diff --git a/README.md b/README.md index 8916969..867c96a 100644 --- a/README.md +++ b/README.md @@ -104,6 +104,7 @@ MCP 客户端需要先执行 `initialize` 握手,拿到 `mcp-session-id` 后 | `list_connections` | 列出所有连接配置 | | `get_connection` | 按名称读取连接配置 | | `delete_connection` | 删除连接配置 | +| `update_connection` | 部分更新一个已存在的连接 (PATCH 语义, 只改提供的字段) | | `write_job_file` | 将 LLM 已生成的 PySpark 代码写入服务端文件 | | `read_job_file` | 读取已存在的 PySpark 脚本内容 | | `update_job_file` | 覆盖更新已存在的 PySpark 脚本 | @@ -120,6 +121,8 @@ MCP 客户端需要先执行 `initialize` 握手,拿到 `mcp-session-id` 后 | `get_external_job_status` | 查询**非本服务提交**的外部 YARN application 状态(按 `application_id` + `connection_name`) | | `get_external_job_result` | 查询外部 YARN application 终态结果视图 | | `get_external_job_logs` | 拉取外部 YARN application 的聚合日志 | +| `list_applications` | 列出 YARN 上所有应用(按 `state` / `queue` / `limit` 过滤),绕过 JobStore | +| `fetch_url` | 代理 HTTP GET 到集群内网 URL (host 受 `Connection.url_allowlist` glob allowlist 约束, 空则全拒) | ### Files MCP 工具 diff --git a/spark_executor/core/connection_store.py b/spark_executor/core/connection_store.py index 6f54bb5..9adbd29 100644 --- a/spark_executor/core/connection_store.py +++ b/spark_executor/core/connection_store.py @@ -77,6 +77,29 @@ class ConnectionStore: self._dump(records) logger.info(f"connection saved name={conn.name} master={conn.master}") + def update(self, name: str, **fields) -> Connection: + """Apply `fields` to the existing Connection identified by `name` and persist. + + PATCH semantics: only fields explicitly passed in `fields` are changed. + Use Pydantic's `model_copy(update=fields)` to apply the patch. + + Raises KeyError if no Connection with `name` exists. + Raises pydantic.ValidationError if the patched Connection is invalid + (e.g. `master='http://...'` fails the master validator). + """ + with self._lock: + records = self._load() + if name not in records: + raise KeyError(f"Connection not found: {name}") + existing = records[name] + patched = Connection.model_validate(existing.model_copy(update=fields).model_dump()) + records[name] = patched + self._dump(records) + logger.info( + f"connection updated name={name} fields={sorted(fields.keys())}" + ) + return patched + def delete(self, name: str) -> bool: with self._lock: records = self._load() diff --git a/spark_executor/core/job_writer.py b/spark_executor/core/job_writer.py index 27df059..796e2a2 100644 --- a/spark_executor/core/job_writer.py +++ b/spark_executor/core/job_writer.py @@ -14,7 +14,7 @@ Resolution order for the output directory: """ import os import secrets -from datetime import datetime +from datetime import datetime, timezone from common.config import settings from common.logging import logger @@ -63,7 +63,7 @@ def write_job_file(code: str, jobs_dir: str | None = None) -> str: effective_dir = resolve_jobs_dir(jobs_dir) os.makedirs(effective_dir, exist_ok=True) - stamp = datetime.utcnow().strftime("%Y%m%d%H%M%S") + stamp = datetime.now(timezone.utc).strftime("%Y%m%d%H%M%S") name = f"job_{stamp}_{secrets.token_hex(3)}.py" path = os.path.join(effective_dir, name) abs_path = os.path.abspath(path) diff --git a/spark_executor/core/pending_store.py b/spark_executor/core/pending_store.py index 5bfccde..0436d54 100644 --- a/spark_executor/core/pending_store.py +++ b/spark_executor/core/pending_store.py @@ -37,7 +37,7 @@ class PendingStore: return Path(self._data_dir) / self._dir_name def _date_path(self, created_at: datetime) -> Path: - # created_at is datetime.utcnow(), so the shard date is a UTC date. + # created_at is timezone-aware UTC, so the shard date is a UTC date. return self.dir_path / f"{created_at.date().isoformat()}.json" def _legacy_path(self) -> Path: diff --git a/spark_executor/core/yarn_client.py b/spark_executor/core/yarn_client.py index 0bb0b56..f86e3f5 100644 --- a/spark_executor/core/yarn_client.py +++ b/spark_executor/core/yarn_client.py @@ -121,13 +121,14 @@ def _base_url(yarn_rm_url: str | None) -> str: def _request(method: str, url: str, *, json_body: dict | None = None, - timeout: float = 30.0, verify: bool | str = True, - auth: httpx.Auth | None = None) -> httpx.Response: + params: dict[str, str] | None = None, timeout: float = 30.0, + verify: bool | str = True, auth: httpx.Auth | None = None) -> httpx.Response: headers = {"Accept": "application/json"} logger.debug(f"YARN {method} {url}" + (f" body={json_body}" if json_body else "")) try: resp = httpx.request( - method, url, json=json_body, headers=headers, timeout=timeout, verify=verify, auth=auth + method, url, json=json_body, params=params, headers=headers, + timeout=timeout, verify=verify, auth=auth ) except httpx.HTTPError as exc: logger.error(f"YARN {method} {url} failed: {exc}") @@ -257,3 +258,51 @@ def kill_application(application_id: str, config: YarnClientConfig) -> None: logger.error(f"YARN PUT {url} -> {resp.status_code}: {resp.text[:500]}") raise YarnError(f"YARN kill returned HTTP {resp.status_code}: {resp.text}") logger.info(f"YARN kill {application_id} -> ok") + + +def list_applications( + config: YarnClientConfig, + *, + state: str | None = None, + queue: str | None = None, + limit: int | None = None, +) -> list[dict]: + """List YARN applications, optionally filtered. + + YARN endpoint: GET /ws/v1/cluster/apps?state=...&queue=...&limit=... + + Filters: + - state: YARN application state. Common values: + "NEW", "NEW_SAVING", "SUBMITTED", "ACCEPTED", "RUNNING", + "FINISHED", "FAILED", "KILLED". + Note: "FINISHED" is the umbrella state covering SUCCEEDED/FAILED/KILLED. + - queue: YARN queue name + - limit: cap on number of returned apps (YARN has no pagination; + callers that need a full enumeration should make multiple + calls with state=... filters or accept the cap) + + Returns a list of YARN app dicts (each with id, name, user, queue, + state, finalStatus, applicationType, startedTime, finishedTime, + trackingUrl, progress, etc). Empty list if no apps match. + + Raises YarnError on transport / 4xx / 5xx. + """ + params: dict[str, str] = {} + if state is not None: + params["state"] = state + if queue is not None: + params["queue"] = queue + if limit is not None: + params["limit"] = str(limit) + + url = f"{_base_url(config.yarn_rm_url)}/ws/v1/cluster/apps" + resp = _request("GET", url, params=params, + verify=config.verify_for_httpx(), + auth=config.auth_for_httpx()) + if resp.status_code >= 400: + raise YarnError( + f"YARN list applications failed: {resp.status_code} {resp.text[:200]}" + ) + data = resp.json() + apps_container = data.get("apps") or {} + return apps_container.get("app", []) or [] diff --git a/spark_executor/models.py b/spark_executor/models.py index 798d72b..9e35af8 100644 --- a/spark_executor/models.py +++ b/spark_executor/models.py @@ -40,6 +40,35 @@ class SubmitResult(BaseModel): tracking_url: str | None = None +class FetchUrlResult(BaseModel): + url: str + status_code: int + content_type: str + body: str + + +class ApplicationSummary(BaseModel): + """A YARN application summary from /ws/v1/cluster/apps. + + Field names are mapped from the YARN JSON keys to clearer + snake_case names by the tool function. Unused YARN fields + (memorySeconds, vcoreSeconds, preemptedResource*, etc.) are + not exposed — the LLM doesn't need them. + """ + application_id: str + name: str + user: str + queue: str + state: str + final_status: str | None = None + application_type: str | None = None + application_tags: str = "" + started_time: int = 0 + finished_time: int = 0 + tracking_url: str | None = None + progress: float | None = None + + class Connection(BaseModel): name: str # Defaults to "yarn" because that's the literal string spark-submit wants @@ -61,6 +90,21 @@ class Connection(BaseModel): auth_principal: str | None = None auth_keytab: str | None = None + url_allowlist: list[str] = Field( + default_factory=list, + description=( + "List of fnmatch glob patterns for hosts the fetch_url tool may access. " + "The list is mandatory-opt-in: an empty list (the default) denies all " + "hosts, so you must populate it before fetch_url can access any URL. " + "Useful for clusters whose hostnames do NOT share a common suffix — " + "e.g. single-label hosts like 'ccam1'-'ccam99' (configure ['ccam*']) " + "or HDFS namenode on a different subdomain ('*.hadoop.internal'). " + "Patterns are matched against the URL host only (no port, no path). " + "fnmatch rules apply: '*' does NOT match '.', so 'ccam*' matches " + "'ccam50' but not 'ccam50.evil.com'." + ), + ) + @field_validator("master") @classmethod def _check_master(cls, v: str) -> str: diff --git a/spark_executor/server.py b/spark_executor/server.py index 3e156c2..4b26128 100644 --- a/spark_executor/server.py +++ b/spark_executor/server.py @@ -11,6 +11,7 @@ from spark_executor.tools.connections import ( get_connection, list_connections, save_connection, + update_connection, ) from spark_executor.tools.write_job import write_job_file from spark_executor.tools.job_file import read_job_file, update_job_file @@ -20,7 +21,9 @@ from spark_executor.tools.external_jobs import ( get_external_job_logs, get_external_job_status, get_external_job_result, + list_applications, ) +from spark_executor.tools.fetch_url import fetch_url from spark_executor.tools.requests import ( ConnectionNameRequest, EmptyRequest, @@ -28,12 +31,15 @@ from spark_executor.tools.requests import ( ExternalJobLogsRequest, ExternalJobStatusRequest, ExternalJobResultRequest, + ListApplicationsRequest, + FetchUrlRequest, GetJobLogsRequest, JobIdRequest, PendingIdRequest, PrepareSubmitJobRequest, ReadJobFileRequest, SaveConnectionRequest, + UpdateConnectionRequest, UpdateJobFileRequest, UpdatePendingJobRequest, ) @@ -342,6 +348,34 @@ def _get_external_job_result(req: ExternalJobResultRequest): return get_external_job_result(req.application_id, req.connection_name) +@app.post( + "/list_applications", + operation_id="list_applications", + summary="List YARN applications on a cluster, optionally filtered", + description=( + "Query YARN's /ws/v1/cluster/apps endpoint through the named " + "Connection, returning a list of ApplicationSummary records. " + "Bypasses the local JobStore — useful for enumerating apps that " + "were not submitted through this service.\n\n" + "**Filters:** state (YARN state, e.g. 'RUNNING', 'FINISHED', " + "'FAILED'), queue (YARN queue name), limit (default 100, max ~10000). " + "YARN has no offset-based pagination, so for large clusters combine " + "state/queue filters to scope the result. The `FINISHED` state " + "covers SUCCEEDED/FAILED/KILLED.\n\n" + "Returns an empty list if no apps match. The Connection's auth_type " + "/ auth_user / auth_password / ssl_verify / ssl_ca_bundle are reused " + "for the request." + ), +) +def _list_applications(req: ListApplicationsRequest): + return list_applications( + req.connection_name, + state=req.state, + queue=req.queue, + limit=req.limit, + ) + + # --- Connection management tools --- @app.post( @@ -356,12 +390,20 @@ def _get_external_job_result(req: ExternalJobResultRequest): "yarn_rm_url is required for get_external_* to work** — spark-submit " "can discover the RM for submissions, but direct YARN REST queries " "need an explicit URL. Saving with an existing name overwrites the " - "record in place (no version history)." + "record in place (no version history). When the name already exists, " + "only the provided fields are changed (PATCH semantics); omitted " + "fields keep their previous values." ), ) def _save_connection(req: SaveConnectionRequest): - # exclude_none so we don't overwrite the function's default with explicit None - return save_connection(**req.model_dump(exclude_none=True)) + fields = req.model_dump(exclude_none=True) + name = fields.pop("name") + try: + get_connection(name) + except KeyError: + # exclude_none so we don't overwrite the function's default with explicit None + return save_connection(name=name, **fields) + return update_connection(name=name, **fields) @app.post( @@ -394,6 +436,28 @@ def _get_connection(req: ConnectionNameRequest): return get_connection(req.name) +@app.post( + "/update_connection", + operation_id="update_connection", + summary="Update an existing connection's fields", + description=( + "Apply a partial update (PATCH) to an existing Connection record. " + "Only the fields you provide are changed; the rest are kept as-is. " + "The `name` is the immutable identifier (use delete_connection + " + "save_connection to rename).\n\n" + "To CLEAR an optional field (e.g. remove `yarn_rm_url`), use " + "delete_connection followed by save_connection with the field omitted. " + "This tool cannot clear fields — only replace them.\n\n" + "Returns the full updated Connection record. 404 if no Connection " + "with the given name exists." + ), +) +def _update_connection(req: UpdateConnectionRequest): + fields = req.model_dump(exclude_none=True) + fields.pop("name", None) # name is the identity, not a field to patch + return update_connection(name=req.name, **fields) + + @app.post( "/delete_connection", operation_id="delete_connection", @@ -467,3 +531,26 @@ def _read_job_file(req: ReadJobFileRequest): ) def _update_job_file(req: UpdateJobFileRequest): return update_job_file(req.script_path, req.content) + +# --- HTTP fetch proxy (host allowlist via Connection.yarn_rm_url) --- + +@app.post( + "/fetch_url", + operation_id="fetch_url", + summary="Fetch a URL on the cluster's network and return the body", + description=( + "Proxy an HTTP GET to a URL on the cluster's network, returning the " + "response body. Useful when the agent is on a different network from " + "the cluster and cannot reach YARN tracking pages, Spark History " + "Server, or NodeManager web UIs directly.\n\n" + "**Security constraints:** the URL host must match one of the fnmatch " + "glob patterns in the named Connection's url_allowlist. An empty or " + "omitted allowlist denies every host. There are no scheme or IP-literal " + "guardrails — the allowlist is the only gate — so keep it tight. The " + "Connection's saved auth is reused, so the agent does not need cluster " + "credentials.\n\n" + "30s timeout, redirects followed." + ), +) +def _fetch_url(req: FetchUrlRequest): + return fetch_url(req.url, req.connection_name) diff --git a/spark_executor/tools/connections.py b/spark_executor/tools/connections.py index ae381f8..4c631a7 100644 --- a/spark_executor/tools/connections.py +++ b/spark_executor/tools/connections.py @@ -4,47 +4,94 @@ @Author :tao.chen """ from common.logging import logger -from spark_executor.core.connection_store import ConnectionStore, store +from spark_executor.core.connection_store import store from spark_executor.models import Connection +_UNSET = object() + + def save_connection( *, name: str, master: str, - deploy_mode: str = "cluster", - yarn_rm_url: str | None = None, - spark_conf: dict[str, str] | None = None, - ssl_verify: bool | None = None, - ssl_ca_bundle: str | None = None, - auth_type: str = "none", - auth_user: str | None = None, - auth_password: str | None = None, - auth_principal: str | None = None, - auth_keytab: str | None = None, -) -> dict[str, str]: + deploy_mode: str = _UNSET, # type: ignore[assignment] + yarn_rm_url: str | None = _UNSET, # type: ignore[assignment] + spark_conf: dict[str, str] | None = _UNSET, # type: ignore[assignment] + ssl_verify: bool | None = _UNSET, # type: ignore[assignment] + ssl_ca_bundle: str | None = _UNSET, # type: ignore[assignment] + auth_type: str = _UNSET, # type: ignore[assignment] + auth_user: str | None = _UNSET, # type: ignore[assignment] + auth_password: str | None = _UNSET, # type: ignore[assignment] + auth_principal: str | None = _UNSET, # type: ignore[assignment] + auth_keytab: str | None = _UNSET, # type: ignore[assignment] + url_allowlist: list[str] | None = _UNSET, # type: ignore[assignment] +) -> dict[str, object]: logger.debug( - f"save_connection enter name={name} master={master} deploy_mode={deploy_mode} " - f"yarn_rm_url={yarn_rm_url} spark_conf_keys={list((spark_conf or {}).keys())}" - ) - conn = Connection( - name=name, - master=master, - deploy_mode=deploy_mode, - yarn_rm_url=yarn_rm_url, - spark_conf=spark_conf or {}, - ssl_verify=ssl_verify, - ssl_ca_bundle=ssl_ca_bundle, - auth_type=auth_type, - auth_user=auth_user, - auth_password=auth_password, - auth_principal=auth_principal, - auth_keytab=auth_keytab, + f"save_connection enter name={name} master={master} " + f"spark_conf_keys={list((spark_conf if isinstance(spark_conf, dict) else {}).keys())}" ) + existing = store.get(name) + if existing is not None: + fields: dict[str, object] = {"master": master} + if deploy_mode is not _UNSET: + fields["deploy_mode"] = deploy_mode + if yarn_rm_url is not _UNSET: + fields["yarn_rm_url"] = yarn_rm_url + if spark_conf is not _UNSET: + fields["spark_conf"] = spark_conf or {} + if ssl_verify is not _UNSET: + fields["ssl_verify"] = ssl_verify + if ssl_ca_bundle is not _UNSET: + fields["ssl_ca_bundle"] = ssl_ca_bundle + if auth_type is not _UNSET: + fields["auth_type"] = auth_type + if auth_user is not _UNSET: + fields["auth_user"] = auth_user + if auth_password is not _UNSET: + fields["auth_password"] = auth_password + if auth_principal is not _UNSET: + fields["auth_principal"] = auth_principal + if auth_keytab is not _UNSET: + fields["auth_keytab"] = auth_keytab + if url_allowlist is not _UNSET: + fields["url_allowlist"] = url_allowlist or [] + return update_connection(name, **fields) + new_fields: dict[str, object] = { + "name": name, + "master": master, + "deploy_mode": deploy_mode if deploy_mode is not _UNSET else "cluster", + "yarn_rm_url": yarn_rm_url if yarn_rm_url is not _UNSET else None, + "spark_conf": (spark_conf if spark_conf is not _UNSET else None) or {}, + "ssl_verify": ssl_verify if ssl_verify is not _UNSET else None, + "ssl_ca_bundle": ssl_ca_bundle if ssl_ca_bundle is not _UNSET else None, + "auth_type": auth_type if auth_type is not _UNSET else "none", + "auth_user": auth_user if auth_user is not _UNSET else None, + "auth_password": auth_password if auth_password is not _UNSET else None, + "auth_principal": auth_principal if auth_principal is not _UNSET else None, + "auth_keytab": auth_keytab if auth_keytab is not _UNSET else None, + "url_allowlist": (url_allowlist if url_allowlist is not _UNSET else None) or [], + } + conn = Connection(**new_fields) store.save(conn) return {"name": name, "status": "SAVED"} +def update_connection(name: str, **fields) -> dict[str, object]: + """Update an existing Connection's mutable fields. + + `name` is the identifier (immutable). PATCH semantics: only the fields + you pass are changed. To clear an optional field (e.g. `yarn_rm_url`), + use `delete_connection(name=...)` followed by `save_connection(...)`. + + Mutable fields: master, deploy_mode, yarn_rm_url, spark_conf, + ssl_verify, ssl_ca_bundle, auth_type, auth_user, auth_password, + auth_principal, auth_keytab, url_allowlist. + """ + logger.debug(f"update_connection enter name={name} fields={sorted(fields.keys())}") + return store.update(name, **fields).model_dump() + + def list_connections() -> list[dict[str, object]]: logger.debug("list_connections enter") return [c.model_dump() for c in store.list_all()] diff --git a/spark_executor/tools/external_jobs.py b/spark_executor/tools/external_jobs.py index eb60329..1527317 100644 --- a/spark_executor/tools/external_jobs.py +++ b/spark_executor/tools/external_jobs.py @@ -10,8 +10,13 @@ YARN application_id and the name of a saved Connection. import json from common.logging import logger -from spark_executor.core.yarn_client import YarnClientConfig, get_application_status, get_application_logs -from spark_executor.models import JobStatus, JobResult +from spark_executor.core.yarn_client import ( + YarnClientConfig, + get_application_logs, + get_application_status, + list_applications as list_applications_yarn, # alias to avoid collision +) +from spark_executor.models import ApplicationSummary, JobResult, JobStatus from spark_executor.tools.connections import store as conn_store @@ -63,3 +68,53 @@ def get_external_job_result(application_id: str, connection_name: str) -> JobRes ) logger.info(f"get_external_job_result ok application_id={application_id} connection_name={connection_name} state={state}") return result + + +def list_applications( + connection_name: str, + state: str | None = None, + queue: str | None = None, + limit: int = 100, +) -> list[ApplicationSummary]: + """List YARN applications on the named cluster, optionally filtered. + + Bypasses the local JobStore (this is for apps not submitted through + this service). The YARN ResourceManager REST endpoint + /ws/v1/cluster/apps is queried through the connection's auth/SSL + config. + + Defaults: limit=100 (YARN has no offset-based pagination, so large + clusters should use state/queue filters to scope the result). + """ + logger.debug( + f"list_applications enter connection_name={connection_name} " + f"state={state} queue={queue} limit={limit}" + ) + conn = conn_store.get(connection_name) + if conn is None: + raise KeyError(f"Connection not found: {connection_name}") + config = YarnClientConfig.from_connection(conn) + raw_apps = list_applications_yarn( + config, state=state, queue=queue, limit=limit + ) + summaries = [ + ApplicationSummary( + application_id=app.get("id", ""), + name=app.get("name", ""), + user=app.get("user", ""), + queue=app.get("queue", ""), + state=app.get("state", ""), + final_status=app.get("finalStatus"), + application_type=app.get("applicationType"), + application_tags=app.get("applicationTags", ""), + started_time=app.get("startedTime", 0), + finished_time=app.get("finishedTime", 0), + tracking_url=app.get("trackingUrl"), + progress=app.get("progress"), + ) + for app in raw_apps + ] + logger.info( + f"list_applications ok connection_name={connection_name} count={len(summaries)}" + ) + return summaries diff --git a/spark_executor/tools/fetch_url.py b/spark_executor/tools/fetch_url.py new file mode 100644 index 0000000..9615372 --- /dev/null +++ b/spark_executor/tools/fetch_url.py @@ -0,0 +1,122 @@ +# coding=utf-8 +""" +@Time :2026/7/9 +@Author :tao.chen + +Generic HTTP GET proxy for the agent. Lets the agent fetch URLs on the +cluster's network when it cannot reach those hosts directly. Security: +the host is checked against an explicit fnmatch glob allowlist configured +on the connection (`url_allowlist`). Empty or missing allowlist means no +URL access; the allowlist is the only gate — scheme and IP-literal checks +are intentionally NOT performed. Reuses the connection's saved auth/SSL +config so the agent doesn't need cluster credentials. +""" +import fnmatch +from urllib.parse import urlparse, urljoin + +import httpx + +from common.logging import logger +from spark_executor.core.yarn_client import YarnClientConfig +from spark_executor.models import FetchUrlResult +from spark_executor.tools.connections import store as conn_store + +_REQUEST_TIMEOUT_SECONDS = 30 + +_MAX_REDIRECTS = 3 +_REDIRECT_STATUSES = frozenset({301, 302, 303, 307, 308}) + + +def _host_matches_any_glob(host: str, allowlist: list[str]) -> bool: + """True if `host` matches any of the fnmatch glob patterns. + + fnmatch is case-sensitive on Linux (our deployment target). `*` in a + pattern does NOT cross '.' boundaries, so 'ccam*' matches 'ccam50' but + NOT 'ccam50.evil.com' (which is what we want — domain boundary matters + for SSRF). + """ + host_labels = host.split(".") + for p in allowlist: + pat_labels = p.split(".") + if len(pat_labels) != len(host_labels): + continue + if all(fnmatch.fnmatchcase(h, pat) for h, pat in zip(host_labels, pat_labels)): + return True + return False + + +def _validate_url_host(url: str, allowlist: list[str] | None) -> None: + """Reject the URL unless its host matches a glob in `allowlist`. + + The only check. The url_allowlist is the single source of truth for + what fetch_url is allowed to access — no scheme, IP-literal, or + "must be the same as yarn_rm_url" guardrails. The caller is + responsible for writing a tight allowlist. + + The only structural check: the URL must have a host (otherwise + the glob match has nothing to test). Anything else is the + allowlist's job. + """ + host = urlparse(url).hostname + if not host: + raise ValueError(f"URL has no host: {url!r}") + if _host_matches_any_glob(host, allowlist or []): + return + raise ValueError( + f"URL host {host!r} is not in Connection.url_allowlist {allowlist!r}. " + f"Add the host pattern to url_allowlist (or use a broader glob) " + f"and try again." + ) + + +def fetch_url(url: str, connection_name: str) -> FetchUrlResult: + """Proxy an HTTP GET to url using the auth/SSL settings of connection_name.""" + logger.debug(f"fetch_url enter url={url} connection_name={connection_name}") + conn = conn_store.get(connection_name) + if conn is None: + raise KeyError(f"Connection not found: {connection_name}") + + _validate_url_host(url, conn.url_allowlist) + + config = YarnClientConfig.from_connection(conn) + auth = config.auth_for_httpx() + verify = config.verify_for_httpx() + + for hop in range(_MAX_REDIRECTS + 1): + resp = httpx.get( + url, + auth=auth, + verify=verify, + timeout=_REQUEST_TIMEOUT_SECONDS, + follow_redirects=False, + ) + if resp.status_code not in _REDIRECT_STATUSES: + logger.info( + f"fetch_url ok url={url} connection_name={connection_name} " + f"status_code={resp.status_code} " + f"content_type={resp.headers.get('content-type', '')} " + f"body_chars={len(resp.text)}" + ) + + return FetchUrlResult( + url=url, + status_code=resp.status_code, + content_type=resp.headers.get("content-type", ""), + body=resp.text, + ) + + location = resp.headers.get("Location") or resp.headers.get("location") + if not location: + raise ValueError( + f"HTTP GET returned {resp.status_code} with no Location header at {url}" + ) + next_url = urljoin(url, location) + logger.debug( + f"fetch_url redirect url={url} status={resp.status_code} to={next_url}" + ) + _validate_url_host(next_url, conn.url_allowlist) + url = next_url + + raise ValueError( + f"fetch_url exceeded {_MAX_REDIRECTS} redirects (last url: {url})" + ) diff --git a/spark_executor/tools/requests.py b/spark_executor/tools/requests.py index d9fbee7..9cb1f0f 100644 --- a/spark_executor/tools/requests.py +++ b/spark_executor/tools/requests.py @@ -122,6 +122,18 @@ class SaveConnectionRequest(BaseModel): ), ) + url_allowlist: list[str] | None = Field( + default=None, + description=( + "Optional list of fnmatch glob patterns for hosts the fetch_url tool " + "may access. See Connection.url_allowlist for full semantics. " + "Example for single-label host clusters: ['ccam*'] allows any host " + "starting with 'ccam' (ccam1, ccam2, ..., ccam99). If omitted, None, " + "or empty, the saved connection will have url_allowlist=[] (the " + "default), meaning fetch_url will reject every URL until the list is " + "populated via update_connection." + ), + ) class PrepareSubmitJobRequest(BaseModel): connection: str = Field( @@ -378,3 +390,126 @@ class ExternalJobResultRequest(BaseModel): ..., description="Name of a saved Connection pointing at the YARN cluster.", ) + +class FetchUrlRequest(BaseModel): + url: str = Field( + ..., + description=( + "Absolute URL to fetch. The only access control is the named " + "Connection's url_allowlist: the URL host must match one of the " + "fnmatch glob patterns in that list. An empty or omitted allowlist " + "denies every host. IP literals and non-HTTP schemes are allowed " + "if and only if they are matched by the allowlist. The Connection's " + "saved auth is reused for the outbound request — the agent does not " + "need cluster credentials." + ), + ) + connection_name: str = Field( + ..., + description=( + "Name of a saved Connection (see list_connections). The " + "Connection's yarn_rm_url defines the allowed host domain. " + "The Connection's auth_type / auth_user / auth_password / " + "ssl_verify / ssl_ca_bundle are reused for the request." + ), + ) + + +class UpdateConnectionRequest(BaseModel): + name: str = Field( + ..., + description=( + "Name of the existing Connection to update (immutable identifier). " + "If you want to rename, use delete_connection + save_connection." + ), + ) + master: str | None = Field( + default=None, + description="New Spark master URL. Omit to keep current.", + ) + deploy_mode: str | None = Field( + default=None, + description="New Spark deploy mode ('cluster' or 'client'). Omit to keep current.", + ) + yarn_rm_url: str | None = Field( + default=None, + description="New YARN ResourceManager REST base URL. Omit to keep current.", + ) + spark_conf: dict[str, str] | None = Field( + default=None, + description="Replacement Spark conf dict (not merged with existing). Omit to keep current.", + ) + ssl_verify: bool | None = Field( + default=None, + description="New SSL verify setting. Omit to keep current.", + ) + ssl_ca_bundle: str | None = Field( + default=None, + description="New SSL CA bundle path. Omit to keep current.", + ) + auth_type: str | None = Field( + default=None, + description="New auth mode. Omit to keep current.", + ) + auth_user: str | None = Field( + default=None, + description="New auth username. Omit to keep current.", + ) + auth_password: str | None = Field( + default=None, + description="New auth password. Omit to keep current.", + ) + auth_principal: str | None = Field( + default=None, + description="New Kerberos principal. Omit to keep current.", + ) + auth_keytab: str | None = Field( + default=None, + description="New Kerberos keytab path. Omit to keep current.", + ) + url_allowlist: list[str] | None = Field( + default=None, + description=( + "Replacement url_allowlist list (not merged). Omit to keep current. " + "Pass an empty list to deny all hosts (the default for new connections). " + "Example: ['ccam*'] allows ccam1-ccam99." + ), + ) + + +class ListApplicationsRequest(BaseModel): + connection_name: str = Field( + ..., + description=( + "Name of a saved Connection (see list_connections) pointing at " + "the YARN cluster to query." + ), + ) + state: str | None = Field( + default=None, + description=( + "Optional YARN application state filter. One of: 'NEW', " + "'NEW_SAVING', 'SUBMITTED', 'ACCEPTED', 'RUNNING', 'FINISHED', " + "'FAILED', 'KILLED'. 'FINISHED' is the umbrella state covering " + "SUCCEEDED/FAILED/KILLED. None = no state filter (returns all " + "states up to `limit`)." + ), + ) + queue: str | None = Field( + default=None, + description=( + "Optional YARN queue name filter (e.g. 'default', 'prod'). " + "None = no queue filter." + ), + ) + limit: int = Field( + default=100, + ge=1, + le=10000, + description=( + "Maximum number of applications to return. YARN has no " + "offset-based pagination, so for large clusters use state/queue " + "filters to scope the result. Max 10000 in practice (YARN's own " + "limit on the limit param)." + ), + ) diff --git a/spark_executor/tools/submit.py b/spark_executor/tools/submit.py index 10ff531..04c2bda 100644 --- a/spark_executor/tools/submit.py +++ b/spark_executor/tools/submit.py @@ -6,9 +6,8 @@ import os import secrets import uuid -from datetime import datetime +from datetime import datetime, timezone -from common.config import settings from common.logging import logger from common.sql_guard import validate_pyspark_code from spark_executor.core.connection_store import store as conn_store @@ -124,7 +123,7 @@ def prepare_submit_job( num_executors=num_executors, spark_conf=dict(conn.spark_conf), extra_args=dict(extra_args or {}), - created_at=datetime.utcnow(), + created_at=datetime.now(timezone.utc), status="PENDING", ) pending_store.save(pending) @@ -248,7 +247,7 @@ def confirm_submit_job(*, pending_id: str) -> SubmitResult: application_id=application_id, script_path=pending.script_path, queue=pending.queue, - submit_time=datetime.utcnow(), + submit_time=datetime.now(timezone.utc), connection=pending.connection, yarn_rm_url=pending.yarn_rm_url, ) @@ -282,7 +281,7 @@ def confirm_submit_job(*, pending_id: str) -> SubmitResult: application_id=application_id, script_path=pending.script_path, queue=pending.queue, - submit_time=datetime.utcnow(), + submit_time=datetime.now(timezone.utc), connection=pending.connection, yarn_rm_url=pending.yarn_rm_url, ) diff --git a/tests/integration/test_mcp_routes.py b/tests/integration/test_mcp_routes.py index 3b555b8..df69baf 100644 --- a/tests/integration/test_mcp_routes.py +++ b/tests/integration/test_mcp_routes.py @@ -64,12 +64,15 @@ def test_seventeen_tool_routes_registered(): assert "/update_pending_job" in paths -def test_twenty_tool_routes_registered(): +def test_twenty_three_tool_routes_registered(): paths = {r.path for r in app.routes} for path in ( "/get_external_job_logs", "/get_external_job_status", "/get_external_job_result", + "/list_applications", + "/fetch_url", + "/update_connection", ): assert path in paths, f"missing MCP tool route: {path}" diff --git a/tests/unit/test_connection_tools.py b/tests/unit/test_connection_tools.py index 3b73af2..dad74c7 100644 --- a/tests/unit/test_connection_tools.py +++ b/tests/unit/test_connection_tools.py @@ -1,4 +1,5 @@ # coding=utf-8 +import json from pathlib import Path import pytest @@ -120,3 +121,77 @@ def test_delete_connection_returns_status(_fresh_store): def test_delete_connection_unknown_raises(_fresh_store): with pytest.raises(KeyError): connections.delete_connection("missing") + + +def test_save_connection_preserves_url_allowlist_on_existing_record(_fresh_store): + connections.save_connection(name="prod", master="yarn", url_allowlist=["ccam*"]) + out = connections.save_connection(name="prod", master="spark://new:7077") + assert out["master"] == "spark://new:7077" + assert out["url_allowlist"] == ["ccam*"] + + +def test_update_connection_changes_specified_field(_fresh_store): + connections.save_connection(name="prod", master="yarn") + out = connections.update_connection(name="prod", master="spark://new:7077") + assert out["master"] == "spark://new:7077" + assert _fresh_store.get("prod").master == "spark://new:7077" + + +def test_update_connection_keeps_omitted_fields(_fresh_store): + connections.save_connection( + name="prod", + master="yarn", + deploy_mode="client", + yarn_rm_url="http://rm:8088", + ) + out = connections.update_connection(name="prod", deploy_mode="cluster") + assert out["deploy_mode"] == "cluster" + assert out["master"] == "yarn" + assert out["yarn_rm_url"] == "http://rm:8088" + + +def test_update_connection_with_no_fields_is_noop(_fresh_store): + connections.save_connection(name="prod", master="yarn", deploy_mode="client") + out = connections.update_connection(name="prod") + assert out["master"] == "yarn" + assert out["deploy_mode"] == "client" + + +def test_update_connection_changes_url_allowlist(_fresh_store): + connections.save_connection(name="prod", master="yarn") + out = connections.update_connection(name="prod", url_allowlist=["ccam*"]) + assert out["url_allowlist"] == ["ccam*"] + + +def test_update_connection_replaces_not_merges_dict(_fresh_store): + connections.save_connection(name="prod", master="yarn", spark_conf={"a": "1"}) + out = connections.update_connection(name="prod", spark_conf={"b": "2"}) + assert out["spark_conf"] == {"b": "2"} + + +def test_update_connection_replaces_not_merges_url_allowlist(_fresh_store): + connections.save_connection( + name="prod", + master="yarn", + url_allowlist=["ccam*"], + ) + out = connections.update_connection(name="prod", url_allowlist=["nm*"]) + assert out["url_allowlist"] == ["nm*"] + + +def test_update_connection_raises_for_unknown_name(_fresh_store): + with pytest.raises(KeyError, match="Connection not found"): + connections.update_connection(name="missing", master="yarn") + + +def test_update_connection_validates_patched_connection(_fresh_store): + connections.save_connection(name="prod", master="yarn") + with pytest.raises(ValueError, match="master must be"): + connections.update_connection(name="prod", master="http://bad") + + +def test_update_connection_persists_to_disk(_fresh_store, tmp_path): + connections.save_connection(name="prod", master="yarn") + connections.update_connection(name="prod", master="spark://new:7077") + raw = json.loads((tmp_path / "connections.json").read_text()) + assert raw["prod"]["master"] == "spark://new:7077" diff --git a/tests/unit/test_external_jobs.py b/tests/unit/test_external_jobs.py index dc67df5..7b22e21 100644 --- a/tests/unit/test_external_jobs.py +++ b/tests/unit/test_external_jobs.py @@ -5,8 +5,10 @@ from unittest.mock import patch import pytest from spark_executor.core import connection_store -from spark_executor.models import Connection +from spark_executor.core.yarn_client import YarnError +from spark_executor.models import ApplicationSummary, Connection from spark_executor.tools import connections, external_jobs +from spark_executor.tools.requests import ListApplicationsRequest def _fresh_stores(): @@ -17,6 +19,13 @@ def _fresh_stores(): external_jobs.conn_store = store +@pytest.fixture +def fresh_stores(tmp_path, monkeypatch): + """Reset connection store singletons to an isolated tmp_path.""" + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + _fresh_stores() + + def test_get_external_job_logs_returns_tailed(): _fresh_stores() external_jobs.conn_store.save( @@ -131,3 +140,166 @@ def test_get_external_job_result_raises_when_connection_missing(): external_jobs.get_external_job_result( application_id="application_1", connection_name="missing" ) + + +def test_list_applications_returns_summaries(fresh_stores): + external_jobs.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088") + ) + with patch( + "spark_executor.tools.external_jobs.list_applications_yarn", + return_value=[ + { + "id": "application_1", + "name": "app-one", + "user": "alice", + "queue": "default", + "state": "RUNNING", + "finalStatus": "UNDEFINED", + "applicationType": "SPARK", + "applicationTags": "tag1", + "startedTime": 1000, + "finishedTime": 0, + "trackingUrl": "http://rm:8088/proxy/application_1", + "progress": 75.0, + }, + { + "id": "application_2", + "name": "app-two", + "user": "bob", + "queue": "research", + "state": "FINISHED", + "finalStatus": "SUCCEEDED", + "applicationType": "SPARK", + "applicationTags": "", + "startedTime": 2000, + "finishedTime": 3000, + "trackingUrl": "http://rm:8088/proxy/application_2", + "progress": 100.0, + }, + ], + ) as m: + out = external_jobs.list_applications("prod") + assert len(out) == 2 + assert all(isinstance(item, ApplicationSummary) for item in out) + assert out[0].application_id == "application_1" + assert out[1].application_id == "application_2" + args = m.call_args.args + assert args[0].yarn_rm_url == "http://rm:8088" + + +def test_list_applications_passes_state_filter(fresh_stores): + external_jobs.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088") + ) + with patch( + "spark_executor.tools.external_jobs.list_applications_yarn", + return_value=[], + ) as m: + external_jobs.list_applications("prod", state="RUNNING") + assert m.call_args.kwargs == {"state": "RUNNING", "queue": None, "limit": 100} + + +def test_list_applications_passes_queue_filter(fresh_stores): + external_jobs.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088") + ) + with patch( + "spark_executor.tools.external_jobs.list_applications_yarn", + return_value=[], + ) as m: + external_jobs.list_applications("prod", queue="research") + assert m.call_args.kwargs == {"state": None, "queue": "research", "limit": 100} + + +def test_list_applications_passes_limit_filter(fresh_stores): + external_jobs.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088") + ) + with patch( + "spark_executor.tools.external_jobs.list_applications_yarn", + return_value=[], + ) as m: + external_jobs.list_applications("prod", limit=50) + assert m.call_args.kwargs == {"state": None, "queue": None, "limit": 50} + + +def test_list_applications_default_limit_is_100(fresh_stores): + external_jobs.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088") + ) + with patch( + "spark_executor.tools.external_jobs.list_applications_yarn", + return_value=[], + ) as m: + external_jobs.list_applications("prod") + assert m.call_args.kwargs == {"state": None, "queue": None, "limit": 100} + assert ListApplicationsRequest(connection_name="prod").limit == 100 + + +def test_list_applications_returns_empty_list_when_no_apps(fresh_stores): + external_jobs.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088") + ) + with patch( + "spark_executor.tools.external_jobs.list_applications_yarn", + return_value=[], + ): + out = external_jobs.list_applications("prod") + assert out == [] + + +def test_list_applications_raises_for_missing_connection(fresh_stores): + with pytest.raises(KeyError, match="Connection not found"): + external_jobs.list_applications("missing") + + +def test_list_applications_maps_yarn_json_to_summary(fresh_stores): + external_jobs.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088") + ) + yarn_app = { + "id": "application_42", + "name": "mapped-app", + "user": "carol", + "queue": "prod", + "state": "ACCEPTED", + "finalStatus": "UNDEFINED", + "applicationType": "MAPREDUCE", + "applicationTags": "batch", + "startedTime": 12345, + "finishedTime": 0, + "trackingUrl": "http://rm:8088/proxy/application_42", + "progress": 12.5, + } + with patch( + "spark_executor.tools.external_jobs.list_applications_yarn", + return_value=[yarn_app], + ): + out = external_jobs.list_applications("prod") + assert len(out) == 1 + summary = out[0] + assert summary.application_id == "application_42" + assert summary.name == "mapped-app" + assert summary.user == "carol" + assert summary.queue == "prod" + assert summary.state == "ACCEPTED" + assert summary.final_status == "UNDEFINED" + assert summary.application_type == "MAPREDUCE" + assert summary.application_tags == "batch" + assert summary.started_time == 12345 + assert summary.finished_time == 0 + assert summary.tracking_url == "http://rm:8088/proxy/application_42" + assert summary.progress == 12.5 + + +def test_list_applications_raises_on_404(fresh_stores): + external_jobs.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm:8088") + ) + with patch( + "spark_executor.tools.external_jobs.list_applications_yarn", + side_effect=YarnError("YARN list applications failed: 404"), + ): + with pytest.raises(YarnError, match="YARN list applications failed"): + external_jobs.list_applications("prod") diff --git a/tests/unit/test_fetch_url.py b/tests/unit/test_fetch_url.py new file mode 100644 index 0000000..affe030 --- /dev/null +++ b/tests/unit/test_fetch_url.py @@ -0,0 +1,329 @@ +# coding=utf-8 +from pathlib import Path +from unittest.mock import patch + +import httpx +import pytest +from pydantic import ValidationError + +from spark_executor.core import connection_store +from spark_executor.core.connection_store import ConnectionStore +from spark_executor.models import Connection +from spark_executor.tools import connections, fetch_url +from spark_executor.tools.requests import ListApplicationsRequest + + +def _fresh_stores(tmp_path, monkeypatch): + """Reset connection store singletons for a single test.""" + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + store = ConnectionStore() + monkeypatch.setattr(connection_store, "store", store) + connections.store = store + fetch_url.conn_store = store + + +@pytest.fixture +def fresh_stores(tmp_path: Path, monkeypatch): + """Reset connection store singletons for a single test.""" + _fresh_stores(tmp_path, monkeypatch) + yield + + +def test_fetch_url_returns_body_and_status(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="prod", + master="yarn", + yarn_rm_url="http://rm.prod.internal:8088", + url_allowlist=["*.prod.internal"], + ) + ) + resp = httpx.Response(200, text="hello", headers={"content-type": "text/html"}) + with patch( + "spark_executor.tools.fetch_url.httpx.get", return_value=resp + ) as m: + out = fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "prod") + assert out.url == "http://nm01.prod.internal:8042/node" + assert out.status_code == 200 + assert out.content_type == "text/html" + assert out.body == "hello" + assert "truncated" not in fetch_url.FetchUrlResult.model_fields + assert m.call_count == 1 + + +def test_fetch_url_raises_for_missing_connection(fresh_stores): + with pytest.raises(KeyError, match="Connection not found"): + fetch_url.fetch_url("http://rm.prod.internal:8088/", "missing") + + +def test_fetch_url_passes_auth_from_connection(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="auth", + master="yarn", + yarn_rm_url="http://rm.prod.internal:8088", + auth_type="basic", + auth_user="u", + auth_password="p", + url_allowlist=["*.prod.internal"], + ) + ) + resp = httpx.Response(200, text="ok") + with patch( + "spark_executor.tools.fetch_url.httpx.get", return_value=resp + ) as m: + fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "auth") + auth = m.call_args.kwargs["auth"] + assert isinstance(auth, httpx.BasicAuth) + import base64 + + creds = base64.b64decode(auth._auth_header.split()[1]).decode() + assert creds == "u:p" + + +def test_fetch_url_passes_ssl_verify_from_connection(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="insecure", + master="yarn", + yarn_rm_url="http://rm.prod.internal:8088", + ssl_verify=False, + url_allowlist=["*.prod.internal"], + ) + ) + resp = httpx.Response(200, text="ok") + with patch( + "spark_executor.tools.fetch_url.httpx.get", return_value=resp + ) as m: + fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "insecure") + assert m.call_args.kwargs["verify"] is False + + +def test_fetch_url_allows_host_matching_glob_pattern(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="ccam", + master="yarn", + yarn_rm_url="http://ccam1:8088", + url_allowlist=["ccam*"], + ) + ) + resp = httpx.Response(200, text="hello") + with patch( + "spark_executor.tools.fetch_url.httpx.get", return_value=resp + ) as m: + out = fetch_url.fetch_url("http://ccam50:8088/foo", "ccam") + assert out.status_code == 200 + assert out.body == "hello" + assert m.call_count == 1 + + +def test_fetch_url_allows_host_matching_any_of_multiple_globs(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="prod", + master="yarn", + yarn_rm_url="http://rm.prod.internal:8088", + url_allowlist=["ccam*", "*.prod.internal"], + ) + ) + resp = httpx.Response(200, text="hello") + with patch( + "spark_executor.tools.fetch_url.httpx.get", return_value=resp + ) as m: + out = fetch_url.fetch_url( + "http://history.prod.internal:18080/api/v1/info", "prod" + ) + assert out.status_code == 200 + assert out.body == "hello" + assert m.call_count == 1 + + +def test_fetch_url_glob_does_not_match_unrelated_host(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="ccam", + master="yarn", + yarn_rm_url="http://ccam1:8088", + url_allowlist=["ccam*"], + ) + ) + with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): + fetch_url.fetch_url("http://evil.com/foo", "ccam") + + +def test_fetch_url_glob_does_not_cross_dot_boundary(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="ccam", + master="yarn", + yarn_rm_url="http://ccam1:8088", + url_allowlist=["ccam*"], + ) + ) + with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): + fetch_url.fetch_url("http://ccam50.evil.com/", "ccam") + + +def test_fetch_url_glob_match_does_not_require_suffix_overlap(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="prod", + master="yarn", + yarn_rm_url="http://rm:8088", + url_allowlist=["ccam*"], + ) + ) + resp = httpx.Response(200, text="hello") + with patch( + "spark_executor.tools.fetch_url.httpx.get", return_value=resp + ) as m: + out = fetch_url.fetch_url("http://ccam50:8088/", "prod") + assert out.status_code == 200 + assert out.body == "hello" + assert m.call_count == 1 + + +def test_fetch_url_accepts_ip_literal_when_in_url_allowlist(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="prod", + master="yarn", + url_allowlist=["*.*.*.*"], + ) + ) + resp = httpx.Response(200, text="hello") + with patch( + "spark_executor.tools.fetch_url.httpx.get", return_value=resp + ) as m: + out = fetch_url.fetch_url("http://10.0.0.1/secret", "prod") + assert out.status_code == 200 + assert out.body == "hello" + assert m.call_count == 1 + + +def test_fetch_url_accepts_https_when_in_url_allowlist(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="prod", + master="yarn", + url_allowlist=["ccam*.example.com"], + ) + ) + resp = httpx.Response(200, text="hello") + with patch( + "spark_executor.tools.fetch_url.httpx.get", return_value=resp + ) as m: + out = fetch_url.fetch_url("https://ccam1.example.com/secure", "prod") + assert out.status_code == 200 + assert out.body == "hello" + assert m.call_count == 1 + + +def test_fetch_url_rejects_when_url_has_no_host(fresh_stores): + with pytest.raises(ValueError, match="URL has no host"): + fetch_url._validate_url_host("", ["*"]) + with pytest.raises(ValueError, match="URL has no host"): + fetch_url._validate_url_host("not-a-url", ["*"]) + + +def test_fetch_url_rejects_when_url_allowlist_empty(fresh_stores): + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", url_allowlist=[]) + ) + with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): + fetch_url.fetch_url("http://anything.com/", "prod") + + +def test_fetch_url_error_message_mentions_url_allowlist(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="prod", + master="yarn", + yarn_rm_url="http://rm.prod.internal:8088", + ) + ) + with pytest.raises(ValueError, match="url_allowlist"): + fetch_url.fetch_url("http://evil.com/foo", "prod") + + +def test_fetch_url_omitted_url_allowlist_defaults_to_empty_and_rejects(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="prod", + master="yarn", + yarn_rm_url="http://rm.prod.internal", + ) + ) + with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): + fetch_url.fetch_url("http://nm.prod.internal/", "prod") + + +def test_fetch_url_rejects_redirect_to_disallowed_host(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="ccam", + master="yarn", + yarn_rm_url="http://ccam1:8088", + url_allowlist=["ccam*"], + ) + ) + redirect = httpx.Response(302, headers={"Location": "http://evil.com/"}) + with patch( + "spark_executor.tools.fetch_url.httpx.get", + return_value=redirect, + ) as m: + with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): + fetch_url.fetch_url("http://ccam50/foo", "ccam") + assert m.call_count == 1 + + +def test_fetch_url_follows_redirect_to_allowed_host(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="prod", + master="yarn", + yarn_rm_url="http://rm.prod.internal:8088", + url_allowlist=["*.internal"], + ) + ) + redirect = httpx.Response( + 302, headers={"Location": "http://other.internal/"} + ) + final = httpx.Response(200, text="ok") + with patch( + "spark_executor.tools.fetch_url.httpx.get", + side_effect=[redirect, final], + ) as m: + out = fetch_url.fetch_url("http://foo.internal/", "prod") + assert out.status_code == 200 + assert out.body == "ok" + assert m.call_count == 2 + assert m.call_args_list[1].args[0] == "http://other.internal/" + + +def test_fetch_url_rejects_redirect_to_ip_literal_not_in_allowlist(fresh_stores): + fetch_url.conn_store.save( + Connection( + name="ccam", + master="yarn", + yarn_rm_url="http://ccam1:8088", + url_allowlist=["ccam*"], + ) + ) + redirect = httpx.Response(302, headers={"Location": "http://10.0.0.1/"}) + with patch( + "spark_executor.tools.fetch_url.httpx.get", + return_value=redirect, + ) as m: + with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): + fetch_url.fetch_url("http://ccam50/foo", "ccam") + assert m.call_count == 1 + + +def test_list_applications_request_limit_bounds(): + assert ListApplicationsRequest(connection_name="prod", limit=10000).limit == 10000 + with pytest.raises(ValidationError): + ListApplicationsRequest(connection_name="prod", limit=0) + with pytest.raises(ValidationError): + ListApplicationsRequest(connection_name="prod", limit=10001)