refactor(fetch_url): trust the allowlist, rename to url_allowlist

Two changes:

1) Drop every check except the allowlist lookup.

  Old _validate_url_host did: scheme check, host-presence check,
  IP-literal check, empty-allowlist check, then glob match.

  New _validate_url_host does: parse host, return on glob match,
  raise on miss. That's it. The only remaining structural check is
  'the URL must have a host' (otherwise the glob has nothing to
  test against).

  Security implication: scheme (file://, gopher://, ftp://) and
  IP literals (10.0.0.1, ::1) are NO LONGER rejected by the
  validator. The allowlist is the single source of truth. If the
  user writes ['*.*.*.*'], they have opted in to 4-label hosts
  including IP literals; if they write ['ccam*'], they get ccam1-
  ccam99 and nothing else. The default ['ccam*'] / [] pattern is
  tight by construction.

  Removed: import ipaddress, the scheme/IP rejection branches, the
  'allowlist empty' explicit branch (the empty list naturally
  matches nothing).

2) Rename allowed_url_hosts -> url_allowlist.

  The previous name was a verbose double-negative ('allowed ... hosts').
  The new name is short, modern (allowlist > whitelist), and matches
  the pattern of the field (URL hosts allowed). Renamed in:
    - Connection (models.py)
    - SaveConnectionRequest, UpdateConnectionRequest, FetchUrlRequest
      (requests.py)
    - _validate_url_host, _host_matches_any_glob parameters
      (fetch_url.py)
    - save_connection / update_connection call sites and error
      messages (connections.py, fetch_url.py)
    - Route descriptions (server.py)
    - All test files
    - README

  No backward-compat alias: the field was added in 0291f36 and
  hasn't shipped, so no production migration. Local dev data
  (data/connections.json, gitignored) with allowed_url_hosts set
  will be silently dropped by Pydantic v2 (default for extra
  fields is ignore) — those connections lose their allowlist and
  fetch_url will reject everything until re-saved.

Test cleanup:
  - Removed 9 obsolete tests (scheme/IP/suffix rejection)
  - Renamed allowed_url_hosts -> url_allowlist in 11 surviving tests
  - Added 4 new tests documenting the 'allowlist is the only gate'
    model: IP literal accepted, HTTPS accepted, no-host rejected,
    empty allowlist rejected, error message mentions url_allowlist

  -1 obsolete test, net -4 from 386 -> 382 tests passing.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Claude
2026-07-09 12:09:41 +08:00
co-authored by Claude Fable 5
parent 6d7387b022
commit a5b9539663
8 changed files with 138 additions and 186 deletions
+1 -1
View File
@@ -121,7 +121,7 @@ MCP 客户端需要先执行 `initialize` 握手,拿到 `mcp-session-id` 后
| `get_external_job_status` | 查询**非本服务提交**的外部 YARN application 状态(按 `application_id` + `connection_name` | | `get_external_job_status` | 查询**非本服务提交**的外部 YARN application 状态(按 `application_id` + `connection_name` |
| `get_external_job_result` | 查询外部 YARN application 终态结果视图 | | `get_external_job_result` | 查询外部 YARN application 终态结果视图 |
| `get_external_job_logs` | 拉取外部 YARN application 的聚合日志 | | `get_external_job_logs` | 拉取外部 YARN application 的聚合日志 |
| `fetch_url` | 代理 HTTP GET 到集群内网 URL (host 受 `Connection.allowed_url_hosts` glob allowlist 约束, 空则全拒) | | `fetch_url` | 代理 HTTP GET 到集群内网 URL (host 受 `Connection.url_allowlist` glob allowlist 约束, 空则全拒) |
### Files MCP 工具 ### Files MCP 工具
+1 -1
View File
@@ -69,7 +69,7 @@ class Connection(BaseModel):
auth_principal: str | None = None auth_principal: str | None = None
auth_keytab: str | None = None auth_keytab: str | None = None
allowed_url_hosts: list[str] = Field( url_allowlist: list[str] = Field(
default_factory=list, default_factory=list,
description=( description=(
"List of fnmatch glob patterns for hosts the fetch_url tool may access. " "List of fnmatch glob patterns for hosts the fetch_url tool may access. "
+6 -7
View File
@@ -505,13 +505,12 @@ def _update_job_file(req: UpdateJobFileRequest):
"response body. Useful when the agent is on a different network from " "response body. Useful when the agent is on a different network from "
"the cluster and cannot reach YARN tracking pages, Spark History " "the cluster and cannot reach YARN tracking pages, Spark History "
"Server, or NodeManager web UIs directly.\n\n" "Server, or NodeManager web UIs directly.\n\n"
"**Security constraints:** the URL host must share at least 2 labels " "**Security constraints:** the URL host must match one of the fnmatch "
"of suffix with the named Connection's yarn_rm_url host (e.g. if " "glob patterns in the named Connection's url_allowlist. An empty or "
"yarn_rm_url is 'rm.prod.internal:8088', you may fetch " "omitted allowlist denies every host. There are no scheme or IP-literal "
"'http://nm01.prod.internal:8042/...' but NOT 'http://evil.com/...'). " "guardrails — the allowlist is the only gate — so keep it tight. The "
"IP literals (10.0.0.1, ::1) and non-http(s) schemes (file://, " "Connection's saved auth is reused, so the agent does not need cluster "
"gopher://) are rejected. The Connection's saved auth is reused, so " "credentials.\n\n"
"the agent does not need cluster credentials.\n\n"
"**Limits:** response body capped at 1 MB (truncated=true if larger), " "**Limits:** response body capped at 1 MB (truncated=true if larger), "
"30s timeout, redirects followed." "30s timeout, redirects followed."
), ),
+11 -11
View File
@@ -19,10 +19,10 @@ def save_connection(
ssl_ca_bundle: str | None = None, ssl_ca_bundle: str | None = None,
auth_type: str = "none", auth_type: str = "none",
auth_user: str | None = None, auth_user: str | None = None,
auth_password: str | None = None, auth_password: str | None = None,
auth_principal: str | None = None, auth_principal: str | None = None,
auth_keytab: str | None = None, auth_keytab: str | None = None,
allowed_url_hosts: list[str] | None = None, url_allowlist: list[str] | None = None,
) -> dict[str, str]: ) -> dict[str, str]:
logger.debug( logger.debug(
f"save_connection enter name={name} master={master} deploy_mode={deploy_mode} " f"save_connection enter name={name} master={master} deploy_mode={deploy_mode} "
@@ -39,10 +39,10 @@ def save_connection(
auth_type=auth_type, auth_type=auth_type,
auth_user=auth_user, auth_user=auth_user,
auth_password=auth_password, auth_password=auth_password,
auth_principal=auth_principal, auth_principal=auth_principal,
auth_keytab=auth_keytab, auth_keytab=auth_keytab,
allowed_url_hosts=allowed_url_hosts or [], url_allowlist=url_allowlist or [],
) )
store.save(conn) store.save(conn)
return {"name": name, "status": "SAVED"} return {"name": name, "status": "SAVED"}
@@ -54,9 +54,9 @@ def update_connection(name: str, **fields) -> dict[str, object]:
you pass are changed. To clear an optional field (e.g. `yarn_rm_url`), you pass are changed. To clear an optional field (e.g. `yarn_rm_url`),
use `delete_connection(name=...)` followed by `save_connection(...)`. use `delete_connection(name=...)` followed by `save_connection(...)`.
Mutable fields: master, deploy_mode, yarn_rm_url, spark_conf, Mutable fields: master, deploy_mode, yarn_rm_url, spark_conf,
ssl_verify, ssl_ca_bundle, auth_type, auth_user, auth_password, ssl_verify, ssl_ca_bundle, auth_type, auth_user, auth_password,
auth_principal, auth_keytab, allowed_url_hosts. auth_principal, auth_keytab, url_allowlist.
""" """
logger.debug(f"update_connection enter name={name} fields={sorted(fields.keys())}") logger.debug(f"update_connection enter name={name} fields={sorted(fields.keys())}")
return store.update(name, **fields).model_dump() return store.update(name, **fields).model_dump()
+23 -40
View File
@@ -6,11 +6,11 @@
Generic HTTP GET proxy for the agent. Lets the agent fetch URLs on the 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: cluster's network when it cannot reach those hosts directly. Security:
the host is checked against an explicit fnmatch glob allowlist configured the host is checked against an explicit fnmatch glob allowlist configured
on the connection (`allowed_url_hosts`). Empty or missing allowlist means on the connection (`url_allowlist`). Empty or missing allowlist means no
no URL access; IP literals and non-HTTP schemes are rejected. Reuses the URL access; the allowlist is the only gate — scheme and IP-literal checks
connection's saved auth/SSL config so the agent doesn't need cluster credentials. are intentionally NOT performed. Reuses the connection's saved auth/SSL
config so the agent doesn't need cluster credentials.
""" """
import ipaddress
import fnmatch import fnmatch
from urllib.parse import urlparse from urllib.parse import urlparse
@@ -25,7 +25,7 @@ _MAX_BODY_BYTES = 1_000_000 # 1 MB cap on response body
_REQUEST_TIMEOUT_SECONDS = 30 _REQUEST_TIMEOUT_SECONDS = 30
def _host_matches_any_glob(host: str, patterns: list[str]) -> bool: def _host_matches_any_glob(host: str, allowlist: list[str]) -> bool:
"""True if `host` matches any of the fnmatch glob patterns. """True if `host` matches any of the fnmatch glob patterns.
fnmatch is case-sensitive on Linux (our deployment target). `*` in a fnmatch is case-sensitive on Linux (our deployment target). `*` in a
@@ -34,7 +34,7 @@ def _host_matches_any_glob(host: str, patterns: list[str]) -> bool:
for SSRF). for SSRF).
""" """
host_labels = host.split(".") host_labels = host.split(".")
for p in patterns: for p in allowlist:
pat_labels = p.split(".") pat_labels = p.split(".")
if len(pat_labels) != len(host_labels): if len(pat_labels) != len(host_labels):
continue continue
@@ -43,47 +43,30 @@ def _host_matches_any_glob(host: str, patterns: list[str]) -> bool:
return False return False
def _validate_url_host( def _validate_url_host(url: str, allowlist: list[str] | None) -> None:
url: str, """Reject the URL unless its host matches a glob in `allowlist`.
allowed_hosts: list[str] | None,
) -> None:
"""Raise ValueError if the URL is not allowed to be fetched.
Allowed only if the URL host matches one of the fnmatch glob patterns in The only check. The url_allowlist is the single source of truth for
`allowed_hosts`. An empty or missing allowlist rejects everything. 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.
""" """
parsed = urlparse(url) host = urlparse(url).hostname
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: if not host:
raise ValueError(f"URL has no host: {url!r}") raise ValueError(f"URL has no host: {url!r}")
# IP literal check if _host_matches_any_glob(host, allowlist or []):
try:
ipaddress.ip_address(host)
raise ValueError(
f"URL host {host!r} is an IP literal — IP targets are not allowed. "
f"Use a hostname on the cluster network."
)
except ValueError as e:
if "IP literal" in str(e):
raise
# not an IP, continue
if not allowed_hosts:
raise ValueError(
"allowed_url_hosts is empty — set Connection.allowed_url_hosts to "
"allow specific hosts before calling fetch_url (e.g. ['ccam*'] "
"for ccam1-ccam99 or ['*.prod.internal'] for a subdomain)."
)
if _host_matches_any_glob(host, allowed_hosts):
return return
raise ValueError( raise ValueError(
f"URL host {host!r} is not in Connection.allowed_url_hosts " f"URL host {host!r} is not in Connection.url_allowlist {allowlist!r}. "
f"{allowed_hosts!r}. Reject this fetch to prevent SSRF." 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: def fetch_url(url: str, connection_name: str) -> FetchUrlResult:
"""Proxy an HTTP GET to url using the auth/SSL settings of connection_name.""" """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}") logger.debug(f"fetch_url enter url={url} connection_name={connection_name}")
@@ -91,7 +74,7 @@ def fetch_url(url: str, connection_name: str) -> FetchUrlResult:
if conn is None: if conn is None:
raise KeyError(f"Connection not found: {connection_name}") raise KeyError(f"Connection not found: {connection_name}")
_validate_url_host(url, conn.allowed_url_hosts) _validate_url_host(url, conn.url_allowlist)
config = YarnClientConfig.from_connection(conn) config = YarnClientConfig.from_connection(conn)
resp = httpx.get( resp = httpx.get(
+13 -14
View File
@@ -119,17 +119,17 @@ class SaveConnectionRequest(BaseModel):
"for 'kinit -kt' workflows. The service does NOT auto-initialize " "for 'kinit -kt' workflows. The service does NOT auto-initialize "
"from the keytab — you must `kinit -kt <auth_keytab> <auth_principal>` " "from the keytab — you must `kinit -kt <auth_keytab> <auth_principal>` "
"yourself before calling the tools." "yourself before calling the tools."
), ),
) )
allowed_url_hosts: list[str] | None = Field( url_allowlist: list[str] | None = Field(
default=None, default=None,
description=( description=(
"Optional list of fnmatch glob patterns for hosts the fetch_url tool " "Optional list of fnmatch glob patterns for hosts the fetch_url tool "
"may access. See Connection.allowed_url_hosts for full semantics. " "may access. See Connection.url_allowlist for full semantics. "
"Example for single-label host clusters: ['ccam*'] allows any host " "Example for single-label host clusters: ['ccam*'] allows any host "
"starting with 'ccam' (ccam1, ccam2, ..., ccam99). If omitted, None, " "starting with 'ccam' (ccam1, ccam2, ..., ccam99). If omitted, None, "
"or empty, the saved connection will have allowed_url_hosts=[] (the " "or empty, the saved connection will have url_allowlist=[] (the "
"default), meaning fetch_url will reject every URL until the list is " "default), meaning fetch_url will reject every URL until the list is "
"populated via update_connection." "populated via update_connection."
), ),
@@ -395,13 +395,12 @@ class FetchUrlRequest(BaseModel):
url: str = Field( url: str = Field(
..., ...,
description=( description=(
"Absolute http:// or https:// URL to fetch. The host must share " "Absolute URL to fetch. The only access control is the named "
"at least 2 labels of suffix with the named Connection's " "Connection's url_allowlist: the URL host must match one of the "
"yarn_rm_url host (e.g. if yarn_rm_url is 'rm.prod.internal:8088', " "fnmatch glob patterns in that list. An empty or omitted allowlist "
"you may fetch 'http://nm01.prod.internal:8042/...' but NOT " "denies every host. IP literals and non-HTTP schemes are allowed "
"'http://evil.com/...' or 'http://10.0.0.1/...'). IP literals " "if and only if they are matched by the allowlist. The Connection's "
"and non-http(s) schemes are rejected. The Connection's saved " "saved auth is reused for the outbound request — the agent does not "
"auth is reused for the outbound request — the agent does not "
"need cluster credentials." "need cluster credentials."
), ),
) )
@@ -468,10 +467,10 @@ class UpdateConnectionRequest(BaseModel):
default=None, default=None,
description="New Kerberos keytab path. Omit to keep current.", description="New Kerberos keytab path. Omit to keep current.",
) )
allowed_url_hosts: list[str] | None = Field( url_allowlist: list[str] | None = Field(
default=None, default=None,
description=( description=(
"Replacement allowed_url_hosts list (not merged). Omit to keep current. " "Replacement url_allowlist list (not merged). Omit to keep current. "
"Pass an empty list to deny all hosts (the default for new connections). " "Pass an empty list to deny all hosts (the default for new connections). "
"Example: ['ccam*'] allows ccam1-ccam99." "Example: ['ccam*'] allows ccam1-ccam99."
), ),
+7 -7
View File
@@ -148,10 +148,10 @@ def test_update_connection_with_no_fields_is_noop(_fresh_store):
assert out["deploy_mode"] == "client" assert out["deploy_mode"] == "client"
def test_update_connection_allows_url_hosts_change(_fresh_store): def test_update_connection_changes_url_allowlist(_fresh_store):
connections.save_connection(name="prod", master="yarn") connections.save_connection(name="prod", master="yarn")
out = connections.update_connection(name="prod", allowed_url_hosts=["ccam*"]) out = connections.update_connection(name="prod", url_allowlist=["ccam*"])
assert out["allowed_url_hosts"] == ["ccam*"] assert out["url_allowlist"] == ["ccam*"]
def test_update_connection_replaces_not_merges_dict(_fresh_store): def test_update_connection_replaces_not_merges_dict(_fresh_store):
@@ -160,14 +160,14 @@ def test_update_connection_replaces_not_merges_dict(_fresh_store):
assert out["spark_conf"] == {"b": "2"} assert out["spark_conf"] == {"b": "2"}
def test_update_connection_replaces_not_merges_list(_fresh_store): def test_update_connection_replaces_not_merges_url_allowlist(_fresh_store):
connections.save_connection( connections.save_connection(
name="prod", name="prod",
master="yarn", master="yarn",
allowed_url_hosts=["ccam*"], url_allowlist=["ccam*"],
) )
out = connections.update_connection(name="prod", allowed_url_hosts=["nm*"]) out = connections.update_connection(name="prod", url_allowlist=["nm*"])
assert out["allowed_url_hosts"] == ["nm*"] assert out["url_allowlist"] == ["nm*"]
def test_update_connection_raises_for_unknown_name(_fresh_store): def test_update_connection_raises_for_unknown_name(_fresh_store):
+76 -105
View File
@@ -26,14 +26,14 @@ def fresh_stores(tmp_path: Path, monkeypatch):
_fresh_stores(tmp_path, monkeypatch) _fresh_stores(tmp_path, monkeypatch)
yield yield
def test_fetch_url_returns_body_and_status(tmp_path, monkeypatch):
_fresh_stores(tmp_path, monkeypatch) def test_fetch_url_returns_body_and_status(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="prod", name="prod",
master="yarn", master="yarn",
yarn_rm_url="http://rm.prod.internal:8088", yarn_rm_url="http://rm.prod.internal:8088",
allowed_url_hosts=["*.prod.internal"], url_allowlist=["*.prod.internal"],
) )
) )
resp = httpx.Response(200, text="hello", headers={"content-type": "text/html"}) resp = httpx.Response(200, text="hello", headers={"content-type": "text/html"})
@@ -47,89 +47,18 @@ def test_fetch_url_returns_body_and_status(tmp_path, monkeypatch):
assert m.call_count == 1 assert m.call_count == 1
def test_fetch_url_raises_for_missing_connection(tmp_path, monkeypatch): def test_fetch_url_raises_for_missing_connection(fresh_stores):
_fresh_stores(tmp_path, monkeypatch)
with pytest.raises(KeyError, match="Connection not found"): with pytest.raises(KeyError, match="Connection not found"):
fetch_url.fetch_url("http://rm.prod.internal:8088/", "missing") fetch_url.fetch_url("http://rm.prod.internal:8088/", "missing")
def test_fetch_url_rejects_non_http_scheme(tmp_path, monkeypatch): def test_fetch_url_truncates_body_over_1mb(fresh_stores):
_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( fetch_url.conn_store.save(
Connection( Connection(
name="prod", name="prod",
master="yarn", master="yarn",
yarn_rm_url="http://rm.prod.internal:8088", yarn_rm_url="http://rm.prod.internal:8088",
allowed_url_hosts=["*.prod.internal"], url_allowlist=["*.prod.internal"],
)
)
with pytest.raises(ValueError, match="is not in Connection.allowed_url_hosts"):
fetch_url.fetch_url("http://evil.com/foo", "prod")
def test_fetch_url_rejects_unrelated_single_label_host(tmp_path, monkeypatch):
"""Without a common suffix rule, a single-label RM host no longer helps."""
_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="allowed_url_hosts is empty"):
fetch_url.fetch_url("http://other-rm/", "prod")
def test_fetch_url_rejects_when_allowed_url_hosts_empty(tmp_path, monkeypatch):
_fresh_stores(tmp_path, monkeypatch)
fetch_url.conn_store.save(
Connection(name="prod", master="yarn", allowed_url_hosts=[])
)
with pytest.raises(ValueError, match="allowed_url_hosts is empty"):
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",
allowed_url_hosts=["*.prod.internal"],
) )
) )
big = "x" * (fetch_url._MAX_BODY_BYTES + 1) big = "x" * (fetch_url._MAX_BODY_BYTES + 1)
@@ -140,8 +69,7 @@ def test_fetch_url_truncates_body_over_1mb(tmp_path, monkeypatch):
assert len(out.body) == fetch_url._MAX_BODY_BYTES assert len(out.body) == fetch_url._MAX_BODY_BYTES
def test_fetch_url_passes_auth_from_connection(tmp_path, monkeypatch): def test_fetch_url_passes_auth_from_connection(fresh_stores):
_fresh_stores(tmp_path, monkeypatch)
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="auth", name="auth",
@@ -150,7 +78,7 @@ def test_fetch_url_passes_auth_from_connection(tmp_path, monkeypatch):
auth_type="basic", auth_type="basic",
auth_user="u", auth_user="u",
auth_password="p", auth_password="p",
allowed_url_hosts=["*.prod.internal"], url_allowlist=["*.prod.internal"],
) )
) )
resp = httpx.Response(200, text="ok") resp = httpx.Response(200, text="ok")
@@ -159,19 +87,19 @@ def test_fetch_url_passes_auth_from_connection(tmp_path, monkeypatch):
auth = m.call_args.kwargs["auth"] auth = m.call_args.kwargs["auth"]
assert isinstance(auth, httpx.BasicAuth) assert isinstance(auth, httpx.BasicAuth)
import base64 import base64
creds = base64.b64decode(auth._auth_header.split()[1]).decode() creds = base64.b64decode(auth._auth_header.split()[1]).decode()
assert creds == "u:p" assert creds == "u:p"
def test_fetch_url_passes_ssl_verify_from_connection(tmp_path, monkeypatch): def test_fetch_url_passes_ssl_verify_from_connection(fresh_stores):
_fresh_stores(tmp_path, monkeypatch)
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="insecure", name="insecure",
master="yarn", master="yarn",
yarn_rm_url="http://rm.prod.internal:8088", yarn_rm_url="http://rm.prod.internal:8088",
ssl_verify=False, ssl_verify=False,
allowed_url_hosts=["*.prod.internal"], url_allowlist=["*.prod.internal"],
) )
) )
resp = httpx.Response(200, text="ok") resp = httpx.Response(200, text="ok")
@@ -180,14 +108,13 @@ def test_fetch_url_passes_ssl_verify_from_connection(tmp_path, monkeypatch):
assert m.call_args.kwargs["verify"] is False assert m.call_args.kwargs["verify"] is False
def test_fetch_url_follows_redirects(tmp_path, monkeypatch): def test_fetch_url_follows_redirects(fresh_stores):
_fresh_stores(tmp_path, monkeypatch)
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="prod", name="prod",
master="yarn", master="yarn",
yarn_rm_url="http://rm.prod.internal:8088", yarn_rm_url="http://rm.prod.internal:8088",
allowed_url_hosts=["*.prod.internal"], url_allowlist=["*.prod.internal"],
) )
) )
resp = httpx.Response(200, text="ok") resp = httpx.Response(200, text="ok")
@@ -195,13 +122,14 @@ def test_fetch_url_follows_redirects(tmp_path, monkeypatch):
fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "prod") fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "prod")
assert m.call_args.kwargs["follow_redirects"] is True assert m.call_args.kwargs["follow_redirects"] is True
def test_fetch_url_allows_host_matching_glob_pattern(fresh_stores): def test_fetch_url_allows_host_matching_glob_pattern(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="ccam", name="ccam",
master="yarn", master="yarn",
yarn_rm_url="http://ccam1:8088", yarn_rm_url="http://ccam1:8088",
allowed_url_hosts=["ccam*"], url_allowlist=["ccam*"],
) )
) )
resp = httpx.Response(200, text="hello") resp = httpx.Response(200, text="hello")
@@ -211,53 +139,59 @@ def test_fetch_url_allows_host_matching_glob_pattern(fresh_stores):
assert out.body == "hello" assert out.body == "hello"
assert m.call_count == 1 assert m.call_count == 1
def test_fetch_url_allows_host_matching_any_of_multiple_globs(fresh_stores): def test_fetch_url_allows_host_matching_any_of_multiple_globs(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="prod", name="prod",
master="yarn", master="yarn",
yarn_rm_url="http://rm.prod.internal:8088", yarn_rm_url="http://rm.prod.internal:8088",
allowed_url_hosts=["ccam*", "*.prod.internal"], url_allowlist=["ccam*", "*.prod.internal"],
) )
) )
resp = httpx.Response(200, text="hello") resp = httpx.Response(200, text="hello")
with patch("spark_executor.tools.fetch_url.httpx.get", return_value=resp) as m: 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") out = fetch_url.fetch_url(
"http://history.prod.internal:18080/api/v1/info", "prod"
)
assert out.status_code == 200 assert out.status_code == 200
assert out.body == "hello" assert out.body == "hello"
assert m.call_count == 1 assert m.call_count == 1
def test_fetch_url_glob_does_not_match_unrelated_host(fresh_stores): def test_fetch_url_glob_does_not_match_unrelated_host(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="ccam", name="ccam",
master="yarn", master="yarn",
yarn_rm_url="http://ccam1:8088", yarn_rm_url="http://ccam1:8088",
allowed_url_hosts=["ccam*"], url_allowlist=["ccam*"],
) )
) )
with pytest.raises(ValueError, match="is not in Connection.allowed_url_hosts"): with pytest.raises(ValueError, match="is not in Connection.url_allowlist"):
fetch_url.fetch_url("http://evil.com/foo", "ccam") fetch_url.fetch_url("http://evil.com/foo", "ccam")
def test_fetch_url_glob_does_not_cross_dot_boundary(fresh_stores): def test_fetch_url_glob_does_not_cross_dot_boundary(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="ccam", name="ccam",
master="yarn", master="yarn",
yarn_rm_url="http://ccam1:8088", yarn_rm_url="http://ccam1:8088",
allowed_url_hosts=["ccam*"], url_allowlist=["ccam*"],
) )
) )
with pytest.raises(ValueError, match="is not in Connection.allowed_url_hosts"): with pytest.raises(ValueError, match="is not in Connection.url_allowlist"):
fetch_url.fetch_url("http://ccam50.evil.com/", "ccam") fetch_url.fetch_url("http://ccam50.evil.com/", "ccam")
def test_fetch_url_glob_match_does_not_require_suffix_overlap(fresh_stores): def test_fetch_url_glob_match_does_not_require_suffix_overlap(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="prod", name="prod",
master="yarn", master="yarn",
yarn_rm_url="http://rm:8088", yarn_rm_url="http://rm:8088",
allowed_url_hosts=["ccam*"], url_allowlist=["ccam*"],
) )
) )
resp = httpx.Response(200, text="hello") resp = httpx.Response(200, text="hello")
@@ -267,30 +201,55 @@ def test_fetch_url_glob_match_does_not_require_suffix_overlap(fresh_stores):
assert out.body == "hello" assert out.body == "hello"
assert m.call_count == 1 assert m.call_count == 1
def test_fetch_url_empty_allowed_hosts_rejects_everything(fresh_stores):
def test_fetch_url_accepts_ip_literal_when_in_url_allowlist(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="prod", name="prod",
master="yarn", master="yarn",
yarn_rm_url="http://rm.prod.internal", url_allowlist=["*.*.*.*"],
allowed_url_hosts=[],
) )
) )
with pytest.raises(ValueError, match="allowed_url_hosts is empty"): resp = httpx.Response(200, text="hello")
fetch_url.fetch_url("http://nm.prod.internal/", "prod") 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_omitted_allowed_url_hosts_defaults_to_empty_and_rejects(fresh_stores):
def test_fetch_url_accepts_https_when_in_url_allowlist(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
name="prod", name="prod",
master="yarn", master="yarn",
yarn_rm_url="http://rm.prod.internal", url_allowlist=["ccam*.example.com"],
) )
) )
with pytest.raises(ValueError, match="allowed_url_hosts is empty"): resp = httpx.Response(200, text="hello")
fetch_url.fetch_url("http://nm.prod.internal/", "prod") 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_error_message_hints_at_allowed_url_hosts(fresh_stores):
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( fetch_url.conn_store.save(
Connection( Connection(
name="prod", name="prod",
@@ -298,5 +257,17 @@ def test_fetch_url_error_message_hints_at_allowed_url_hosts(fresh_stores):
yarn_rm_url="http://rm.prod.internal:8088", yarn_rm_url="http://rm.prod.internal:8088",
) )
) )
with pytest.raises(ValueError, match="set Connection.allowed_url_hosts"): with pytest.raises(ValueError, match="url_allowlist"):
fetch_url.fetch_url("http://evil.com/foo", "prod") 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")