Merge pull request #37778 from BerriAI/litellm_decrease_anys_opus5

chore(typing): clear Any seams across 47 files, ratchet basedpyright ceilings -3,302
This commit is contained in:
Mateo Wang 2026-08-29 21:48:11 -07:00 • committed by GitHub
commit 5e4b3838aa
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
52 changed files with 1128 additions and 590 deletions

View file

@ -1,9 +1,9 @@
{
"reportAny": {
"limit": 17270
"limit": 16389
},
"reportArgumentType": {
"limit": 2538
"limit": 2229
},
"reportAssignmentType": {
"limit": 319
@ -24,7 +24,7 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 5485
"limit": 5242
},
"reportFunctionMemberAccess": {
"limit": 7
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5658
"limit": 5614
},
"reportMissingTypeArgument": {
"limit": 15425
"limit": 15356
},
"reportMissingTypeStubs": {
"limit": 40
@ -72,7 +72,7 @@
"limit": 0
},
"reportOptionalMemberAccess": {
"limit": 1055
"limit": 0
},
"reportOptionalOperand": {
"limit": 0
@ -90,7 +90,7 @@
"limit": 8
},
"reportReturnType": {
"limit": 213
"limit": 181
},
"reportTypedDictNotRequiredAccess": {
"limit": 25
@ -99,25 +99,25 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44526
"limit": 44389
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38721
"limit": 38500
},
"reportUnknownParameterType": {
"limit": 19778
"limit": 19673
},
"reportUnknownVariableType": {
"limit": 30290
"limit": 30092
},
"reportUnnecessaryCast": {
"limit": 117
"limit": 111
},
"reportUnnecessaryComparison": {
"limit": 697
"limit": 695
},
"reportUnnecessaryContains": {
"limit": 5

View file

@ -7,7 +7,7 @@ GET - /audit/{id} - Get audit log by id
GET - /audit - Get all audit logs
"""
from typing import TYPE_CHECKING, Final, Optional
from typing import TYPE_CHECKING, Final
#### AUDIT LOGGING ####
from fastapi import APIRouter, Depends, HTTPException, Query
@ -58,33 +58,33 @@ async def get_audit_logs(
page: int = Query(1, ge=1),
page_size: int = Query(10, ge=1, le=100),
# Filter parameters
changed_by: Optional[str] = Query(
changed_by: str | None = Query(
None, description="Filter by user or system that performed the action"
),
changed_by_api_key: Optional[str] = Query(
changed_by_api_key: str | None = Query(
None, description="Filter by API key hash that performed the action"
),
action: Optional[str] = Query(
action: str | None = Query(
None, description="Filter by action type (create, update, delete)"
),
table_name: Optional[str] = Query(
table_name: str | None = Query(
None, description="Filter by table name that was modified"
),
object_id: Optional[str] = Query(
object_id: str | None = Query(
None, description="Filter by ID of the object that was modified"
),
start_date: Optional[str] = Query(None, description="Filter logs after this date"),
end_date: Optional[str] = Query(None, description="Filter logs before this date"),
object_team_id: Optional[str] = Query(
start_date: str | None = Query(None, description="Filter logs after this date"),
end_date: str | None = Query(None, description="Filter logs before this date"),
object_team_id: str | None = Query(
None,
description="Filter by team_id present in before_value or updated_values JSON (PostgreSQL only)",
),
object_key_hash: Optional[str] = Query(
object_key_hash: str | None = Query(
None,
description="Filter by token (key hash) present in before_value or updated_values JSON (PostgreSQL only)",
),
# Sorting parameters
sort_by: Optional[str] = Query(
sort_by: str | None = Query(
None,
description="Column to sort by (e.g. 'updated_at', 'action', 'table_name')",
),

View file

@ -17,7 +17,7 @@ until they're actually needed.
import importlib
import sys
from collections.abc import Callable
from collections.abc import Callable, Mapping
from types import ModuleType
from typing import TYPE_CHECKING, Any, Final, cast
@ -57,10 +57,11 @@ from ._lazy_imports_registry import (
)
if TYPE_CHECKING:
import httpx
from tiktoken import Encoding
def get_litellm_globals() -> dict:
def get_litellm_globals() -> dict[str, object]:
"""
Get the globals dictionary of the litellm module.
@ -70,7 +71,7 @@ def get_litellm_globals() -> dict:
return sys.modules["litellm"].__dict__
def _get_utils_globals() -> dict:
def _get_utils_globals() -> dict[str, object]:
"""
Get the globals dictionary of the utils module.
@ -80,6 +81,11 @@ def _get_utils_globals() -> dict:
return sys.modules["litellm.utils"].__dict__
def _get_module_level_client_timeout(litellm_globals: Mapping[str, Any]) -> "float | httpx.Timeout | None":
"""Read the configured `litellm.request_timeout` used for the module level http clients."""
return litellm_globals.get("request_timeout")
# These are special lazy loaders for things that are used internally
# They're separate from the main lazy import system because they have specific use cases
@ -435,8 +441,8 @@ def _lazy_import_http_handlers(name: str) -> object:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
# Get timeout from module config (if set)
timeout = _globals.get("request_timeout")
params: Final = {"timeout": timeout, "client_alias": "module level aclient"}
async_timeout: Final = _get_module_level_client_timeout(_globals)
params: Final = {"timeout": async_timeout, "client_alias": "module level aclient"}
# Create the client instance
provider_id: Final = cast(Any, "litellm_module_level_client")
@ -453,8 +459,8 @@ def _lazy_import_http_handlers(name: str) -> object:
# Create a sync HTTP client
from litellm.llms.custom_httpx.http_handler import HTTPHandler
timeout = _globals.get("request_timeout")
sync_client: Final = HTTPHandler(timeout=timeout)
sync_timeout: Final = _get_module_level_client_timeout(_globals)
sync_client: Final = HTTPHandler(timeout=sync_timeout)
# Cache it
_globals["module_level_client"] = sync_client

View file

@ -17,6 +17,7 @@ RedisSemanticCache since those are backend agnostic.
import asyncio
import hashlib
import os
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Final
@ -64,7 +65,7 @@ class ValkeySemanticCache(RedisSemanticCache):
async_client: AsyncRedis | None = None,
embedding_max_input_tokens: int | None = None,
embedding_timeout: float | None = None,
**kwargs: Any,
**kwargs: object,
):
if similarity_threshold is None:
raise ValueError("similarity_threshold must be provided, passed None")
@ -87,11 +88,13 @@ class ValkeySemanticCache(RedisSemanticCache):
self.key_prefix = f"{self.index_name}:"
self._index_dim: int | None = None
resolved_url = None
if sync_client is None or async_client is None:
resolved_url = redis_url or self._build_valkey_url(host, port, password, ssl)
self.sync_client = sync_client if sync_client is not None else Redis.from_url(resolved_url)
self.async_client = async_client if async_client is not None else AsyncRedis.from_url(resolved_url)
if sync_client is not None and async_client is not None:
self.sync_client = sync_client
self.async_client = async_client
else:
resolved_url: Final = redis_url or self._build_valkey_url(host, port, password, ssl)
self.sync_client = sync_client if sync_client is not None else Redis.from_url(resolved_url)
self.async_client = async_client if async_client is not None else AsyncRedis.from_url(resolved_url)
print_verbose(f"Valkey semantic-cache initializing index - {self.index_name}")
@ -118,7 +121,7 @@ class ValkeySemanticCache(RedisSemanticCache):
return hashlib.sha256(str(key).encode("utf-8")).hexdigest()
@staticmethod
def _embedding_to_bytes(embedding: list[float]) -> bytes:
def _embedding_to_bytes(embedding: Sequence[float]) -> bytes:
return pack_vector(embedding)
def _index_schema(self, dim: int) -> tuple[TagField, VectorField]:
@ -192,7 +195,9 @@ class ValkeySemanticCache(RedisSemanticCache):
def _doc_key(self, key: str) -> str:
return f"{self.key_prefix}{self._scope_tag(key)}:{uuid.uuid4()}"
def _doc_mapping(self, key: str, prompt: str, value_str: str, embedding: list[float]) -> dict:
def _doc_mapping(
self, key: str, prompt: str, value_str: str, embedding: Sequence[float]
) -> Mapping[str | bytes, str | bytes]:
return {
self.CACHE_KEY_FIELD_NAME: self._scope_tag(key),
self.PROMPT_FIELD_NAME: prompt,
@ -208,30 +213,49 @@ class ValkeySemanticCache(RedisSemanticCache):
)
return Query(query_string).return_fields(self.RESPONSE_FIELD_NAME, self.DISTANCE_FIELD_NAME).dialect(2)
async def _async_search(self, key: str, embedding: Sequence[float]) -> object:
"""Run the KNN query on the async client, stopping the untyped search surface here."""
return await self.async_client.ft(self.index_name).search(
self._knn_query(key),
query_params={"vec": self._embedding_to_bytes(embedding)}, # pyright: ignore[reportArgumentType] # redis stubs omit bytes; KNN vectors are raw bytes at runtime
)
@classmethod
def _first_hit(cls, search_result: Any) -> _ValkeyCacheHit | None:
docs: Final = getattr(search_result, "docs", [])
def _first_hit(cls, search_result: object) -> _ValkeyCacheHit | None:
docs: Final[Sequence[object]] = getattr(search_result, "docs", [])
if not docs:
return None
doc: Final = docs[0]
response_field: Final[object] = getattr(doc, cls.RESPONSE_FIELD_NAME)
distance_field: Final[str | bytes | float] = getattr(doc, cls.DISTANCE_FIELD_NAME)
return _ValkeyCacheHit(
response=str(getattr(doc, cls.RESPONSE_FIELD_NAME)),
distance=float(getattr(doc, cls.DISTANCE_FIELD_NAME)),
response=str(response_field),
distance=float(distance_field),
)
def _resolve_hit(self, hit: _ValkeyCacheHit | None, key: str, **kwargs: Any) -> Any:
@staticmethod
def _record_similarity(kwargs: dict[str, Any], similarity: float) -> None:
"""Stamp the semantic-similarity score onto the request metadata carried in ``kwargs``."""
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
@staticmethod
def _embedding_metadata(kwargs: dict[str, Any]) -> dict[str, Any] | None:
"""The request metadata forwarded to the embedding call."""
return kwargs.get("metadata")
def _resolve_hit(self, hit: _ValkeyCacheHit | None, key: str, **kwargs: object) -> object:
if hit is None:
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
self._record_similarity(kwargs, 0.0)
return None
similarity: Final = 1 - hit.distance
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
self._record_similarity(kwargs, similarity)
if similarity < self.similarity_threshold:
return None
return self._get_cache_logic(cached_response=hit.response)
def set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
def set_cache(self, key: str, value: object, **kwargs: object) -> None:
print_verbose(f"Valkey semantic-cache set_cache, kwargs: {kwargs}")
try:
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
@ -250,12 +274,12 @@ class ValkeySemanticCache(RedisSemanticCache):
except Exception as e:
print_verbose(f"Error in Valkey semantic-cache set_cache: {e}")
def get_cache(self, key: str, **kwargs: Any) -> Any:
def get_cache(self, key: str, **kwargs: object) -> object:
print_verbose(f"Valkey semantic-cache get_cache, kwargs: {kwargs}")
try:
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
if prompt is None:
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
self._record_similarity(kwargs, 0.0)
return None
embedding: Final = self._get_embedding(prompt)
@ -263,14 +287,14 @@ class ValkeySemanticCache(RedisSemanticCache):
search_result: Final = self.sync_client.ft(self.index_name).search(
self._knn_query(key),
query_params={"vec": self._embedding_to_bytes(embedding)},
query_params={"vec": self._embedding_to_bytes(embedding)}, # pyright: ignore[reportArgumentType] # redis stubs omit bytes; KNN vectors are raw bytes at runtime
)
return self._resolve_hit(self._first_hit(search_result), key, **kwargs)
except Exception as e:
print_verbose(f"Error in Valkey semantic-cache get_cache: {e}")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
self._record_similarity(kwargs, 0.0)
async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None:
print_verbose(f"Async Valkey semantic-cache set_cache, kwargs: {kwargs}")
try:
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
@ -278,7 +302,7 @@ class ValkeySemanticCache(RedisSemanticCache):
print_verbose("No prompt provided for semantic caching")
return
embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
embedding: Final = await self._get_async_embedding(prompt, metadata=self._embedding_metadata(kwargs))
await self._ensure_index_async(len(embedding))
doc_key: Final = self._doc_key(key)
@ -289,31 +313,28 @@ class ValkeySemanticCache(RedisSemanticCache):
except Exception as e:
print_verbose(f"Error in async Valkey semantic-cache set_cache: {e}")
async def async_get_cache(self, key: str, **kwargs: Any) -> Any:
async def async_get_cache(self, key: str, **kwargs: object) -> object:
print_verbose(f"Async Valkey semantic-cache get_cache, kwargs: {kwargs}")
try:
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
if prompt is None:
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
self._record_similarity(kwargs, 0.0)
return None
embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
embedding: Final = await self._get_async_embedding(prompt, metadata=self._embedding_metadata(kwargs))
await self._ensure_index_async(len(embedding))
search_result: Final = await self.async_client.ft(self.index_name).search(
self._knn_query(key),
query_params={"vec": self._embedding_to_bytes(embedding)},
)
search_result: Final[object] = await self._async_search(key, embedding)
return self._resolve_hit(self._first_hit(search_result), key, **kwargs)
except Exception as e:
print_verbose(f"Error in async Valkey semantic-cache get_cache: {e}")
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
self._record_similarity(kwargs, 0.0)
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs: Any) -> None:
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, object]], **kwargs: object) -> None:
try:
await asyncio.gather(*[self.async_set_cache(key, value, **kwargs) for key, value in cache_list])
except Exception as e:
print_verbose(f"Error in Valkey semantic-cache async_set_cache_pipeline: {e}")
async def _index_info(self) -> dict:
async def _index_info(self) -> Mapping[str, object]:
return await self.async_client.ft(self.index_name).info()

View file

@ -163,15 +163,11 @@ class BitBucketClient:
response.raise_for_status()
data: Final[BitBucketSrcListing] = response.json()
files: Final[list[str]] = []
for item in data.get("values", []):
if item.get("type") == "commit_file":
file_path = item.get("path", "")
if file_path.endswith(file_extension):
files.append(file_path)
return files
return [
file_path
for item in data.get("values", [])
if item.get("type") == "commit_file" and (file_path := item.get("path", "")).endswith(file_extension)
]
except Exception as e:
# Check if it's an HTTP error

View file

@ -7,7 +7,10 @@ litellm_content_retrieve tool calls server-side via the typed agentic loop plan.
import time
import uuid
from typing import TYPE_CHECKING, Any, ClassVar, Final, cast
from collections.abc import Mapping, Sequence
from typing import Any, ClassVar, Final, Protocol, cast
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_logger
from litellm.compression import compress
@ -22,13 +25,23 @@ from litellm.types.integrations.custom_logger import (
)
from litellm.types.utils import CallTypes
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
LITELLM_CONTENT_RETRIEVE_TOOL_NAME: Final = "litellm_content_retrieve"
_CACHE_TTL_SECONDS: Final = 15 * 60
class _AgenticLoopParams(TypedDict, total=False):
"""The ``agentic_loop_params`` entry the agentic loop driver records on the logging object."""
model: ReadOnly[str]
class _AgenticLoopLoggingObj(Protocol):
"""Logging object view exposing the untyped call details this handler reads."""
@property
def model_call_details(self) -> Mapping[str, _AgenticLoopParams]: ...
def _compression_savings_from_counts(
original_tokens: object, compressed_tokens: object
) -> CompressionSavingsMetadata | None:
@ -83,7 +96,7 @@ class CompressionInterceptionLogger(CustomLogger):
compression_trigger: int = 200_000,
compression_target: int | None = None,
embedding_model: str | None = None,
embedding_model_params: dict[str, Any] | None = None,
embedding_model_params: dict[str, object] | None = None,
):
super().__init__()
self.enabled = enabled
@ -106,7 +119,7 @@ class CompressionInterceptionLogger(CustomLogger):
@staticmethod
def initialize_from_proxy_config(
litellm_settings: dict[str, Any],
callback_specific_params: dict[str, Any],
callback_specific_params: Mapping[str, object],
) -> "CompressionInterceptionLogger":
compression_params: CompressionInterceptionConfig = {}
if "compression_interception_params" in litellm_settings:
@ -120,7 +133,9 @@ class CompressionInterceptionLogger(CustomLogger):
)
return CompressionInterceptionLogger.from_config_yaml(compression_params)
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
async def async_pre_call_deployment_hook(
self, kwargs: dict[str, Any], call_type: CallTypes | None
) -> dict[str, object] | None:
if not self.enabled:
return None
if call_type is not None and call_type != CallTypes.anthropic_messages:
@ -150,7 +165,7 @@ class CompressionInterceptionLogger(CustomLogger):
cache: Final = cast(dict[str, str], compressed.get("cache", {}))
skip_reason: Final = cast(str | None, compressed.get("compression_skipped_reason"))
compressed_tools: Final = cast(list[dict[str, Any]], compressed.get("tools", []))
compressed_tools: Final = cast(list[dict[str, object]], compressed.get("tools", []))
# Only mutate kwargs when compression actually produced a result.
# If compression was a no-op (below trigger, invalid tool sequence, etc.),
@ -161,7 +176,7 @@ class CompressionInterceptionLogger(CustomLogger):
kwargs["messages"] = compressed["messages"]
if compressed_tools:
kwargs["tools"] = self._merge_tools(
existing_tools=cast(list[dict[str, Any]] | None, kwargs.get("tools")),
existing_tools=cast(list[dict[str, object]] | None, kwargs.get("tools")),
compressed_tools=compressed_tools,
)
call_id = cast(str | None, kwargs.get("litellm_call_id"))
@ -194,14 +209,14 @@ class CompressionInterceptionLogger(CustomLogger):
async def async_should_run_agentic_loop(
self,
response: Any,
response: object,
model: str,
messages: list[dict],
tools: list[dict] | None,
messages: Sequence[Mapping[str, object]],
tools: Sequence[Mapping[str, object]] | None,
stream: bool,
custom_llm_provider: str,
kwargs: dict,
) -> tuple[bool, dict]:
kwargs: Mapping[str, object],
) -> tuple[bool, dict[str, object]]:
if not self.enabled:
return False, {}
if not self._has_retrieval_tool(tools):
@ -219,19 +234,19 @@ class CompressionInterceptionLogger(CustomLogger):
async def async_build_agentic_loop_plan(
self,
tools: dict,
tools: Mapping[str, object],
model: str,
messages: list[dict],
response: Any,
anthropic_messages_provider_config: Any,
anthropic_messages_optional_request_params: dict,
logging_obj: "LiteLLMLoggingObj | None",
messages: list[dict[str, object]],
response: object,
anthropic_messages_provider_config: object,
anthropic_messages_optional_request_params: Mapping[str, object],
logging_obj: _AgenticLoopLoggingObj | None,
stream: bool,
kwargs: dict,
kwargs: Mapping[str, object],
) -> AgenticLoopPlan:
self._prune_expired_cache()
tool_calls: Final = cast(list[dict[str, Any]], tools.get("tool_calls", []))
thinking_blocks: Final = cast(list[dict[str, Any]], tools.get("thinking_blocks", []))
tool_calls: Final = cast(list[dict[str, object]], tools.get("tool_calls", []))
thinking_blocks: Final = cast(list[dict[str, object]], tools.get("thinking_blocks", []))
call_id: Final = self._resolve_call_id(logging_obj=logging_obj, kwargs=kwargs)
cache: Final = self._get_cache(call_id=call_id)
@ -274,7 +289,7 @@ class CompressionInterceptionLogger(CustomLogger):
full_model_name = model
if logging_obj is not None:
agentic_params: Final = logging_obj.model_call_details.get("agentic_loop_params", {})
full_model_name = cast(str, agentic_params.get("model", model))
full_model_name = agentic_params.get("model", model)
request_patch: Final = AgenticLoopRequestPatch(
model=full_model_name,
@ -309,15 +324,15 @@ class CompressionInterceptionLogger(CustomLogger):
return {}
return cache_entry[0]
def _resolve_call_id(self, logging_obj: Any, kwargs: dict[str, Any]) -> str | None:
def _resolve_call_id(self, logging_obj: _AgenticLoopLoggingObj | None, kwargs: Mapping[str, object]) -> str | None:
if logging_obj is not None:
logging_call_id: Final = getattr(logging_obj, "litellm_call_id", None)
if isinstance(logging_call_id, str) and logging_call_id:
return logging_call_id
kwargs_call_id: Final = kwargs.get("litellm_call_id")
return cast(str | None, kwargs_call_id if isinstance(kwargs_call_id, str) else None)
return kwargs_call_id if isinstance(kwargs_call_id, str) else None
def _resolve_retrieval_content(self, tool_call: dict[str, Any], cache: dict[str, str]) -> str:
def _resolve_retrieval_content(self, tool_call: Mapping[str, object], cache: Mapping[str, str]) -> str:
raw_input: Final = tool_call.get("input", {})
key = ""
if isinstance(raw_input, dict):
@ -328,7 +343,9 @@ class CompressionInterceptionLogger(CustomLogger):
return cache[key]
return f"[compressed content key '{key}' not found]"
def _extract_retrieval_tool_calls(self, response: Any) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
def _extract_retrieval_tool_calls(
self, response: object
) -> tuple[list[dict[str, object]], list[dict[str, object]]]:
if isinstance(response, dict):
content = response.get("content", [])
else:
@ -337,8 +354,8 @@ class CompressionInterceptionLogger(CustomLogger):
if not isinstance(content, list):
return [], []
tool_calls: Final[list[dict[str, Any]]] = []
thinking_blocks: Final[list[dict[str, Any]]] = []
tool_calls: Final[list[dict[str, object]]] = []
thinking_blocks: Final[list[dict[str, object]]] = []
for block in content:
if isinstance(block, dict):
@ -385,13 +402,13 @@ class CompressionInterceptionLogger(CustomLogger):
return tool_calls, thinking_blocks
def _prepare_followup_kwargs(self, kwargs: dict[str, Any]) -> dict[str, Any]:
def _prepare_followup_kwargs(self, kwargs: Mapping[str, object]) -> dict[str, object]:
internal_keys: Final = {"litellm_logging_obj"}
return {
k: v for k, v in kwargs.items() if not k.startswith("_compression_interception") and k not in internal_keys
}
def _has_retrieval_tool(self, tools: Any) -> bool:
def _has_retrieval_tool(self, tools: object) -> bool:
if not isinstance(tools, list):
return False
for tool in tools:
@ -407,9 +424,9 @@ class CompressionInterceptionLogger(CustomLogger):
def _merge_tools(
self,
existing_tools: list[dict[str, Any]] | None,
compressed_tools: list[dict[str, Any]],
) -> list[dict[str, Any]]:
existing_tools: Sequence[Mapping[str, object]] | None,
compressed_tools: Sequence[Mapping[str, object]],
) -> list[Mapping[str, object]]:
merged: Final = list(existing_tools or [])
if self._has_retrieval_tool(merged):
return merged

View file

@ -2,7 +2,7 @@
# On success, logs events to Promptlayer
import re
import traceback
from collections.abc import AsyncGenerator, Mapping
from collections.abc import AsyncGenerator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional
from pydantic import BaseModel
@ -123,11 +123,11 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
return []
callbacks: Final = AllCallbacks()
callback_info: Final = getattr(callbacks, lookup_name, None)
callback_info: Final[object] = getattr(callbacks, lookup_name, None)
if callback_info is None:
return []
params: Final = getattr(callback_info, "litellm_callback_params", None)
params: Final[Sequence[str] | None] = getattr(callback_info, "litellm_callback_params", None)
if not params:
return []
@ -851,7 +851,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
- Converting to string and then truncating the logged content catches this
2. We want to avoid modifying the original `messages`, `response`, and `error_str` in the logging payload since these are in kwargs and could be returned to the user
"""
field_value: Final = standard_logging_object.get(field_name)
field_value: Final[object] = standard_logging_object.get(field_name)
if field_value:
str_value: Final = str(field_value)
if len(str_value) > max_length:
@ -1005,8 +1005,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
• Keep untyped or text content.
• Recursively redact inline base64 blobs in *any* string field, at any depth.
"""
raw_messages: Final[Any] = payload.get("messages", [])
messages: Final[list[Any]] = raw_messages if isinstance(raw_messages, list) else []
raw_messages: Final[object] = payload.get("messages", [])
messages: Final[list[object]] = raw_messages if isinstance(raw_messages, list) else []
verbose_logger.debug("[CustomLogger] Stripping base64 from %s messages", len(messages))
if messages:
@ -1037,8 +1037,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
• Keep untyped or text content.
• Recursively redact inline base64 blobs in *any* string field, at any depth.
"""
raw_messages: Final[Any] = payload.get("messages", [])
messages: Final[list[Any]] = raw_messages if isinstance(raw_messages, list) else []
raw_messages: Final[object] = payload.get("messages", [])
messages: Final[list[object]] = raw_messages if isinstance(raw_messages, list) else []
verbose_logger.debug("[CustomLogger] Stripping base64 from %s messages", len(messages))
if messages:
@ -1059,7 +1059,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
value: Any,
depth: int = 0,
max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER,
) -> Any:
) -> object:
"""Recursively redact inline base64 from any nested structure with a max recursion depth limit."""
if depth > max_depth:
verbose_logger.warning("[CustomLogger] Max recursion depth %s reached while redacting base64", max_depth)
@ -1090,16 +1090,16 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
def _process_messages(
self,
messages: list[Any],
messages: list[object],
max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER,
) -> list[dict[str, Any]]:
filtered_messages: Final[list[dict[str, Any]]] = []
) -> list[dict[str, object]]:
filtered_messages: Final[list[dict[str, object]]] = []
for msg in messages:
if not isinstance(msg, dict):
continue
contents: Any = msg.get("content")
contents: object = msg.get("content")
if isinstance(contents, list):
cleaned: list[Any] = []
cleaned: list[object] = []
for c in contents:
if self._should_keep_content(content=c):
cleaned.append(self._redact_base64(value=c, max_depth=max_depth))

View file

@ -17,12 +17,12 @@ def build_trace_payload(
end_time: datetime,
input_data: Any,
output_data: Any,
metadata: dict[str, Any],
metadata: dict[str, object],
tags: list[str],
thread_id: str | None,
) -> types.TracePayload:
"""Build a complete trace payload."""
trace_name: Final = response_obj.get("object", "unknown type")
trace_name: Final[str] = response_obj.get("object", "unknown type")
return types.TracePayload(
project_name=project_name,
@ -47,7 +47,7 @@ def build_span_payload(
end_time: datetime,
input_data: Any,
output_data: Any,
metadata: dict[str, Any],
metadata: dict[str, object],
tags: list[str],
usage: dict[str, int],
provider: str | None = None,
@ -56,9 +56,9 @@ def build_span_payload(
"""Build a complete span payload."""
span_id: Final = utils.create_uuid7()
model: Final = response_obj.get("model", "unknown-model")
obj_type: Final = response_obj.get("object", "unknown-object")
created: Final = response_obj.get("created", 0)
model: Final[str] = response_obj.get("model", "unknown-model")
obj_type: Final[str] = response_obj.get("object", "unknown-object")
created: Final[int] = response_obj.get("created", 0)
span_name: Final = f"{model}_{obj_type}_{created}"
_logging.verbose_logger.debug("OpikLogger creating span with id %s for trace %s", span_id, trace_id)

View file

@ -10,7 +10,7 @@ import asyncio
import math
import uuid
from collections.abc import AsyncIterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
from typing import TYPE_CHECKING, Any, Final, TypedDict, TypeVar, cast
from typing_extensions import ReadOnly
@ -77,6 +77,10 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b
# ``web_search_tool_result`` blocks to inject into the final response.
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks"
_RESPONSE_CONTENT_FIELD: Final = "content"
_ResponseT: Final = TypeVar("_ResponseT")
class _PlanMetadataView(TypedDict):
websearch_native_blocks: Sequence[Mapping[str, object]] | None
@ -948,17 +952,17 @@ class WebSearchInterceptionLogger(CustomLogger):
)
@staticmethod
def _inject_native_blocks(response: Any, native_blocks: Sequence[Mapping[str, object]]) -> Any:
def _inject_native_blocks(response: _ResponseT, native_blocks: Sequence[Mapping[str, object]]) -> _ResponseT:
"""Prepend native blocks to response content, dict or object form."""
if not native_blocks:
return response
if isinstance(response, dict):
existing = response.get("content") or []
response["content"] = list(native_blocks) + list(existing)
existing = response.get(_RESPONSE_CONTENT_FIELD) or []
response[_RESPONSE_CONTENT_FIELD] = list(native_blocks) + list(existing)
return response
existing = getattr(response, "content", None) or []
existing = getattr(response, _RESPONSE_CONTENT_FIELD, None) or []
try:
response.content = list(native_blocks) + list(existing)
setattr(response, _RESPONSE_CONTENT_FIELD, list(native_blocks) + list(existing))
except (AttributeError, TypeError):
# Object refused write — fall through and leave the response
# untouched rather than crash the request.

View file

@ -2,6 +2,8 @@ import datetime
from collections.abc import Mapping
from typing import Any, Final
import httpx
from litellm.constants import LITELLM_DETAILED_TIMING
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.llm_response_utils.get_api_base import get_api_base
@ -59,11 +61,7 @@ class ResponseMetadata:
@property
def supports_response_time(self) -> bool:
"""Check if response type supports timing metrics"""
return (
isinstance(self.result, ModelResponse)
or isinstance(self.result, EmbeddingResponse)
or isinstance(self.result, TranscriptionResponse)
)
return isinstance(self.result, (ModelResponse, EmbeddingResponse, TranscriptionResponse))
def set_hidden_params(self, logging_obj: LiteLLMLoggingObject, model: str | None, kwargs: dict) -> None:
"""Set hidden parameters on the response"""
@ -79,7 +77,7 @@ class ResponseMetadata:
result=self.result, litellm_model_name=model, router_model_id=model_id
),
"additional_headers": process_response_headers(
self._get_value_from_hidden_params("additional_headers") or {},
self._get_additional_headers_from_hidden_params() or {},
preserve_litellm_internal_headers=True,
),
"litellm_model_name": model,
@ -98,12 +96,12 @@ class ResponseMetadata:
for key, value in new_params.items():
setattr(self._hidden_params, key, value)
def _get_value_from_hidden_params(self, key: str) -> Any | None:
"""Get value from hidden params - handles when self._hidden_params is a dict or HiddenParams object"""
def _get_additional_headers_from_hidden_params(self) -> httpx.Headers | dict[str, str] | None:
"""Get `additional_headers` from hidden params - handles when self._hidden_params is a dict or HiddenParams object"""
if isinstance(self._hidden_params, dict):
return self._hidden_params.get(key, None)
return self._hidden_params.get("additional_headers", None)
elif isinstance(self._hidden_params, HiddenParams):
return getattr(self._hidden_params, key, None)
return getattr(self._hidden_params, "additional_headers", None)
def set_timing_metrics(
self,
@ -129,7 +127,7 @@ class ResponseMetadata:
#########################################################
# 2. Add callback processing duration
#########################################################
callback_duration_ms: Final = getattr(logging_obj, "callback_duration_ms", None)
callback_duration_ms: Final[float | None] = getattr(logging_obj, "callback_duration_ms", None)
if callback_duration_ms is not None:
self._update_hidden_params(
{
@ -142,17 +140,17 @@ class ResponseMetadata:
#########################################################
llm_api_duration_ms: Final = logging_obj.model_call_details.get("llm_api_duration_ms")
if LITELLM_DETAILED_TIMING and llm_api_duration_ms is not None:
detailed: Final[dict] = {
detailed: Final[dict[str, float]] = {
"timing_llm_api_ms": round(llm_api_duration_ms, 4),
}
# message copy time from Logging.__init__()
msg_copy_ms: Final = getattr(logging_obj, "message_copy_duration_ms", None)
msg_copy_ms: Final[float | None] = getattr(logging_obj, "message_copy_duration_ms", None)
if msg_copy_ms is not None:
detailed["timing_message_copy_ms"] = round(msg_copy_ms, 4)
# pre-processing = time from request start to LLM API call start
api_call_start: Final = logging_obj.model_call_details.get("api_call_start_time")
api_call_start: Final[datetime.datetime | None] = logging_obj.model_call_details.get("api_call_start_time")
if api_call_start is not None and start_time is not None:
pre_ms: Final = (api_call_start - start_time).total_seconds() * 1000
detailed["timing_pre_processing_ms"] = round(pre_ms, 4)

View file

@ -93,7 +93,7 @@ def print_verbose(print_statement: object):
@dataclass(frozen=True, slots=True)
class _ProviderChunkParsed:
response_obj: dict[str, Any]
response_obj: dict[str, object]
@dataclass(frozen=True, slots=True)
@ -1288,7 +1288,7 @@ class CustomStreamWrapper:
for key, value in anthropic_response_obj["provider_specific_fields"].items():
setattr(model_response, key, value)
response_obj = cast(dict[str, Any], anthropic_response_obj)
response_obj = cast(dict[str, object], anthropic_response_obj)
elif self.model == "replicate" or self.custom_llm_provider == "replicate":
response_obj = self.handle_replicate_chunk(chunk)
completion_obj["content"] = response_obj["text"]
@ -1444,7 +1444,7 @@ class CustomStreamWrapper:
if not isinstance(chunk, str):
raise ValueError(f"chunk is not a string: {chunk}")
response_obj = cast(
dict[str, Any],
dict[str, object],
litellm.CodestralTextCompletionConfig()._chunk_parser(chunk),
)
completion_obj["content"] = response_obj["text"]
@ -2551,7 +2551,7 @@ def calculate_total_usage(chunks: list[ModelResponse]) -> Usage:
prompt_tokens: int = 0
completion_tokens: int = 0
latest_usage_chunk = None
latest_usage_chunk: Usage | Mapping[str, int] | None = None
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
completion_tokens_details: CompletionTokensDetailsWrapper | None = None
cache_creation_token_details: CacheCreationTokenDetails | None = None

View file

@ -417,7 +417,7 @@ def _extract_redirect_url(response: httpx.Response, request_url: str) -> str:
return str(httpx.URL(request_url).join(location))
def safe_get(client: Any, url: str, **kwargs: Any) -> Any:
def safe_get(client: Any, url: str, **kwargs: Any) -> httpx.Response:
"""
Fetch a user-supplied URL with SSRF protection on every redirect hop.
@ -460,7 +460,7 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any:
raise SSRFError("Too many redirects")
async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any:
async def async_safe_get(client: Any, url: str, **kwargs: Any) -> httpx.Response:
"""Async version of safe_get."""
if not getattr(litellm, "user_url_validation", True):
kwargs.setdefault("follow_redirects", True)

View file

@ -1,9 +1,11 @@
import json
import time
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, cast
import httpx
from httpx import Headers, Response
from typing_extensions import ReadOnly, TypedDict
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
@ -21,6 +23,29 @@ else:
LoggingClass = Any
class AnthropicBatchRequestCounts(TypedDict, total=False):
"""The ``request_counts`` object of an Anthropic Message Batch."""
processing: ReadOnly[int]
succeeded: ReadOnly[int]
errored: ReadOnly[int]
canceled: ReadOnly[int]
expired: ReadOnly[int]
class AnthropicMessageBatch(TypedDict, total=False):
"""The fields of an Anthropic Message Batch that map onto an OpenAI Batch."""
id: ReadOnly[str]
processing_status: ReadOnly[str]
created_at: ReadOnly[str | None]
ended_at: ReadOnly[str | None]
expires_at: ReadOnly[str | None]
cancel_initiated_at: ReadOnly[str | None]
archived_at: ReadOnly[str | None]
request_counts: ReadOnly[AnthropicBatchRequestCounts]
class AnthropicBatchesConfig(BaseBatchesConfig):
def __init__(self):
from ..chat.transformation import AnthropicConfig
@ -85,7 +110,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
create_batch_data: CreateBatchRequest,
optional_params: dict,
litellm_params: dict,
) -> bytes | str | dict[str, Any]:
) -> bytes | str | dict[str, object]:
"""
Transform the batch creation request to Anthropic format.
@ -135,7 +160,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
batch_id: str,
optional_params: dict,
litellm_params: dict,
) -> bytes | str | dict[str, Any]:
) -> bytes | str | dict[str, object]:
"""
Transform batch retrieval request for Anthropic.
@ -154,7 +179,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
) -> LiteLLMBatch:
"""Transform Anthropic MessageBatch retrieval response to LiteLLM format."""
try:
response_data: Final = raw_response.json()
response_data: Final[AnthropicMessageBatch] = raw_response.json()
except Exception as e:
raise ValueError(f"Failed to parse Anthropic batch response: {e}")
@ -163,18 +188,20 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
processing_status: Final = response_data.get("processing_status", "in_progress")
# Map Anthropic processing_status to OpenAI status
status_mapping: dict[
str,
Literal[
"validating",
"failed",
"in_progress",
"finalizing",
"completed",
"expired",
"cancelling",
"cancelled",
],
status_mapping: Final[
Mapping[
str,
Literal[
"validating",
"failed",
"in_progress",
"finalizing",
"completed",
"expired",
"cancelling",
"cancelled",
],
]
] = {
"in_progress": "in_progress",
"canceling": "cancelling",
@ -281,7 +308,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
if not line:
continue
try:
response_json = json.loads(line)
response_json: Mapping[str, Mapping[str, dict[str, object]]] = json.loads(line)
# Update model_response with the parsed JSON
completion_response = response_json["result"]["message"]
transformed_response = self.anthropic_chat_config.transform_parsed_response(

View file

@ -160,7 +160,7 @@ class LiteLLMMessagesToResponsesAPIHandler:
top_k: int | None = None,
top_p: float | None = None,
output_format: AnthropicOutputSchema | None = None,
**kwargs,
**kwargs: object,
) -> AnthropicMessagesResponse | AsyncIterator[bytes]:
responses_kwargs: Final = _build_responses_kwargs(
max_tokens=max_tokens,
@ -214,7 +214,7 @@ class LiteLLMMessagesToResponsesAPIHandler:
top_p: float | None = None,
output_format: AnthropicOutputSchema | None = None,
_is_async: bool = False,
**kwargs,
**kwargs: object,
) -> (
AnthropicMessagesResponse
| AsyncIterator[bytes]

View file

@ -94,7 +94,7 @@ class BedrockRealtime(BaseAWSLLM):
aws_sts_endpoint: str | None = None,
aws_bedrock_runtime_endpoint: str | None = None,
aws_external_id: str | None = None,
**kwargs,
**kwargs: object,
):
"""
Establish bidirectional streaming connection with Bedrock Nova Sonic.
@ -166,13 +166,16 @@ class BedrockRealtime(BaseAWSLLM):
)
bedrock_client: Final = BedrockRuntimeClient(config=config)
async def open_bidirectional_stream() -> BedrockBidirectionalStream:
return await bedrock_client.invoke_model_with_bidirectional_stream(
InvokeModelWithBidirectionalStreamOperationInput(model_id=model)
)
transformation_config: Final = BedrockRealtimeConfig()
try:
# Initialize the bidirectional stream
bedrock_stream: Final = await bedrock_client.invoke_model_with_bidirectional_stream(
InvokeModelWithBidirectionalStreamOperationInput(model_id=model)
)
bedrock_stream: Final = await open_bidirectional_stream()
verbose_proxy_logger.debug("Bedrock Realtime: Bidirectional stream established")
@ -243,10 +246,11 @@ class BedrockRealtime(BaseAWSLLM):
InvokeModelWithBidirectionalStreamInputChunk,
)
def build_input_chunk(payload: bytes) -> object:
return InvokeModelWithBidirectionalStreamInputChunk(value=BidirectionalInputPayloadPart(bytes_=payload))
async def send_to_bedrock(bedrock_message: str) -> None:
event: Final = InvokeModelWithBidirectionalStreamInputChunk(
value=BidirectionalInputPayloadPart(bytes_=bedrock_message.encode("utf-8"))
)
event: Final = build_input_chunk(bedrock_message.encode("utf-8"))
await bedrock_stream.input_stream.send(event)
verbose_proxy_logger.debug("Bedrock Realtime: Sent to Bedrock: %s", bedrock_message[:200])

View file

@ -49,10 +49,10 @@ class ChatGPTToolCallNormalizer:
def __getattr__(self, name: str) -> object:
return getattr(self._stream, name)
def __iter__(self):
def __iter__(self) -> "ChatGPTToolCallNormalizer":
return self
def __aiter__(self):
def __aiter__(self) -> "ChatGPTToolCallNormalizer":
return self
def __next__(self) -> ModelResponseStream:

View file

@ -2,13 +2,16 @@
CompactifAI chat completion transformation
"""
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
import httpx
from typing_extensions import ReadOnly, TypedDict
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.openai.common_utils import OpenAIError
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
@ -23,6 +26,18 @@ else:
LiteLLMLoggingObj = Any
class CompactifAIResponseFields(TypedDict, total=False):
"""The chat completion fields of a CompactifAI response body."""
id: ReadOnly[str]
choices: ReadOnly[Sequence[Mapping[str, object]]]
created: ReadOnly[int]
model: ReadOnly[str]
system_fingerprint: ReadOnly[str | None]
usage: ReadOnly[Mapping[str, object]]
object: ReadOnly[str]
class CompactifAIChatConfig(OpenAIGPTConfig):
"""
Configuration class for CompactifAI chat completions.
@ -47,10 +62,10 @@ class CompactifAIChatConfig(OpenAIGPTConfig):
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: list,
optional_params: dict,
litellm_params: dict,
request_data: Mapping[str, object],
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
encoding: "tiktoken.Encoding | None",
api_key: str | None = None,
json_mode: bool | None = None,
@ -81,14 +96,18 @@ class CompactifAIChatConfig(OpenAIGPTConfig):
message["content"] = tool_calls[0]["function"].get("arguments", "")
message["tool_calls"] = None
returned_response: Final = ModelResponse(**response_json)
response_fields: Final[CompactifAIResponseFields] = response_json
returned_response: Final = ModelResponse(**response_fields)
# Set model name with provider prefix
returned_response.model = f"compactifai/{model}"
return returned_response
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException:
def get_error_class(
self, error_message: str, status_code: int, headers: dict[str, str] | httpx.Headers
) -> BaseLLMException:
"""
Get the appropriate error class for CompactifAI errors.
Since CompactifAI is OpenAI-compatible, we use OpenAI error handling.

View file

@ -6,11 +6,12 @@ endpoint defined in endpoints.json, eliminating the need for individual handler
"""
import json
from collections.abc import Coroutine
from collections.abc import Coroutine, Mapping, Sequence
from pathlib import Path
from typing import TYPE_CHECKING, Any, Final
import httpx
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
@ -32,26 +33,58 @@ if TYPE_CHECKING:
from litellm.llms.base_llm.containers.transformation import BaseContainerConfig
class EndpointConfig(TypedDict):
"""One endpoint entry of ``litellm/containers/endpoints.json``."""
name: ReadOnly[str]
async_name: ReadOnly[str]
path: ReadOnly[str]
method: ReadOnly[str]
path_params: ReadOnly[Sequence[str]]
query_params: ReadOnly[Sequence[str]]
response_type: ReadOnly[str]
is_multipart: NotRequired[ReadOnly[bool]]
returns_binary: NotRequired[ReadOnly[bool]]
class EndpointsConfig(TypedDict):
"""The parsed ``litellm/containers/endpoints.json`` document."""
endpoints: ReadOnly[Sequence[EndpointConfig]]
class ContainerErrorDetail(TypedDict, total=False):
"""The ``error`` object of a container API error body."""
message: ReadOnly[str]
class ContainerResponseBody(TypedDict, total=False):
"""The fields this handler reads off a container API JSON body."""
error: ReadOnly[ContainerErrorDetail]
_ContainerResponseModel = ContainerFileListResponse | ContainerFileObject | DeleteContainerFileResponse
# Response type mapping
RESPONSE_TYPES: Final[dict[str, type]] = {
RESPONSE_TYPES: Final[Mapping[str, type[_ContainerResponseModel]]] = {
"ContainerFileListResponse": ContainerFileListResponse,
"ContainerFileObject": ContainerFileObject,
"DeleteContainerFileResponse": DeleteContainerFileResponse,
}
ContainerEndpointResponse = (
ContainerFileListResponse | ContainerFileObject | DeleteContainerFileResponse | bytes | dict[str, object]
)
ContainerEndpointResponse = _ContainerResponseModel | bytes | ContainerResponseBody
def _load_endpoints_config() -> dict:
def _load_endpoints_config() -> EndpointsConfig:
"""Load the endpoints configuration from JSON file."""
config_path: Final = Path(__file__).parent.parent.parent / "containers" / "endpoints.json"
with open(config_path) as f:
return json.load(f)
def _get_endpoint_config(endpoint_name: str) -> dict | None:
def _get_endpoint_config(endpoint_name: str) -> EndpointConfig | None:
"""Get config for a specific endpoint by name."""
config: Final = _load_endpoints_config()
for endpoint in config["endpoints"]:
@ -60,10 +93,15 @@ def _get_endpoint_config(endpoint_name: str) -> dict | None:
return None
def _response_model(response_type_name: str) -> type[_ContainerResponseModel] | None:
"""The pydantic model a container endpoint's ``response_type`` names."""
return RESPONSE_TYPES.get(response_type_name)
def _build_url(
api_base: str,
path_template: str,
path_params: dict[str, str],
path_params: Mapping[str, object],
) -> str:
"""Build the full URL by substituting path parameters.
@ -93,16 +131,12 @@ def _build_url(
def _build_query_params(
query_param_names: list,
kwargs: dict[str, Any],
) -> dict[str, str]:
query_param_names: Sequence[str],
kwargs: Mapping[str, object],
) -> dict[str, object]:
"""Build query parameters from kwargs."""
params: Final = {}
for param_name in query_param_names:
value = kwargs.get(param_name)
if value is not None:
params[param_name] = str(value) if not isinstance(value, str) else value
return params
supplied: Final = ((param_name, kwargs.get(param_name)) for param_name in query_param_names)
return {name: value if isinstance(value, str) else str(value) for name, value in supplied if value is not None}
def _error_message_from_response(response: httpx.Response) -> str:
@ -136,24 +170,24 @@ def _transform_response(
if returns_binary:
return response.content
response_json: Final = response.json()
response_json: Final[ContainerResponseBody] = response.json()
if "error" in response_json:
raise BaseLLMException(
status_code=response.status_code,
message=response_json.get("error", {}).get("message", str(response_json)),
message=response_json["error"].get("message", str(response_json)),
headers=dict(response.headers),
)
response_type: Final = RESPONSE_TYPES.get(response_type_name)
response_type: Final = _response_model(response_type_name)
if response_type:
return response_type(**response_json)
return response_type.model_validate(response_json)
return response_json
def _prepare_multipart_file_upload(
file: Any,
headers: dict[str, Any],
) -> tuple:
headers: dict[str, object],
) -> tuple[dict[str, tuple[str, bytes, str]], dict[str, object]]:
"""
Prepare file and headers for multipart upload.
@ -178,6 +212,52 @@ def _prepare_multipart_file_upload(
return files, headers_copy
def _request_headers(
container_provider_config: "BaseContainerConfig",
extra_headers: dict[str, object] | None,
litellm_params: GenericLiteLLMParams,
) -> dict[str, object]:
"""The provider auth headers for a container request."""
return container_provider_config.validate_environment(
headers=extra_headers or {},
api_key=litellm_params.get("api_key", None),
)
def _request_api_base(
container_provider_config: "BaseContainerConfig",
litellm_params: GenericLiteLLMParams,
) -> str:
"""The provider base URL for a container request."""
return container_provider_config.get_complete_url(
api_base=litellm_params.get("api_base", None),
litellm_params=dict(litellm_params),
)
def _sync_http_client(
client: HTTPHandler | AsyncHTTPHandler | None,
litellm_params: GenericLiteLLMParams,
) -> HTTPHandler:
"""The sync HTTP client for a container request, reusing the caller's when usable."""
if client is None or not isinstance(client, HTTPHandler):
return _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)})
return client
def _async_http_client(
client: HTTPHandler | AsyncHTTPHandler | None,
litellm_params: GenericLiteLLMParams,
) -> AsyncHTTPHandler:
"""The async HTTP client for a container request, reusing the caller's when usable."""
if client is None or not isinstance(client, AsyncHTTPHandler):
return get_async_httpx_client(
llm_provider=litellm.LlmProviders.OPENAI,
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
)
return client
class GenericContainerHandler:
"""
Generic handler for container file API endpoints.
@ -192,13 +272,13 @@ class GenericContainerHandler:
container_provider_config: "BaseContainerConfig",
litellm_params: GenericLiteLLMParams,
logging_obj: "LiteLLMLoggingObj",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout = 600,
_is_async: bool = False,
client: HTTPHandler | AsyncHTTPHandler | None = None,
**kwargs,
) -> Any | Coroutine[Any, Any, Any]:
**kwargs: object,
) -> Any | Coroutine[object, object, Any]:
"""
Generic handler for any container file endpoint.
@ -245,11 +325,11 @@ class GenericContainerHandler:
container_provider_config: "BaseContainerConfig",
litellm_params: GenericLiteLLMParams,
logging_obj: "LiteLLMLoggingObj",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout = 600,
client: HTTPHandler | AsyncHTTPHandler | None = None,
**kwargs,
**kwargs: object,
) -> Any:
"""Synchronous request handler."""
endpoint_config: Final = _get_endpoint_config(endpoint_name)
@ -257,23 +337,14 @@ class GenericContainerHandler:
raise ValueError(f"Unknown endpoint: {endpoint_name}")
# Get HTTP client
if client is None or not isinstance(client, HTTPHandler):
http_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)})
else:
http_client = client
http_client: Final = _sync_http_client(client, litellm_params)
# Build request
headers = container_provider_config.validate_environment(
headers=extra_headers or {},
api_key=litellm_params.get("api_key", None),
)
headers = _request_headers(container_provider_config, extra_headers, litellm_params)
if extra_headers:
headers.update(extra_headers)
api_base: Final = container_provider_config.get_complete_url(
api_base=litellm_params.get("api_base", None),
litellm_params=dict(litellm_params),
)
api_base: Final = _request_api_base(container_provider_config, litellm_params)
# Build URL with path params
path_params: Final = {p: kwargs.get(p, "") for p in endpoint_config.get("path_params", [])}
@ -334,11 +405,11 @@ class GenericContainerHandler:
container_provider_config: "BaseContainerConfig",
litellm_params: GenericLiteLLMParams,
logging_obj: "LiteLLMLoggingObj",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout = 600,
client: HTTPHandler | AsyncHTTPHandler | None = None,
**kwargs,
**kwargs: object,
) -> Any:
"""Asynchronous request handler."""
endpoint_config: Final = _get_endpoint_config(endpoint_name)
@ -346,26 +417,14 @@ class GenericContainerHandler:
raise ValueError(f"Unknown endpoint: {endpoint_name}")
# Get HTTP client
if client is None or not isinstance(client, AsyncHTTPHandler):
http_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.OPENAI,
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
)
else:
http_client = client
http_client: Final = _async_http_client(client, litellm_params)
# Build request
headers = container_provider_config.validate_environment(
headers=extra_headers or {},
api_key=litellm_params.get("api_key", None),
)
headers = _request_headers(container_provider_config, extra_headers, litellm_params)
if extra_headers:
headers.update(extra_headers)
api_base: Final = container_provider_config.get_complete_url(
api_base=litellm_params.get("api_base", None),
litellm_params=dict(litellm_params),
)
api_base: Final = _request_api_base(container_provider_config, litellm_params)
# Build URL with path params
path_params: Final = {p: kwargs.get(p, "") for p in endpoint_config.get("path_params", [])}

View file

@ -9,7 +9,7 @@ import threading
import time
from collections.abc import AsyncIterable, Callable, Iterable, Mapping
from http.cookiejar import CookieJar, DefaultCookiePolicy
from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional, TypeAlias, TypedDict
from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, Optional, TypeAlias, TypedDict
import certifi
import httpx
@ -447,7 +447,7 @@ def _safe_read_response(response: httpx.Response, timeout: float | None = None)
return b""
def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> None:
def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> NoReturn:
"""Raise a MaskedHTTPStatusError for sync HTTP handlers."""
if stream:
try:
@ -467,7 +467,7 @@ def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> None:
raise MaskedHTTPStatusError(e, message=_text, text=_text) from None
async def _raise_masked_async_error(e: httpx.HTTPStatusError, stream: bool) -> None:
async def _raise_masked_async_error(e: httpx.HTTPStatusError, stream: bool) -> NoReturn:
"""Raise a MaskedHTTPStatusError for async HTTP handlers."""
if stream:
try:

View file

@ -8,7 +8,7 @@ Talks to e2b's REST API directly over httpx (no e2b SDK dependency):
"""
import json
from typing import Final, cast
from typing import Final
import httpx
@ -68,13 +68,10 @@ class E2BSandboxConfig(BaseSandboxConfig):
if metadata:
body["metadata"] = metadata
response: Final = cast(
httpx.Response,
await self._http(client).post(
url=f"{base}/sandboxes",
headers={"X-API-Key": key, "Content-Type": "application/json"},
json=body,
),
response: Final = await self._http(client).post(
url=f"{base}/sandboxes",
headers={"X-API-Key": key, "Content-Type": "application/json"},
json=body,
)
data: Final = response.json()
@ -117,14 +114,11 @@ class E2BSandboxConfig(BaseSandboxConfig):
headers["E2B-Traffic-Access-Token"] = traffic_token
url: Final = f"https://{JUPYTER_PORT}-{handle.id}.{handle.domain}/execute"
response: Final = cast(
httpx.Response,
await self._http(client).post(
url=url,
headers=headers,
json={"code": code, "context_id": None, "env_vars": env_vars},
stream=True,
),
response: Final = await self._http(client).post(
url=url,
headers=headers,
json={"code": code, "context_id": None, "env_vars": env_vars},
stream=True,
)
lines: Final = await self._read_capped_lines(response)
return self._parse_lines(lines)
@ -142,12 +136,9 @@ class E2BSandboxConfig(BaseSandboxConfig):
key: Final = api_key or handle._hidden_params.get("api_key") or self.validate_environment()
base: Final = api_base or handle._hidden_params.get("api_base") or E2B_API_BASE
try:
response: Final = cast(
httpx.Response,
await self._http(client).delete(
url=f"{base}/sandboxes/{handle.id}",
headers={"X-API-Key": key},
),
response: Final = await self._http(client).delete(
url=f"{base}/sandboxes/{handle.id}",
headers={"X-API-Key": key},
)
except httpx.HTTPStatusError as e:
if e.response.status_code == 404:

View file

@ -6,11 +6,25 @@ import json
import os
import re
import threading
from typing import Any, Final
from collections.abc import Callable
from typing import Any, Final, Protocol
from urllib.parse import urlsplit
import litellm
from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig
from litellm.types.llms.openai import AllMessageValues
class _GDCHAudienceCredentials(Protocol):
"""A GDCH service account credential already bound to an audience, ready to mint a bearer token."""
@property
def valid(self) -> bool: ...
@property
def token(self) -> str: ...
def refresh(self, request: object) -> None: ...
class GDCGeminiConfig(OpenAILikeChatConfig):
@ -21,7 +35,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
def __init__(self, **kwargs: Any) -> None:
super().__init__(**kwargs)
self._creds_lock = threading.Lock()
self._gdch_creds_cache: dict = {}
self._gdch_creds_cache: dict[tuple[str, str], _GDCHAudienceCredentials] = {}
def get_supported_openai_params(self, model: str) -> list:
return [
@ -110,7 +124,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
return f"{api_base}/v1/projects/{project}/locations/{location}/chat/completions"
def _read_env_bool(self, val: Any, env_var: str, default: bool = True) -> bool | str:
def _read_env_bool(self, val: bool | str | None, env_var: str, default: bool = True) -> bool | str:
def _parse(s: str) -> bool | str:
cleaned: Final = s.strip().lower()
if cleaned in ("false", "0", "no", "off"):
@ -129,7 +143,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
return default
return _parse(_env_val)
def _fetch_auth(self, gdch_creds: Any, ssl_verify: bool | str) -> None:
def _fetch_auth(self, gdch_creds: _GDCHAudienceCredentials, ssl_verify: bool | str) -> None:
import requests
from google.auth.transport import requests as auth_requests
@ -138,13 +152,24 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
auth_request: Final = auth_requests.Request(session=auth_session)
gdch_creds.refresh(auth_request)
def _cached_fetch_token(self, creds: Any, audience: str, ssl_verify: bool | str, api_key: str | None = None) -> str:
def _with_gdch_audience(self, creds: object, audience: str) -> _GDCHAudienceCredentials:
"""The credential rebound to ``audience``, which GDCH requires before a token refresh."""
bind_audience: Final[Callable[[str], _GDCHAudienceCredentials] | None] = getattr(
creds, "with_gdch_audience", None
)
if bind_audience is None:
raise AttributeError("GDC credentials must expose with_gdch_audience to be bound to a request audience")
return bind_audience(audience)
def _cached_fetch_token(
self, creds: object, audience: str, ssl_verify: bool | str, api_key: str | None = None
) -> str:
# Key cache by both audience and credential identity to prevent cross-caller contamination
cache_key: Final = (audience.rstrip("/"), api_key or str(id(creds)))
with self._creds_lock:
if cache_key not in self._gdch_creds_cache:
self._gdch_creds_cache[cache_key] = creds.with_gdch_audience(audience.rstrip("/"))
self._gdch_creds_cache[cache_key] = self._with_gdch_audience(creds, audience.rstrip("/"))
gdch_creds: Final = self._gdch_creds_cache[cache_key]
@ -155,7 +180,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
return token
def _load_creds_from_key(self, api_key: str) -> tuple[Any, bool]:
def _load_creds_from_key(self, api_key: str) -> tuple[object | None, bool]:
import google.auth
try:
@ -175,7 +200,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
self,
headers: dict,
model: str,
messages: list[Any],
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: str | None = None,
@ -230,7 +255,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
if self._read_env_bool(litellm_params.get("gdc_token_caching"), "GDC_TOKEN_CACHING", default=False):
token = self._cached_fetch_token(creds, audience, ssl_verify, api_key)
else:
gdch_creds: Final = creds.with_gdch_audience(audience)
gdch_creds: Final = self._with_gdch_audience(creds, audience)
self._fetch_auth(gdch_creds, ssl_verify)
token = gdch_creds.token
headers["Authorization"] = f"Bearer {token}"
@ -252,7 +277,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
def transform_request(
self,
model: str,
messages: list[Any],
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,

View file

@ -4,10 +4,11 @@ Transformation logic from Cohere's /v1/rerank format to Infinity's `/v1/rerank`
Why separate file? Make it easy to see how transformation works
"""
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from typing import Final
import httpx
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
from litellm._uuid import uuid
@ -26,6 +27,31 @@ from litellm.types.rerank import (
from ..common_utils import InfinityError
class _InfinityRerankUsage(TypedDict, extra_items=ReadOnly[int]):
"""The token counters Infinity reports in the ``usage`` block of a rerank response."""
class _InfinityRerankResult(TypedDict):
"""One scored document in an Infinity ``/v1/rerank`` response."""
index: ReadOnly[int]
relevance_score: ReadOnly[float]
document: ReadOnly[str]
class _InfinityRerankResponse(TypedDict):
"""The JSON body returned by Infinity's ``/v1/rerank`` endpoint."""
id: ReadOnly[NotRequired[str]]
usage: ReadOnly[NotRequired[_InfinityRerankUsage]]
results: ReadOnly[Sequence[_InfinityRerankResult]]
def _parse_rerank_response(raw_response: httpx.Response) -> _InfinityRerankResponse:
"""Read the untyped JSON body of an Infinity rerank response."""
return raw_response.json()
class InfinityRerankConfig(CohereRerankConfig):
def get_complete_url(
self,
@ -82,7 +108,7 @@ class InfinityRerankConfig(CohereRerankConfig):
No transformation required, Infinity follows Cohere API response format
"""
try:
raw_response_json: Final = raw_response.json()
raw_response_json: Final = _parse_rerank_response(raw_response)
except Exception:
raise InfinityError(message=raw_response.text, status_code=raw_response.status_code)

View file

@ -13,12 +13,57 @@ Generated files are returned directly in the response - no separate storage need
import base64
import json
from collections.abc import Sequence
from enum import Enum
from typing import Any, Final
from typing import Any, Final, Protocol
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_logger
class _ToolCallFunction(Protocol):
"""Function payload of an assistant tool call."""
name: str | None
arguments: str
class _ToolCall(Protocol):
"""Tool call requested by the assistant on a chat completion choice."""
id: str
function: _ToolCallFunction
class _AssistantMessage(Protocol):
"""Assistant message carried by a chat completion choice."""
content: str | None
tool_calls: Sequence[_ToolCall] | None
class _CompletionChoice(Protocol):
"""Single choice of a chat completion response."""
finish_reason: str
message: _AssistantMessage
class _SandboxFile(TypedDict):
"""File generated inside the sandbox during a code execution run."""
name: ReadOnly[str]
mime_type: ReadOnly[str]
content_base64: ReadOnly[str]
class _CodeExecutionArguments(TypedDict):
"""Arguments the model passes to the `litellm_code_execution` tool."""
code: NotRequired[ReadOnly[str]]
class LiteLLMInternalTools(str, Enum):
"""
Enum for internal LiteLLM tools that are injected into requests.
@ -30,7 +75,7 @@ class LiteLLMInternalTools(str, Enum):
CODE_EXECUTION = "litellm_code_execution"
def get_litellm_code_execution_tool() -> dict[str, Any]:
def get_litellm_code_execution_tool() -> dict[str, object]:
"""
Returns the litellm_code_execution tool definition in OpenAI format.
@ -51,7 +96,7 @@ def get_litellm_code_execution_tool() -> dict[str, Any]:
}
def get_litellm_code_execution_tool_anthropic() -> dict[str, Any]:
def get_litellm_code_execution_tool_anthropic() -> dict[str, object]:
"""
Returns the litellm_code_execution tool definition in Anthropic/messages API format.
@ -103,7 +148,7 @@ class CodeExecutionHandler:
skill_files: dict[str, bytes],
skill_id: str | None = None,
**kwargs,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Execute an LLM call with automatic code execution handling.
@ -134,8 +179,8 @@ class CodeExecutionHandler:
)
current_messages: Final = list(messages)
generated_files: Final[list[dict[str, Any]]] = [] # Files returned directly
execution_results: Final[list[dict]] = []
generated_files: Final[list[dict[str, object]]] = [] # Files returned directly
execution_results: Final[list[dict[str, object]]] = []
executor: Final = SkillsSandboxExecutor(timeout=self.sandbox_timeout)
response: Any = None # Initialize to avoid possibly unbound error
@ -151,11 +196,12 @@ class CodeExecutionHandler:
**kwargs,
)
assistant_message = response.choices[0].message
stop_reason = response.choices[0].finish_reason
choice: _CompletionChoice = response.choices[0]
assistant_message = choice.message
stop_reason: str = choice.finish_reason
# Build assistant message for conversation history
assistant_msg_dict: dict[str, Any] = {
assistant_msg_dict: dict[str, object] = {
"role": "assistant",
"content": assistant_message.content,
}
@ -190,8 +236,8 @@ class CodeExecutionHandler:
if tool_name == LiteLLMInternalTools.CODE_EXECUTION.value:
# Execute code in sandbox
try:
args = json.loads(tool_call.function.arguments)
code = args.get("code", "")
args: _CodeExecutionArguments = json.loads(tool_call.function.arguments)
code: str = args.get("code", "")
verbose_logger.debug("CodeExecutionHandler: Executing code (%s chars)", len(code))
@ -202,13 +248,15 @@ class CodeExecutionHandler:
verbose_logger.debug("CodeExecutionHandler: Execution result: %s", exec_result)
sandbox_files: Sequence[_SandboxFile] = exec_result["files"]
execution_results.append(
{
"iteration": iteration,
"success": exec_result["success"],
"output": exec_result["output"],
"error": exec_result["error"],
"files": [f["name"] for f in exec_result["files"]],
"files": [f["name"] for f in sandbox_files],
}
)
@ -216,9 +264,9 @@ class CodeExecutionHandler:
tool_result = exec_result["output"] or ""
# Collect generated files (returned directly, no storage)
if exec_result["files"]:
if sandbox_files:
tool_result += "\n\nGenerated files:"
for f in exec_result["files"]:
for f in sandbox_files:
file_content = base64.b64decode(f["content_base64"])
# Add to generated files list (returned in response)
generated_files.append(

View file

@ -4,16 +4,32 @@ Ollama /chat/completion calls handled in llm_http_handler.py
[TODO]: migrate embeddings to a base handler as well.
"""
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Any, Final, Protocol, TypedDict
from typing_extensions import NotRequired, ReadOnly
import litellm
from litellm.types.utils import EmbeddingResponse
class TokenEncoder(Protocol):
"""The tokenizer surface used to estimate prompt tokens."""
def encode(self, text: str, /) -> Sequence[int]: ...
class OllamaEmbeddingResponse(TypedDict):
"""Body of an Ollama ``/api/embed`` response."""
embeddings: ReadOnly[list[list[float]]]
prompt_eval_count: ReadOnly[NotRequired[int]]
def _prepare_ollama_embedding_payload(
model: str, prompts: list[str], optional_params: dict[str, Any]
) -> dict[str, Any]:
data: Final[dict[str, Any]] = {"model": model, "input": prompts}
model: str, prompts: list[str], optional_params: Mapping[str, object]
) -> dict[str, object]:
data: Final[dict[str, object]] = {"model": model, "input": prompts}
special_optional_params: Final = ["truncate", "options", "keep_alive", "dimensions"]
for k, v in optional_params.items():
@ -27,12 +43,12 @@ def _prepare_ollama_embedding_payload(
def _process_ollama_embedding_response(
response_json: dict,
response_json: OllamaEmbeddingResponse,
prompts: list[str],
model: str,
model_response: EmbeddingResponse,
logging_obj: Any,
encoding: Any,
encoding: TokenEncoder | None,
) -> EmbeddingResponse:
output_data: Final = []
embeddings: Final[list[list[float]]] = response_json["embeddings"]
@ -72,7 +88,7 @@ async def ollama_aembeddings(
model_response: EmbeddingResponse,
optional_params: dict,
logging_obj: Any,
encoding: Any,
encoding: TokenEncoder | None,
):
if not api_base.endswith("/api/embed"):
api_base += "/api/embed"
@ -80,7 +96,7 @@ async def ollama_aembeddings(
data: Final = _prepare_ollama_embedding_payload(model, prompts, optional_params)
response: Final = await litellm.module_level_aclient.post(url=api_base, json=data)
response_json: Final = response.json()
response_json: Final[OllamaEmbeddingResponse] = response.json()
return _process_ollama_embedding_response(
response_json=response_json,
@ -99,7 +115,7 @@ def ollama_embeddings(
optional_params: dict,
model_response: EmbeddingResponse,
logging_obj: Any,
encoding: Any = None,
encoding: TokenEncoder | None = None,
):
if not api_base.endswith("/api/embed"):
api_base += "/api/embed"
@ -107,7 +123,7 @@ def ollama_embeddings(
data: Final = _prepare_ollama_embedding_payload(model, prompts, optional_params)
response: Final = litellm.module_level_client.post(url=api_base, json=data)
response_json: Final = response.json()
response_json: Final[OllamaEmbeddingResponse] = response.json()
return _process_ollama_embedding_response(
response_json=response_json,

View file

@ -1,6 +1,8 @@
from typing import TYPE_CHECKING, Any, Final
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
@ -11,9 +13,11 @@ from litellm.secret_managers.main import get_secret_str
from litellm.types.containers.main import (
ContainerCreateOptionalRequestParams,
ContainerFileListResponse,
ContainerFileObject,
ContainerListResponse,
ContainerObject,
DeleteContainerResult,
ExpiresAfter,
)
from litellm.types.router import GenericLiteLLMParams
@ -32,6 +36,46 @@ else:
BaseLLMException = Any
class OpenAIContainerPayload(TypedDict):
"""The JSON body OpenAI returns for a single container."""
id: ReadOnly[str]
object: ReadOnly[Literal["container"]]
created_at: ReadOnly[int]
status: ReadOnly[str]
expires_after: ReadOnly[ExpiresAfter | None]
last_active_at: ReadOnly[int | None]
name: ReadOnly[str | None]
class OpenAIContainerListPayload(TypedDict):
"""The JSON body OpenAI returns for a page of containers."""
object: ReadOnly[Literal["list"]]
data: ReadOnly[list[ContainerObject]]
first_id: ReadOnly[str | None]
last_id: ReadOnly[str | None]
has_more: ReadOnly[bool]
class OpenAIContainerDeletedPayload(TypedDict):
"""The JSON body OpenAI returns for a deleted container."""
id: ReadOnly[str]
object: ReadOnly[Literal["container.deleted"]]
deleted: ReadOnly[bool]
class OpenAIContainerFileListPayload(TypedDict):
"""The JSON body OpenAI returns for a page of container files."""
object: ReadOnly[Literal["list"]]
data: ReadOnly[list[ContainerFileObject]]
first_id: ReadOnly[str | None]
last_id: ReadOnly[str | None]
has_more: ReadOnly[bool]
class OpenAIContainerConfig(BaseContainerConfig):
"""Configuration class for OpenAI container API."""
@ -87,7 +131,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
def transform_container_create_request(
self,
name: str,
container_create_optional_request_params: dict,
container_create_optional_request_params: Mapping[str, object],
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> dict:
@ -111,7 +155,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> ContainerObject:
"""Transform the OpenAI container creation response."""
response_data: Final = raw_response.json()
response_data: Final[OpenAIContainerPayload] = raw_response.json()
# Transform the response data
container_obj: Final = ContainerObject(**response_data)
@ -140,7 +184,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
after: str | None = None,
limit: int | None = None,
order: str | None = None,
extra_query: dict[str, Any] | None = None,
extra_query: Mapping[str, object] | None = None,
) -> tuple[str, dict]:
"""Transform the container list request for OpenAI API.
@ -151,7 +195,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
url: Final = api_base
# Prepare query parameters
params: Final = {}
params: Final[dict[str, object]] = {}
if after is not None:
params["after"] = after
if limit is not None:
@ -171,7 +215,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> ContainerListResponse:
"""Transform the OpenAI container list response."""
response_data: Final = raw_response.json()
response_data: Final[OpenAIContainerListPayload] = raw_response.json()
# Transform the response data
container_list: Final = ContainerListResponse(**response_data)
@ -191,7 +235,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}")
# No additional data needed for GET request
data: Final[dict[str, Any]] = {}
data: Final[dict[str, object]] = {}
return url, data
@ -201,7 +245,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> ContainerObject:
"""Transform the OpenAI container retrieve response."""
response_data: Final = raw_response.json()
response_data: Final[OpenAIContainerPayload] = raw_response.json()
# Transform the response data
container_obj: Final = ContainerObject(**response_data)
@ -224,7 +268,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}")
# No data needed for DELETE request
data: Final[dict[str, Any]] = {}
data: Final[dict[str, object]] = {}
return url, data
@ -234,7 +278,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> DeleteContainerResult:
"""Transform the OpenAI container delete response."""
response_data: Final = raw_response.json()
response_data: Final[OpenAIContainerDeletedPayload] = raw_response.json()
# Transform the response data
delete_result: Final = DeleteContainerResult(**response_data)
@ -250,7 +294,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
after: str | None = None,
limit: int | None = None,
order: str | None = None,
extra_query: dict[str, Any] | None = None,
extra_query: Mapping[str, object] | None = None,
) -> tuple[str, dict]:
"""Transform the container file list request for OpenAI API.
@ -262,7 +306,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}/files")
# Prepare query parameters
params: Final[dict[str, Any]] = {}
params: Final[dict[str, object]] = {}
if after is not None:
params["after"] = after
if limit is not None:
@ -282,7 +326,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
logging_obj: LiteLLMLoggingObj,
) -> ContainerFileListResponse:
"""Transform the OpenAI container file list response."""
response_data: Final = raw_response.json()
response_data: Final[OpenAIContainerFileListPayload] = raw_response.json()
# Transform the response data
file_list: Final = ContainerFileListResponse(**response_data)
@ -308,7 +352,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}/files/{encoded_file_id}/content")
# No query parameters needed
params: Final[dict[str, Any]] = {}
params: Final[dict[str, object]] = {}
return url, params

View file

@ -1,7 +1,7 @@
import asyncio
import json
import time
from typing import Final, cast
from typing import Final
import httpx
@ -86,13 +86,10 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
secure_access=secure_access,
)
response: Final = cast(
httpx.Response,
await self._http(client).post(
url=f"{base}/sandboxes",
headers=self._lifecycle_headers(key),
json=body,
),
response: Final = await self._http(client).post(
url=f"{base}/sandboxes",
headers=self._lifecycle_headers(key),
json=body,
)
data: Final = response.json()
sandbox_id: Final = str(data["id"])
@ -182,12 +179,9 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
base: Final = str(handle._hidden_params.get("api_base") or self._api_base(api_base))
key: Final = self._api_key(api_key=api_key, handle=handle)
try:
response: Final = cast(
httpx.Response,
await self._http(client).delete(
url=f"{base}/sandboxes/{handle.id}",
headers=self._lifecycle_headers(key),
),
response: Final = await self._http(client).delete(
url=f"{base}/sandboxes/{handle.id}",
headers=self._lifecycle_headers(key),
)
except httpx.HTTPStatusError as e:
if e.response.status_code == 404:
@ -245,12 +239,9 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
) -> None:
deadline: Final = time.monotonic() + ready_timeout
while True:
response = cast(
httpx.Response,
await self._http(client).get(
url=f"{api_base}/sandboxes/{sandbox_id}",
headers=headers,
),
response = await self._http(client).get(
url=f"{api_base}/sandboxes/{sandbox_id}",
headers=headers,
)
data = response.json()
state = self._sandbox_state(data)
@ -306,13 +297,10 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
use_server_proxy: bool,
client: AsyncHTTPHandler | None,
) -> tuple[str, dict[str, str]]:
response: Final = cast(
httpx.Response,
await self._http(client).get(
url=f"{api_base}/sandboxes/{sandbox_id}/endpoints/{OPEN_SANDBOX_EXECD_PORT}",
headers=headers,
params={"use_server_proxy": use_server_proxy},
),
response: Final = await self._http(client).get(
url=f"{api_base}/sandboxes/{sandbox_id}/endpoints/{OPEN_SANDBOX_EXECD_PORT}",
headers=headers,
params={"use_server_proxy": use_server_proxy},
)
data: Final = response.json()
endpoint: Final = data.get("endpoint")
@ -329,15 +317,12 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
client: AsyncHTTPHandler | None,
) -> list[str]:
timeout: Final = httpx.Timeout(connect=30.0, read=None, write=30.0, pool=None)
response: Final = cast(
httpx.Response,
await self._http(client).post(
url=url,
headers=headers,
timeout=timeout,
json=body,
stream=True,
),
response: Final = await self._http(client).post(
url=url,
headers=headers,
timeout=timeout,
json=body,
stream=True,
)
return await self._read_capped_lines(response)

View file

@ -97,7 +97,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
def __init__(self):
super().__init__()
def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials:
def get_auth_credentials(self, litellm_params: Mapping[str, object]) -> BaseVectorStoreAuthCredentials:
# Get credentials and project info
vertex_credentials: Final = self.get_vertex_ai_credentials(dict(litellm_params))
vertex_project: Final = self.get_vertex_ai_project(dict(litellm_params))
@ -122,7 +122,9 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
"write": [("POST", "/ragCorpora")],
}
def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict:
def validate_environment(
self, headers: dict[str, str], litellm_params: GenericLiteLLMParams | None
) -> dict[str, str]:
"""
Validate and set up authentication for Vertex AI RAG API
"""
@ -135,7 +137,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
def get_complete_url(
self,
api_base: str | None,
litellm_params: dict,
litellm_params: dict[str, object],
) -> str:
"""
Get the Base endpoint for Vertex AI RAG API

View file

@ -33,6 +33,6 @@ class DomainModel(BaseModel):
return cls(**record.dict())
return cls(**dict(record))
def to_db_dict(self, exclude_unset: bool = False) -> dict[str, Any]:
def to_db_dict(self, exclude_unset: bool = False) -> dict[str, object]:
"""Convert domain model to a dictionary for database operations."""
return self.model_dump(exclude_none=True, exclude_unset=exclude_unset)

View file

@ -1096,8 +1096,7 @@ async def exchange_token_with_server(
headers={"Accept": "application/json", **token_request.headers},
data=token_data,
)
if response is not None:
response.raise_for_status()
response.raise_for_status()
except httpx.HTTPStatusError as exc:
fault: Final = classify_upstream_token_rejection(
exc.response,
@ -1119,11 +1118,6 @@ async def exchange_token_with_server(
)
return _bridge_mint_error_response("invalid_refresh")
return render_token_fault(fault)
if response is None:
raise HTTPException(
status_code=502,
detail="MCP upstream token endpoint returned no response",
)
token_response = response.json()
# Validate token response against server-configured rules before any storage.
@ -1536,16 +1530,10 @@ async def _post_dcr_registration(
headers=headers,
json=register_data,
)
if response is not None:
response.raise_for_status()
response.raise_for_status()
except httpx.HTTPStatusError as exc:
status_code, detail = dcr_fault_detail(classify_upstream_dcr_rejection(exc.response, log_context=server_id))
raise HTTPException(status_code=status_code, detail=detail) from exc
if response is None:
raise HTTPException(
status_code=502,
detail="MCP upstream registration endpoint returned no response",
)
return response

View file

@ -36,7 +36,10 @@ a healed fleet has no null rows and the backfill exits after one query.
import json
from collections import Counter
from typing import Any, Final, Literal
from collections.abc import Mapping, Sequence
from typing import Final, Literal, Protocol
from pydantic import JsonValue
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.db import _decode_oauth_payload, decrypt_credentials
@ -55,9 +58,59 @@ BackfillRule = Literal[
_BACKFILL_AUDIT_ACTOR: Final = "oauth2_flow_backfill"
def _decrypted_credentials(raw_credentials: Any) -> MCPCredentials | None:
class _MCPServerRow(Protocol):
"""The ``LiteLLM_MCPServerTable`` columns this backfill reads."""
@property
def server_id(self) -> str: ...
@property
def authorization_url(self) -> str | None: ...
@property
def registration_url(self) -> str | None: ...
@property
def token_url(self) -> str | None: ...
@property
def credentials(self) -> str | Mapping[str, JsonValue] | None: ...
class _MCPUserCredentialRow(Protocol):
"""The ``LiteLLM_MCPUserCredentials`` columns this backfill reads."""
@property
def server_id(self) -> str: ...
@property
def credential_b64(self) -> str: ...
class _MCPServerTable(Protocol):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_MCPServerRow]: ...
async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, str]) -> object: ...
class _MCPUserCredentialsTable(Protocol):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_MCPUserCredentialRow]: ...
def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable:
"""The MCP server table, typed so the untyped prisma client surface stops here."""
return prisma_client.db.litellm_mcpservertable
def _mcp_user_credentials_table(prisma_client: PrismaClient) -> _MCPUserCredentialsTable:
"""The per-user MCP credential table, typed so the untyped prisma client surface stops here."""
return prisma_client.db.litellm_mcpusercredentials
def _decrypted_credentials(raw_credentials: str | Mapping[str, JsonValue] | None) -> MCPCredentials | None:
if raw_credentials is None:
return None
parsed: JsonValue | Mapping[str, JsonValue]
if isinstance(raw_credentials, str):
try:
parsed = json.loads(raw_credentials)
@ -92,14 +145,14 @@ def classify_null_flow_row(
async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[BackfillRule, int]:
"""Classify every ``auth_type=oauth2`` row whose ``oauth2_flow`` is null; stamp the provable
ones, warn on the ambiguous ones, and return counts per rule."""
null_rows: Final[list[Any]] = await prisma_client.db.litellm_mcpservertable.find_many(
null_rows: Final[Sequence[_MCPServerRow]] = await _mcp_server_table(prisma_client).find_many(
where={"auth_type": "oauth2", "oauth2_flow": None},
)
if not null_rows:
return {}
server_ids: Final = [row.server_id for row in null_rows]
token_rows: Final[list[Any]] = await prisma_client.db.litellm_mcpusercredentials.find_many(
token_rows: Final[Sequence[_MCPUserCredentialRow]] = await _mcp_user_credentials_table(prisma_client).find_many(
where={"server_id": {"in": server_ids}},
)
server_ids_with_oauth_tokens: Final[set[str]] = {
@ -141,7 +194,7 @@ async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[Backfi
stamped_flows: Final = {flow for _, (flow, _) in classified if flow is not None}
for stamped_flow in stamped_flows:
server_ids_for_flow = [row.server_id for row, (row_flow, _) in classified if row_flow == stamped_flow]
await prisma_client.db.litellm_mcpservertable.update_many(
await _mcp_server_table(prisma_client).update_many(
where={"server_id": {"in": server_ids_for_flow}, "oauth2_flow": None},
data={"oauth2_flow": stamped_flow, "updated_by": _BACKFILL_AUDIT_ACTOR},
)

View file

@ -19,9 +19,9 @@ Implements the client-credentials behavior contract for the v2 resolver:
identity.
The token-endpoint POST is injected (``M2MTokenEndpointPost``) so the grant orchestration is
testable without a live IdP; ``post_client_credentials_grant`` is the httpx edge and the one
place the untyped response boundary is contained. Failures are values: the source returns
``Result[OAuthToken, CredError]``; only the httpx edge touches exceptions.
testable without a live IdP; ``post_client_credentials_grant`` is the httpx edge. Failures are
values: the source returns ``Result[OAuthToken, CredError]``; only the httpx edge touches
exceptions.
"""
from __future__ import annotations
@ -95,18 +95,17 @@ async def post_client_credentials_grant(
) -> TokenEndpointOutcome:
"""POST the grant to the token endpoint and classify the transport outcome.
The httpx edge: litellm's handler is partially typed (and raises ``HTTPStatusError`` itself on
a 4xx/5xx), so the untyped boundary is contained here and every field the caller reads comes
out of a validated ``TokenEndpointOutcome``.
The httpx edge: litellm's handler raises ``HTTPStatusError`` itself on a 4xx/5xx, and every
field the caller reads comes out of a validated ``TokenEndpointOutcome``.
"""
from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 # defer heavy handler import to call time
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # handler is partially typed
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # handler factory params are coarsely typed
)
from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 # deferred with the handler import
try:
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
response = await client.post( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # handler is partially typed
response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # handler params are coarsely typed
url, headers={"Accept": "application/json", **headers}, data=form
)
except httpx.HTTPStatusError as status_err:
@ -114,8 +113,6 @@ async def post_client_credentials_grant(
return TokenEndpointDenied(status_code=status_code, detail=f"token endpoint returned HTTP {status_code}")
except Exception as exc: # noqa: BLE001 # any transport failure is the same outcome: unreachable
return TokenEndpointUnreachable(detail=str(exc))
if not isinstance(response, httpx.Response):
return TokenEndpointUnreachable(detail="token endpoint returned no response")
try:
body: Final = _TOKEN_BODY_ADAPTER.validate_json(response.content)
except ValidationError:

View file

@ -111,9 +111,6 @@ class TokenEndpointClient:
return Error(
CredError.of_upstream_unavailable("token exchange failed: token endpoint returned a non-JSON response")
)
if raw is None:
verbose_proxy_logger.warning("MCP token endpoint %s returned no response", endpoint)
return Error(CredError.of_upstream_unavailable("token exchange failed: no response from token endpoint"))
try:
parsed: Final = _TokenEndpointResponse.model_validate(raw)
except ValidationError:
@ -199,7 +196,7 @@ def _cache_ttl_seconds(expires_in: int | None) -> int:
)
async def _post_form(endpoint: str, data: dict[str, str]) -> object | None:
async def _post_form(endpoint: str, data: dict[str, str]) -> object:
# litellm's httpx handler and httpx.Response are only partially typed; the token endpoint
# returns a JSON object that `_TokenEndpointResponse` validates, so the untyped boundary is
# contained here. A non-2xx raises `httpx.HTTPStatusError`, an unreachable endpoint raises
@ -208,8 +205,6 @@ async def _post_form(endpoint: str, data: dict[str, str]) -> object | None:
# each to a CredError.
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore[reportUnknownVariableType] # litellm http handler is untyped
response = await client.post(endpoint, data=data) # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType] # litellm http handler is untyped
if response is None:
return None
response.raise_for_status()
return response.json() # pyright: ignore[reportAny] # untyped JSON; validated by _TokenEndpointResponse in fetch

View file

@ -1,5 +1,9 @@
import json
from typing import Final
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Final, Protocol
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -7,18 +11,73 @@ from litellm.proxy.utils import PrismaClient
from litellm.repositories.table_repositories import MCPToolsetRepository
from litellm.types.mcp_server.mcp_toolset import (
MCPToolset,
MCPToolsetTool,
NewMCPToolsetRequest,
UpdateMCPToolsetRequest,
)
def _toolset_from_row(row) -> MCPToolset:
class MCPToolsetFields(TypedDict):
"""The ``MCPToolset`` constructor keywords a toolset row expands into."""
toolset_id: ReadOnly[str]
toolset_name: ReadOnly[str]
description: NotRequired[ReadOnly[str | None]]
tools: NotRequired[ReadOnly[list[MCPToolsetTool]]]
created_at: NotRequired[ReadOnly[datetime | None]]
created_by: NotRequired[ReadOnly[str | None]]
updated_at: NotRequired[ReadOnly[datetime | None]]
updated_by: NotRequired[ReadOnly[str | None]]
class MCPToolsetRowData(TypedDict):
"""A toolset table row, whose ``tools`` column is stored as JSON."""
toolset_id: ReadOnly[str]
toolset_name: ReadOnly[str]
description: NotRequired[ReadOnly[str | None]]
tools: NotRequired[ReadOnly[str | list[MCPToolsetTool]]]
created_at: NotRequired[ReadOnly[datetime | None]]
created_by: NotRequired[ReadOnly[str | None]]
updated_at: NotRequired[ReadOnly[datetime | None]]
updated_by: NotRequired[ReadOnly[str | None]]
class MCPToolsetRow(Protocol):
"""A row of the toolset table, as the prisma client returns it."""
def model_dump(self) -> MCPToolsetRowData: ...
class MCPToolsetTable(Protocol):
"""The prisma table actions this module runs against the toolset table."""
async def create(self, data: Mapping[str, object]) -> MCPToolsetRow: ...
async def find_unique(self, where: Mapping[str, object]) -> MCPToolsetRow | None: ...
async def find_first(self, where: Mapping[str, object]) -> MCPToolsetRow | None: ...
async def find_many(self, where: Mapping[str, object]) -> Sequence[MCPToolsetRow]: ...
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> MCPToolsetRow: ...
async def delete(self, where: Mapping[str, object]) -> MCPToolsetRow: ...
def _toolset_table(prisma_client: PrismaClient) -> MCPToolsetTable:
"""The toolset table actions of the prisma client."""
return MCPToolsetRepository(prisma_client).table
def _toolset_from_row(row: MCPToolsetRow) -> MCPToolset:
data: Final = row.model_dump()
tools = data.get("tools") or []
if isinstance(tools, str):
tools = json.loads(tools)
data["tools"] = tools
return MCPToolset(**data)
tools: Final = data.get("tools") or []
resolved: Final[MCPToolsetFields] = {
**data,
"tools": json.loads(tools) if isinstance(tools, str) else tools,
}
return MCPToolset(**resolved)
async def create_mcp_toolset(
@ -31,7 +90,7 @@ async def create_mcp_toolset(
data_dict["tools"] = json.dumps(data_dict.get("tools", []))
data_dict["created_by"] = touched_by
data_dict["updated_by"] = touched_by
row: Final = await MCPToolsetRepository(prisma_client).table.create(data=data_dict)
row: Final = await _toolset_table(prisma_client).create(data=data_dict)
return _toolset_from_row(row)
@ -39,7 +98,7 @@ async def get_mcp_toolset(
prisma_client: PrismaClient,
toolset_id: str,
) -> MCPToolset | None:
row: Final = await MCPToolsetRepository(prisma_client).table.find_unique(where={"toolset_id": toolset_id})
row: Final = await _toolset_table(prisma_client).find_unique(where={"toolset_id": toolset_id})
if row is None:
return None
return _toolset_from_row(row)
@ -47,13 +106,11 @@ async def get_mcp_toolset(
async def list_mcp_toolsets(
prisma_client: PrismaClient,
toolset_ids: list[str] | None = None,
) -> list[MCPToolset]:
toolset_ids: Sequence[str] | None = None,
) -> Sequence[MCPToolset]:
try:
where = {}
if toolset_ids is not None:
where = {"toolset_id": {"in": toolset_ids}}
rows: Final = await MCPToolsetRepository(prisma_client).table.find_many(where=where)
where: Final[Mapping[str, object]] = {} if toolset_ids is None else {"toolset_id": {"in": toolset_ids}}
rows: Final = await _toolset_table(prisma_client).find_many(where=where)
return [_toolset_from_row(r) for r in rows]
except Exception as e:
verbose_proxy_logger.warning("litellm.proxy._experimental.mcp_server.toolset_db::list_mcp_toolsets - %s", e)
@ -64,7 +121,7 @@ async def get_mcp_toolset_by_name(
prisma_client: PrismaClient,
toolset_name: str,
) -> MCPToolset | None:
row: Final = await MCPToolsetRepository(prisma_client).table.find_first(where={"toolset_name": toolset_name})
row: Final = await _toolset_table(prisma_client).find_first(where={"toolset_name": toolset_name})
if row is None:
return None
return _toolset_from_row(row)
@ -80,7 +137,7 @@ async def update_mcp_toolset(
data_dict["tools"] = json.dumps(data_dict["tools"])
data_dict["updated_by"] = touched_by
try:
row: Final = await MCPToolsetRepository(prisma_client).table.update(
row: Final = await _toolset_table(prisma_client).update(
where={"toolset_id": data.toolset_id},
data=data_dict,
)
@ -98,7 +155,7 @@ async def delete_mcp_toolset(
toolset_id: str,
) -> MCPToolset | None:
try:
row: Final = await MCPToolsetRepository(prisma_client).table.delete(where={"toolset_id": toolset_id})
row: Final = await _toolset_table(prisma_client).delete(where={"toolset_id": toolset_id})
except Exception as e:
from prisma.errors import RecordNotFoundError

View file

@ -401,12 +401,10 @@ class AgentRegistry:
The patched agent
"""
try:
existing_row: Final = await AgentsRepository(prisma_client).table.find_unique(
where={"agent_id": agent_id} # mutable-ok: prisma filters are plain dicts
)
if existing_row is None:
existing_record: Final = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
if existing_record is None:
raise Exception(f"Agent with ID {agent_id} not found")
existing_agent: Final = dict(existing_row)
existing_agent: Final[Mapping[str, object]] = dict(existing_record)
augment_agent: Final = {**existing_agent, **agent}
update_data: Final[dict[str, object]] = {}
@ -433,7 +431,7 @@ class AgentRegistry:
update_data["extra_headers"] = extra_headers_value if extra_headers_value is not None else []
if agent.get("object_permission") is not None:
agent_copy: Final = dict(augment_agent)
existing_object_permission_id: Final = existing_agent.get("object_permission_id")
existing_object_permission_id: Final = existing_record.object_permission_id
object_permission_id: Final = await handle_update_object_permission_common(
agent_copy,
existing_object_permission_id,

View file

@ -34,7 +34,7 @@ def teams():
"""Manage teams and team assignments"""
def display_teams_table(teams: list[dict[str, Any]]) -> None:
def display_teams_table(teams: Sequence[dict[str, Any]]) -> None:
"""Display teams in a formatted table"""
console: Final = Console()

View file

@ -1,6 +1,6 @@
import copy
import os
from collections.abc import Callable, Iterable
from collections.abc import Callable, Iterable, Mapping
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias
@ -525,8 +525,8 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset(
def sanitize_openai_provider_metadata(
metadata: dict[str, Any] | None,
) -> dict[str, str] | None:
metadata: Mapping[str, object] | None,
) -> Mapping[str, object] | None:
"""
Keep only provider-safe OpenAI metadata entries (string keys -> string values).
@ -644,7 +644,7 @@ def process_callback(_callback: str, callback_type: str, environment_variables:
return {"name": _callback, "variables": env_vars_dict, "type": callback_type}
def normalize_callback_names(callbacks: Iterable[Any]) -> list[Any]:
def normalize_callback_names(callbacks: Iterable[object] | None) -> list[object]:
if callbacks is None:
return []
return [c.lower() if isinstance(c, str) else c for c in callbacks]
@ -674,7 +674,7 @@ def decrypt_callback_vars(metadata: Any) -> Any:
return _transform_callback_vars(metadata, _decrypt_or_passthrough)
def _transform_callback_vars(metadata: Any, transform: Callable[[str, Any], Any]) -> Any:
def _transform_callback_vars(metadata: object, transform: Callable[[str, Any], Any]) -> object:
if not isinstance(metadata, dict):
return metadata
out: Final = copy.deepcopy(metadata)
@ -704,7 +704,7 @@ def is_sensitive_callback_key(
return _CALLBACK_VAR_MASKER.is_sensitive_key(key)
def _encrypt_if_plaintext(key: str, value: Any) -> Any:
def _encrypt_if_plaintext(key: str, value: object) -> object:
if not isinstance(value, str) or not value:
return value
if not is_sensitive_callback_key(key):
@ -725,7 +725,7 @@ def _encrypt_if_plaintext(key: str, value: Any) -> Any:
return value
def _decrypt_or_passthrough(key: str, value: Any) -> Any:
def _decrypt_or_passthrough(key: str, value: object) -> object:
if not isinstance(value, str) or not value:
return value
if not value.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX):

View file

@ -40,30 +40,31 @@ class UserApiKeyCache(DualCache):
@overload
def get_cache(
self,
key: Any,
parent_otel_span: Any = None,
key: object,
parent_otel_span: object = None,
local_only: bool = False,
*,
model_type: type[T],
**kwargs: Any,
**kwargs: object,
) -> T | None: ...
@overload
def get_cache(
self,
key: Any,
parent_otel_span: Any = None,
key: object,
parent_otel_span: object = None,
local_only: bool = False,
**kwargs: Any,
model_type: None = None,
**kwargs: object,
) -> Any: ...
def get_cache(
self,
key,
parent_otel_span=None,
key: object,
parent_otel_span: object = None,
local_only: bool = False,
model_type: type[BaseModel] | None = None,
**kwargs,
**kwargs: object,
) -> Any | BaseModel | None:
if model_type is None and "model_type" in kwargs:
model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
@ -85,30 +86,31 @@ class UserApiKeyCache(DualCache):
@overload
async def async_get_cache(
self,
key: Any,
parent_otel_span: Any = None,
key: object,
parent_otel_span: object = None,
local_only: bool = False,
*,
model_type: type[T],
**kwargs: Any,
**kwargs: object,
) -> T | None: ...
@overload
async def async_get_cache(
self,
key: Any,
parent_otel_span: Any = None,
key: object,
parent_otel_span: object = None,
local_only: bool = False,
**kwargs: Any,
model_type: None = None,
**kwargs: object,
) -> Any: ...
async def async_get_cache(
self,
key,
parent_otel_span=None,
key: object,
parent_otel_span: object = None,
local_only: bool = False,
model_type: type[BaseModel] | None = None,
**kwargs,
**kwargs: object,
) -> Any | BaseModel | None:
if model_type is None and "model_type" in kwargs:
model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
@ -129,17 +131,17 @@ class UserApiKeyCache(DualCache):
return None
return decoded
def set_cache(self, key, value, local_only: bool = False, **kwargs):
def set_cache(self, key: object, value: object, local_only: bool = False, **kwargs: object):
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
payload: Final = CacheCodec.serialize(value, model_type=model_type)
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs)
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs):
async def async_set_cache(self, key: object, value: object, local_only: bool = False, **kwargs: object):
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
payload: Final = CacheCodec.serialize(value, model_type=model_type)
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs) -> None:
async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs: object) -> None:
"""
Batch writes with the same Codec boundary as ``async_set_cache`` without
``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged.

View file

@ -748,7 +748,7 @@ class CompresrGuardrail(CustomGuardrail):
}
try:
raw_response: HttpxResponse | None = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
url=url,
json=payload,
headers=self._request_headers(),
@ -778,11 +778,11 @@ class CompresrGuardrail(CustomGuardrail):
{"detail": str(e)},
)
return None
if raw_response is None or not 200 <= raw_response.status_code < 300:
if not 200 <= raw_response.status_code < 300:
self._handle_compress_failure(
"Compresr compression service returned an error",
{
"status_code": getattr(raw_response, "status_code", None),
"status_code": raw_response.status_code,
"body": _safe_response_text(raw_response),
},
)

View file

@ -433,7 +433,7 @@ class HeadroomGuardrail(CustomGuardrail):
payload["model"] = model
try:
raw_response: HttpxResponse | None = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
url=f"{self.headroom_api_base}/v1/compress",
json=payload,
headers=self._request_headers(),
@ -458,16 +458,6 @@ class HeadroomGuardrail(CustomGuardrail):
False,
{},
)
if raw_response is None:
return (
self._handle_compress_failure(
messages,
"Headroom compression service returned no response",
{},
),
False,
{},
)
response: Final[HttpxResponse] = raw_response
if response.status_code != 200:
@ -580,7 +570,7 @@ class HeadroomGuardrail(CustomGuardrail):
params["query"] = query
try:
raw_response: HttpxResponse | None = await self.async_handler.get( # pyright: ignore[reportUnknownMemberType]
raw_response: HttpxResponse = await self.async_handler.get( # pyright: ignore[reportUnknownMemberType]
url=f"{self.headroom_api_base}/v1/retrieve/{hash_value}",
params=params,
headers=self._request_headers(),
@ -589,7 +579,7 @@ class HeadroomGuardrail(CustomGuardrail):
verbose_proxy_logger.warning("Headroom: retrieve failed for hash=%s: %s", hash_value, e)
return f"[Headroom: retrieval failed for hash={hash_value}]"
if raw_response is None or raw_response.status_code == 404:
if raw_response.status_code == 404:
return f"[Headroom: hash={hash_value} not found or expired]"
if raw_response.status_code != 200:

View file

@ -2,7 +2,7 @@ from __future__ import annotations
import os
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict
from typing import TYPE_CHECKING, Final, Literal, Protocol
from urllib.parse import urlparse
from uuid import uuid4
@ -11,7 +11,7 @@ import requests
from fastapi import HTTPException
from httpx import HTTPStatusError
from requests.auth import HTTPBasicAuth
from typing_extensions import ReadOnly
from typing_extensions import ReadOnly, TypedDict, Unpack
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
@ -40,14 +40,21 @@ if TYPE_CHECKING:
_AUTH_TIMEOUT_SECONDS: Final[float] = 30.0
class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
"""Base-class constructor options carried by this guardrail's forwarded keyword arguments."""
guardrail_name: ReadOnly[str | None]
supported_event_hooks: list[GuardrailEventHooks] | None
class _HiddenlayerEvaluation(TypedDict, total=False):
action: str
threat_level: str
action: ReadOnly[str]
threat_level: ReadOnly[str]
class _HiddenlayerAnalysisEntry(TypedDict, total=False):
name: str
detected: bool
name: ReadOnly[str]
detected: ReadOnly[bool]
class _HiddenlayerModifiedMessage(TypedDict):
@ -59,9 +66,9 @@ class _HiddenlayerModifiedSide(TypedDict):
class _HiddenlayerResponse(TypedDict, total=False):
evaluation: _HiddenlayerEvaluation
analysis: Sequence[_HiddenlayerAnalysisEntry]
modified_data: Mapping[str, _HiddenlayerModifiedSide]
evaluation: ReadOnly[_HiddenlayerEvaluation]
analysis: ReadOnly[Sequence[_HiddenlayerAnalysisEntry]]
modified_data: ReadOnly[Mapping[str, _HiddenlayerModifiedSide]]
class _ProxyServerRequest(TypedDict, total=False):
@ -194,7 +201,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
api_key: str | None = None,
api_base: str | None = None,
auth_url: str | None = None,
**kwargs: Any,
**kwargs: Unpack[_CustomGuardrailOptions],
) -> None:
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
self.hiddenlayer_client_id = api_id or os.getenv("HIDDENLAYER_CLIENT_ID")
@ -399,7 +406,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
api_key: str | None = None,
api_base: str | None = None,
auth_url: str | None = None,
**kwargs: Any,
**kwargs: Unpack[_CustomGuardrailOptions],
) -> None:
self.hiddenlayer_client_id = api_id or os.getenv("HIDDENLAYER_CLIENT_ID")
self.hiddenlayer_client_secret = api_key or os.getenv("HIDDENLAYER_CLIENT_SECRET")
@ -530,7 +537,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
self,
payload: _HiddenlayerV2Payload,
input_type: Literal["request", "response"],
hl_headers: dict[str, str],
hl_headers: Mapping[str, str],
) -> httpx.Response:
if input_type == "request":
path = "detection/v2/request-evaluations"

View file

@ -7,8 +7,9 @@
import enum
import json
import os
from collections.abc import Callable, Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, cast
from urllib.parse import urlparse
from litellm._logging import verbose_proxy_logger
@ -36,6 +37,8 @@ _AIDR_SCAN_ENDPOINT: Final = "/litellm/guardrail"
_INTERVENED_INPUT_FIELDS: Final = ("texts", "images", "tools", "tool_calls")
_DEFAULT_API_BASE_HOSTNAME: Final = urlparse(_DEFAULT_API_BASE).hostname
_GuardrailJsonResponse: TypeAlias = Exception | str | dict[str, object]
_KEYS_DUPLICATING_SCAN_INPUTS: Final = ("messages", "input")
_LOGGING_KEYS_DUPLICATING_SCAN_INPUTS: Final = _KEYS_DUPLICATING_SCAN_INPUTS + (
"additional_args",
@ -119,7 +122,7 @@ class NomaV2Guardrail(CustomGuardrail):
def _resolve_action_from_response(
self,
response_json: dict,
response_json: Mapping[str, object],
) -> _Action:
action: Final = response_json.get("action")
if isinstance(action, str):
@ -165,10 +168,11 @@ class NomaV2Guardrail(CustomGuardrail):
@staticmethod
def _sanitize_payload_for_transport(payload: dict) -> dict:
def _default(obj: Any) -> Any:
if hasattr(obj, "model_dump"):
def _default(obj: object) -> object:
model_dump: Final[Callable[[], Mapping[str, object]] | None] = getattr(obj, "model_dump", None)
if model_dump is not None:
try:
return obj.model_dump()
return model_dump()
except Exception:
pass
return str(obj)
@ -178,7 +182,7 @@ class NomaV2Guardrail(CustomGuardrail):
except (ValueError, TypeError):
json_str = safe_dumps(payload)
safe_payload: Final = safe_json_loads(json_str, default={})
safe_payload: Final[object] = safe_json_loads(json_str, default={})
if safe_payload == {} and payload:
verbose_proxy_logger.warning(
"Noma v2 guardrail: payload serialization failed, falling back to empty payload"
@ -196,7 +200,7 @@ class NomaV2Guardrail(CustomGuardrail):
async def _call_noma_scan(
self,
payload: dict,
) -> dict:
) -> dict[str, object]:
headers: Final[dict[str, str]] = {"Content-Type": "application/json"}
authorization_header: Final = self._get_authorization_header()
if authorization_header:
@ -215,7 +219,7 @@ class NomaV2Guardrail(CustomGuardrail):
response.text,
)
response.raise_for_status()
response_json: Final = response.json()
response_json: Final[dict[str, object]] = response.json()
verbose_proxy_logger.debug(
"Noma v2 AIDR response parsed: %s",
json.dumps(response_json, default=str),
@ -227,7 +231,7 @@ class NomaV2Guardrail(CustomGuardrail):
request_data: dict,
start_time: datetime,
guardrail_status: GuardrailStatus,
guardrail_json_response: Any,
guardrail_json_response: _GuardrailJsonResponse,
) -> None:
end_time: Final = datetime.now()
duration: Final = (end_time - start_time).total_seconds()
@ -270,11 +274,11 @@ class NomaV2Guardrail(CustomGuardrail):
) -> GenericGuardrailAPIInputs:
start_time: Final = datetime.now()
guardrail_status: GuardrailStatus = "success"
guardrail_json_response: Any = {}
guardrail_json_response: _GuardrailJsonResponse = {}
dynamic_params = self.get_guardrail_dynamic_request_body_params(request_data)
if not isinstance(dynamic_params, dict):
dynamic_params = {}
response_json: dict | None = None
response_json: dict[str, object] | None = None
# Per-request dynamic params can override configured application context.
application_id = self._get_non_empty_str(dynamic_params.get("application_id"))

View file

@ -197,14 +197,11 @@ class RepelloAIGuardrail(CustomGuardrail):
repelloai_response: RepelloAIAnalyzeResponse | None = None
try:
verbose_proxy_logger.debug("RepelloAI Argus request: %s", request)
raw_response: HttpxResponse | None = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
response: Final[HttpxResponse] = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
url=endpoint,
headers={"X-API-Key": self.repelloai_api_key},
json=request,
)
if raw_response is None:
raise ValueError("RepelloAI Argus returned no response")
response: Final[HttpxResponse] = raw_response
self._raise_for_config_error(response)
response.raise_for_status()
try:

View file

@ -1,4 +1,6 @@
from typing import Final
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Final, Protocol
from fastapi import APIRouter, Depends, HTTPException, Query
@ -18,7 +20,59 @@ from litellm.repositories.table_repositories import JWTKeyMappingRepository
router: Final = APIRouter()
def _to_response(mapping) -> JWTKeyMappingResponse:
class _JWTKeyMappingRecord(Protocol):
"""A ``LiteLLM_JWTKeyMapping`` row, viewed through the columns these endpoints read."""
@property
def id(self) -> str: ...
@property
def jwt_claim_name(self) -> str: ...
@property
def jwt_claim_value(self) -> str: ...
@property
def description(self) -> str | None: ...
@property
def is_active(self) -> bool: ...
@property
def created_at(self) -> datetime: ...
@property
def updated_at(self) -> datetime: ...
@property
def created_by(self) -> str | None: ...
@property
def updated_by(self) -> str | None: ...
class _JWTKeyMappingTable(Protocol):
"""The Prisma table actions these endpoints issue against the JWT key mapping table."""
async def create(self, *, data: Mapping[str, object]) -> _JWTKeyMappingRecord: ...
async def find_unique(self, *, where: Mapping[str, object]) -> _JWTKeyMappingRecord | None: ...
async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> _JWTKeyMappingRecord: ...
async def delete(self, *, where: Mapping[str, object]) -> _JWTKeyMappingRecord | None: ...
async def find_many(self, *, skip: int, take: int, order: Mapping[str, str]) -> Sequence[_JWTKeyMappingRecord]: ...
async def count(self) -> int: ...
def _mapping_table(prisma_client: object) -> _JWTKeyMappingTable:
"""View the JWT key mapping repository's untyped Prisma table through the actions used here."""
return JWTKeyMappingRepository(prisma_client).table
def _to_response(mapping: _JWTKeyMappingRecord) -> JWTKeyMappingResponse:
"""Convert a Prisma mapping object to a safe response (no hashed token)."""
return JWTKeyMappingResponse(
id=mapping.id,
@ -62,7 +116,7 @@ async def create_jwt_key_mapping(
if data.description is not None:
create_data["description"] = data.description
new_mapping: Final = await JWTKeyMappingRepository(prisma_client).table.create(data=create_data)
new_mapping: Final = await _mapping_table(prisma_client).create(data=create_data)
# Invalidate cache
cache_key: Final = f"jwt_key_mapping:{data.jwt_claim_name}:{data.jwt_claim_value}"
@ -110,7 +164,7 @@ async def update_jwt_key_mapping(
try:
# Get old mapping for cache invalidation
old_mapping: Final = await JWTKeyMappingRepository(prisma_client).table.find_unique(where={"id": data.id})
old_mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": data.id})
if old_mapping is None:
raise HTTPException(status_code=404, detail="Mapping not found")
@ -118,9 +172,7 @@ async def update_jwt_key_mapping(
cache_key = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}"
await user_api_key_cache.async_delete_cache(cache_key)
updated_mapping: Final = await JWTKeyMappingRepository(prisma_client).table.update(
where={"id": data.id}, data=update_data
)
updated_mapping: Final = await _mapping_table(prisma_client).update(where={"id": data.id}, data=update_data)
if updated_mapping is None:
raise HTTPException(status_code=404, detail="Mapping not found")
@ -162,7 +214,7 @@ async def delete_jwt_key_mapping(
try:
# Get old mapping for cache invalidation
old_mapping: Final = await JWTKeyMappingRepository(prisma_client).table.find_unique(where={"id": data.id})
old_mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": data.id})
if old_mapping is None:
raise HTTPException(status_code=404, detail="Mapping not found")
@ -170,7 +222,7 @@ async def delete_jwt_key_mapping(
cache_key: Final = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}"
await user_api_key_cache.async_delete_cache(cache_key)
await JWTKeyMappingRepository(prisma_client).table.delete(where={"id": data.id})
await _mapping_table(prisma_client).delete(where={"id": data.id})
return {"status": "success"}
except HTTPException:
raise
@ -198,12 +250,12 @@ async def list_jwt_key_mappings(
try:
skip: Final = (page - 1) * size
mappings: Final = await JWTKeyMappingRepository(prisma_client).table.find_many(
mappings: Final = await _mapping_table(prisma_client).find_many(
skip=skip,
take=size,
order={"created_at": "desc"},
)
total_count: Final = await JWTKeyMappingRepository(prisma_client).table.count()
total_count: Final = await _mapping_table(prisma_client).count()
return {
"mappings": [_to_response(m) for m in mappings],
"total_count": total_count,
@ -235,7 +287,7 @@ async def info_jwt_key_mapping(
raise HTTPException(status_code=500, detail="Database not connected")
try:
mapping: Final = await JWTKeyMappingRepository(prisma_client).table.find_unique(where={"id": id})
mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": id})
if mapping is None:
raise HTTPException(status_code=404, detail="Mapping not found")
return _to_response(mapping)

View file

@ -10,8 +10,8 @@ https://platform.openai.com/docs/api-reference/responses-streaming
import asyncio
import json
from collections.abc import Sequence
from typing import TYPE_CHECKING, Final, TypedDict, cast
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Final, TypedDict
from fastapi import Request, Response
from fastapi.responses import StreamingResponse
@ -38,6 +38,36 @@ class _StreamOutputItem(TypedDict, total=False):
content: ReadOnly[Sequence[_StreamContentPart | None]]
class _StreamTerminalResponse(TypedDict, total=False):
"""Fields of the ``response`` payload carried by a terminal streaming event."""
status: ReadOnly[ResponsesAPIStatus]
tool_choice: ReadOnly[object]
model: ReadOnly[str]
instructions: ReadOnly[str]
temperature: ReadOnly[float]
top_p: ReadOnly[float]
max_output_tokens: ReadOnly[int]
previous_response_id: ReadOnly[str]
truncation: ReadOnly[str]
parallel_tool_calls: ReadOnly[bool]
user: ReadOnly[str]
store: ReadOnly[bool]
output: ReadOnly[Sequence[_StreamOutputItem]]
class _StreamEvent(TypedDict, total=False):
"""One decoded ``data:`` frame of an OpenAI Responses streaming body."""
type: ReadOnly[str]
item: ReadOnly[_StreamOutputItem]
item_id: ReadOnly[str]
part: ReadOnly[_StreamContentPart]
content_index: ReadOnly[int]
delta: ReadOnly[str]
response: ReadOnly[_StreamTerminalResponse]
async def background_streaming_task(
polling_id: str,
data,
@ -139,7 +169,7 @@ async def background_streaming_task(
None # Will be set by response.completed/failed/incomplete/cancelled
)
terminal_error = None
_event_to_status: Final = {
_event_to_status: Final[Mapping[str, ResponsesAPIStatus]] = {
"response.completed": "completed",
"response.failed": "failed",
"response.incomplete": "incomplete",
@ -180,7 +210,7 @@ async def background_streaming_task(
break
try:
event = json.loads(chunk_data)
event: _StreamEvent = json.loads(chunk_data)
event_type = event.get("type", "")
# Process different event types based on OpenAI streaming spec
@ -288,12 +318,9 @@ async def background_streaming_task(
# Terminal event - extract all ResponsesAPIResponse fields
# https://platform.openai.com/docs/api-reference/responses-streaming
response_data = event.get("response", {})
terminal_status = cast(
ResponsesAPIStatus,
response_data.get(
"status",
_event_to_status.get(event_type, "completed"),
),
terminal_status = response_data.get(
"status",
_event_to_status.get(event_type, "completed"),
)
# Extract error for failed and incomplete responses

View file

@ -2,8 +2,9 @@
Search Tool Registry for managing search tool configurations.
"""
from collections.abc import Iterator, Mapping, Sequence
from datetime import datetime, timezone
from typing import Final
from typing import Final, Protocol
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -13,6 +14,40 @@ from litellm.repositories.table_repositories import SearchToolsRepository
from litellm.types.search import SearchTool
class SearchToolRecord(Protocol):
search_tool_id: str
search_tool_name: str
created_at: datetime
updated_at: datetime
def __iter__(self) -> Iterator[tuple[str, object]]: ...
class SearchToolTableClient(Protocol):
async def create(self, data: Mapping[str, object]) -> SearchToolRecord: ...
async def find_unique(self, where: Mapping[str, object]) -> SearchToolRecord | None: ...
async def find_many(self, order: Mapping[str, str] | None = None) -> Sequence[SearchToolRecord]: ...
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> SearchToolRecord: ...
async def delete(self, where: Mapping[str, object]) -> SearchToolRecord: ...
class _SearchToolsRepositoryView(Protocol):
@property
def table(self) -> SearchToolTableClient: ...
def _search_tools_table_of(repository: _SearchToolsRepositoryView) -> SearchToolTableClient:
return repository.table
def _search_tools_table(prisma_client: PrismaClient) -> SearchToolTableClient:
return _search_tools_table_of(SearchToolsRepository(prisma_client))
class SearchToolRegistry:
"""
Handles adding, removing, and getting search tools in DB + in memory.
@ -22,7 +57,7 @@ class SearchToolRegistry:
pass
@staticmethod
def _convert_prisma_to_dict(prisma_obj) -> dict:
def _convert_prisma_to_dict(prisma_obj: SearchToolRecord) -> dict:
"""
Convert Prisma result to dict with datetime objects as ISO format strings.
@ -35,9 +70,9 @@ class SearchToolRegistry:
result: Final = dict(prisma_obj)
# Convert datetime objects to ISO format strings
if "created_at" in result and result["created_at"]:
result["created_at"] = result["created_at"].isoformat()
result["created_at"] = prisma_obj.created_at.isoformat()
if "updated_at" in result and result["updated_at"]:
result["updated_at"] = result["updated_at"].isoformat()
result["updated_at"] = prisma_obj.updated_at.isoformat()
return result
###########################################################
@ -61,7 +96,7 @@ class SearchToolRegistry:
search_tool_info: Final[str] = safe_dumps(search_tool.get("search_tool_info", {}))
# Create search tool in DB
created_search_tool: Final = await SearchToolsRepository(prisma_client).table.create(
created_search_tool: Final = await _search_tools_table(prisma_client).create(
data={
"search_tool_name": search_tool_name,
"litellm_params": litellm_params,
@ -95,7 +130,7 @@ class SearchToolRegistry:
"""
try:
# Get search tool before deletion for response
existing_tool: Final = await SearchToolsRepository(prisma_client).table.find_unique(
existing_tool: Final = await _search_tools_table(prisma_client).find_unique(
where={"search_tool_id": search_tool_id}
)
@ -103,7 +138,7 @@ class SearchToolRegistry:
raise Exception(f"Search tool with ID {search_tool_id} not found")
# Delete from DB
await SearchToolsRepository(prisma_client).table.delete(where={"search_tool_id": search_tool_id})
await _search_tools_table(prisma_client).delete(where={"search_tool_id": search_tool_id})
return {
"message": f"Search tool {search_tool_id} deleted successfully",
@ -131,7 +166,7 @@ class SearchToolRegistry:
search_tool_info: Final[str] = safe_dumps(search_tool.get("search_tool_info", {}))
# Update in DB
updated_search_tool: Final = await SearchToolsRepository(prisma_client).table.update(
updated_search_tool: Final = await _search_tools_table(prisma_client).update(
where={"search_tool_id": search_tool_id},
data={
"search_tool_name": search_tool_name,
@ -163,7 +198,7 @@ class SearchToolRegistry:
try:
search_tools_from_db: Final = await call_with_db_reconnect_retry(
prisma_client,
lambda: SearchToolsRepository(prisma_client).table.find_many(
lambda: _search_tools_table(prisma_client).find_many(
order={"created_at": "desc"},
),
reason="get_all_search_tools_from_db_lookup_failure",
@ -194,7 +229,7 @@ class SearchToolRegistry:
Search tool configuration or None if not found
"""
try:
search_tool: Final = await SearchToolsRepository(prisma_client).table.find_unique(
search_tool: Final = await _search_tools_table(prisma_client).find_unique(
where={"search_tool_id": search_tool_id}
)
@ -222,7 +257,7 @@ class SearchToolRegistry:
Search tool configuration or None if not found
"""
try:
search_tool: Final = await SearchToolsRepository(prisma_client).table.find_unique(
search_tool: Final = await _search_tools_table(prisma_client).find_unique(
where={"search_tool_name": search_tool_name}
)

View file

@ -1,11 +1,45 @@
from collections.abc import Sequence
from typing import Any, Final
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
from litellm.types.utils import ModelResponse
from litellm.types.vector_stores import (
VectorStoreResultContent,
VectorStoreSearchResponse,
)
from litellm.types.vector_stores import VectorStoreSearchResponse
class _ResultContentView(TypedDict):
"""Content entry carried by a vector store search result."""
type: ReadOnly[NotRequired[str]]
text: ReadOnly[str]
class _SearchResultView(TypedDict):
"""Vector store search result, as far as :class:`RAGQuery` reads it."""
content: ReadOnly[NotRequired[Sequence[_ResultContentView]]]
text: ReadOnly[NotRequired[str]]
class _SearchDataView(TypedDict):
results: ReadOnly[Sequence[_SearchResultView]]
class _ContextChunksView(TypedDict):
chunks: ReadOnly[Sequence[_SearchResultView | str | None]]
class _RerankResultView(TypedDict):
index: ReadOnly[NotRequired[int]]
class _RerankResultsView(TypedDict):
results: ReadOnly[Sequence[_RerankResultView]]
class _MessageView(TypedDict):
message: ReadOnly[object]
class RAGQuery:
@ -42,9 +76,10 @@ class RAGQuery:
"""
context_content = RAGQuery.CONTENT_PREFIX_STRING
for chunk in context_chunks:
chunks: Final[_ContextChunksView] = {"chunks": context_chunks}
for chunk in chunks["chunks"]:
if isinstance(chunk, dict):
result_content: list[VectorStoreResultContent] | None = chunk.get("content")
result_content: Sequence[_ResultContentView] | None = chunk.get("content")
if result_content:
for content_item in result_content:
content_text: str | None = content_item.get("text")
@ -64,14 +99,15 @@ class RAGQuery:
def add_search_results_to_response(
response: ModelResponse,
search_results: VectorStoreSearchResponse,
rerank_results: Any | None = None,
rerank_results: object = None,
) -> ModelResponse:
"""
Add search results to the response choices.
"""
if hasattr(response, "choices") and response.choices:
for choice in response.choices:
message = getattr(choice, "message", None)
message_view: _MessageView = {"message": getattr(choice, "message", None)}
message = message_view["message"]
if message is not None:
# Get existing provider_specific_fields or create new dict
provider_fields = getattr(message, "provider_specific_fields", None) or {}
@ -91,7 +127,8 @@ class RAGQuery:
) -> list[str | dict[str, Any]]:
"""Extract text documents from vector store search response."""
documents: Final[list[str | dict[str, Any]]] = []
for result in search_response.get("data", []):
search_data: Final[_SearchDataView] = {"results": search_response.get("data", [])}
for result in search_data["results"]:
content_list = result.get("content", [])
for content in content_list:
if content.get("type") == "text" and content.get("text"):
@ -99,11 +136,13 @@ class RAGQuery:
return documents
@staticmethod
def get_top_chunks_from_rerank(search_response: Any, rerank_response: Any) -> list[Any]:
def get_top_chunks_from_rerank(search_response: Any, rerank_response: Any) -> list[_SearchResultView]:
"""Get the original search results corresponding to the top reranked results."""
top_chunks: Final = []
original_results: Final = search_response.get("data", [])
for result in rerank_response.get("results", []):
top_chunks: Final[list[_SearchResultView]] = []
search_data: Final[_SearchDataView] = {"results": search_response.get("data", [])}
original_results: Final = search_data["results"]
reranked: Final[_RerankResultsView] = {"results": rerank_response.get("results", [])}
for result in reranked["results"]:
index = result.get("index")
if index is not None and index < len(original_results):
top_chunks.append(original_results[index])

View file

@ -40,7 +40,7 @@ def record_to_dict(record: DbRecord) -> Mapping[str, object]:
class BaseRepository(ABC, Generic[T]):
"""Abstract base class for all repositories."""
def __init__(self, prisma_client: Any): # any-ok: PrismaClient is an untyped runtime wrapper
def __init__(self, prisma_client: object):
self._prisma_client = prisma_client
@property

View file

@ -3,9 +3,9 @@ Team repository for database operations on LiteLLM_TeamTable.
"""
import json
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, Protocol
from pydantic import TypeAdapter
@ -21,6 +21,25 @@ if TYPE_CHECKING:
from prisma import Prisma
from prisma import models as prisma_models
class _TeamArrays(Protocol):
"""The string array columns of a team row, which the domain model leaves untyped."""
@property
def members(self) -> Sequence[str]: ...
@property
def admins(self) -> Sequence[str]: ...
@property
def models(self) -> Sequence[str]: ...
def _team_arrays(team: LiteLLM_TeamTable) -> _TeamArrays:
"""View a team's untyped list columns as sequences of ids."""
return team
_MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member])
_JSON_ENCODED_TEAM_FIELDS: Final = (
"metadata",
@ -80,8 +99,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
)
if not rows:
return None
raw_value: Final = rows[0]["members_with_roles"]
parsed: Final = json.loads(raw_value) if isinstance(raw_value, str) else raw_value
raw_value: Final[object] = rows[0]["members_with_roles"]
parsed: Final[object] = json.loads(raw_value) if isinstance(raw_value, str) else raw_value
if not parsed:
return []
return _MEMBERS_WITH_ROLES_ADAPTER.validate_python(parsed)
@ -315,7 +334,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
if team is None:
return None
members: Final = [m for m in team.members if m != user_id]
members: Final = [m for m in _team_arrays(team).members if m != user_id]
return await self.update(team_id, {"members": members}, id_field="team_id")
async def add_admin(self, team_id: str, user_id: str) -> LiteLLM_TeamTable | None:
@ -340,7 +359,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
if team is None:
return None
admins: Final = [a for a in team.admins if a != user_id]
admins: Final = [a for a in _team_arrays(team).admins if a != user_id]
return await self.update(team_id, {"admins": admins}, id_field="team_id")
async def add_models(self, team_id: str, models: list[str]) -> LiteLLM_TeamTable | None:
@ -365,5 +384,5 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
if team is None:
return None
current_models: Final = [m for m in team.models if m not in models]
current_models: Final = [m for m in _team_arrays(team).models if m not in models]
return await self.update(team_id, {"models": current_models}, id_field="team_id")

View file

@ -1,12 +1,12 @@
{
"ANN001": {
"limit": 3012
"limit": 2998
},
"ANN002": {
"limit": 71
},
"ANN003": {
"limit": 827
"limit": 809
},
"ANN201": {
"limit": 2003
@ -15,7 +15,7 @@
"limit": 845
},
"ANN204": {
"limit": 702
"limit": 698
},
"ANN205": {
"limit": 112
@ -24,7 +24,7 @@
"limit": 133
},
"ANN401": {
"limit": 654
"limit": 597
},
"ASYNC230": {
"limit": 11
@ -195,7 +195,7 @@
"limit": 22
},
"SIM101": {
"limit": 58
"limit": 56
},
"SIM102": {
"limit": 315
@ -231,7 +231,7 @@
"limit": 5
},
"TID251": {
"limit": 1116
"limit": 1111
},
"TRY002": {
"limit": 524
@ -246,7 +246,7 @@
"limit": 113
},
"TRY300": {
"limit": 857
"limit": 855
},
"UP028": {
"limit": 2

View file

@ -310,6 +310,25 @@ class TestGDCGeminiConfig:
api_base=TEST_API_BASE,
)
def test_validate_environment_credentials_missing_audience_binding_are_named(self):
config = GDCGeminiConfig()
creds_without_audience_binding = MagicMock(spec=[])
with patch(
"google.auth.load_credentials_from_dict",
return_value=(creds_without_audience_binding, None),
):
with pytest.raises(AttributeError, match="must expose with_gdch_audience"):
config.validate_environment(
headers={},
model=TEST_MODEL,
messages=[],
optional_params={},
litellm_params={"vertex_project": TEST_PROJECT},
api_key=TEST_API_KEY,
api_base=TEST_API_BASE,
)
def test_validate_environment_string_false_disables_token_caching(self):
config = GDCGeminiConfig()
mock_creds = MagicMock()

View file

@ -224,20 +224,6 @@ async def test_fetch_invalid_json_maps_to_upstream_unavailable():
assert "idp.example.com" not in result.error.summary
@pytest.mark.asyncio
async def test_fetch_none_response_is_upstream_unavailable():
with patch(_PATCH_TARGET, return_value=_client(None)):
result = await TokenEndpointClient().fetch(
_ENDPOINT,
_CLIENT_ID,
{"grant_type": "g"},
ClientSecretAuth(client_secret=SecretStr("s")),
)
assert isinstance(result, Error)
assert result.error.tag == "upstream_unavailable"
@pytest.mark.asyncio
async def test_fetch_missing_access_token_is_upstream_unavailable():
bad = MagicMock()
@ -275,21 +261,6 @@ async def test_fetch_http_error_does_not_leak_endpoint_url():
assert "idp.example.com" not in result.error.summary
@pytest.mark.asyncio
async def test_fetch_none_response_does_not_leak_endpoint_url():
with patch(_PATCH_TARGET, return_value=_client(None)):
result = await TokenEndpointClient().fetch(
_ENDPOINT,
_CLIENT_ID,
{"grant_type": "g"},
ClientSecretAuth(client_secret=SecretStr("s")),
)
assert isinstance(result, Error)
assert _ENDPOINT not in result.error.summary
assert "idp.example.com" not in result.error.summary
@pytest.mark.asyncio
async def test_fetch_missing_access_token_does_not_leak_endpoint_url():
bad = MagicMock()

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22704
"limit": 22642
},
"LIT002": {
"limit": 26854
"limit": 26834
},
"LIT003": {
"limit": 269
@ -15,7 +15,7 @@
"limit": 0
},
"LIT006": {
"limit": 1063
"limit": 1041
},
"LIT007": {
"limit": 0
@ -27,12 +27,12 @@
"limit": 0
},
"LIT010": {
"limit": 16564
"limit": 16550
},
"LIT011": {
"limit": 5577
"limit": 5576
},
"LIT012": {
"limit": 4506
"limit": 4496
}
}