Initial public ModelForge release
This commit is contained in:
@@ -0,0 +1,154 @@
|
||||
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
|
||||
Reference in New Issue
Block a user