diff --git a/tests/test_main_observability.py b/tests/test_main_observability.py index 3b39a59..27f0434 100644 --- a/tests/test_main_observability.py +++ b/tests/test_main_observability.py @@ -175,6 +175,63 @@ async def asgi_fetch(*args, **kwargs): self.assertEqual(emitted[-1]["status_code"], 413) self.assertEqual(emitted[-1]["outcome"], "client_error") + def test_every_bridge_call_caps_the_request_body_at_the_submit_limit(self): + # Covers the three asgi.fetch call sites: cached GET miss, uncacheable + # GET, and POST. Each must hand the bridge the same cap run_example + # enforces, or the friendly "too large" page becomes unreachable. + main = self.import_main() + bridge_kwargs = [] + + class FakeCache: + async def match(self, key): + return None + + async def put(self, key, value): + pass + + class FakeJsRequest: + @staticmethod + def new(url, init=None): + return url + + async def asgi_fetch(*args, **kwargs): + bridge_kwargs.append(kwargs) + response = SimpleNamespace(status=200, headers=_Headers()) + response.clone = lambda: response + return response + + main.caches = SimpleNamespace(default=FakeCache()) + main.JsRequest = FakeJsRequest + main.asgi.fetch = asgi_fetch + main.observability.emit = lambda value, env=None: None + worker = main.Default() + worker.env = SimpleNamespace() + + for calls, (method, path, cache) in enumerate( + [ + ("GET", "/examples/values", "miss"), + ("GET", "/layout-options/a", "bypass"), + ("POST", "/examples/values", "bypass"), + ], + start=1, + ): + with self.subTest(method=method, path=path): + event = {"cache": "bypass"} + main.observability.event_from_worker_request = ( + lambda *args, event=event, **kwargs: event + ) + request = SimpleNamespace( + method=method, + url=f"https://www.pythonbyexample.dev{path}", + js_object=SimpleNamespace(), + ) + asyncio.run(worker.fetch(request)) + self.assertEqual(event["cache"], cache) + self.assertEqual(len(bridge_kwargs), calls) + self.assertEqual( + bridge_kwargs[-1].get("max_body_bytes"), main.MAX_SUBMITTED_BODY_BYTES + ) + def test_dynamic_response_reader_cancels_stream_above_output_cap(self): main = self.import_main() cancelled = [] diff --git a/tests/test_worker_asgi_bridge_scope.py b/tests/test_worker_asgi_bridge_scope.py index df3451a..d3f3ad7 100644 --- a/tests/test_worker_asgi_bridge_scope.py +++ b/tests/test_worker_asgi_bridge_scope.py @@ -1,6 +1,5 @@ import asyncio import importlib -import pathlib import sys import types import unittest @@ -8,8 +7,6 @@ from typing import ClassVar from urllib.parse import urlparse -ROOT = pathlib.Path(__file__).resolve().parents[1] - class WorkerAsgiBridgeScopeTests(unittest.TestCase): def setUp(self): @@ -243,6 +240,26 @@ async def app(scope, receive, send): # pragma: no cover - must not run self.assertEqual(response.status, 413) self.assertFalse(ran["app"], "oversize body must be rejected before the app runs") + def test_fetch_forwards_body_cap_to_request_processing(self): + ran = {"app": False} + + async def app(scope, receive, send): # pragma: no cover - must not run + ran["app"] = True + + async def start_application(_app): + pass + + self.bridge._app_lifespans.clear() + self.bridge.start_application = start_application + req = self._fake_request( + url="https://x.dev/examples/values", + headers={"content-type": "text/plain"}, + body_chunks=[b"toolong"], + ) + response = asyncio.run(self.bridge.fetch(app, req, SimpleNamespace(), max_body_bytes=3)) + self.assertEqual(response.status, 413) + self.assertFalse(ran["app"], "fetch must apply the cap before the app runs") + def test_request_to_scope_accepts_state_without_request_globals_or_extra_scope(self): class FakeRequest(self.Request): method = "POST" @@ -347,14 +364,6 @@ async def scenario(): asyncio.run(scenario()) self.assertEqual(calls, {"startup": 1, "request": 2}) - def test_bridge_has_pre_asgi_body_cap_hook(self): - bridge_source = (ROOT / "src" / "worker_asgi_bridge.py").read_text() - main_source = (ROOT / "src" / "main.py").read_text() - - self.assertIn("max_body_bytes", bridge_source) - self.assertIn("body_bytes > max_body_bytes", bridge_source) - self.assertIn("max_body_bytes=MAX_SUBMITTED_BODY_BYTES", main_source) - if __name__ == "__main__": unittest.main()