diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 5af669591f7..06cbfd4fc04 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -2432,7 +2432,7 @@ class Logging(LiteLLMLoggingBaseClass): await invalidate_baseline_cache(self, reason, completed=completed) def _build_standard_logging_payload( - self, init_response_obj: object, start_time: Any, end_time: Any + self, init_response_obj: object, start_time: dt_object, end_time: dt_object ) -> StandardLoggingPayload | None: """Build StandardLoggingPayload and accumulate its construction time.""" _start: Final = time.time() @@ -2732,7 +2732,7 @@ class Logging(LiteLLMLoggingBaseClass): def success_handler( self, - result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) + result: object = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) start_time: datetime.datetime | None = None, end_time: datetime.datetime | None = None, cache_hit: bool | None = None, @@ -3171,7 +3171,7 @@ class Logging(LiteLLMLoggingBaseClass): async def async_success_handler( self, - result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) + result: object = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) start_time: datetime.datetime | None = None, end_time: datetime.datetime | None = None, cache_hit: bool | None = None, @@ -3189,7 +3189,7 @@ class Logging(LiteLLMLoggingBaseClass): async def _async_success_handler_body( self, - result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) + result: object = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) start_time: datetime.datetime | None = None, end_time: datetime.datetime | None = None, cache_hit: bool | None = None, @@ -4296,7 +4296,7 @@ class Logging(LiteLLMLoggingBaseClass): ) return result - def _handle_a2a_response_logging(self, result: Any) -> Any: + def _handle_a2a_response_logging(self, result: Any) -> object: """ Handles logging for A2A (Agent-to-Agent) responses. @@ -5705,7 +5705,7 @@ class StandardLoggingPayloadSetup: @staticmethod def get_standard_logging_metadata( - metadata: dict[str, Any] | None, + metadata: Mapping[str, object] | None, litellm_params: dict | None = None, prompt_integration: str | None = None, applied_guardrails: list[str] | None = None, diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index fa4650aec4f..946d19c028f 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -816,7 +816,7 @@ class CustomStreamWrapper: self, completion_obj: dict[str, Any], model_response: ModelResponseStream, - response_obj: dict[str, Any], + response_obj: Mapping[str, object], ) -> bool: if ( "content" in completion_obj diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 60d12337447..8d65aa7b0ca 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -201,6 +201,7 @@ if TYPE_CHECKING: from aiohttp import ClientSession from websockets.asyncio.client import ClientConnection + from litellm.google_genai.streaming_iterator import AsyncGoogleGenAIGenerateContentStreamingIterator from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer @@ -209,6 +210,7 @@ if TYPE_CHECKING: ) from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.google_genai.main import GenerateContentResponse from litellm.types.llms.openai_evals import ( CancelEvalResponse, CancelRunResponse, @@ -401,7 +403,8 @@ def _decoded_body_headers(response: httpx.Response) -> httpx.Headers: `aiter_bytes` yields the decoded body, so the upstream transfer headers only describe the bytes on the wire when no content-encoding was applied. """ - if response.headers.get("content-encoding", "identity").lower() == "identity": + headers: Final[Mapping[str, str]] = response.headers + if headers.get("content-encoding", "identity").lower() == "identity": return response.headers return httpx.Headers( [ @@ -3291,7 +3294,8 @@ class BaseLLMHTTPHandler: """ if upload_url_location == "headers": # Google Cloud Storage style - URL in X-Goog-Upload-URL header - upload_url = response.headers.get("X-Goog-Upload-URL") + upload_headers: Final[Mapping[str, str]] = response.headers + upload_url = upload_headers.get("X-Goog-Upload-URL") return upload_url, None else: # Response body style (e.g., Manus, S3 presigned URLs) @@ -11594,7 +11598,7 @@ class BaseLLMHTTPHandler: stream: bool = False, litellm_metadata: dict[str, object] | None = None, system_instruction: object | None = None, - ) -> Any: + ) -> "AsyncGoogleGenAIGenerateContentStreamingIterator | GenerateContentResponse": """ Async version of the generate content handler. Uses async HTTP client to make requests. diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 941ec4ad419..b5f32d57061 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -1571,12 +1571,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): gemini_call_id = part["functionCall"].get("id") if is_function_call is True: - function_dict: dict[str, Any] = dict(_function_chunk) - if thought_signature: - if "provider_specific_fields" not in function_dict: - function_dict["provider_specific_fields"] = {} - function_dict["provider_specific_fields"]["thought_signature"] = thought_signature - function = cast(ChatCompletionToolCallFunctionChunk, function_dict) + function = ( + {**_function_chunk, "provider_specific_fields": {"thought_signature": thought_signature}} + if thought_signature + else {**_function_chunk} + ) else: _tool_response_chunk: ChatCompletionToolCallChunk = { "id": f"call_{uuid.uuid4().hex[:28]}", diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 3e3b53f387d..490d3072955 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -28,7 +28,6 @@ from contextlib import asynccontextmanager from dataclasses import dataclass, replace from functools import lru_cache from itertools import chain, groupby -from operator import itemgetter from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse @@ -6988,7 +6987,7 @@ class MCPServerManager: ) return { server_id: list(dict.fromkeys(chain.from_iterable(tools for _, tools in group))) - for server_id, group in groupby(sorted(expanded, key=itemgetter(0)), key=itemgetter(0)) + for server_id, group in groupby(sorted(expanded, key=lambda pair: pair[0]), key=lambda pair: pair[0]) } def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index f4330ad6aa9..ddf9cace8b9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -474,7 +474,7 @@ class XecGuardGuardrail(CustomGuardrail): return "\n".join(text_parts) or None @staticmethod - def _extract_choice_content(choice: Any) -> Any: + def _extract_choice_content(choice: Any) -> object: if hasattr(choice, "message"): message = choice.message elif isinstance(choice, dict): diff --git a/litellm/utils.py b/litellm/utils.py index 13a46840431..09b5067339d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3595,7 +3595,7 @@ def get_optional_params_image_gen( passed_params.pop("provider_config", None) passed_params.pop("drop_params", None) drop_params = normalize_drop_params(drop_params) - additional_drop_params = passed_params.pop("additional_drop_params", None) + passed_params.pop("additional_drop_params", None) passed_params.pop("kwargs") special_params: Final[Mapping[str, object]] = kwargs for k, v in special_params.items(): @@ -4434,11 +4434,12 @@ def get_optional_params( store: bool | None = None, prompt_cache_key: str | None = None, base_model: str | None = None, - **kwargs, + **kwargs: object, ): drop_params = normalize_drop_params(drop_params) # rebind-ok: config and DB deployments pass "true" as a string passed_params: Final = locals().copy() - special_params: Final = passed_params.pop("kwargs") + passed_params.pop("kwargs") + special_params: Final = kwargs # Remove base_model from passed_params so it doesn't interfere with # non_default_params / _check_valid_arg — it's a routing hint, not an # OpenAI param. diff --git a/tests/integration/mcp/test_mcp_tool_permission_merge.py b/tests/integration/mcp/test_mcp_tool_permission_merge.py new file mode 100644 index 00000000000..40e57a6cd21 --- /dev/null +++ b/tests/integration/mcp/test_mcp_tool_permission_merge.py @@ -0,0 +1,50 @@ +import uuid +from typing import Final + +from integration._support.client import Gateway, eventually +from integration._support.mcp import ( + call_tool, + mcp_peer, + register_mcp, + tool_names, +) + + +def test_tool_permissions_merge_when_keys_resolve_to_same_server(gateway: Gateway) -> None: + with mcp_peer() as first, mcp_peer() as second, gateway.scenario() as scenario: + shared_alias: Final = "merge" + uuid.uuid4().hex[:8] + other_alias: Final = "other" + uuid.uuid4().hex[:8] + first_id: Final = register_mcp(scenario, first, shared_alias) + second_id: Final = register_mcp(scenario, second, other_alias) + key: Final = scenario.key( + object_permission={ + "mcp_servers": [first_id, second_id], + "mcp_tool_permissions": { + shared_alias: ["add"], + first_id: ["multiply", "add"], + second_id: ["add"], + }, + } + ) + + first_names: Final = eventually( + lambda: tool_names(gateway, key, first_id), + lambda names: set(names) != set(), + seconds=15, + ) + assert set(first_names) == {"add", "multiply"}, first_names + assert set(tool_names(gateway, key, second_id)) == {"add"} + + first.drain() + add: Final = call_tool(gateway, key, first_id, first_names["add"], {"a": 1, "b": 2}) + assert add.status_code == 200 and add.json()["isError"] is False, add.text + assert add.json()["content"][0]["text"] == "3" + multiply: Final = call_tool(gateway, key, first_id, first_names["multiply"], {"a": 2, "b": 3}) + assert multiply.status_code == 200 and multiply.json()["isError"] is False, multiply.text + assert multiply.json()["content"][0]["text"] == "6" + fail_name: Final = f"{shared_alias}-fail" + denied: Final = call_tool(gateway, key, first_id, fail_name, {}) + assert denied.status_code == 403, denied.text + detail: Final = denied.json()["detail"]["error"] + assert "is not allowed for your key/team" in detail and "fail" in detail, detail + assert len(tuple(item for item in first.drain() if item["body"].get("method") == "tools/call")) == 2 diff --git a/tests/integration/observability/test_xecguard_wire.py b/tests/integration/observability/test_xecguard_wire.py new file mode 100644 index 00000000000..df80ca83577 --- /dev/null +++ b/tests/integration/observability/test_xecguard_wire.py @@ -0,0 +1,84 @@ +import json +import uuid +from pathlib import Path +from typing import Final + +import yaml +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def test_xecguard_post_call_scan_reaches_vendor_and_call_succeeds(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "xecguard" + uuid.uuid4().hex + + def vendor(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == "/xecguard/v1/scan" + assert request.headers["authorization"] == "Bearer synthetic-xecguard-key" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == "xecguard_v2" + assert body["scan_type"] in ("input", "response") + assert any(message.get("content") == "hi" for message in body.get("messages", [])), body + return Reply(body=json.dumps({"decision": "SAFE", "violations": []}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/chat/completions" + return Reply( + body=json.dumps( + { + "id": "chatcmpl-xec", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "permitted"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + } + ).encode() + ) + + with wire_server(vendor) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "xecguard", + "mode": "post_call", + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-xecguard-key", + }, + } + ] + path: Final = tmp_path / "xecguard.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=upstream.url, + api_key="synthetic-openai-key", + ) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": "hi"}], + }, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "permitted" + scans: Final = tuple(request for request in policy.drain() if request.target == "/xecguard/v1/scan") + assert scans, "post-call xecguard scan never reached the vendor" + assert len(upstream.drain()) == 1 diff --git a/tests/integration/providers/test_image_gen_drop_params_wire.py b/tests/integration/providers/test_image_gen_drop_params_wire.py new file mode 100644 index 00000000000..7addfefeb36 --- /dev/null +++ b/tests/integration/providers/test_image_gen_drop_params_wire.py @@ -0,0 +1,49 @@ +import json +from typing import Final + +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def test_image_generation_additional_drop_params_reaches_provider_body(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == "/images/generations" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert "style" not in body, body + assert body["model"] == "dall-e-3" + assert body["prompt"] == "a scripted cat" + assert body["size"] == "1024x1024" + return Reply( + body=json.dumps( + { + "created": 1700000000, + "data": [{"b64_json": "aW1n", "revised_prompt": None, "url": None}], + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/dall-e-3", + api_base=wire.url, + api_key="synthetic-image-key", + additional_drop_params=["style"], + ) + response: Final = gateway.client.post( + "/v1/images/generations", + json={ + "model": model, + "prompt": "a scripted cat", + "size": "1024x1024", + "style": "vivid", + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) + assert response.status_code == 200, response.text + assert response.json()["data"][0]["b64_json"] == "aW1n" + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/images/generations")] diff --git a/tests/integration/providers/test_openai_stream_text_usage_wire.py b/tests/integration/providers/test_openai_stream_text_usage_wire.py new file mode 100644 index 00000000000..735d1a904bb --- /dev/null +++ b/tests/integration/providers/test_openai_stream_text_usage_wire.py @@ -0,0 +1,81 @@ +import json +from collections.abc import Mapping +from typing import Final + +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + +_IDENTITY: Final = "chatcmpl-stream-usage" + + +def _frame(delta: Mapping[str, JsonValue], finish: str | None = None) -> bytes: + return ( + b"data: " + + json.dumps( + { + "id": _IDENTITY, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish}], + } + ).encode() + + b"\n\n" + ) + + +def test_streaming_chat_assembles_text_and_final_usage(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.target == "/chat/completions" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["stream"] is True, body + assert body["stream_options"]["include_usage"] is True, body + usage: Final = json.dumps( + { + "id": _IDENTITY, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [], + "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, + } + ) + return Reply( + content_type="text/event-stream", + chunks=[ + _frame({"role": "assistant", "content": "Hello "}), + _frame({"content": "world"}), + _frame({}, finish="stop"), + b"data: " + usage.encode() + b"\n\n", + b"data: [DONE]\n\n", + ], + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(api_base=wire.url) + response: Final = gateway.client.post( + "/chat/completions", + json={ + "model": model, + "stream": True, + "stream_options": {"include_usage": True}, + "messages": [{"role": "user", "content": "hi"}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) + assert response.status_code == 200, response.text + chunks: Final = tuple( + json.loads(line[6:]) + for line in response.text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + text: Final = "".join(choice["delta"].get("content", "") for chunk in chunks for choice in chunk["choices"]) + assert text == "Hello world" + usages: Final = tuple(chunk["usage"] for chunk in chunks if chunk.get("usage")) + assert len(usages) == 1 + assert usages[0]["prompt_tokens"] == 11 and usages[0]["completion_tokens"] == 4 + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")] diff --git a/tests/integration/providers/test_vertex_gemini_function_call_wire.py b/tests/integration/providers/test_vertex_gemini_function_call_wire.py new file mode 100644 index 00000000000..fef5e31f9c7 --- /dev/null +++ b/tests/integration/providers/test_vertex_gemini_function_call_wire.py @@ -0,0 +1,229 @@ +import json +from typing import Final + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, Scenario +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gemini-3.7-flash" +_PROJECT: Final = "scripted-project" +_LOCATION: Final = "us-central1" +_MODEL_PATH: Final = f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/google/models/{_BACKEND}" +_SIGNATURE: Final = "sig-4f2a" +_ARGS: Final = {"city": "Paris"} +_FUNCTIONS: Final = [ + { + "name": "get_weather", + "description": "Return the weather for a city", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + } +] +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def _service_account_json(token_url: str) -> str: + private_key: Final = ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() + ) + return json.dumps( + { + "type": "service_account", + "project_id": _PROJECT, + "private_key_id": "scripted", + "private_key": private_key, + "client_email": f"scripted@{_PROJECT}.iam.gserviceaccount.com", + "client_id": "0", + "auth_uri": f"{token_url}/_oauth/authorize", + "token_uri": f"{token_url}/_oauth/token", + } + ) + + +def _candidate(*, with_signature: bool) -> dict[str, JsonValue]: + part: Final = { + "functionCall": {"name": "get_weather", "args": _ARGS, "id": "fc-1"}, + **({"thoughtSignature": _SIGNATURE} if with_signature else {}), + } + return { + "candidates": [ + { + "content": {"role": "model", "parts": [part]}, + "finishReason": "STOP", + } + ], + "usageMetadata": {"promptTokenCount": 11, "candidatesTokenCount": 7, "totalTokenCount": 18}, + "modelVersion": _BACKEND, + } + + +def _model(gateway: Gateway, scenario: Scenario, wire_url: str) -> str: + return scenario.model( + model=f"vertex_ai/{_BACKEND}", + api_base=f"{wire_url}{_MODEL_PATH}", + api_key=None, + vertex_project=_PROJECT, + vertex_location=_LOCATION, + vertex_credentials=_service_account_json(gateway.upstream_url.rstrip("/")), + ) + + +def _non_streaming_call(gateway: Gateway, model: str) -> dict[str, JsonValue]: + response: Final = gateway.client.post( + "/v1/chat/completions", + json={ + "model": model, + "functions": _FUNCTIONS, + "messages": [{"role": "user", "content": "weather?"}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) + assert response.status_code == 200, response.text + return response.json() + + +def _streaming_call(gateway: Gateway, model: str) -> tuple[dict[str, JsonValue], ...]: + with gateway.client.stream( + "POST", + "/v1/chat/completions", + json={ + "model": model, + "functions": _FUNCTIONS, + "messages": [{"role": "user", "content": "weather?"}], + "stream": True, + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) as response: + assert response.status_code == 200, response.read() + lines: Final = tuple(line for line in response.iter_lines() if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", lines[-3:] + return tuple(_JSON_OBJECT.validate_json(line.removeprefix("data: ").encode()) for line in lines[:-1]) + + +def _function_call_of(response: dict[str, JsonValue]) -> dict[str, JsonValue]: + message: Final = response["choices"][0]["message"] + assert isinstance(message, dict) + call: Final = message["function_call"] + assert isinstance(call, dict) + return call + + +def test_vertex_gemini_function_call_thought_signature_is_returned_non_streaming(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:generateContent" + return Reply(body=json.dumps(_candidate(with_signature=True)).encode()) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + call: Final = _function_call_of(_non_streaming_call(gateway, model)) + assert call["name"] == "get_weather" + assert json.loads(str(call["arguments"])) == _ARGS + assert call.get("provider_specific_fields") == {"thought_signature": _SIGNATURE} + + +def test_vertex_gemini_function_call_without_signature_has_no_provider_fields_non_streaming( + gateway: Gateway, +) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:generateContent" + return Reply(body=json.dumps(_candidate(with_signature=False)).encode()) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + call: Final = _function_call_of(_non_streaming_call(gateway, model)) + assert call["name"] == "get_weather" + assert json.loads(str(call["arguments"])) == _ARGS + assert "provider_specific_fields" not in call + assert "thought_signature" not in json.dumps(call) + + +def test_vertex_gemini_function_call_thought_signature_is_returned_streaming(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:streamGenerateContent?alt=sse" + payload: Final = json.dumps(_candidate(with_signature=True)) + return Reply(content_type="text/event-stream", chunks=[f"data: {payload}\n\n".encode()]) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + chunks: Final = _streaming_call(gateway, model) + function_calls: Final = tuple( + choice["delta"]["function_call"] + for chunk in chunks + for choice in chunk.get("choices", ()) + if choice.get("delta", {}).get("function_call") + ) + assert function_calls, "no function_call delta received" + merged: Final = "".join(str(call.get("arguments", "")) for call in function_calls) + assert json.loads(merged) == _ARGS + assert function_calls[-1].get("provider_specific_fields") == {"thought_signature": _SIGNATURE} + + +def test_vertex_gemini_function_call_without_signature_has_no_provider_fields_streaming( + gateway: Gateway, +) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:streamGenerateContent?alt=sse" + payload: Final = json.dumps(_candidate(with_signature=False)) + return Reply(content_type="text/event-stream", chunks=[f"data: {payload}\n\n".encode()]) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + chunks: Final = _streaming_call(gateway, model) + function_calls: Final = tuple( + choice["delta"]["function_call"] + for chunk in chunks + for choice in chunk.get("choices", ()) + if choice.get("delta", {}).get("function_call") + ) + assert function_calls, "no function_call delta received" + assert all("provider_specific_fields" not in call for call in function_calls) + assert "thought_signature" not in json.dumps(function_calls) + + +def test_vertex_gemini_kwargs_extra_param_reaches_generation_config(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:generateContent" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["generationConfig"]["top_k"] == 3, body + return Reply( + body=json.dumps( + { + "candidates": [ + { + "content": {"role": "model", "parts": [{"text": "done"}]}, + "finishReason": "STOP", + } + ], + "usageMetadata": {"promptTokenCount": 4, "candidatesTokenCount": 2, "totalTokenCount": 6}, + "modelVersion": _BACKEND, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + response: Final = gateway.client.post( + "/v1/chat/completions", + json={ + "model": model, + "top_k": 3, + "messages": [{"role": "user", "content": "hi"}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "done" diff --git a/tests/integration/spend/test_chaos_burst_spend_once.py b/tests/integration/spend/test_chaos_burst_spend_once.py new file mode 100644 index 00000000000..77b08b1d559 --- /dev/null +++ b/tests/integration/spend/test_chaos_burst_spend_once.py @@ -0,0 +1,56 @@ +import uuid +from concurrent.futures import ThreadPoolExecutor +from typing import Final + +import httpx +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows + +_BURST: Final = 24 + + +def test_burst_with_partial_upstream_failures_logs_each_success_once(gateway: Gateway) -> None: + with ( + httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream, + gateway.scenario() as scenario, + ): + provider_model: Final = f"burst-{uuid.uuid4().hex}" + model: Final = scenario.model(model=f"openai/{provider_model}", input_cost_per_token=0, output_cost_per_token=0) + statuses: Final = [500] + [200, 200, 200] * (_BURST // 4 + 2) + + def remove_script() -> None: + response: Final = upstream.delete(f"/__scripts/{provider_model}") + assert response.status_code in (200, 404), response.text + + scenario.cleanups.callback(remove_script) + configured: Final = upstream.post(f"/__scripts/{provider_model}", json={"statuses": statuses}) + assert configured.status_code == 200, configured.text + upstream.get("/__observations").raise_for_status() + + def attempt(index: int) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"burst {index}"}]}, + ) + + with ThreadPoolExecutor(max_workers=_BURST) as pool: + responses: Final = tuple(pool.map(attempt, range(_BURST))) + + succeeded: Final = tuple(response.json()["id"] for response in responses if response.status_code == 200) + assert len(succeeded) > 0, [response.status_code for response in responses] + assert len(set(succeeded)) == len(succeeded), "duplicate response id in burst" + assert all(response.status_code in (200, 429, 500) for response in responses), [ + response.status_code for response in responses + ] + + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s)', + (list(succeeded),), + ), + lambda values: len(values) == len(succeeded), + seconds=90, + ) + landed: Final = [row["request_id"] for row in rows] + assert sorted(landed) == sorted(succeeded), "a successful burst id did not land exactly once"