mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
7f95b5f361
commit
9dfa42dcde
13 changed files with 574 additions and 22 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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]}",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
50
tests/integration/mcp/test_mcp_tool_permission_merge.py
Normal file
50
tests/integration/mcp/test_mcp_tool_permission_merge.py
Normal 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
|
||||
84
tests/integration/observability/test_xecguard_wire.py
Normal file
84
tests/integration/observability/test_xecguard_wire.py
Normal 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
|
||||
|
|
@ -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")]
|
||||
|
|
@ -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")]
|
||||
|
|
@ -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"
|
||||
56
tests/integration/spend/test_chaos_burst_spend_once.py
Normal file
56
tests/integration/spend/test_chaos_burst_spend_once.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue