# 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 _MAX_BODY_BYTES = 1_000_000 # 1 MB cap on response body _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 _read_streamed_body(resp: httpx.Response) -> tuple[bytes, bool]: """Stream the response body, keeping at most ``_MAX_BODY_BYTES`` bytes. Returns ``(body_bytes, truncated)``. Stops reading as soon as the cap is exceeded so that a multi-gigabyte response from an allowlisted host cannot OOM the service. """ chunks: list[bytes] = [] total = 0 truncated = False for chunk in resp.iter_bytes(): if total + len(chunk) <= _MAX_BODY_BYTES: chunks.append(chunk) total += len(chunk) else: if total < _MAX_BODY_BYTES: chunks.append(chunk[: _MAX_BODY_BYTES - total]) total = _MAX_BODY_BYTES truncated = True break return b"".join(chunks), truncated 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): with httpx.stream( "GET", url, auth=auth, verify=verify, timeout=_REQUEST_TIMEOUT_SECONDS, follow_redirects=False, ) as resp: if resp.status_code not in _REDIRECT_STATUSES: body_bytes, truncated = _read_streamed_body(resp) encoding = resp.encoding or "utf-8" body = body_bytes.decode(encoding, errors="replace") 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, ) 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})" )