refactor: drop fetch_url body cap and modernize datetime usage

- fetch_url: remove the 1MB body cap and the streaming helper. The
  manual redirect loop with allowlist re-check (the SSRF fix) is kept
  intact; the body is now read in full via resp.text. Drop the now-
  meaningless `truncated` field from FetchUrlResult and the tests that
  asserted on it. Switched from httpx.stream() back to httpx.get()
  for the redirect loop — cleaner without the body cap.

- datetime: replace deprecated datetime.utcnow() with
  datetime.now(timezone.utc) in submit.py (3 sites) and
  core/job_writer.py (1 site). Update the stale comment in
  core/pending_store.py that referenced the old call.

- Clean up an unused `from common.config import settings` import in
  submit.py that ruff flagged.

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
Claude
2026-07-09 14:35:16 +08:00
co-authored by Claude
parent 7d4e512cb0
commit deafb5a26b
7 changed files with 50 additions and 158 deletions
+2 -2
View File
@@ -14,7 +14,7 @@ Resolution order for the output directory:
""" """
import os import os
import secrets import secrets
from datetime import datetime from datetime import datetime, timezone
from common.config import settings from common.config import settings
from common.logging import logger 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) effective_dir = resolve_jobs_dir(jobs_dir)
os.makedirs(effective_dir, exist_ok=True) 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" name = f"job_{stamp}_{secrets.token_hex(3)}.py"
path = os.path.join(effective_dir, name) path = os.path.join(effective_dir, name)
abs_path = os.path.abspath(path) abs_path = os.path.abspath(path)
+1 -1
View File
@@ -37,7 +37,7 @@ class PendingStore:
return Path(self._data_dir) / self._dir_name return Path(self._data_dir) / self._dir_name
def _date_path(self, created_at: datetime) -> Path: 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" return self.dir_path / f"{created_at.date().isoformat()}.json"
def _legacy_path(self) -> Path: def _legacy_path(self) -> Path:
-1
View File
@@ -45,7 +45,6 @@ class FetchUrlResult(BaseModel):
status_code: int status_code: int
content_type: str content_type: str
body: str body: str
truncated: bool = False
class ApplicationSummary(BaseModel): class ApplicationSummary(BaseModel):
-1
View File
@@ -549,7 +549,6 @@ def _update_job_file(req: UpdateJobFileRequest):
"guardrails — the allowlist is the only gate — so keep it tight. The " "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 " "Connection's saved auth is reused, so the agent does not need cluster "
"credentials.\n\n" "credentials.\n\n"
"**Limits:** response body capped at 1 MB (truncated=true if larger), "
"30s timeout, redirects followed." "30s timeout, redirects followed."
), ),
) )
+27 -57
View File
@@ -21,7 +21,6 @@ from spark_executor.core.yarn_client import YarnClientConfig
from spark_executor.models import FetchUrlResult from spark_executor.models import FetchUrlResult
from spark_executor.tools.connections import store as conn_store 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 _REQUEST_TIMEOUT_SECONDS = 30
_MAX_REDIRECTS = 3 _MAX_REDIRECTS = 3
@@ -70,29 +69,6 @@ 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: 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}")
@@ -107,45 +83,39 @@ def fetch_url(url: str, connection_name: str) -> FetchUrlResult:
verify = config.verify_for_httpx() verify = config.verify_for_httpx()
for hop in range(_MAX_REDIRECTS + 1): for hop in range(_MAX_REDIRECTS + 1):
with httpx.stream( resp = httpx.get(
"GET",
url, url,
auth=auth, auth=auth,
verify=verify, verify=verify,
timeout=_REQUEST_TIMEOUT_SECONDS, timeout=_REQUEST_TIMEOUT_SECONDS,
follow_redirects=False, follow_redirects=False,
) as resp: )
if resp.status_code not in _REDIRECT_STATUSES: if resp.status_code not in _REDIRECT_STATUSES:
body_bytes, truncated = _read_streamed_body(resp) logger.info(
encoding = resp.encoding or "utf-8" f"fetch_url ok url={url} connection_name={connection_name} "
body = body_bytes.decode(encoding, errors="replace") f"status_code={resp.status_code} "
f"content_type={resp.headers.get('content-type', '')} "
logger.info( f"body_chars={len(resp.text)}"
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 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( raise ValueError(
f"fetch_url exceeded {_MAX_REDIRECTS} redirects (last url: {url})" f"fetch_url exceeded {_MAX_REDIRECTS} redirects (last url: {url})"
+4 -5
View File
@@ -6,9 +6,8 @@
import os import os
import secrets import secrets
import uuid import uuid
from datetime import datetime from datetime import datetime, timezone
from common.config import settings
from common.logging import logger from common.logging import logger
from common.sql_guard import validate_pyspark_code from common.sql_guard import validate_pyspark_code
from spark_executor.core.connection_store import store as conn_store from spark_executor.core.connection_store import store as conn_store
@@ -124,7 +123,7 @@ def prepare_submit_job(
num_executors=num_executors, num_executors=num_executors,
spark_conf=dict(conn.spark_conf), spark_conf=dict(conn.spark_conf),
extra_args=dict(extra_args or {}), extra_args=dict(extra_args or {}),
created_at=datetime.utcnow(), created_at=datetime.now(timezone.utc),
status="PENDING", status="PENDING",
) )
pending_store.save(pending) pending_store.save(pending)
@@ -248,7 +247,7 @@ def confirm_submit_job(*, pending_id: str) -> SubmitResult:
application_id=application_id, application_id=application_id,
script_path=pending.script_path, script_path=pending.script_path,
queue=pending.queue, queue=pending.queue,
submit_time=datetime.utcnow(), submit_time=datetime.now(timezone.utc),
connection=pending.connection, connection=pending.connection,
yarn_rm_url=pending.yarn_rm_url, yarn_rm_url=pending.yarn_rm_url,
) )
@@ -282,7 +281,7 @@ def confirm_submit_job(*, pending_id: str) -> SubmitResult:
application_id=application_id, application_id=application_id,
script_path=pending.script_path, script_path=pending.script_path,
queue=pending.queue, queue=pending.queue,
submit_time=datetime.utcnow(), submit_time=datetime.now(timezone.utc),
connection=pending.connection, connection=pending.connection,
yarn_rm_url=pending.yarn_rm_url, yarn_rm_url=pending.yarn_rm_url,
) )
+16 -91
View File
@@ -13,18 +13,6 @@ from spark_executor.tools import connections, fetch_url
from spark_executor.tools.requests import ListApplicationsRequest from spark_executor.tools.requests import ListApplicationsRequest
def _stream_cm(resp):
"""Wrap a response (or fake response) in a context manager for httpx.stream."""
class _CM:
def __enter__(self):
return resp
def __exit__(self, exc_type, exc, tb):
return False
return _CM()
def _fresh_stores(tmp_path, monkeypatch): def _fresh_stores(tmp_path, monkeypatch):
"""Reset connection store singletons for a single test.""" """Reset connection store singletons for a single test."""
monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path))
@@ -52,14 +40,14 @@ def test_fetch_url_returns_body_and_status(fresh_stores):
) )
resp = httpx.Response(200, text="hello", headers={"content-type": "text/html"}) resp = httpx.Response(200, text="hello", headers={"content-type": "text/html"})
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(resp) "spark_executor.tools.fetch_url.httpx.get", return_value=resp
) as m: ) as m:
out = fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "prod") out = fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "prod")
assert out.url == "http://nm01.prod.internal:8042/node" assert out.url == "http://nm01.prod.internal:8042/node"
assert out.status_code == 200 assert out.status_code == 200
assert out.content_type == "text/html" assert out.content_type == "text/html"
assert out.body == "hello" assert out.body == "hello"
assert out.truncated is False assert "truncated" not in fetch_url.FetchUrlResult.model_fields
assert m.call_count == 1 assert m.call_count == 1
@@ -68,25 +56,6 @@ def test_fetch_url_raises_for_missing_connection(fresh_stores):
fetch_url.fetch_url("http://rm.prod.internal:8088/", "missing") fetch_url.fetch_url("http://rm.prod.internal:8088/", "missing")
def test_fetch_url_truncates_body_over_1mb(fresh_stores):
fetch_url.conn_store.save(
Connection(
name="prod",
master="yarn",
yarn_rm_url="http://rm.prod.internal:8088",
url_allowlist=["*.prod.internal"],
)
)
big = "x" * (fetch_url._MAX_BODY_BYTES + 1)
resp = httpx.Response(200, text=big)
with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(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(fresh_stores): def test_fetch_url_passes_auth_from_connection(fresh_stores):
fetch_url.conn_store.save( fetch_url.conn_store.save(
Connection( Connection(
@@ -101,7 +70,7 @@ def test_fetch_url_passes_auth_from_connection(fresh_stores):
) )
resp = httpx.Response(200, text="ok") resp = httpx.Response(200, text="ok")
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(resp) "spark_executor.tools.fetch_url.httpx.get", return_value=resp
) as m: ) as m:
fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "auth") fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "auth")
auth = m.call_args.kwargs["auth"] auth = m.call_args.kwargs["auth"]
@@ -124,7 +93,7 @@ def test_fetch_url_passes_ssl_verify_from_connection(fresh_stores):
) )
resp = httpx.Response(200, text="ok") resp = httpx.Response(200, text="ok")
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(resp) "spark_executor.tools.fetch_url.httpx.get", return_value=resp
) as m: ) as m:
fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "insecure") fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "insecure")
assert m.call_args.kwargs["verify"] is False assert m.call_args.kwargs["verify"] is False
@@ -141,7 +110,7 @@ def test_fetch_url_allows_host_matching_glob_pattern(fresh_stores):
) )
resp = httpx.Response(200, text="hello") resp = httpx.Response(200, text="hello")
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(resp) "spark_executor.tools.fetch_url.httpx.get", return_value=resp
) as m: ) as m:
out = fetch_url.fetch_url("http://ccam50:8088/foo", "ccam") out = fetch_url.fetch_url("http://ccam50:8088/foo", "ccam")
assert out.status_code == 200 assert out.status_code == 200
@@ -160,7 +129,7 @@ def test_fetch_url_allows_host_matching_any_of_multiple_globs(fresh_stores):
) )
resp = httpx.Response(200, text="hello") resp = httpx.Response(200, text="hello")
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(resp) "spark_executor.tools.fetch_url.httpx.get", return_value=resp
) as m: ) as m:
out = fetch_url.fetch_url( out = fetch_url.fetch_url(
"http://history.prod.internal:18080/api/v1/info", "prod" "http://history.prod.internal:18080/api/v1/info", "prod"
@@ -207,7 +176,7 @@ def test_fetch_url_glob_match_does_not_require_suffix_overlap(fresh_stores):
) )
resp = httpx.Response(200, text="hello") resp = httpx.Response(200, text="hello")
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(resp) "spark_executor.tools.fetch_url.httpx.get", return_value=resp
) as m: ) as m:
out = fetch_url.fetch_url("http://ccam50:8088/", "prod") out = fetch_url.fetch_url("http://ccam50:8088/", "prod")
assert out.status_code == 200 assert out.status_code == 200
@@ -225,7 +194,7 @@ def test_fetch_url_accepts_ip_literal_when_in_url_allowlist(fresh_stores):
) )
resp = httpx.Response(200, text="hello") resp = httpx.Response(200, text="hello")
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(resp) "spark_executor.tools.fetch_url.httpx.get", return_value=resp
) as m: ) as m:
out = fetch_url.fetch_url("http://10.0.0.1/secret", "prod") out = fetch_url.fetch_url("http://10.0.0.1/secret", "prod")
assert out.status_code == 200 assert out.status_code == 200
@@ -243,7 +212,7 @@ def test_fetch_url_accepts_https_when_in_url_allowlist(fresh_stores):
) )
resp = httpx.Response(200, text="hello") resp = httpx.Response(200, text="hello")
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", return_value=_stream_cm(resp) "spark_executor.tools.fetch_url.httpx.get", return_value=resp
) as m: ) as m:
out = fetch_url.fetch_url("https://ccam1.example.com/secure", "prod") out = fetch_url.fetch_url("https://ccam1.example.com/secure", "prod")
assert out.status_code == 200 assert out.status_code == 200
@@ -301,8 +270,8 @@ def test_fetch_url_rejects_redirect_to_disallowed_host(fresh_stores):
) )
redirect = httpx.Response(302, headers={"Location": "http://evil.com/"}) redirect = httpx.Response(302, headers={"Location": "http://evil.com/"})
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", "spark_executor.tools.fetch_url.httpx.get",
return_value=_stream_cm(redirect), return_value=redirect,
) as m: ) as m:
with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): with pytest.raises(ValueError, match="is not in Connection.url_allowlist"):
fetch_url.fetch_url("http://ccam50/foo", "ccam") fetch_url.fetch_url("http://ccam50/foo", "ccam")
@@ -323,14 +292,14 @@ def test_fetch_url_follows_redirect_to_allowed_host(fresh_stores):
) )
final = httpx.Response(200, text="ok") final = httpx.Response(200, text="ok")
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", "spark_executor.tools.fetch_url.httpx.get",
side_effect=[_stream_cm(redirect), _stream_cm(final)], side_effect=[redirect, final],
) as m: ) as m:
out = fetch_url.fetch_url("http://foo.internal/", "prod") out = fetch_url.fetch_url("http://foo.internal/", "prod")
assert out.status_code == 200 assert out.status_code == 200
assert out.body == "ok" assert out.body == "ok"
assert m.call_count == 2 assert m.call_count == 2
assert m.call_args_list[1].args[1] == "http://other.internal/" 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): def test_fetch_url_rejects_redirect_to_ip_literal_not_in_allowlist(fresh_stores):
@@ -344,58 +313,14 @@ def test_fetch_url_rejects_redirect_to_ip_literal_not_in_allowlist(fresh_stores)
) )
redirect = httpx.Response(302, headers={"Location": "http://10.0.0.1/"}) redirect = httpx.Response(302, headers={"Location": "http://10.0.0.1/"})
with patch( with patch(
"spark_executor.tools.fetch_url.httpx.stream", "spark_executor.tools.fetch_url.httpx.get",
return_value=_stream_cm(redirect), return_value=redirect,
) as m: ) as m:
with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): with pytest.raises(ValueError, match="is not in Connection.url_allowlist"):
fetch_url.fetch_url("http://ccam50/foo", "ccam") fetch_url.fetch_url("http://ccam50/foo", "ccam")
assert m.call_count == 1 assert m.call_count == 1
def test_fetch_url_streams_body_and_truncates_without_buffering_full(fresh_stores):
"""The response body must be read incrementally; .text must not be accessed."""
class _FakeStreamResponse:
def __init__(self, chunks):
self.status_code = 200
self.headers = httpx.Headers({"content-type": "text/plain"})
self.encoding = "utf-8"
self._chunks = chunks
def iter_bytes(self):
yield from self._chunks
@property
def text(self):
raise AssertionError("response.text should not be accessed")
fetch_url.conn_store.save(
Connection(
name="prod",
master="yarn",
url_allowlist=["*.prod.internal"],
)
)
chunk_size = 400_000
chunks = [
b"a" * chunk_size,
b"b" * chunk_size,
b"c" * chunk_size,
]
fake = _FakeStreamResponse(chunks)
with patch(
"spark_executor.tools.fetch_url.httpx.stream",
return_value=_stream_cm(fake),
) as m:
out = fetch_url.fetch_url("http://nm01.prod.internal/big", "prod")
assert m.call_count == 1
assert out.truncated is True
assert len(out.body) == fetch_url._MAX_BODY_BYTES
assert out.body.startswith("a" * chunk_size)
# Only the first 200 KB of the third chunk were consumed before truncation.
assert out.body[-200_000:] == "c" * 200_000
def test_list_applications_request_limit_bounds(): def test_list_applications_request_limit_bounds():
assert ListApplicationsRequest(connection_name="prod", limit=10000).limit == 10000 assert ListApplicationsRequest(connection_name="prod", limit=10000).limit == 10000
with pytest.raises(ValidationError): with pytest.raises(ValidationError):