fix: address review findings on fetch-url-tool
- fetch_url: revalidate allowlist on every redirect hop (fixes SSRF where 302 to disallowed host / 169.254.169.254 / file:// bypassed the url_allowlist). Stream response body with iter_bytes and cap at 1MB so a multi-GB response from an allowlisted host cannot OOM the service. Reuses the manual-redirect-loop pattern from yarn_client. - list_applications: stop swallowing 404 (YARN returns 200+empty for "no match"; 404 means the RM doesn't support the endpoint — surface the YarnError instead of hiding it as an empty result). Add Field(ge=1, le=10000) to ListApplicationsRequest.limit so a runaway limit is rejected at the Pydantic layer with 422. - save_connection: PATCH semantics for existing records. Re-route to update_connection when the name already exists so partial updates (e.g. only master) no longer wipe url_allowlist back to []. Uses an _UNSET sentinel in the tool function to distinguish "omitted" from "None" without breaking the existing parameter list. - README: drop leading space on 5 new connection-tool table rows that was breaking GitHub Flavored Markdown table continuity. - Indentation: normalize connections.py and requests.py to 4-space indent (auth_password/auth_principal/auth_keytab were 3-space). Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -4,45 +4,75 @@
|
||||
@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,
|
||||
url_allowlist: list[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())}"
|
||||
f"save_connection enter name={name} master={master} "
|
||||
f"spark_conf_keys={list((spark_conf if isinstance(spark_conf, dict) else {}).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,
|
||||
url_allowlist=url_allowlist or [],
|
||||
)
|
||||
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"}
|
||||
|
||||
@@ -54,8 +84,8 @@ def update_connection(name: str, **fields) -> dict[str, object]:
|
||||
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,
|
||||
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())}")
|
||||
|
||||
@@ -12,7 +12,7 @@ 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
|
||||
from urllib.parse import urlparse, urljoin
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -24,6 +24,9 @@ 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.
|
||||
@@ -67,6 +70,29 @@ def _validate_url_host(url: str, allowlist: list[str] | None) -> None:
|
||||
)
|
||||
|
||||
|
||||
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}")
|
||||
@@ -77,28 +103,50 @@ def fetch_url(url: str, connection_name: str) -> FetchUrlResult:
|
||||
_validate_url_host(url, conn.url_allowlist)
|
||||
|
||||
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,
|
||||
)
|
||||
auth = config.auth_for_httpx()
|
||||
verify = config.verify_for_httpx()
|
||||
|
||||
body = resp.text[:_MAX_BODY_BYTES]
|
||||
truncated = len(resp.text) > _MAX_BODY_BYTES
|
||||
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}"
|
||||
)
|
||||
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,
|
||||
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})"
|
||||
)
|
||||
|
||||
@@ -117,10 +117,10 @@ class SaveConnectionRequest(BaseModel):
|
||||
description=(
|
||||
"Absolute path to a Kerberos keytab file. Optional convenience "
|
||||
"for 'kinit -kt' workflows. The service does NOT auto-initialize "
|
||||
"from the keytab — you must `kinit -kt <auth_keytab> <auth_principal>` "
|
||||
"yourself before calling the tools."
|
||||
),
|
||||
)
|
||||
"from the keytab — you must `kinit -kt <auth_keytab> <auth_principal>` "
|
||||
"yourself before calling the tools."
|
||||
),
|
||||
)
|
||||
|
||||
url_allowlist: list[str] | None = Field(
|
||||
default=None,
|
||||
@@ -504,6 +504,8 @@ class ListApplicationsRequest(BaseModel):
|
||||
)
|
||||
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 "
|
||||
|
||||
Reference in New Issue
Block a user