From e9bf2cfd01a2c99721a09c40978249e15153613f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 21:03:16 +0000 Subject: [PATCH] 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> --- litellm/harness/endpoint.py | 18 +- litellm/harness/sync.py | 2 +- litellm/integrations/galileo.py | 2 +- .../vector_store_pre_call_hook.py | 13 +- .../prompt_templates/factory.py | 11 +- .../claude_code/harness/transformation.py | 14 +- litellm/llms/ollama/chat/transformation.py | 2 +- litellm/proxy/auth/user_api_key_auth.py | 16 +- litellm/proxy/proxy_server.py | 2 +- .../test_audit_any_sweep_auth.py | 384 ++++++++++++++++++ .../test_audit_any_sweep_harness_endpoint.py | 128 ++++++ .../test_audit_any_sweep_vector_store.py | 346 ++++++++++++++++ .../routing/test_audit_any_sweep_chaos.py | 203 +++++++++ 13 files changed, 1107 insertions(+), 34 deletions(-) create mode 100644 tests/integration/authorization/test_audit_any_sweep_auth.py create mode 100644 tests/integration/providers/test_audit_any_sweep_harness_endpoint.py create mode 100644 tests/integration/providers/test_audit_any_sweep_vector_store.py create mode 100644 tests/integration/routing/test_audit_any_sweep_chaos.py diff --git a/litellm/harness/endpoint.py b/litellm/harness/endpoint.py index 21e789ed1fc..192c97533ea 100644 --- a/litellm/harness/endpoint.py +++ b/litellm/harness/endpoint.py @@ -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)) diff --git a/litellm/harness/sync.py b/litellm/harness/sync.py index cdcf251af51..9787800833b 100644 --- a/litellm/harness/sync.py +++ b/litellm/harness/sync.py @@ -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 diff --git a/litellm/integrations/galileo.py b/litellm/integrations/galileo.py index 010f8ad8ef2..21906b0d996 100644 --- a/litellm/integrations/galileo.py +++ b/litellm/integrations/galileo.py @@ -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): diff --git a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py index 216749eda6c..80530622e36 100644 --- a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py +++ b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py @@ -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 {} diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 0c48b7c1c2a..fde687a41ab 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -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": diff --git a/litellm/llms/claude_code/harness/transformation.py b/litellm/llms/claude_code/harness/transformation.py index a85968f78be..ea3cdd67ffb 100644 --- a/litellm/llms/claude_code/harness/transformation.py +++ b/litellm/llms/claude_code/harness/transformation.py @@ -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) diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index cb3080e6534..dc3705b0fe5 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -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) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 5395d817e33..cb39801a8e0 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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). diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8e0dbfa7308..89f19a16da4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, diff --git a/tests/integration/authorization/test_audit_any_sweep_auth.py b/tests/integration/authorization/test_audit_any_sweep_auth.py new file mode 100644 index 00000000000..078f781ee36 --- /dev/null +++ b/tests/integration/authorization/test_audit_any_sweep_auth.py @@ -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 diff --git a/tests/integration/providers/test_audit_any_sweep_harness_endpoint.py b/tests/integration/providers/test_audit_any_sweep_harness_endpoint.py new file mode 100644 index 00000000000..5d3f438a217 --- /dev/null +++ b/tests/integration/providers/test_audit_any_sweep_harness_endpoint.py @@ -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 diff --git a/tests/integration/providers/test_audit_any_sweep_vector_store.py b/tests/integration/providers/test_audit_any_sweep_vector_store.py new file mode 100644 index 00000000000..269d8b90c4a --- /dev/null +++ b/tests/integration/providers/test_audit_any_sweep_vector_store.py @@ -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 diff --git a/tests/integration/routing/test_audit_any_sweep_chaos.py b/tests/integration/routing/test_audit_any_sweep_chaos.py new file mode 100644 index 00000000000..7f5e0b8b1c4 --- /dev/null +++ b/tests/integration/routing/test_audit_any_sweep_chaos.py @@ -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)