from __future__ import annotations import json from dataclasses import dataclass import pytest from starlette.responses import JSONResponse from starlette.types import Message, Receive, Scope, Send from modelforge_api.api.request_limits import ( MAX_CONSECUTIVE_EMPTY_REQUEST_EVENTS, MAX_REQUEST_BODY_EVENTS, MAX_TOTAL_EMPTY_REQUEST_EVENTS, RequestBodyLimitMiddleware, ) @dataclass(frozen=True, slots=True) class LimitResult: status_code: int body: dict[str, object] receive_calls: int async def _consume_body(scope: Scope, receive: Receive, send: Send) -> None: received_bytes = 0 while True: message = await receive() if message["type"] == "http.disconnect": await JSONResponse( status_code=400, content={"control": "disconnect"}, )(scope, receive, send) return received_bytes += len(message.get("body", b"")) if not message.get("more_body", False): break await JSONResponse(status_code=200, content={"received_bytes": received_bytes})( scope, receive, send ) async def _exercise(messages: list[Message], *, limit_bytes: int = 65_536) -> LimitResult: receive_calls = 0 async def receive() -> Message: nonlocal receive_calls if receive_calls >= len(messages): raise AssertionError("middleware received beyond the supplied event budget") message = messages[receive_calls] receive_calls += 1 return message sent: list[Message] = [] async def send(message: Message) -> None: sent.append(message) scope: Scope = { "type": "http", "asgi": {"version": "3.0"}, "http_version": "1.1", "method": "POST", "scheme": "http", "path": "/bounded", "raw_path": b"/bounded", "query_string": b"", "root_path": "", "headers": [], "client": ("test", 1), "server": ("test", 80), "state": { "correlation_id": "limit-test", "request_body_limit_bytes": limit_bytes, }, } await RequestBodyLimitMiddleware(_consume_body)(scope, receive, send) start = next(message for message in sent if message["type"] == "http.response.start") body = b"".join( message.get("body", b"") for message in sent if message["type"] == "http.response.body" ) return LimitResult( status_code=int(start["status"]), body=json.loads(body), receive_calls=receive_calls, ) @pytest.mark.asyncio async def test_consecutive_empty_request_events_stop_at_the_fixed_cutoff() -> None: messages: list[Message] = [ {"type": "http.request", "body": b"", "more_body": True} for _ in range(MAX_CONSECUTIVE_EMPTY_REQUEST_EVENTS + 1) ] messages.append({"type": "http.request", "body": b"unreachable", "more_body": False}) result = await _exercise(messages) assert result.status_code == 400 assert result.body["error"]["code"] == "request_body_progress_exhausted" assert result.receive_calls == MAX_CONSECUTIVE_EMPTY_REQUEST_EVENTS + 1 @pytest.mark.asyncio async def test_total_empty_request_events_are_bounded_even_with_intermittent_progress() -> None: messages: list[Message] = [] for _ in range(MAX_TOTAL_EMPTY_REQUEST_EVENTS + 1): messages.append({"type": "http.request", "body": b"", "more_body": True}) messages.append({"type": "http.request", "body": b"x", "more_body": True}) messages.append({"type": "http.request", "body": b"unreachable", "more_body": False}) result = await _exercise(messages) assert result.status_code == 400 assert result.body["error"]["code"] == "request_body_progress_exhausted" assert result.receive_calls == (MAX_TOTAL_EMPTY_REQUEST_EVENTS * 2) + 1 @pytest.mark.asyncio async def test_total_request_event_count_bounds_tiny_nonempty_progress() -> None: messages: list[Message] = [ {"type": "http.request", "body": b"x", "more_body": True} for _ in range(MAX_REQUEST_BODY_EVENTS + 1) ] messages.append({"type": "http.request", "body": b"unreachable", "more_body": False}) result = await _exercise(messages) assert result.status_code == 400 assert result.body["error"]["code"] == "request_body_progress_exhausted" assert result.receive_calls == MAX_REQUEST_BODY_EVENTS + 1 @pytest.mark.asyncio async def test_finite_empty_events_normal_chunks_and_disconnect_remain_valid_controls() -> None: finite_empty = [ {"type": "http.request", "body": b"", "more_body": True} for _ in range(MAX_CONSECUTIVE_EMPTY_REQUEST_EVENTS) ] completed = await _exercise( [ *finite_empty, {"type": "http.request", "body": b"ab", "more_body": True}, {"type": "http.request", "body": b"cd", "more_body": False}, ] ) disconnected = await _exercise([{"type": "http.disconnect"}]) assert completed.status_code == 200 assert completed.body == {"received_bytes": 4} assert completed.receive_calls == MAX_CONSECUTIVE_EMPTY_REQUEST_EVENTS + 2 assert disconnected.status_code == 400 assert disconnected.body == {"control": "disconnect"} assert disconnected.receive_calls == 1