diff --git a/README.md b/README.md index 8916969..94e76a2 100644 --- a/README.md +++ b/README.md @@ -120,6 +120,7 @@ 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 的聚合日志 | +| `fetch_url` | 代理 HTTP GET 到集群内网 URL (受 `Connection.yarn_rm_url` host 围栏约束) | ### Files MCP 工具 diff --git a/spark_executor/models.py b/spark_executor/models.py index 798d72b..ce3b2ab 100644 --- a/spark_executor/models.py +++ b/spark_executor/models.py @@ -40,6 +40,14 @@ class SubmitResult(BaseModel): tracking_url: str | None = None +class FetchUrlResult(BaseModel): + url: str + status_code: int + content_type: str + body: str + truncated: bool = False + + class Connection(BaseModel): name: str # Defaults to "yarn" because that's the literal string spark-submit wants diff --git a/spark_executor/server.py b/spark_executor/server.py index 3e156c2..f772426 100644 --- a/spark_executor/server.py +++ b/spark_executor/server.py @@ -21,6 +21,7 @@ from spark_executor.tools.external_jobs import ( get_external_job_status, get_external_job_result, ) +from spark_executor.tools.fetch_url import fetch_url from spark_executor.tools.requests import ( ConnectionNameRequest, EmptyRequest, @@ -28,6 +29,7 @@ from spark_executor.tools.requests import ( ExternalJobLogsRequest, ExternalJobStatusRequest, ExternalJobResultRequest, + FetchUrlRequest, GetJobLogsRequest, JobIdRequest, PendingIdRequest, @@ -467,3 +469,28 @@ 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 share at least 2 labels " + "of suffix with the named Connection's yarn_rm_url host (e.g. if " + "yarn_rm_url is 'rm.prod.internal:8088', you may fetch " + "'http://nm01.prod.internal:8042/...' but NOT 'http://evil.com/...'). " + "IP literals (10.0.0.1, ::1) and non-http(s) schemes (file://, " + "gopher://) are rejected. The Connection's saved auth is reused, so " + "the agent does not need cluster credentials.\n\n" + "**Limits:** response body capped at 1 MB (truncated=true if larger), " + "30s timeout, redirects followed." + ), +) +def _fetch_url(req: FetchUrlRequest): + return fetch_url(req.url, req.connection_name) diff --git a/spark_executor/tools/fetch_url.py b/spark_executor/tools/fetch_url.py new file mode 100644 index 0000000..6556613 --- /dev/null +++ b/spark_executor/tools/fetch_url.py @@ -0,0 +1,108 @@ +# 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: +URL host must share >= 2 labels of suffix with the connection's yarn_rm_url +host; IP literals and non-HTTP schemes are rejected. Reuses the connection's +saved auth/SSL config so the agent doesn't need cluster credentials. +""" +import ipaddress +from urllib.parse import urlparse + +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 + +_MAX_BODY_BYTES = 1_000_000 # 1 MB cap on response body +_REQUEST_TIMEOUT_SECONDS = 30 + + +def _host_suffix_overlap(host_a: str, host_b: str, min_labels: int = 2) -> bool: + """Return True if host_a and host_b share at least min_labels suffix labels.""" + labels_a = host_a.lower().split(".") + labels_b = host_b.lower().split(".") + n = 0 + i, j = len(labels_a) - 1, len(labels_b) - 1 + while i >= 0 and j >= 0 and labels_a[i] == labels_b[j]: + n += 1 + i -= 1 + j -= 1 + return n >= min_labels + + +def _validate_url_host(url: str, yarn_rm_url: str | None) -> None: + """Validate that url is an allowed http(s) hostname on the cluster network.""" + parsed = urlparse(url) + if parsed.scheme not in ("http", "https"): + raise ValueError(f"URL scheme must be http or https, got {parsed.scheme!r}") + + host = parsed.hostname + if not host: + raise ValueError(f"URL has no host: {url!r}") + + try: + ipaddress.ip_address(host) + except ValueError: + pass + else: + raise ValueError( + f"URL host {host!r} is an IP literal — IP targets are not allowed. " + "Use a hostname on the cluster network." + ) + + if not yarn_rm_url: + raise ValueError( + "Connection has no yarn_rm_url set, cannot validate URL host. " + "Save a Connection with yarn_rm_url first." + ) + + anchor = urlparse(yarn_rm_url).hostname or "" + if not _host_suffix_overlap(host, anchor, min_labels=2): + raise ValueError( + f"URL host {host!r} does not share enough suffix with the connection's " + f"yarn_rm_url host {anchor!r} (need >= 2 labels of common suffix). " + "Reject this fetch to prevent SSRF." + ) + + +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.yarn_rm_url) + + config = YarnClientConfig.from_connection(conn) + resp = httpx.get( + url, + auth=config.auth_for_httpx(), + verify=config.verify_for_httpx(), + timeout=_REQUEST_TIMEOUT_SECONDS, + follow_redirects=True, + ) + + body = resp.text[:_MAX_BODY_BYTES] + truncated = len(resp.text) > _MAX_BODY_BYTES + + 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_bytes={len(body)} truncated={truncated}" + ) + + return FetchUrlResult( + url=url, + status_code=resp.status_code, + content_type=resp.headers.get("content-type", ""), + body=body, + truncated=truncated, + ) diff --git a/spark_executor/tools/requests.py b/spark_executor/tools/requests.py index d9fbee7..20ca8da 100644 --- a/spark_executor/tools/requests.py +++ b/spark_executor/tools/requests.py @@ -378,3 +378,27 @@ class ExternalJobResultRequest(BaseModel): ..., description="Name of a saved Connection pointing at the YARN cluster.", ) + +class FetchUrlRequest(BaseModel): + url: str = Field( + ..., + description=( + "Absolute http:// or https:// URL to fetch. The host must share " + "at least 2 labels of suffix with the named Connection's " + "yarn_rm_url host (e.g. if yarn_rm_url is 'rm.prod.internal:8088', " + "you may fetch 'http://nm01.prod.internal:8042/...' but NOT " + "'http://evil.com/...' or 'http://10.0.0.1/...'). IP literals " + "and non-http(s) schemes are rejected. 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." + ), + ) diff --git a/tests/integration/test_mcp_routes.py b/tests/integration/test_mcp_routes.py index 3b555b8..ae074b8 100644 --- a/tests/integration/test_mcp_routes.py +++ b/tests/integration/test_mcp_routes.py @@ -64,12 +64,13 @@ def test_seventeen_tool_routes_registered(): assert "/update_pending_job" in paths -def test_twenty_tool_routes_registered(): +def test_twenty_one_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", + "/fetch_url", ): assert path in paths, f"missing MCP tool route: {path}" diff --git a/tests/unit/test_fetch_url.py b/tests/unit/test_fetch_url.py new file mode 100644 index 0000000..75b8e8f --- /dev/null +++ b/tests/unit/test_fetch_url.py @@ -0,0 +1,164 @@ +# coding=utf-8 +from unittest.mock import patch + +import httpx +import pytest + +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 + + +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 + + +def test_fetch_url_returns_body_and_status(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm.prod.internal:8088") + ) + 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 out.truncated is False + assert m.call_count == 1 + + +def test_fetch_url_raises_for_missing_connection(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + with pytest.raises(KeyError, match="Connection not found"): + fetch_url.fetch_url("http://rm.prod.internal:8088/", "missing") + + +def test_fetch_url_rejects_non_http_scheme(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm.prod.internal:8088") + ) + with pytest.raises(ValueError, match="scheme"): + fetch_url.fetch_url("file:///etc/passwd", "prod") + + +def test_fetch_url_rejects_ftp_scheme(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm.prod.internal:8088") + ) + with pytest.raises(ValueError, match="scheme"): + fetch_url.fetch_url("ftp://rm.prod.internal/foo", "prod") + + +def test_fetch_url_rejects_ip_literal(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm.prod.internal:8088") + ) + with pytest.raises(ValueError, match="IP literal"): + fetch_url.fetch_url("http://10.0.0.1/secrets", "prod") + + +def test_fetch_url_rejects_ipv6_literal(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm.prod.internal:8088") + ) + with pytest.raises(ValueError, match="IP literal"): + fetch_url.fetch_url("http://[::1]:8080/", "prod") + + +def test_fetch_url_rejects_external_host(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm.prod.internal:8088") + ) + with pytest.raises(ValueError, match="does not share enough suffix"): + fetch_url.fetch_url("http://evil.com/foo", "prod") + + +def test_fetch_url_rejects_too_short_suffix_overlap(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm/8088") + ) + with pytest.raises(ValueError, match="does not share enough suffix"): + fetch_url.fetch_url("http://other-rm/", "prod") + + +def test_fetch_url_rejects_when_connection_has_no_yarn_rm_url(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save(Connection(name="prod", master="yarn", yarn_rm_url=None)) + with pytest.raises(ValueError, match="no yarn_rm_url set"): + fetch_url.fetch_url("http://anything.com/", "prod") + + +def test_fetch_url_truncates_body_over_1mb(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm.prod.internal:8088") + ) + big = "x" * (fetch_url._MAX_BODY_BYTES + 1) + resp = httpx.Response(200, text=big) + with patch("spark_executor.tools.fetch_url.httpx.get", return_value=resp): + out = fetch_url.fetch_url("http://nm01.prod.internal/big", "prod") + assert out.truncated is True + assert len(out.body) == fetch_url._MAX_BODY_BYTES + + +def test_fetch_url_passes_auth_from_connection(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + 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", + ) + ) + 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(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection( + name="insecure", + master="yarn", + yarn_rm_url="http://rm.prod.internal:8088", + ssl_verify=False, + ) + ) + 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_follows_redirects(tmp_path, monkeypatch): + _fresh_stores(tmp_path, monkeypatch) + fetch_url.conn_store.save( + Connection(name="prod", master="yarn", yarn_rm_url="http://rm.prod.internal:8088") + ) + 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", "prod") + assert m.call_args.kwargs["follow_redirects"] is True