refactor(types): replace Any with proven types in 9 files (#44389)

* refactor(types): replace Any with proven types in 13 files

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(types): keep provider error paths for malformed prefetch and poll JSON

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(types): keep main's login body parsing

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(types): keep main's Copilot auth and budget alert typing

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): audit cells for Any sweep 20261003_2

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): load audit video deployments at proxy start

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(types): keep main's video-edit prefetch handling

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): fix audit chaos Responses cells

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): probe every route after audit worker kill

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-03 21:03:16 +00:00 • committed by GitHub
parent fe683ea139
commit e9bf2cfd01
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1107 additions and 34 deletions

View file

@ -22,6 +22,7 @@ from typing import TYPE_CHECKING, Any, Final, Protocol
import httpx
import openai
from pydantic import ConfigDict, TypeAdapter
import litellm
from litellm.constants import (
@ -46,6 +47,9 @@ if TYPE_CHECKING:
from uvicorn import Server
verbose_logger: Final = logging.getLogger("LiteLLM")
_JSON_VALUE_ADAPTER: Final = TypeAdapter(object, config=ConfigDict(strict=True))
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True))
_PORT_ADAPTER: Final = TypeAdapter(int, config=ConfigDict(strict=True))
MISSING_DEPS_MESSAGE = "litellm.harness needs starlette and uvicorn: pip install starlette uvicorn"
@ -207,11 +211,11 @@ class SSEUsageParser:
if not payload or payload == b"[DONE]":
return
try:
event = json.loads(payload)
event: Final[object] = _JSON_VALUE_ADAPTER.validate_python(json.loads(payload))
except ValueError:
return
if isinstance(event, Mapping):
self.absorb(event)
self.absorb(_JSON_OBJECT_ADAPTER.validate_python(event))
def absorb(self, event: Mapping[str, object]) -> None:
event_type = event.get("type")
@ -434,7 +438,7 @@ class ModelEndpoint:
except BaseException:
await self.stop()
raise
self.port = self._server.servers[0].sockets[0].getsockname()[1]
self.port = _PORT_ADAPTER.validate_python(self._server.servers[0].sockets[0].getsockname()[1])
async def stop(self) -> None:
if self._server is not None:
@ -540,11 +544,12 @@ class ModelEndpoint:
if not self._authorized(request):
return self._unauthorized()
try:
body = json.loads(await request.body())
parsed_body: Final[object] = _JSON_VALUE_ADAPTER.validate_python(json.loads(await request.body()))
except ValueError as e:
return self._error(e, 400)
if not isinstance(body, dict):
if not isinstance(parsed_body, dict):
return self._error(ValueError("request body must be a JSON object"), 400)
body: Final = _JSON_OBJECT_ADAPTER.validate_python(parsed_body)
route = route_of(request.url.path)
if self.gateway is not None:
return await self._forward(request, route, body)
@ -615,7 +620,8 @@ class ModelEndpoint:
tokens = (parser.input_tokens, parser.output_tokens)
else:
try:
tokens = usage_from_body(json.loads(collected))
body: Final[object] = _JSON_VALUE_ADAPTER.validate_python(json.loads(collected))
tokens = usage_from_body(body)
except ValueError:
tokens = (0, 0)
self._record(model, tokens[0], tokens[1], header_cost(upstream.headers))

View file

@ -186,7 +186,7 @@ class Session:
def history(
self,
) -> list[dict[str, Any]]: # mutable-ok: public API returns OpenAI-format message dicts from the handler
) -> list[dict[str, object]]: # mutable-ok: public API returns OpenAI-format message dicts from the handler
return run_sync(self._inner.history(), "history")
@property

View file

@ -492,7 +492,7 @@ class GalileoObserve(CustomLogger):
@staticmethod
def _get_chat_content_for_galileo(response_obj: litellm.ModelResponse) -> object:
if response_obj.choices and len(response_obj.choices) > 0:
message: Final = response_obj["choices"][0]["message"]
message: Final = response_obj.choices[0].message
if hasattr(message, "json"):
message_json: Final[object] = message.json()
if isinstance(message_json, str):

View file

@ -27,7 +27,7 @@ from litellm.types.llms.openai import (
ResponsesAPIResponse,
)
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.utils import CallTypes, StandardCallbackDynamicParams
from litellm.types.utils import CallTypes, LLMResponseTypes, ModelResponse, StandardCallbackDynamicParams
from litellm.types.vector_stores import (
LiteLLM_ManagedVectorStore,
VectorStoreSearchFailure,
@ -46,6 +46,7 @@ else:
SEARCH_FAILURES_FIELD: Final = "vector_store_search_failures"
_DEFAULT_FAILURE_MODE: Final[VectorStoreSearchFailureMode] = "annotate"
_FAILURE_MODE_ADAPTER: Final = TypeAdapter(VectorStoreSearchFailureMode)
_OBJECT_ADAPTER: Final = TypeAdapter(object)
_STR_KEYED_ADAPTER: Final = TypeAdapter(dict[str, object])
_GUARDRAIL_KEYS_THE_PROXY_MERGES_INTO_METADATA: Final = frozenset(
{"guardrails", "guardrail_config", "policies", "include_guardrail_response"}
@ -269,7 +270,9 @@ class VectorStorePreCallHook(CustomLogger):
verbose_logger.debug("No query found in messages for vector store search")
return None
request_litellm_params: Final = litellm_logging_obj.model_call_details.get("litellm_params", {})
request_litellm_params: Final = _OBJECT_ADAPTER.validate_python(
litellm_logging_obj.model_call_details.get("litellm_params", {})
)
request_metadata: Final = (
request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {}
)
@ -415,9 +418,9 @@ class VectorStorePreCallHook(CustomLogger):
async def async_post_call_success_deployment_hook(
self,
request_data: dict,
response: Any,
response: LLMResponseTypes,
call_type: CallTypes | None,
) -> Any | None:
) -> LLMResponseTypes | None:
"""
Add search results to the response after successful LLM call.
@ -451,7 +454,7 @@ class VectorStorePreCallHook(CustomLogger):
return response
# Add search results to response object
if hasattr(response, "choices") and response.choices:
if isinstance(response, ModelResponse) and response.choices:
for choice in response.choices:
if hasattr(choice, "message") and choice.message:
provider_fields = getattr(choice.message, "provider_specific_fields", None) or {}

View file

@ -1741,11 +1741,14 @@ def _find_server_tool_result(
)
_AnthropicToolInvokeItem: TypeAlias = AnthropicMessagesToolUseParam | dict[str, Any]
def convert_to_anthropic_tool_invoke(
tool_calls: list[ChatCompletionAssistantToolCall],
web_search_results: Sequence[object] | None = None,
tool_results: Sequence[object] | None = None,
) -> list[AnthropicMessagesToolUseParam | dict[str, Any]]:
) -> list[_AnthropicToolInvokeItem]:
"""
OpenAI tool invokes:
{
@ -2565,9 +2568,9 @@ def anthropic_messages_pt(
# Group tool invoke results into (server_tool_use, result) pairs
# and separate regular tool_use blocks
server_tool_groups: list[list[Any]] = []
regular_tool_uses: list[Any] = []
_current_group: list[Any] = []
server_tool_groups: list[list[_AnthropicToolInvokeItem]] = []
regular_tool_uses: list[_AnthropicToolInvokeItem] = []
_current_group: list[_AnthropicToolInvokeItem] = []
for item in tool_invoke_results:
item_type = item.get("type", "") if isinstance(item, dict) else getattr(item, "type", "")
if item_type == "server_tool_use":

View file

@ -129,7 +129,7 @@ class ClaudeCodeStreamState:
result_text: str | None = None
is_error: bool = False
errors: Sequence[str] = ()
structured_output: Any | None = None
structured_output: object = None
@property
def final_text(self) -> str:
@ -157,7 +157,7 @@ def _stringify_block(block: object) -> str:
return json.dumps(block, ensure_ascii=False)
def _message_blocks(event: Mapping[str, Any]) -> Sequence[Any]:
def _message_blocks(event: Mapping[str, object]) -> Sequence[Any]:
message: Final = event.get("message")
content: Final = message.get("content") if isinstance(message, Mapping) else None
if isinstance(content, str):
@ -187,7 +187,7 @@ def _assistant_block_events(block: Mapping[str, Any], state: ClaudeCodeStreamSta
return ()
def _assistant_events(event: Mapping[str, Any], state: ClaudeCodeStreamState) -> Sequence[Event]:
def _assistant_events(event: Mapping[str, object], state: ClaudeCodeStreamState) -> Sequence[Event]:
if event.get("parent_tool_use_id"):
return event_list() # subagent traffic
message: Final = event.get("message")
@ -201,7 +201,7 @@ def _is_tool_result(block: object) -> bool:
return isinstance(block, dict) and block.get("type") == "tool_result"
def _user_events(event: Mapping[str, Any]) -> Sequence[Event]:
def _user_events(event: Mapping[str, object]) -> Sequence[Event]:
if event.get("parent_tool_use_id"):
return event_list()
return event_list(
@ -217,7 +217,7 @@ def _user_events(event: Mapping[str, Any]) -> Sequence[Event]:
)
def _system_events(event: Mapping[str, Any], state: ClaudeCodeStreamState) -> Sequence[Event]:
def _system_events(event: Mapping[str, object], state: ClaudeCodeStreamState) -> Sequence[Event]:
subtype = event.get("subtype")
if subtype == "init" and event.get("session_id"):
state.session_id = str(event["session_id"])
@ -253,7 +253,7 @@ def turn_error_message(state: ClaudeCodeStreamState, exit_code: int, stderr_tail
return f"{message}\nstderr:\n{tail}" if tail else message
def build_system_prompt(instructions: str | None, output_schema: Mapping[str, Any] | None) -> str | None:
def build_system_prompt(instructions: str | None, output_schema: Mapping[str, object] | None) -> str | None:
schema_part: Final = (
STRUCTURED_OUTPUT_INSTRUCTION.format(schema=json.dumps(output_schema)) if output_schema is not None else None
)
@ -355,7 +355,7 @@ class ClaudeCodeHarnessConfig(BaseCLIHarnessConfig):
def create_stream_state(self) -> ClaudeCodeStreamState:
return ClaudeCodeStreamState()
def transform_stream_line(self, line: Mapping[str, Any], state: ClaudeCodeStreamState) -> Sequence[Event]:
def transform_stream_line(self, line: Mapping[str, object], state: ClaudeCodeStreamState) -> Sequence[Event]:
kind = line.get("type")
if kind == "assistant":
return _assistant_events(line, state)

View file

@ -117,7 +117,7 @@ class OllamaChatConfig(BaseConfig):
system: str | None = None,
template: str | None = None,
) -> None:
locals_: Final = locals().copy()
locals_: Final[dict[str, object]] = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)

View file

@ -13,7 +13,7 @@ import re
import secrets
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import Any, Final, NamedTuple, Protocol, Union, cast
from typing import Final, NamedTuple, Protocol, Union, cast
import fastapi
import orjson
@ -215,7 +215,7 @@ def _get_model_from_request_context(
request_data: dict,
route: str,
request: Request | None,
llm_router: Any | None = None,
llm_router: litellm.Router | None = None,
team_id: str | None = None,
) -> str | list[str] | None:
return get_model_from_request(
@ -520,8 +520,8 @@ def _get_bearer_token_or_received_api_key(api_key: str) -> str:
def _routing_selector_matches_claim(
selector_value: Any | None,
claim_value: Any | None,
selector_value: object,
claim_value: object,
*,
split_space_delimited: bool = False,
) -> bool:
@ -661,7 +661,7 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
# ``websocket.url``, which Starlette reconstructs from the (poisonable)
# Host header. Carry the ASGI scope's path / root_path so the lookup
# never reaches the fallback.
synthetic_scope: Final[dict[str, Any]] = {
synthetic_scope: Final[dict[str, object]] = {
"type": "http",
"method": "GET",
"query_string": ws_scope.get("query_string", b""),
@ -3211,7 +3211,7 @@ async def _reserve_budget_after_common_checks(
user_api_key_auth_obj: UserAPIKeyAuth,
request_data: dict,
route: str,
llm_router: Any | None,
llm_router: litellm.Router | None,
team_object: LiteLLM_TeamTableCachedObj | None,
user_object: LiteLLM_UserTable | None,
prisma_client: PrismaClient | None,
@ -3257,7 +3257,7 @@ def _should_skip_budget_checks(
request_data: dict,
route: str,
request: Request | None,
llm_router: Any | None,
llm_router: litellm.Router | None,
team_id: str | None = None,
) -> bool:
model: Final = _get_model_from_request_context(
@ -3760,7 +3760,7 @@ async def _enforce_key_and_fallback_model_access(
route: str,
request: Request | None,
llm_model_list: list | None,
llm_router: Any | None,
llm_router: litellm.Router | None,
) -> None:
"""
Key-level model allowlist and client fallbacks (same as standard auth).

View file

@ -9441,7 +9441,7 @@ def _format_fallback_metadata_sse_event(
def _restamp_streaming_chunk_model(
*,
chunk: Any,
chunk: object,
requested_model_from_client: str,
request_data: dict,
model_mismatch_logged: bool,

View file

@ -0,0 +1,384 @@
from __future__ import annotations
import json
import uuid
from typing import Final, Literal
import anthropic
import httpx
import openai
import pytest
import websockets
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from pydantic import JsonValue
def _marker() -> str:
return uuid.uuid4().hex
def _register_scenario(gateway: Gateway, scenario_id: str, response: dict[str, JsonValue]) -> None:
with httpx.Client(timeout=5, trust_env=False) as client:
result: Final = client.post(
f"{gateway.upstream_url}/__scenarios",
json={"scenario_id": scenario_id, "response": response},
)
assert result.status_code == 200, result.text
_ANTHROPIC_MESSAGE: Final = {
"id": "msg_$UNIQUE_ID",
"type": "message",
"role": "assistant",
"model": "claude-3-5-sonnet-20241022",
"content": [{"type": "text", "text": "scripted reply"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 9, "output_tokens": 5},
}
_ANTHROPIC_EVENTS: Final = (
'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_$UNIQUE_ID","type":"message","role":"assistant","model":"claude-3-5-sonnet-20241022","content":[],"usage":{"input_tokens":9,"output_tokens":0}}}',
'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}',
'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"scripted reply"}}',
'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}',
'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":5}}',
'event: message_stop\ndata: {"type":"message_stop"}',
)
_RESPONSES_BODY: Final = {
"id": "resp_$UNIQUE_ID",
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "openai/gpt-4o-mini",
"output": [{"type": "message", "content": [{"type": "output_text", "text": "done"}]}],
"usage": {"input_tokens": 3, "output_tokens": 4, "total_tokens": 7},
}
_RESPONSES_EVENTS: Final = (
'event: response.created\ndata: {"type":"response.created","response":{"id":"$UNIQUE_ID","object":"response","status":"in_progress"}}',
'event: response.completed\ndata: {"type":"response.completed","response":{"id":"$UNIQUE_ID","object":"response","status":"completed","output":[{"type":"message","content":[{"type":"output_text","text":"done"}]}],"usage":{"input_tokens":3,"output_tokens":4,"total_tokens":7}}}',
)
def _anthropic_model(gateway: Gateway, *, stream: bool) -> str:
scenario_id: Final = "audit-anthropic-stream" if stream else "audit-anthropic"
if stream:
_register_scenario(gateway, scenario_id, {"content_type": "text/event-stream", "frames": _ANTHROPIC_EVENTS})
else:
_register_scenario(
gateway,
scenario_id,
{
"content_type": "application/x-routed",
"routes": {"POST /v1/messages": {"content_type": "application/json", "body": _ANTHROPIC_MESSAGE}},
},
)
name: Final = f"integration-{uuid.uuid4().hex}"
created: Final = gateway.post(
"/model/new",
{
"model_name": name,
"litellm_params": {
"model": "anthropic/claude-3-5-sonnet-20241022",
"api_key": "integration-provider-key",
"api_base": f"{gateway.upstream_url}/{scenario_id}",
},
},
)
assert isinstance(created["model_info"], dict) and created["model_info"]["id"], created
return name
def _responses_model(gateway: Gateway, *, stream: bool) -> str:
scenario_id: Final = "audit-resp-stream" if stream else "audit-resp"
if stream:
_register_scenario(gateway, scenario_id, {"content_type": "text/event-stream", "frames": _RESPONSES_EVENTS})
else:
_register_scenario(
gateway,
scenario_id,
{
"content_type": "application/x-routed",
"routes": {"POST /responses": {"content_type": "application/json", "body": _RESPONSES_BODY}},
},
)
name: Final = f"integration-{uuid.uuid4().hex}"
created: Final = gateway.post(
"/model/new",
{
"model_name": name,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "integration-provider-key",
"api_base": f"{gateway.upstream_url}/{scenario_id}",
},
},
)
assert isinstance(created["model_info"], dict) and created["model_info"]["id"], created
return name
def _openai_client(gateway: Gateway, key: str) -> openai.OpenAI:
return openai.OpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=key,
http_client=httpx.Client(trust_env=False),
)
def _openai_async_client(gateway: Gateway, key: str) -> openai.AsyncOpenAI:
return openai.AsyncOpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=key,
http_client=httpx.AsyncClient(trust_env=False),
)
def _anthropic_client(gateway: Gateway, key: str) -> anthropic.Anthropic:
return anthropic.Anthropic(
base_url=str(gateway.client.base_url).rstrip("/"),
api_key=key,
http_client=httpx.Client(trust_env=False),
)
def _anthropic_async_client(gateway: Gateway, key: str) -> anthropic.AsyncAnthropic:
return anthropic.AsyncAnthropic(
base_url=str(gateway.client.base_url).rstrip("/"),
api_key=key,
http_client=httpx.AsyncClient(trust_env=False),
)
def _sse_chunk_models(text: str) -> tuple[str, ...]:
models: list[str] = []
for line in text.splitlines():
if not line.startswith("data: ") or line == "data: [DONE]":
continue
value = json.loads(line.removeprefix("data: "))
if isinstance(value, dict) and isinstance(value.get("model"), str):
models.append(value["model"])
return tuple(models)
@pytest.mark.parametrize("client_kind", ("httpx", "openai", "openai_async"))
@pytest.mark.parametrize("stream", (False, True))
def test_allowed_key_chat_completions_restamps_alias(
gateway: Gateway, client_kind: Literal["httpx", "openai", "openai_async"], stream: bool
) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
body: Final = {
"model": model,
"messages": [{"role": "user", "content": _marker()}],
"stream": stream,
**({"stream_options": {"include_usage": True}} if stream else {}),
}
if client_kind == "httpx":
response: Final = gateway.request("POST", "/v1/chat/completions", body, key=key)
assert response.status_code == 200, response.text
if stream:
models: Final = _sse_chunk_models(response.text)
assert models and all(chunk_model == model for chunk_model in models), response.text
else:
assert response.json()["model"] == model, response.text
elif client_kind == "openai":
client: Final = _openai_client(gateway, key)
completion: Final = client.chat.completions.create(**body)
if stream:
chunk_models: list[str] = []
for chunk in completion:
if chunk.model:
chunk_models.append(chunk.model)
assert chunk_models and all(chunk_model == model for chunk_model in chunk_models)
else:
assert completion.model == model
else:
async def _run() -> None:
client_async: Final = _openai_async_client(gateway, key)
completion_async: Final = await client_async.chat.completions.create(**body)
if stream:
seen: list[str] = []
async for chunk in completion_async:
if chunk.model:
seen.append(chunk.model)
assert seen and all(chunk_model == model for chunk_model in seen)
else:
assert completion_async.model == model
import asyncio
asyncio.run(_run())
@pytest.mark.parametrize("client_kind", ("httpx", "anthropic", "anthropic_async"))
@pytest.mark.parametrize("stream", (False, True))
def test_allowed_key_messages(gateway: Gateway, client_kind: str, stream: bool) -> None:
model: Final = _anthropic_model(gateway, stream=stream)
with gateway.scenario() as scenario:
key: Final = scenario.key(models=[model])
body: Final = {
"model": model,
"max_tokens": 64,
"messages": [{"role": "user", "content": _marker()}],
"stream": stream,
}
if client_kind == "httpx":
response: Final = gateway.request("POST", "/v1/messages", body, key=key)
assert response.status_code == 200, response.text
if stream:
assert "message_start" in response.text and "message_stop" in response.text, response.text
else:
assert response.json()["type"] == "message", response.text
elif client_kind == "anthropic":
client: Final = _anthropic_client(gateway, key)
if stream:
with client.messages.stream(**{k: v for k, v in body.items() if k != "stream"}) as event_stream:
message: Final = event_stream.get_final_message()
assert message.type == "message" and message.content, message
else:
message = client.messages.create(**{k: v for k, v in body.items() if k != "stream"})
assert message.type == "message" and message.content, message
else:
async def _run() -> None:
client_async: Final = _anthropic_async_client(gateway, key)
params: Final = {k: v for k, v in body.items() if k != "stream"}
if stream:
async with client_async.messages.stream(**params) as event_stream:
message_async: Final = await event_stream.get_final_message()
else:
message_async = await client_async.messages.create(**params)
assert message_async.type == "message" and message_async.content, message_async
import asyncio
asyncio.run(_run())
@pytest.mark.parametrize("client_kind", ("httpx", "openai", "openai_async"))
@pytest.mark.parametrize("stream", (False, True))
def test_allowed_key_responses(gateway: Gateway, client_kind: str, stream: bool) -> None:
model: Final = _responses_model(gateway, stream=stream)
with gateway.scenario() as scenario:
key: Final = scenario.key(models=[model])
body: Final = {
"model": model,
"input": [{"role": "user", "content": [{"type": "input_text", "text": _marker()}]}],
"stream": stream,
}
if client_kind == "httpx":
response: Final = gateway.request("POST", "/v1/responses", body, key=key)
assert response.status_code == 200, response.text
if stream:
assert "response.completed" in response.text, response.text
else:
assert response.json()["object"] == "response", response.text
elif client_kind == "openai":
client: Final = _openai_client(gateway, key)
result: Final = client.responses.create(**body)
if stream:
events: list[str] = []
for event in result:
events.append(event.type)
assert "response.completed" in events, events
else:
assert result.object == "response", result
else:
async def _run() -> None:
client_async: Final = _openai_async_client(gateway, key)
result_async: Final = await client_async.responses.create(**body)
if stream:
events_async: list[str] = []
async for event in result_async:
events_async.append(event.type)
assert "response.completed" in events_async, events_async
else:
assert result_async.object == "response", result_async
import asyncio
asyncio.run(_run())
def test_denied_key_rejected(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
other_model: Final = scenario.model()
key: Final = scenario.key(models=[other_model])
for path, body in (
("/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "x"}]}),
("/v1/messages", {"model": model, "max_tokens": 4, "messages": [{"role": "user", "content": "x"}]}),
("/v1/responses", {"model": model, "input": "x"}),
):
response: Final = gateway.request("POST", path, body, key=key)
assert response.status_code in (401, 403), f"{path}: {response.status_code} {response.text}"
async def test_realtime_websocket_allowed_and_denied(gateway: Gateway) -> None:
rt_scenario: Final = f"auditrt{uuid.uuid4().hex[:8]}"
_register_scenario(
gateway,
rt_scenario,
{
"content_type": "application/x-realtime",
"session_model": "gpt-4o-realtime-preview",
"events": (
{"type": "response.done", "response": {"id": "resp_$UNIQUE_ID", "model": "gpt-4o-realtime-preview"}},
),
},
)
name: Final = f"integration-{uuid.uuid4().hex}"
created: Final = gateway.post(
"/model/new",
{
"model_name": name,
"litellm_params": {
"model": "openai/gpt-4o-realtime-preview",
"api_key": rt_scenario,
"api_base": gateway.upstream_url,
},
},
)
assert isinstance(created["model_info"], dict) and created["model_info"]["id"], created
with gateway.scenario() as scenario:
key: Final = scenario.key(models=[name])
other_model: Final = scenario.model()
denied_key: Final = scenario.key(models=[other_model])
url: Final = f"{str(gateway.client.base_url).replace('http://', 'ws://')}/v1/realtime?model={name}"
async with websockets.connect(url, additional_headers={"Authorization": f"Bearer {key}"}) as websocket:
session: Final = json.loads(await websocket.recv())
assert session["type"] == "session.created", session
assert session["session"]["model"] == "gpt-4o-realtime-preview", session
await websocket.send(json.dumps({"type": "response.create"}))
event: Final = json.loads(await websocket.recv())
assert event["type"] == "response.done", event
with pytest.raises(websockets.exceptions.InvalidHandshake):
async with websockets.connect(
url, additional_headers={"Authorization": f"Bearer {denied_key}"}
) as denied_socket:
await denied_socket.recv()
def test_max_budget_allows_then_denies_and_records_spend(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model], max_budget=0.0000001)
body: Final = {"model": model, "messages": [{"role": "user", "content": _marker()}]}
first: Final = gateway.request("POST", "/v1/chat/completions", body, key=key)
assert first.status_code == 200, first.text
request_id: Final = first.json()["id"]
rows: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (request_id,)),
lambda values: len(values) == 1,
seconds=70,
)
assert rows, first.text
second: Final = gateway.request("POST", "/v1/chat/completions", body, key=key)
assert second.status_code == 422, second.text
assert "budget_exceeded" in second.text, second.text

View file

@ -0,0 +1,128 @@
from __future__ import annotations
import uuid
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway
from pydantic import JsonValue
from litellm.harness.context import GatewayTarget
from litellm.harness.endpoint import ModelEndpoint
from litellm.harness.types import Harness
def _register_scenario(gateway: Gateway, scenario_id: str, response: dict[str, JsonValue]) -> None:
with httpx.Client(timeout=5, trust_env=False) as client:
result: Final = client.post(
f"{gateway.upstream_url}/__scenarios",
json={"scenario_id": scenario_id, "response": response},
)
assert result.status_code == 200, result.text
def _client(endpoint: ModelEndpoint) -> httpx.AsyncClient:
return httpx.AsyncClient(
base_url=endpoint.url,
headers={"Authorization": f"Bearer {endpoint.token}"},
trust_env=False,
timeout=15,
)
_CHAT_BODY: Final = {
"id": "chatcmpl-$UNIQUE_ID",
"object": "chat.completion",
"created": 1,
"model": "scripted-model",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 11, "completion_tokens": 7, "total_tokens": 18},
}
_RESPONSES_BODY: Final = {
"id": "resp_$UNIQUE_ID",
"object": "response",
"status": "completed",
"model": "scripted-model",
"output": [{"type": "message", "content": [{"type": "output_text", "text": "hi"}]}],
"usage": {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8},
}
async def test_endpoint_forwards_json_object_bodies(gateway: Gateway) -> None:
scenario_id: Final = f"audithend{uuid.uuid4().hex[:8]}"
_register_scenario(
gateway,
scenario_id,
{
"content_type": "application/x-routed",
"routes": {
"POST /v1/chat/completions": {"content_type": "application/json", "body": _CHAT_BODY},
"POST /v1/responses": {"content_type": "application/json", "body": _RESPONSES_BODY},
},
},
)
target: Final = GatewayTarget(api_base=f"{gateway.upstream_url}/{scenario_id}", api_key="scripted")
async with ModelEndpoint(Harness.CODEX, "audit-model", target, api_key="x") as endpoint:
assert isinstance(endpoint.port, int) and endpoint.port > 0
async with _client(endpoint) as client:
chat: Final = await client.post("/v1/chat/completions", json={"messages": []})
assert chat.status_code == 200, chat.text
assert chat.json()["object"] == "chat.completion", chat.text
responses: Final = await client.post("/v1/responses", json={"input": "x"})
assert responses.status_code == 200, responses.text
assert responses.json()["object"] == "response", responses.text
usage: Final = endpoint.usage.snapshot()
assert usage.input_tokens == 11 + 5 and usage.output_tokens == 7 + 3, usage
@pytest.mark.parametrize("raw", (b"[1, 2]", b'"a string"', b"42", b"null"))
async def test_endpoint_rejects_valid_json_non_object(gateway: Gateway, raw: bytes) -> None:
async with ModelEndpoint(Harness.CODEX, "audit-model", None, api_key="x") as endpoint:
async with _client(endpoint) as client:
response: Final = await client.post("/v1/chat/completions", content=raw)
assert response.status_code == 400, response.text
assert "must be a JSON object" in response.text, response.text
async def test_endpoint_rejects_unparseable_body(gateway: Gateway) -> None:
async with ModelEndpoint(Harness.CODEX, "audit-model", None, api_key="x") as endpoint:
async with _client(endpoint) as client:
response: Final = await client.post("/v1/chat/completions", content=b"this is not json{")
assert response.status_code == 400, response.text
async def test_endpoint_records_usage_from_sse_stream(gateway: Gateway) -> None:
scenario_id: Final = f"audithend{uuid.uuid4().hex[:8]}"
_register_scenario(
gateway,
scenario_id,
{
"content_type": "text/event-stream",
"frames": (
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"role":"assistant","content":"hi"},"finish_reason":null}]}',
"data: 42",
'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","choices":[],"usage":{"prompt_tokens":13,"completion_tokens":4}}',
"data: [DONE]",
),
},
)
target: Final = GatewayTarget(api_base=f"{gateway.upstream_url}/{scenario_id}", api_key="scripted")
async with ModelEndpoint(Harness.CODEX, "audit-model", target, api_key="x") as endpoint:
async with _client(endpoint) as client:
response: Final = await client.post("/v1/chat/completions", json={"messages": []})
assert response.status_code == 200, response.text
assert "[DONE]" in response.text, response.text
usage: Final = endpoint.usage.snapshot()
assert usage.input_tokens == 13 and usage.output_tokens == 4, usage
async def test_endpoint_binds_free_port() -> None:
async with ModelEndpoint(Harness.CODEX, "audit-model", None) as endpoint:
port: Final = endpoint.port
assert isinstance(port, int) and port > 0
async with _client(endpoint) as client:
response: Final = await client.get("/v1/models")
assert response.status_code == 200, response.text
assert response.json()["object"] == "list", response.text

View file

@ -0,0 +1,346 @@
from __future__ import annotations
import json
import uuid
from collections.abc import Iterator
from contextlib import ExitStack
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import httpx
import openai
import pytest
import yaml
from integration._support.client import Gateway, gateway_from_environment, object_value
from integration._support.process import owned_proxy
from integration.authorization._guardrail_opt_out import upstream_observations
from pydantic import JsonValue
CONFIG_STORE_ID: Final = "vs_integration_config_store"
FAIL_STORE_ID: Final = "vs_audit_fail_store"
PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml"
def _marker() -> str:
return uuid.uuid4().hex
def _chat_body(model: str, marker: str, store_id: str, *, stream: bool = False) -> dict[str, JsonValue]:
return {
"model": model,
"messages": [{"role": "user", "content": marker}],
"vector_store_ids": [store_id],
"stream": stream,
**({"stream_options": {"include_usage": True}} if stream else {}),
}
def _sdk_chat_body(body: dict[str, JsonValue]) -> dict[str, JsonValue]:
return {
**{key: value for key, value in body.items() if key != "vector_store_ids"},
"extra_body": {"vector_store_ids": body["vector_store_ids"]},
}
def _search_observations(gateway: Gateway, marker: str, store_id: str) -> tuple[dict[str, JsonValue], ...]:
path: Final = f"/vector_stores/{store_id}/search"
return tuple(
observation
for observation in upstream_observations(gateway)
if observation["path"] == path and marker in json.dumps(observation["body"])
)
def _chat_observations(gateway: Gateway, marker: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(
observation
for observation in upstream_observations(gateway)
if observation["path"] == "/v1/chat/completions" and marker in json.dumps(observation["body"])
)
def _message_provider_fields(payload: dict[str, JsonValue]) -> dict[str, JsonValue]:
choices: Final = payload["choices"]
assert isinstance(choices, list) and choices, payload
message: Final = object_value(object_value(choices[0])["message"])
fields: Final = message.get("provider_specific_fields")
assert isinstance(fields, dict), payload
return object_value(fields)
def _sse_data_lines(text: str) -> tuple[dict[str, JsonValue], ...]:
parsed: list[dict[str, JsonValue]] = []
for line in text.splitlines():
if not line.startswith("data: ") or line == "data: [DONE]":
continue
value = json.loads(line.removeprefix("data: "))
if isinstance(value, dict):
parsed.append(value)
return tuple(parsed)
def _chunk_annotated(chunk: dict[str, JsonValue]) -> bool:
choices: Final = chunk.get("choices")
if not isinstance(choices, list) or not choices:
return False
delta: Final = object_value(choices[0]).get("delta")
if not isinstance(delta, dict):
return False
fields: Final = object_value(delta).get("provider_specific_fields")
return "search_results" in json.dumps(fields)
def _openai_client(gateway: Gateway, key: str) -> openai.OpenAI:
return openai.OpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=key,
http_client=httpx.Client(trust_env=False),
)
def _openai_async_client(gateway: Gateway, key: str) -> openai.AsyncOpenAI:
return openai.AsyncOpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=key,
http_client=httpx.AsyncClient(trust_env=False),
)
def _register_scenario(gateway: Gateway, scenario_id: str, response: dict[str, JsonValue]) -> None:
with httpx.Client(timeout=5, trust_env=False) as client:
result: Final = client.post(
f"{gateway.upstream_url}/__scenarios",
json={"scenario_id": scenario_id, "response": response},
)
assert result.status_code == 200, result.text
def test_vector_store_hook_chat_completions_raw_httpx(gateway: Gateway) -> None:
marker: Final = _marker()
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST", "/v1/chat/completions", _chat_body(model, marker, CONFIG_STORE_ID), key=key
)
assert response.status_code == 200, response.text
assert _message_provider_fields(response.json())["search_results"]
assert len(_search_observations(gateway, marker, CONFIG_STORE_ID)) == 1
def test_vector_store_hook_chat_completions_raw_httpx_stream(gateway: Gateway) -> None:
marker: Final = _marker()
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST", "/v1/chat/completions", _chat_body(model, marker, CONFIG_STORE_ID, stream=True), key=key
)
assert response.status_code == 200, response.text
chunks: Final = _sse_data_lines(response.text)
assert any(_chunk_annotated(chunk) for chunk in chunks), response.text
assert len(_search_observations(gateway, marker, CONFIG_STORE_ID)) == 1
@pytest.mark.parametrize("stream", (False, True))
def test_vector_store_hook_chat_completions_openai_sdk(gateway: Gateway, stream: bool) -> None:
marker: Final = _marker()
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
client: Final = _openai_client(gateway, key)
body: Final = _chat_body(model, marker, CONFIG_STORE_ID, stream=stream)
if stream:
with client.chat.completions.stream(
**{k: v for k, v in _sdk_chat_body(body).items() if k != "stream"}
) as event_stream:
for _event in event_stream:
pass
snapshot: Final = event_stream.get_final_completion().model_dump(mode="json")
assert _message_provider_fields(snapshot)["search_results"], snapshot
else:
completion: Final = client.chat.completions.create(**_sdk_chat_body(body))
assert _message_provider_fields(completion.model_dump(mode="json"))["search_results"]
assert len(_search_observations(gateway, marker, CONFIG_STORE_ID)) == 1
@pytest.mark.parametrize("stream", (False, True))
async def test_vector_store_hook_chat_completions_openai_async_sdk(gateway: Gateway, stream: bool) -> None:
marker: Final = _marker()
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
client: Final = _openai_async_client(gateway, key)
body: Final = _chat_body(model, marker, CONFIG_STORE_ID, stream=stream)
if stream:
stream_response: Final = await client.chat.completions.create(**_sdk_chat_body(body))
chunks: list[dict[str, JsonValue]] = []
async for chunk in stream_response:
chunks.append(chunk.model_dump(mode="json"))
assert any(_chunk_annotated(chunk) for chunk in chunks), chunks
else:
completion: Final = await client.chat.completions.create(**_sdk_chat_body(body))
assert _message_provider_fields(completion.model_dump(mode="json"))["search_results"]
assert len(_search_observations(gateway, marker, CONFIG_STORE_ID)) == 1
@pytest.mark.parametrize("stream", (False, True))
def test_vector_store_hook_responses_api(gateway: Gateway, stream: bool) -> None:
marker: Final = _marker()
scenario_id: Final = "auditresp" if not stream else "auditrespstream"
if stream:
_register_scenario(
gateway,
scenario_id,
{
"content_type": "text/event-stream",
"frames": (
'event: response.created\ndata: {"type":"response.created","response":{"id":"$UNIQUE_ID","object":"response","status":"in_progress"}}',
'event: response.completed\ndata: {"type":"response.completed","response":{"id":"$UNIQUE_ID","object":"response","status":"completed","output":[{"type":"message","content":[{"type":"output_text","text":"done"}]}],"usage":{"input_tokens":3,"output_tokens":4,"total_tokens":7}}}',
),
},
)
else:
_register_scenario(
gateway,
scenario_id,
{
"content_type": "application/x-routed",
"routes": {
"POST /responses": {
"content_type": "application/json",
"body": {
"id": "resp_$UNIQUE_ID",
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "openai/gpt-4o-mini",
"output": [{"type": "message", "content": [{"type": "output_text", "text": "done"}]}],
"usage": {"input_tokens": 3, "output_tokens": 4, "total_tokens": 7},
},
}
},
},
)
with gateway.scenario() as scenario:
model: Final = scenario.model(api_base=f"{gateway.upstream_url}/{scenario_id}")
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/responses",
{
"model": model,
"input": marker,
"vector_store_ids": [CONFIG_STORE_ID],
"stream": stream,
},
key=key,
)
assert response.status_code == 200, response.text
if stream:
assert "response.completed" in response.text
else:
payload: Final = response.json()
assert payload["object"] == "response" and payload["id"], payload
assert len(_search_observations(gateway, marker, CONFIG_STORE_ID)) == 1
def test_vector_store_annotation_survives_cache_hit(gateway: Gateway) -> None:
marker: Final = _marker()
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
body: Final = _chat_body(model, marker, CONFIG_STORE_ID)
first: Final = gateway.request("POST", "/v1/chat/completions", body, key=key)
assert first.status_code == 200, first.text
second: Final = gateway.request("POST", "/v1/chat/completions", body, key=key)
assert second.status_code == 200, second.text
assert second.json()["id"] == first.json()["id"]
assert _message_provider_fields(second.json())["search_results"]
assert len(_chat_observations(gateway, marker)) == 1
@dataclass(frozen=True, slots=True)
class FailureModeGateways:
annotate: Gateway
error: Gateway
upstream: Gateway
def _failure_mode_config(directory: Path, mode: str, upstream_url: str) -> Path:
config: Final = yaml.safe_load(PROXY_CONFIG.read_text())
config["litellm_settings"]["vector_store_search_failure_mode"] = mode
config["vector_store_registry"] = [
{
"vector_store_name": "audit-fail-store",
"litellm_params": {
"vector_store_id": FAIL_STORE_ID,
"custom_llm_provider": "openai",
"api_base": f"{upstream_url}/auditnostore",
"api_key": "integration-provider-key",
},
}
]
path: Final = directory / f"proxy_vector_store_{mode}.yaml"
path.write_text(yaml.safe_dump(config))
return path
@pytest.fixture(scope="module")
def failure_gateways(tmp_path_factory: pytest.TempPathFactory) -> Iterator[FailureModeGateways]:
with gateway_from_environment() as upstream_gateway, ExitStack() as stack:
gateways: dict[str, Gateway] = {}
for mode in ("annotate", "error"):
directory: Final = tmp_path_factory.mktemp(f"audit_vs_{mode}")
config: Final = _failure_mode_config(directory, mode, upstream_gateway.upstream_url)
owned: Final = stack.enter_context(owned_proxy(upstream_gateway, directory, {}, config=config, workers=2))
gateways[mode] = owned
yield FailureModeGateways(gateways["annotate"], gateways["error"], upstream_gateway)
def _register_failing_store(gateway: Gateway) -> None:
_register_scenario(
gateway,
"auditnostore",
{
"content_type": "application/x-routed",
"routes": {
f"POST /vector_stores/{FAIL_STORE_ID}/search": {
"content_type": "application/json",
"body": {"error": {"message": "scripted store failure"}},
"status": 500,
}
},
},
)
def test_vector_store_search_failure_annotate_mode(failure_gateways: FailureModeGateways) -> None:
_register_failing_store(failure_gateways.upstream)
marker: Final = _marker()
with failure_gateways.annotate.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
response: Final = failure_gateways.annotate.request(
"POST", "/v1/chat/completions", _chat_body(model, marker, FAIL_STORE_ID), key=key
)
assert response.status_code == 200, response.text
fields: Final = _message_provider_fields(response.json())
failures: Final = fields.get("vector_store_search_failures")
assert isinstance(failures, list) and failures, response.text
assert FAIL_STORE_ID in json.dumps(failures) and "scripted store failure" in json.dumps(failures), failures
def test_vector_store_search_failure_error_mode(failure_gateways: FailureModeGateways) -> None:
_register_failing_store(failure_gateways.upstream)
marker: Final = _marker()
with failure_gateways.error.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
response: Final = failure_gateways.error.request(
"POST", "/v1/chat/completions", _chat_body(model, marker, FAIL_STORE_ID), key=key
)
assert response.status_code != 200, response.text
assert "error" in response.json(), response.text
assert "vector store" in response.text.lower() or FAIL_STORE_ID in response.text, response.text

View file

@ -0,0 +1,203 @@
from __future__ import annotations
import json
import re
import signal
from collections.abc import Mapping
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from types import MappingProxyType
from typing import Final
import httpx
import psutil
from integration._support.client import Gateway, eventually, gateway_from_environment
from integration._support.process import owned_proxy_process
_RESPONSES_JSON_SCENARIO: Final = "audit-chaos-responses"
_RESPONSES_STREAM_SCENARIO: Final = "audit-chaos-responses-stream"
_MESSAGES_JSON_SCENARIO: Final = "audit-chaos-messages"
_MESSAGES_STREAM_SCENARIO: Final = "audit-chaos-messages-stream"
_RESPONSES_BODY: Final = {
"id": "resp_$UNIQUE_ID",
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "openai/gpt-4o-mini",
"output": [{"type": "message", "content": [{"type": "output_text", "text": "done"}]}],
"usage": {"input_tokens": 3, "output_tokens": 4, "total_tokens": 7},
}
_ANTHROPIC_MESSAGE: Final = {
"id": "msg_$UNIQUE_ID",
"type": "message",
"role": "assistant",
"model": "claude-3-5-sonnet-20241022",
"content": [{"type": "text", "text": "scripted reply"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 9, "output_tokens": 5},
}
_ANTHROPIC_EVENTS: Final = (
'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_$UNIQUE_ID","type":"message","role":"assistant","model":"claude-3-5-sonnet-20241022","content":[],"usage":{"input_tokens":9,"output_tokens":0}}}',
'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}',
'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"scripted reply"}}',
'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}',
'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":5}}',
'event: message_stop\ndata: {"type":"message_stop"}',
)
_RESPONSES_EVENTS: Final = (
'event: response.completed\ndata: {"type":"response.completed","response":{"id":"resp_$UNIQUE_ID","object":"response","created_at":1700000000,"status":"completed","model":"openai/gpt-4o-mini","output":[{"type":"message","content":[{"type":"output_text","text":"done"}]}],"usage":{"input_tokens":3,"output_tokens":4,"total_tokens":7}}}',
)
def _register_scenarios(upstream_url: str) -> None:
with httpx.Client(timeout=5, trust_env=False) as client:
for scenario_id, response in (
(
_RESPONSES_JSON_SCENARIO,
{
"content_type": "application/x-routed",
"routes": {"POST /responses": {"content_type": "application/json", "body": _RESPONSES_BODY}},
},
),
(_RESPONSES_STREAM_SCENARIO, {"content_type": "text/event-stream", "frames": _RESPONSES_EVENTS}),
(
_MESSAGES_JSON_SCENARIO,
{
"content_type": "application/x-routed",
"routes": {"POST /v1/messages": {"content_type": "application/json", "body": _ANTHROPIC_MESSAGE}},
},
),
(_MESSAGES_STREAM_SCENARIO, {"content_type": "text/event-stream", "frames": _ANTHROPIC_EVENTS}),
):
result: Final = client.post(
f"{upstream_url}/__scenarios", json={"scenario_id": scenario_id, "response": response}
)
assert result.status_code == 200, result.text
def _stream_response_id(text: str) -> str | None:
for line in text.splitlines():
if not line.startswith("data: ") or line == "data: [DONE]":
continue
value: Final = json.loads(line.removeprefix("data: "))
if isinstance(value, dict):
if isinstance(value.get("id"), str):
return value["id"]
for key_name in ("response", "message"):
nested: Final = value.get(key_name)
if isinstance(nested, dict) and isinstance(nested.get("id"), str):
return nested["id"]
return None
_KINDS: Final = ("chat", "chat-stream", "messages", "messages-stream", "responses", "responses-stream")
_PATHS: Final = MappingProxyType(
{"chat": "/v1/chat/completions", "messages": "/v1/messages", "responses": "/v1/responses"}
)
def _burst_kind(index: int) -> str:
return f"{_KINDS[index % 3 * 2].removesuffix('-stream')}{'-stream' if index % 2 == 0 else ''}"
def _fire(
owned: Gateway, models: Mapping[str, str], key: str, kind: str, marker: str
) -> tuple[str, int | None, str | None]:
stream: Final = kind.endswith("-stream")
path: Final = _PATHS[kind.removesuffix("-stream")]
model: Final = models[f"{path.rsplit('/', 1)[-1]}{'-stream' if stream else ''}"]
if path == "/v1/messages":
body: Final = {
"model": model,
"messages": [{"role": "user", "content": f"burst-{kind}-{marker}"}],
"max_tokens": 64,
"stream": stream,
}
elif path == "/v1/responses":
body = {"model": model, "input": f"burst-{kind}-{marker}", "stream": stream}
else:
body = {
"model": model,
"messages": [{"role": "user", "content": f"burst-{kind}-{marker}"}],
"stream": stream,
**({"stream_options": {"include_usage": True}} if stream else {}),
}
try:
response: Final = owned.request("POST", path, body, key=key)
except Exception:
return kind, None, None
if response.status_code != 200:
return kind, response.status_code, None
if stream:
return kind, 200, _stream_response_id(response.text)
response_id: Final = response.json().get("id")
return kind, 200, response_id if isinstance(response_id, str) else None
_WORKER_PID: Final = re.compile(r"Started server process \[(\d+)\]")
def _worker_pids(log: Path) -> tuple[int, ...]:
return tuple(int(match) for match in _WORKER_PID.findall(log.read_text()))
def test_proxy_survives_worker_kill_mid_burst(tmp_path: Path) -> None:
with gateway_from_environment() as upstream_gateway:
directory: Final = tmp_path
_register_scenarios(upstream_gateway.upstream_url)
with owned_proxy_process(upstream_gateway, directory, {}, workers=2) as owned:
with owned.gateway.scenario() as scenario:
model: Final = scenario.model()
responses_model: Final = scenario.model(
api_base=f"{owned.gateway.upstream_url}/{_RESPONSES_JSON_SCENARIO}"
)
responses_stream_model: Final = scenario.model(
api_base=f"{owned.gateway.upstream_url}/{_RESPONSES_STREAM_SCENARIO}"
)
messages_model: Final = scenario.model(
model="anthropic/claude-3-5-sonnet-20241022",
api_base=f"{owned.gateway.upstream_url}/{_MESSAGES_JSON_SCENARIO}",
)
messages_stream_model: Final = scenario.model(
model="anthropic/claude-3-5-sonnet-20241022",
api_base=f"{owned.gateway.upstream_url}/{_MESSAGES_STREAM_SCENARIO}",
)
key: Final = scenario.key(
models=[model, responses_model, responses_stream_model, messages_model, messages_stream_model]
)
models: Final = MappingProxyType(
{
"completions": model,
"completions-stream": model,
"messages": messages_model,
"messages-stream": messages_stream_model,
"responses": responses_model,
"responses-stream": responses_stream_model,
}
)
workers: Final = eventually(lambda: _worker_pids(owned.log), lambda pids: len(pids) == 2, seconds=30)
records: dict[int, tuple[str, int | None, str | None]] = {}
with ThreadPoolExecutor(max_workers=30) as pool:
futures: Final = [
(index, pool.submit(_fire, owned.gateway, models, key, _burst_kind(index), str(index)))
for index in range(30)
]
psutil.Process(workers[0]).send_signal(signal.SIGKILL)
for index, future in futures:
records[index] = future.result()
proxy_stamped: Final = [
record[2]
for record in records.values()
if record[2] is not None and not record[0].startswith("messages")
]
malformed: Final = {i: r for i, r in records.items() if r[1] == 200 and r[2] is None}
assert not malformed, f"200 responses without a well-formed body: {malformed}"
assert len(proxy_stamped) == len(set(proxy_stamped)), records
assert proxy_stamped, "no burst request completed"
for kind in _KINDS:
_, status, response_id = _fire(owned.gateway, models, key, kind, "probe")
assert status == 200 and response_id is not None, (kind, status, records)