Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 57 additions & 0 deletions tests/test_main_observability.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = []
Expand Down
31 changes: 20 additions & 11 deletions tests/test_worker_asgi_bridge_scope.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,12 @@
import asyncio
import importlib
import pathlib
import sys
import types
import unittest
from types import SimpleNamespace
from typing import ClassVar
from urllib.parse import urlparse

ROOT = pathlib.Path(__file__).resolve().parents[1]


class WorkerAsgiBridgeScopeTests(unittest.TestCase):
def setUp(self):
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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()
Loading