diff --git a/spark_executor/core/job_writer.py b/spark_executor/core/job_writer.py index 27df059..796e2a2 100644 --- a/spark_executor/core/job_writer.py +++ b/spark_executor/core/job_writer.py @@ -14,7 +14,7 @@ Resolution order for the output directory: """ import os import secrets -from datetime import datetime +from datetime import datetime, timezone from common.config import settings 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) 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" path = os.path.join(effective_dir, name) abs_path = os.path.abspath(path) diff --git a/spark_executor/core/pending_store.py b/spark_executor/core/pending_store.py index 5bfccde..0436d54 100644 --- a/spark_executor/core/pending_store.py +++ b/spark_executor/core/pending_store.py @@ -37,7 +37,7 @@ class PendingStore: return Path(self._data_dir) / self._dir_name 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" def _legacy_path(self) -> Path: diff --git a/spark_executor/models.py b/spark_executor/models.py index b389021..9e35af8 100644 --- a/spark_executor/models.py +++ b/spark_executor/models.py @@ -45,7 +45,6 @@ class FetchUrlResult(BaseModel): status_code: int content_type: str body: str - truncated: bool = False class ApplicationSummary(BaseModel): diff --git a/spark_executor/server.py b/spark_executor/server.py index e43efb0..4b26128 100644 --- a/spark_executor/server.py +++ b/spark_executor/server.py @@ -549,7 +549,6 @@ def _update_job_file(req: UpdateJobFileRequest): "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 " "credentials.\n\n" - "**Limits:** response body capped at 1 MB (truncated=true if larger), " "30s timeout, redirects followed." ), ) diff --git a/spark_executor/tools/fetch_url.py b/spark_executor/tools/fetch_url.py index 4813c30..9615372 100644 --- a/spark_executor/tools/fetch_url.py +++ b/spark_executor/tools/fetch_url.py @@ -21,7 +21,6 @@ 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 @@ -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: """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}") @@ -107,45 +83,39 @@ def fetch_url(url: str, connection_name: str) -> FetchUrlResult: verify = config.verify_for_httpx() for hop in range(_MAX_REDIRECTS + 1): - with httpx.stream( - "GET", + resp = httpx.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}" + ) + if resp.status_code not in _REDIRECT_STATUSES: + 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_chars={len(resp.text)}" ) - _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( f"fetch_url exceeded {_MAX_REDIRECTS} redirects (last url: {url})" diff --git a/spark_executor/tools/submit.py b/spark_executor/tools/submit.py index 10ff531..04c2bda 100644 --- a/spark_executor/tools/submit.py +++ b/spark_executor/tools/submit.py @@ -6,9 +6,8 @@ import os import secrets import uuid -from datetime import datetime +from datetime import datetime, timezone -from common.config import settings from common.logging import logger from common.sql_guard import validate_pyspark_code from spark_executor.core.connection_store import store as conn_store @@ -124,7 +123,7 @@ def prepare_submit_job( num_executors=num_executors, spark_conf=dict(conn.spark_conf), extra_args=dict(extra_args or {}), - created_at=datetime.utcnow(), + created_at=datetime.now(timezone.utc), status="PENDING", ) pending_store.save(pending) @@ -248,7 +247,7 @@ def confirm_submit_job(*, pending_id: str) -> SubmitResult: application_id=application_id, script_path=pending.script_path, queue=pending.queue, - submit_time=datetime.utcnow(), + submit_time=datetime.now(timezone.utc), connection=pending.connection, yarn_rm_url=pending.yarn_rm_url, ) @@ -282,7 +281,7 @@ def confirm_submit_job(*, pending_id: str) -> SubmitResult: application_id=application_id, script_path=pending.script_path, queue=pending.queue, - submit_time=datetime.utcnow(), + submit_time=datetime.now(timezone.utc), connection=pending.connection, yarn_rm_url=pending.yarn_rm_url, ) diff --git a/tests/unit/test_fetch_url.py b/tests/unit/test_fetch_url.py index 33d556d..affe030 100644 --- a/tests/unit/test_fetch_url.py +++ b/tests/unit/test_fetch_url.py @@ -13,18 +13,6 @@ from spark_executor.tools import connections, fetch_url 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): """Reset connection store singletons for a single test.""" 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"}) 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: 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 "truncated" not in fetch_url.FetchUrlResult.model_fields 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") -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): fetch_url.conn_store.save( Connection( @@ -101,7 +70,7 @@ def test_fetch_url_passes_auth_from_connection(fresh_stores): ) resp = httpx.Response(200, text="ok") 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: fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "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") 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: fetch_url.fetch_url("http://nm01.prod.internal:8042/node", "insecure") 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") 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: out = fetch_url.fetch_url("http://ccam50:8088/foo", "ccam") 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") 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: out = fetch_url.fetch_url( "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") 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: out = fetch_url.fetch_url("http://ccam50:8088/", "prod") 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") 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: out = fetch_url.fetch_url("http://10.0.0.1/secret", "prod") 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") 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: out = fetch_url.fetch_url("https://ccam1.example.com/secure", "prod") 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/"}) with patch( - "spark_executor.tools.fetch_url.httpx.stream", - return_value=_stream_cm(redirect), + "spark_executor.tools.fetch_url.httpx.get", + return_value=redirect, ) as m: with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): 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") with patch( - "spark_executor.tools.fetch_url.httpx.stream", - side_effect=[_stream_cm(redirect), _stream_cm(final)], + "spark_executor.tools.fetch_url.httpx.get", + side_effect=[redirect, final], ) as m: out = fetch_url.fetch_url("http://foo.internal/", "prod") assert out.status_code == 200 assert out.body == "ok" 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): @@ -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/"}) with patch( - "spark_executor.tools.fetch_url.httpx.stream", - return_value=_stream_cm(redirect), + "spark_executor.tools.fetch_url.httpx.get", + return_value=redirect, ) as m: with pytest.raises(ValueError, match="is not in Connection.url_allowlist"): fetch_url.fetch_url("http://ccam50/foo", "ccam") 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(): assert ListApplicationsRequest(connection_name="prod", limit=10000).limit == 10000 with pytest.raises(ValidationError):