155 lines
5.2 KiB
Python
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
|