mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
fe683ea139
commit
e9bf2cfd01
13 changed files with 1107 additions and 34 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
384
tests/integration/authorization/test_audit_any_sweep_auth.py
Normal file
384
tests/integration/authorization/test_audit_any_sweep_auth.py
Normal 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
|
||||
|
|
@ -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
|
||||
346
tests/integration/providers/test_audit_any_sweep_vector_store.py
Normal file
346
tests/integration/providers/test_audit_any_sweep_vector_store.py
Normal 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
|
||||
203
tests/integration/routing/test_audit_any_sweep_chaos.py
Normal file
203
tests/integration/routing/test_audit_any_sweep_chaos.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue