Files
ModelForge/backend/tests/test_request_limits.py

155 lines
5.2 KiB
Python

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