From 2f27eb3f3606be726cf58fe1c3ffcd66c3216ffd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 1 Oct 2026 17:02:10 -0700 Subject: [PATCH] test(bedrock): harden the runtime chat completions audit cells The chaos peer's shared counter and process now come from the same spawn context, since a fork-context Value handed to a spawn-context process raises on Linux. The peer-kill test waits for the first six answers to reach the client before killing the peer instead of counting accepted requests. The Responses wire tests look the spend row up under both the ciphertext id the caller received and the issued id behind it, matching the chaos file's rule for the pre-encryption row --- .../test_bedrock_gpt_responses_native_wire.py | 18 ++++++++---- ..._bedrock_runtime_chat_completions_chaos.py | 29 +++++++++++++------ 2 files changed, 32 insertions(+), 15 deletions(-) diff --git a/tests/integration/providers/test_bedrock_gpt_responses_native_wire.py b/tests/integration/providers/test_bedrock_gpt_responses_native_wire.py index 6a365ac26ba..3d6eb7fcaff 100644 --- a/tests/integration/providers/test_bedrock_gpt_responses_native_wire.py +++ b/tests/integration/providers/test_bedrock_gpt_responses_native_wire.py @@ -58,11 +58,16 @@ def _body(request: Request) -> dict[str, JsonValue]: return _JSON_OBJECT.validate_json(request.body) -def _spend_row(identity: str) -> dict[str, JsonValue]: +# TODO: a Bedrock non-stream /v1/responses spend row can carry the pre-encryption resp_ id instead of the +# ciphertext the caller received, because the spend row id is read from response_obj["id"] before the +# ResponsesIDSecurity hook rewrites it in place; the row is looked up under both ids until that ordering is fixed on +# main +def _spend_row(client_id: str, issued_id: str) -> dict[str, JsonValue]: rows: Final = eventually( lambda: read_rows( - 'SELECT model_group, status, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s', - (identity,), + 'SELECT model_group, status, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" ' + "WHERE request_id = ANY(%s)", + ([client_id, issued_id],), # pyright: ignore[reportArgumentType] # psycopg adapts the list to a text array ), lambda found: len(found) == 1, seconds=70, @@ -85,10 +90,11 @@ def test_openai_sdk_responses_request_is_served_by_the_native_responses_route(ga response: Final = raw.parse() assert response.output_text == answer(marker), raw.text assert response.usage is not None and (response.usage.input_tokens, response.usage.output_tokens) == (30, 5) - assert _issued_id(response.id).upstream == f"resp_upstream_{marker}", response.id + issued: Final = _issued_id(response.id) + assert issued.upstream == f"resp_upstream_{marker}", response.id request: Final = _native_request(wire) assert _body(request) == {"model": GPT, "input": _prompt(marker)}, request.body - assert _spend_row(response.id) == _success_row(model) + assert _spend_row(response.id, issued.issued) == _success_row(model) async def test_async_openai_sdk_responses_stream_is_served_by_the_native_responses_route(gateway: Gateway) -> None: @@ -114,4 +120,4 @@ async def test_async_openai_sdk_responses_stream_is_served_by_the_native_respons assert issued.upstream == f"resp_upstream_{marker}", completed.response.id request: Final = _native_request(wire) assert _body(request) == {"model": GPT, "input": _prompt(marker), "stream": True}, request.body - assert _spend_row(issued.issued) == _success_row(model) + assert _spend_row(completed.response.id, issued.issued) == _success_row(model) diff --git a/tests/integration/providers/test_bedrock_runtime_chat_completions_chaos.py b/tests/integration/providers/test_bedrock_runtime_chat_completions_chaos.py index 030a407bd8d..ef6e0ce60f2 100644 --- a/tests/integration/providers/test_bedrock_runtime_chat_completions_chaos.py +++ b/tests/integration/providers/test_bedrock_runtime_chat_completions_chaos.py @@ -1,6 +1,7 @@ import asyncio import base64 import binascii +import itertools import multiprocessing import os import re @@ -200,6 +201,19 @@ async def _burst( return tuple(result for result in results if isinstance(result, _Served)) +async def _burst_killing_the_peer_once_it_answered( + base_url: str, key: str, model: str, calls: tuple[_Call, ...], peer: _ChildPeer, answered: int +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client: + tasks: Final = tuple(asyncio.create_task(_send(client, key, model, call)) for call in calls) + await asyncio.to_thread(eventually, lambda: peer.received.value, lambda count: count == len(calls), 60) + first: Final = [await finished for finished in itertools.islice(asyncio.as_completed(tasks), answered)] + assert all(item.status == 200 for item in first), [(item.call.marker, item.status) for item in first] + peer.process.kill() + peer.process.join(timeout=10) + return tuple(await asyncio.gather(*tasks)) + + def _calls(count: int, endpoints: tuple[Endpoint, ...], stream: Callable[[int], bool]) -> tuple[_Call, ...]: return tuple( _Call(endpoint=endpoints[index % len(endpoints)], stream=stream(index), marker=uuid.uuid4().hex) @@ -223,10 +237,9 @@ def _accepts_connections(port: int) -> bool: @contextmanager def _child_peer(port: int, answer_first: int) -> Iterator[_ChildPeer]: - received: Final = multiprocessing.Value("i", 0) - process: Final = multiprocessing.get_context("spawn").Process( - target=serve_peer, args=(port, received, answer_first), daemon=True - ) + context: Final = multiprocessing.get_context("spawn") + received: Final = context.Value("i", 0) + process: Final = context.Process(target=serve_peer, args=(port, received, answer_first), daemon=True) process.start() try: eventually(lambda: _accepts_connections(port), bool, seconds=30) @@ -263,11 +276,9 @@ async def test_peer_killed_mid_burst_fails_only_the_held_calls_and_a_restarted_p with gateway.scenario() as scenario: model: Final = _deployment(scenario, f"http://127.0.0.1:{port}") with _child_peer(port, answer_first=6) as peer: - burst: Final = asyncio.create_task(_burst(str(gateway.client.base_url), gateway.key, model, calls)) - await asyncio.to_thread(eventually, lambda: peer.received.value, lambda count: count == 12, 60) - peer.process.kill() - peer.process.join(timeout=10) - served: Final = await burst + served: Final = await _burst_killing_the_peer_once_it_answered( + str(gateway.client.base_url), gateway.key, model, calls, peer, answered=6 + ) succeeded: Final = tuple(item for item in served if item.status == 200) failed: Final = tuple(item for item in served if item.status != 200) assert (len(succeeded), len(failed)) == (6, 6), [(item.call.marker, item.status) for item in served]