refactor(types): replace Any with proven types in 7 files (#43704)

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

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

* refactor(types): revert Any changes that broke existing callers

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

* refactor(types): drop prompt factory helper wrappers

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

* test(integration): cover typing sweep surfaces

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

* test(integration): tighten sweep audit tests

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

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-29 06:12:58 -07:00 • committed by GitHub
parent 7f95b5f361
commit 9dfa42dcde
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 574 additions and 22 deletions

View file

@ -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,

View file

@ -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

View file

@ -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.

View file

@ -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]}",

View file

@ -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:

View file

@ -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):

View file

@ -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.

View file

@ -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

View file

@ -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

View file

@ -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")]

View file

@ -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")]

View file

@ -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"

View file

@ -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"