refactor(types): replace Any with precise types across 73 modules

Narrows reportAny / reportExplicitAny hot spots in provider transformations,
proxy endpoints, integrations and secret managers by introducing TypedDicts,
Protocols and object-typed boundaries instead of Any, then ratchets the
budget ceilings down to match.

reportAny 14765 -> 14076, reportExplicitAny 4493 -> 4128, ANN401 387 -> 307
This commit is contained in:
mateo-berri 2026-09-01 10:59:38 +00:00
parent ec3f8183c3
commit 215bf03373
76 changed files with 1458 additions and 579 deletions

View file

@ -1,6 +1,6 @@
{
"reportAny": {
"limit": 14765
"limit": 14076
},
"reportArgumentType": {
"limit": 2216
@ -24,7 +24,7 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 4493
"limit": 4128
},
"reportFunctionMemberAccess": {
"limit": 7
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5607
"limit": 5601
},
"reportMissingTypeArgument": {
"limit": 15310
"limit": 15306
},
"reportMissingTypeStubs": {
"limit": 40
@ -105,13 +105,13 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38368
"limit": 38350
},
"reportUnknownParameterType": {
"limit": 19633
"limit": 19626
},
"reportUnknownVariableType": {
"limit": 29908
"limit": 29890
},
"reportUnnecessaryCast": {
"limit": 111
@ -123,7 +123,7 @@
"limit": 5
},
"reportUnnecessaryIsInstance": {
"limit": 828
"limit": 826
},
"reportUntypedBaseClass": {
"limit": 0

View file

@ -12,6 +12,7 @@ import hashlib
import json
import time
import traceback
from collections.abc import Mapping
from enum import Enum
from typing import Any, Final
@ -506,7 +507,7 @@ class Cache:
def _get_cache_logic(
self,
cached_result: Any | None,
cached_result: object | None,
max_age: float | None,
):
"""
@ -538,8 +539,8 @@ class Cache:
return cached_result
@staticmethod
def _get_safe_cache_lookup_kwargs(kwargs: dict[str, Any]) -> dict[str, Any]:
cache_lookup_kwargs: Final[dict[str, Any]] = {}
def _get_safe_cache_lookup_kwargs(kwargs: Mapping[str, object]) -> dict[str, object]:
cache_lookup_kwargs: Final[dict[str, object]] = {}
for prompt_kwarg in ("messages", "input"):
if prompt_kwarg in kwargs:
cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg]
@ -552,7 +553,7 @@ class Cache:
@staticmethod
def _update_metadata_from_cache_lookup_kwargs(
original_kwargs: dict[str, Any], cache_lookup_kwargs: dict[str, Any]
original_kwargs: Mapping[str, object], cache_lookup_kwargs: Mapping[str, object]
) -> None:
original_metadata: Final = original_kwargs.get("metadata")
cache_lookup_metadata: Final = cache_lookup_kwargs.get("metadata")

View file

@ -12,7 +12,7 @@ import ast
import asyncio
import json
import os
from typing import TYPE_CHECKING, Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
import litellm
from litellm._logging import print_verbose
@ -39,6 +39,12 @@ if TYPE_CHECKING:
from litellm.router import Router
class _QdrantCollectionDetailsResponse(Protocol):
"""The qdrant `/collections/{name}` response, whose body is kept as an opaque JSON object."""
def json(self) -> dict[str, object]: ...
class QdrantSemanticCache(BaseCache):
CACHE_KEY_FIELD_NAME = "litellm_cache_key"
embedding_max_input_tokens: int | None = None
@ -115,15 +121,15 @@ class QdrantSemanticCache(BaseCache):
raise ValueError(f"Error from qdrant checking if /collections exist {collection_exists.text}")
if collection_exists.json()["result"]["exists"]:
collection_details = self.sync_client.get(
collection_details: _QdrantCollectionDetailsResponse = self.sync_client.get(
url=f"{self.qdrant_api_base}/collections/{self.collection_name}",
headers=self.headers,
)
self.collection_info = collection_details.json()
self.collection_info: dict[str, object] = collection_details.json()
print_verbose(f"Collection already exists.\nCollection details:{self.collection_info}")
self._ensure_cache_key_payload_index()
else:
quantization_params: dict[str, Any]
quantization_params: dict[str, dict[str, object]]
if quantization_config is None or quantization_config == "binary":
quantization_params = {
"binary": {
@ -214,7 +220,7 @@ class QdrantSemanticCache(BaseCache):
resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router),
)
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
def _get_embedding(self, prompt: str, metadata: dict[str, object] | None = None) -> EmbeddingResponse:
"""Embed via the proxy Router when it serves the model, else direct."""
try:
from litellm.proxy.proxy_server import llm_model_list, llm_router
@ -241,7 +247,7 @@ class QdrantSemanticCache(BaseCache):
num_retries=0,
)
async def _get_async_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
async def _get_async_embedding(self, prompt: str, metadata: dict[str, object] | None = None) -> EmbeddingResponse:
try:
from litellm.proxy.proxy_server import llm_model_list, llm_router
except ImportError:

View file

@ -45,14 +45,14 @@ class ResponsesToCompletionBridgeHandler:
return bool(stream)
@staticmethod
def _is_preformatted_cached_chat_stream(result: Any) -> bool:
def _is_preformatted_cached_chat_stream(result: object) -> bool:
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
return isinstance(result, CustomStreamWrapper) and result.custom_llm_provider == "cached_response"
@staticmethod
def _coerce_response_object(
response_obj: Any,
response_obj: object,
hidden_params: dict | None,
) -> "ResponsesAPIResponse":
if isinstance(response_obj, ResponsesAPIResponse):
@ -78,8 +78,8 @@ class ResponsesToCompletionBridgeHandler:
for _ in stream_iter:
pass
completed: Final = getattr(stream_iter, "completed_response", None)
response_obj: Final = getattr(completed, "response", None) if completed else None
completed: Final[object] = getattr(stream_iter, "completed_response", None)
response_obj: Final[object] = getattr(completed, "response", None) if completed else None
if response_obj is None:
raise ValueError("Stream ended without a completed response")
@ -93,8 +93,8 @@ class ResponsesToCompletionBridgeHandler:
async for _ in stream_iter:
pass
completed: Final = getattr(stream_iter, "completed_response", None)
response_obj: Final = getattr(completed, "response", None) if completed else None
completed: Final[object] = getattr(stream_iter, "completed_response", None)
response_obj: Final[object] = getattr(completed, "response", None) if completed else None
if response_obj is None:
raise ValueError("Stream ended without a completed response")
@ -157,7 +157,7 @@ class ResponsesToCompletionBridgeHandler:
def completion(
self, *args, **kwargs
) -> Union[
Coroutine[Any, Any, Union["ModelResponse", "CustomStreamWrapper"]],
Coroutine[None, None, Union["ModelResponse", "CustomStreamWrapper"]],
"ModelResponse",
"CustomStreamWrapper",
]:

View file

@ -52,10 +52,10 @@ class GenerateContentSetupResult(BaseModel):
model_config: ClassVar[ConfigDict] = ConfigDict(arbitrary_types_allowed=True)
model: str
request_body: dict[str, Any]
request_body: dict[str, object]
custom_llm_provider: str
generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig | None
generate_content_config_dict: dict[str, Any]
generate_content_config_dict: dict[str, object]
native_request_fields: dict[str, object]
litellm_params: GenericLiteLLMParams
litellm_logging_obj: LiteLLMLoggingObj
@ -68,7 +68,7 @@ class GenerateContentHelper:
@staticmethod
def mock_generate_content_response(
mock_response: str = "This is a mock response from Google GenAI generate_content.",
) -> dict[str, Any]:
) -> dict[str, object]:
"""Mock response for generate_content for testing purposes"""
return {
"text": mock_response,
@ -239,9 +239,9 @@ async def agenerate_content(
tools: ToolConfigDict | None = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
# LiteLLM specific params,
custom_llm_provider: str | None = None,
@ -307,9 +307,9 @@ def generate_content(
tools: ToolConfigDict | None = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
# LiteLLM specific params,
custom_llm_provider: str | None = None,
@ -397,9 +397,9 @@ async def agenerate_content_stream(
tools: ToolConfigDict | None = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
# LiteLLM specific params,
custom_llm_provider: str | None = None,
@ -492,9 +492,9 @@ def generate_content_stream(
tools: ToolConfigDict | None = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
# LiteLLM specific params,
custom_llm_provider: str | None = None,

View file

@ -3,7 +3,7 @@ import contextvars
import importlib
from collections.abc import Coroutine
from functools import partial
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast, overload
from typing import TYPE_CHECKING, Final, Literal, Optional, cast, overload
if TYPE_CHECKING:
from litellm.images.utils import ImageEditRequestUtils
@ -151,7 +151,7 @@ def image_generation(
*,
aimg_generation: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ImageResponse]:
) -> Coroutine[object, object, ImageResponse]:
...
@ -197,7 +197,7 @@ def image_generation(
api_version: str | None = None,
custom_llm_provider=None,
**kwargs,
) -> ImageResponse | Coroutine[Any, Any, ImageResponse]:
) -> ImageResponse | Coroutine[object, object, ImageResponse]:
"""
Maps the https://api.openai.com/v1/images/generations endpoint.
@ -723,14 +723,14 @@ def image_edit(
user: str | None = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
# LiteLLM specific params,
custom_llm_provider: str | None = None,
**kwargs,
) -> ImageResponse | Coroutine[Any, Any, ImageResponse]:
) -> ImageResponse | Coroutine[object, object, ImageResponse]:
"""
Maps the image edit functionality, similar to OpenAI's images/edits endpoint.
"""
@ -769,7 +769,7 @@ def image_edit(
images: Final = image if isinstance(image, list) else ([image] if image is not None else [])
headers_from_kwargs: Final = kwargs.get("headers")
merged_extra_headers: Final[dict[str, Any]] = {}
merged_extra_headers: Final[dict[str, object]] = {}
if isinstance(headers_from_kwargs, dict):
merged_extra_headers.update(headers_from_kwargs)
if isinstance(extra_headers, dict):
@ -974,9 +974,9 @@ async def aimage_edit(
user: str | None = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
# LiteLLM specific params,
custom_llm_provider: str | None = None,
@ -1044,7 +1044,7 @@ async def aimage_edit(
)
def __getattr__(name: str) -> Any:
def __getattr__(name: str) -> type["ImageEditRequestUtils"]:
"""Lazy import handler for images.main module"""
if name == "ImageEditRequestUtils":
# Lazy load ImageEditRequestUtils to avoid heavy import from images.utils at module load time

View file

@ -545,7 +545,6 @@ class SlackAlerting(CustomBatchLogger):
# Get the appropriate budget alert type handler
budget_alert_class: Final = get_budget_alert_type(type)
_id: Final = budget_alert_class.get_id(user_info)
user_info_json: Final = user_info.model_dump(exclude_none=True)
user_info_str: Final = self._get_user_info_str(user_info)
event_message = budget_alert_class.get_event_message()
@ -575,7 +574,22 @@ class SlackAlerting(CustomBatchLogger):
webhook_event = WebhookEvent(
event=event,
event_message=event_message,
**user_info_json,
spend=user_info.spend,
max_budget=user_info.max_budget,
soft_budget=user_info.soft_budget,
token=user_info.token,
customer_id=user_info.customer_id,
user_id=user_info.user_id,
team_id=user_info.team_id,
team_alias=user_info.team_alias,
organization_id=user_info.organization_id,
user_email=user_info.user_email,
key_alias=user_info.key_alias,
projected_exceeded_date=user_info.projected_exceeded_date,
projected_spend=user_info.projected_spend,
event_group=user_info.event_group,
alert_emails=user_info.alert_emails,
max_budget_alert_emails=user_info.max_budget_alert_emails,
)
await self.send_alert(
message=event_message + "\n\n" + user_info_str,
@ -657,7 +671,7 @@ class SlackAlerting(CustomBatchLogger):
"""
Create a standard message for a budget alert
"""
_all_fields_as_dict: Final = user_info.model_dump(exclude_none=True)
_all_fields_as_dict: Final[dict[str, object]] = user_info.model_dump(exclude_none=True)
_all_fields_as_dict.pop("token")
msg = ""
for k, v in _all_fields_as_dict.items():
@ -1006,7 +1020,7 @@ class SlackAlerting(CustomBatchLogger):
except Exception:
pass
async def model_added_alert(self, model_name: str, litellm_model_name: str, passed_model_info: Any):
async def model_added_alert(self, model_name: str, litellm_model_name: str, passed_model_info: object):
base_model_from_user: Final = getattr(passed_model_info, "base_model", None)
model_info = {}
base_model = ""
@ -1973,7 +1987,7 @@ Model Info:
try:
message = f"`{event_name}`\n"
key_event_dict: Final = key_event.model_dump()
key_event_dict: Final[dict[str, object]] = key_event.model_dump()
# Add Created by information first
message += "*Action Done by:*\n"

View file

@ -3,6 +3,7 @@ BitBucket prompt manager that integrates with LiteLLM's prompt management system
Fetches .prompt files from BitBucket repositories and provides team-based access control.
"""
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
from jinja2 import DictLoader, select_autoescape
@ -65,7 +66,7 @@ class BitBucketTemplateManager:
def __init__(
self,
bitbucket_config: dict[str, Any],
bitbucket_config: Mapping[str, object],
prompt_id: str | None = None,
):
self.bitbucket_config = bitbucket_config
@ -123,7 +124,7 @@ class BitBucketTemplateManager:
template_content = content
# Parse YAML frontmatter
metadata: dict[str, Any] = {}
metadata: dict[str, object] = {}
if frontmatter_str:
try:
import yaml
@ -141,9 +142,9 @@ class BitBucketTemplateManager:
metadata=metadata,
)
def _parse_yaml_basic(self, yaml_str: str) -> dict[str, Any]:
def _parse_yaml_basic(self, yaml_str: str) -> dict[str, object]:
"""Basic YAML parser for simple cases when PyYAML is not available."""
result: Final[dict[str, Any]] = {}
result: Final[dict[str, object]] = {}
for line in yaml_str.split("\n"):
line = line.strip()
if ":" in line and not line.startswith("#"):
@ -162,7 +163,7 @@ class BitBucketTemplateManager:
result[key] = value.strip("\"'")
return result
def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> str:
def render_template(self, template_id: str, variables: Mapping[str, object] | None = None) -> str:
"""Render a template with the given variables."""
if template_id not in self.prompts:
raise ValueError(f"Template '{template_id}' not found")
@ -209,7 +210,7 @@ class BitBucketPromptManager(CustomPromptManagement):
def __init__(
self,
bitbucket_config: dict[str, Any],
bitbucket_config: Mapping[str, object],
prompt_id: str | None = None,
):
self.bitbucket_config = bitbucket_config
@ -234,7 +235,7 @@ class BitBucketPromptManager(CustomPromptManagement):
def get_prompt_template(
self,
prompt_id: str,
prompt_variables: dict[str, Any] | None = None,
prompt_variables: Mapping[str, object] | None = None,
) -> tuple[str, dict[str, Any]]:
"""
Get a prompt template and render it with variables.
@ -267,12 +268,12 @@ class BitBucketPromptManager(CustomPromptManagement):
self,
user_id: str | None,
messages: list[AllMessageValues],
function_call: dict[str, Any] | str | None = None,
litellm_params: dict[str, Any] | None = None,
function_call: Mapping[str, object] | str | None = None,
litellm_params: dict[str, object] | None = None,
prompt_id: str | None = None,
prompt_variables: dict[str, Any] | None = None,
prompt_variables: Mapping[str, object] | None = None,
**kwargs,
) -> tuple[list[AllMessageValues], dict[str, Any] | None]:
) -> tuple[list[AllMessageValues], dict[str, object] | None]:
"""
Pre-call hook that processes the prompt template before making the LLM call.
"""
@ -316,9 +317,9 @@ class BitBucketPromptManager(CustomPromptManagement):
except Exception as e:
# Log error but don't fail the call
import litellm
from litellm._logging import verbose_proxy_logger
litellm._logging.verbose_proxy_logger.error("Error in BitBucket prompt pre_call_hook: %s", e)
verbose_proxy_logger.error("Error in BitBucket prompt pre_call_hook: %s", e)
return messages, litellm_params
def _parse_prompt_to_messages(self, prompt_content: str) -> list[AllMessageValues]:
@ -384,14 +385,14 @@ class BitBucketPromptManager(CustomPromptManagement):
def post_call_hook(
self,
user_id: str | None,
response: Any,
response: object,
input_messages: list[AllMessageValues],
function_call: dict[str, Any] | str | None = None,
litellm_params: dict[str, Any] | None = None,
function_call: Mapping[str, object] | str | None = None,
litellm_params: Mapping[str, object] | None = None,
prompt_id: str | None = None,
prompt_variables: dict[str, Any] | None = None,
prompt_variables: Mapping[str, object] | None = None,
**kwargs,
) -> Any:
) -> object:
"""
Post-call hook for any post-processing after the LLM call.
"""

View file

@ -19,14 +19,29 @@
"""Transform LiteLLM data to CloudZero AnyCost CBF format."""
from datetime import datetime
from typing import Any, Final
from typing import Final, SupportsFloat, SupportsIndex, SupportsInt
import polars as pl
from typing_extensions import Buffer
from ...types.integrations.cloudzero import CBFRecord
from .cz_resource_names import CZEntityType, CZRNGenerator
def _as_int(value: object) -> int:
"""The integer form of a spend table cell, computed the way :func:`int` computes it."""
if isinstance(value, (str, Buffer, SupportsInt, SupportsIndex)):
return int(value)
raise TypeError(f"int() argument must be a string or a number, not {type(value).__name__!r}")
def _as_float(value: object) -> float:
"""The floating point form of a spend table cell, computed the way :func:`float` computes it."""
if isinstance(value, (str, Buffer, SupportsFloat, SupportsIndex)):
return float(value)
raise TypeError(f"float() argument must be a string or a number, not {type(value).__name__!r}")
class CBFTransformer:
"""Transform LiteLLM usage data to CloudZero Billing Format (CBF)."""
@ -82,15 +97,15 @@ class CBFTransformer:
return pl.DataFrame(cbf_data)
def _create_cbf_record(self, row: dict[str, Any]) -> CBFRecord:
def _create_cbf_record(self, row: dict[str, object]) -> CBFRecord:
"""Create a single CBF record from LiteLLM daily spend row."""
# Parse date (daily spend tables use date strings like '2025-04-19')
usage_date: Final = self._parse_date(row.get("date"))
# Calculate total tokens
prompt_tokens: Final = int(row.get("prompt_tokens", 0))
completion_tokens: Final = int(row.get("completion_tokens", 0))
prompt_tokens: Final = _as_int(row.get("prompt_tokens", 0))
completion_tokens: Final = _as_int(row.get("completion_tokens", 0))
total_tokens: Final = prompt_tokens + completion_tokens
# Create CloudZero Resource Name (CZRN) as resource_id
@ -154,7 +169,7 @@ class CBFTransformer:
"time/usage_start": (
usage_date.isoformat() if usage_date else None
), # Required: ISO-formatted UTC datetime
"cost/cost": float(row.get("spend", 0.0)), # Required: billed cost
"cost/cost": _as_float(row.get("spend", 0.0)), # Required: billed cost
"resource/id": resource_id, # CZRN (CloudZero Resource Name)
# Usage metrics for token consumption
"usage/amount": total_tokens, # Numeric value of tokens consumed
@ -187,7 +202,7 @@ class CBFTransformer:
return CBFRecord(cbf_record)
def _parse_date(self, date_str) -> datetime | None:
def _parse_date(self, date_str: object) -> datetime | None:
"""Parse date string from daily spend tables (e.g., '2025-04-19')."""
if date_str is None:
return None

View file

@ -2,6 +2,7 @@ import contextvars
import hashlib
import os
import secrets
from collections.abc import Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args
@ -227,13 +228,13 @@ class CustomGuardrail(CustomLogger):
)
super().__init__(**kwargs)
def render_violation_message(self, default: str, context: dict[str, Any] | None = None) -> str:
def render_violation_message(self, default: str, context: Mapping[str, object] | None = None) -> str:
"""Return a custom violation message if template is configured."""
if not self.violation_message_template:
return default
format_context: Final[dict[str, Any]] = {"default_message": default}
format_context: Final[dict[str, object]] = {"default_message": default}
if context:
format_context.update(context)
try:
@ -661,7 +662,7 @@ class CustomGuardrail(CustomLogger):
value: Final = self._get_admin_metadata(data).get("opted_out_global_guardrails")
return value if isinstance(value, list) else []
def _is_valid_response_type(self, result: Any) -> bool:
def _is_valid_response_type(self, result: object) -> bool:
"""
Check if result is a valid LLMResponseTypes instance.
@ -722,7 +723,7 @@ class CustomGuardrail(CustomLogger):
return None
return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}"
def mark_pre_call_hook_ran(self, data: dict[str, Any]) -> None:
def mark_pre_call_hook_ran(self, data: dict[str, object]) -> None:
"""
Record that this guardrail's ``async_pre_call_hook`` already ran for this
request, so the deployment-level hook does not run it a second time.
@ -747,7 +748,7 @@ class CustomGuardrail(CustomLogger):
return
data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]}
def _pre_call_hook_already_ran(self, data: dict[str, Any]) -> bool:
def _pre_call_hook_already_ran(self, data: dict[str, object]) -> bool:
marker: Final = self._pre_call_marker()
if marker is None:
return False
@ -1170,7 +1171,7 @@ class CustomGuardrail(CustomLogger):
This gets logged on downsteam Langfuse, DataDog, etc.
"""
# Convert None to empty dict to satisfy type requirements
guardrail_response: dict[str, Any] | str = {} if response is None else response
guardrail_response: dict[str, object] | str = {} if response is None else response
# For apply_guardrail functions in custom_code_guardrail scenario,
# simplify the logged response to "allow", "deny", or "mask"

View file

@ -20,10 +20,11 @@ import time
import traceback
from collections.abc import Sequence
from datetime import datetime as datetimeObj
from typing import Any, Final
from typing import TYPE_CHECKING, Any, Final
import httpx
from httpx import Response
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
@ -62,6 +63,18 @@ from litellm.types.utils import StandardLoggingPayload
from ..additional_logging_utils import AdditionalLoggingUtils
if TYPE_CHECKING:
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
class _DatadogLoggingKwargs(TypedDict, total=False):
"""The subset of logging ``kwargs`` that the Datadog payload builder reads."""
standard_logging_object: ReadOnly[StandardLoggingPayload | None]
# max number of logs DD API can accept
@ -87,6 +100,11 @@ def _resolve_dd_batch_size() -> int:
return max(1, min(value, DD_MAX_BATCH_SIZE))
def _span_attribute(span: object, name: str) -> object:
"""Read an optional attribute off whatever span object the active tracer hands back."""
return getattr(span, name, None)
class DataDogLogger(
CustomBatchLogger,
AdditionalLoggingUtils,
@ -271,9 +289,9 @@ class DataDogLogger(
self,
request_data: dict,
original_exception: Exception,
user_api_key_dict: Any,
user_api_key_dict: "UserAPIKeyAuth",
traceback_str: str | None = None,
) -> Any | None:
) -> "HTTPException | None":
"""
Log proxy-level failures (e.g. 401 auth, DB connection errors) to Datadog.
@ -297,7 +315,7 @@ class DataDogLogger(
status_code = int(_code)
# Use project-standard sanitized user context when running in proxy
user_context: dict[str, Any] = {}
user_context: dict[str, object] = {}
try:
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
@ -553,8 +571,8 @@ class DataDogLogger(
def create_datadog_logging_payload(
self,
kwargs: dict | Any,
response_obj: Any,
kwargs: _DatadogLoggingKwargs,
response_obj: object,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> DatadogPayload:
@ -562,8 +580,8 @@ class DataDogLogger(
Helper function to create a datadog payload for logging
Args:
kwargs (Union[dict, Any]): request kwargs
response_obj (Any): llm api response
kwargs: request kwargs, read for its standard logging object
response_obj: llm api response
start_time (datetime.datetime): start time of request
end_time (datetime.datetime): end time of request
@ -625,7 +643,7 @@ class DataDogLogger(
self,
payload: ServiceLoggerPayload,
error: str | None = "",
parent_otel_span: Any | None = None,
parent_otel_span: object = None,
start_time: datetimeObj | float | None = None,
end_time: float | datetimeObj | None = None,
event_metadata: dict | None = None,
@ -659,7 +677,7 @@ class DataDogLogger(
self,
payload: ServiceLoggerPayload,
error: str | None = "",
parent_otel_span: Any | None = None,
parent_otel_span: object = None,
start_time: datetimeObj | float | None = None,
end_time: float | datetimeObj | None = None,
event_metadata: dict | None = None,
@ -696,7 +714,7 @@ class DataDogLogger(
def _create_v0_logging_payload(
self,
kwargs: dict | Any,
kwargs: dict,
response_obj: Any,
start_time: datetime.datetime,
end_time: datetime.datetime,
@ -810,11 +828,11 @@ class DataDogLogger(
if current_span is None:
return None
trace_id: Final = getattr(current_span, "trace_id", None)
trace_id: Final = _span_attribute(current_span, "trace_id")
if trace_id is None:
return None
span_id: Final = getattr(current_span, "span_id", None)
span_id: Final = _span_attribute(current_span, "span_id")
trace_context: Final[dict[str, str]] = {"trace_id": str(trace_id)}
if span_id is not None:
trace_context["span_id"] = str(span_id)

View file

@ -9,6 +9,7 @@ API Reference: https://docs.datadoghq.com/llm_observability/setup/api/?tab=examp
import asyncio
import json
import os
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any, Final, Literal
@ -334,7 +335,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
def _get_response_messages(
self, standard_logging_payload: StandardLoggingPayload, call_type: str | None
) -> list[Any]:
) -> list[object]:
"""
Get the messages from the response object
@ -484,7 +485,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
# Default fallback for unknown or passthrough operations
return "llm"
def _ensure_string_content(self, messages: str | list[Any] | dict[Any, Any] | None) -> list[Any]:
def _ensure_string_content(self, messages: str | Sequence[object] | Mapping[object, object] | None) -> list[object]:
if messages is None:
return []
if isinstance(messages, str):
@ -495,11 +496,11 @@ class DataDogLLMObsLogger(CustomBatchLogger):
return [str(messages.get("content", ""))]
return []
def _get_dd_llm_obs_payload_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, Any]:
def _get_dd_llm_obs_payload_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, object]:
"""
Fields to track in DD LLM Observability metadata from litellm standard logging payload
"""
_metadata: Final[dict[str, Any]] = {
_metadata: Final[dict[str, object]] = {
"model_name": standard_logging_payload.get("model", "unknown"),
"model_provider": standard_logging_payload.get("custom_llm_provider", "unknown"),
"id": standard_logging_payload.get("id", "unknown"),
@ -647,7 +648,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
return spend_metrics
def _process_input_messages_preserving_tool_calls(self, messages: list[Any]) -> list[dict[str, Any]]:
def _process_input_messages_preserving_tool_calls(self, messages: Sequence[object]) -> list[dict[str, object]]:
"""
Process input messages while preserving tool_calls and tool message types.
@ -671,13 +672,13 @@ class DataDogLLMObsLogger(CustomBatchLogger):
return processed
@staticmethod
def _tool_calls_kv_pair(tool_calls: list[dict[str, Any]]) -> dict[str, Any]:
def _tool_calls_kv_pair(tool_calls: list[dict[str, Any]]) -> dict[str, object]:
"""
Extract tool call information into key-value pairs for Datadog metadata.
Similar to OpenTelemetry's implementation but adapted for Datadog's format.
"""
kv_pairs: Final[dict[str, Any]] = {}
kv_pairs: Final[dict[str, object]] = {}
for idx, tool_call in enumerate(tool_calls):
try:
# Extract tool call ID
@ -712,11 +713,11 @@ class DataDogLLMObsLogger(CustomBatchLogger):
return kv_pairs
def _extract_tool_call_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, Any]:
def _extract_tool_call_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, object]:
"""
Extract tool call information from both input messages and response for Datadog metadata.
"""
tool_call_metadata: Final[dict[str, Any]] = {}
tool_call_metadata: Final[dict[str, object]] = {}
try:
# Extract tool calls from input messages

View file

@ -3,12 +3,21 @@ Based on Google's GenAI Kit dotprompt implementation: https://google.github.io/d
"""
import re
from collections.abc import Mapping
from pathlib import Path
from typing import Any, Final
import yaml
from jinja2 import DictLoader, select_autoescape
from jinja2.sandbox import ImmutableSandboxedEnvironment
from typing_extensions import NotRequired, ReadOnly, TypedDict
class _PromptFileJson(TypedDict):
"""JSON form of a .prompt file: rendered template text plus its frontmatter."""
content: ReadOnly[NotRequired[str]]
metadata: ReadOnly[NotRequired[dict[str, object]]]
def strip_version_suffix(prompt_id: str) -> str | None:
@ -167,7 +176,7 @@ class PromptManager:
template_id=prompt_id,
)
def _parse_frontmatter(self, content: str) -> tuple[dict[str, Any], str]:
def _parse_frontmatter(self, content: str) -> tuple[dict[str, object], str]:
"""Parse YAML frontmatter from prompt content."""
# Match YAML frontmatter between --- delimiters
frontmatter_pattern: Final = r"^---\s*\n(.*?)\n---\s*\n(.*)$"
@ -178,7 +187,7 @@ class PromptManager:
template_content = match.group(2)
try:
frontmatter = yaml.safe_load(frontmatter_yaml) or {}
frontmatter: dict[str, object] = yaml.safe_load(frontmatter_yaml) or {}
except yaml.YAMLError as e:
raise ValueError(f"Invalid YAML frontmatter: {e}")
else:
@ -191,7 +200,7 @@ class PromptManager:
def render(
self,
prompt_id: str,
prompt_variables: dict[str, Any] | None = None,
prompt_variables: Mapping[str, object] | None = None,
version: int | None = None,
) -> str:
"""
@ -231,7 +240,7 @@ class PromptManager:
except Exception as e:
raise ValueError(f"Error rendering template '{prompt_id}': {e}")
def _validate_input(self, variables: dict[str, Any], schema: dict[str, Any]) -> None:
def _validate_input(self, variables: Mapping[str, object], schema: Mapping[str, str]) -> None:
"""Basic validation of input variables against schema."""
for field_name, field_type in schema.items():
if field_name in variables:
@ -291,7 +300,7 @@ class PromptManager:
"""Get a list of all available prompt IDs."""
return list(self.prompts.keys())
def get_prompt_metadata(self, prompt_id: str) -> dict[str, Any] | None:
def get_prompt_metadata(self, prompt_id: str) -> dict[str, object] | None:
"""Get metadata for a specific prompt."""
template: Final = self.prompts.get(prompt_id)
return template.metadata if template else None
@ -302,12 +311,12 @@ class PromptManager:
if self.prompt_directory:
self._load_prompts()
def add_prompt(self, prompt_id: str, content: str, metadata: dict[str, Any] | None = None) -> None:
def add_prompt(self, prompt_id: str, content: str, metadata: dict[str, object] | None = None) -> None:
"""Add a prompt template programmatically."""
template: Final = PromptTemplate(content=content, metadata=metadata or {}, template_id=prompt_id)
self.prompts[prompt_id] = template
def prompt_file_to_json(self, file_path: str | Path) -> dict[str, Any]:
def prompt_file_to_json(self, file_path: str | Path) -> _PromptFileJson:
"""Convert a .prompt file to JSON format.
Args:
@ -324,7 +333,7 @@ class PromptManager:
return {"content": template_content.strip(), "metadata": frontmatter}
def json_to_prompt_file(self, prompt_data: dict[str, Any]) -> str:
def json_to_prompt_file(self, prompt_data: _PromptFileJson) -> str:
"""Convert JSON prompt data to .prompt file format.
Args:

View file

@ -6,10 +6,11 @@ import re
import uuid
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone, tzinfo
from typing import Any, Final, TypedDict, cast
from typing import Any, Final, Protocol, cast
import httpx
from pydantic import BaseModel, Field
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
@ -35,6 +36,34 @@ GALILEO_CLOUD_API_BASE_URL: Final = "https://api.galileo.ai"
GALILEO_MAX_IN_MEMORY_RECORDS: Final = 1000
class _GalileoLoginBody(TypedDict):
"""Decoded body of the Galileo login response."""
access_token: ReadOnly[str]
class _GalileoLoginResponse(Protocol):
"""The login call's HTTP response, read for the access token it carries."""
def json(self) -> _GalileoLoginBody: ...
class _JsonResponse(Protocol):
"""An HTTP response read only for whatever JSON body it decodes to."""
def json(self) -> object: ...
def _login_access_token(response: _GalileoLoginResponse) -> str:
"""Read the bearer token out of a Galileo login response body."""
return response.json()["access_token"]
def _decoded_body(response: _JsonResponse) -> object:
"""Decode a response body without asserting anything about its shape."""
return response.json()
class GalileoStandardLoggingFields(TypedDict, total=False):
call_type: str
model: str
@ -156,7 +185,7 @@ class GalileoObserve(CustomLogger):
},
)
galileo_login_response.raise_for_status()
access_token: Final = galileo_login_response.json()["access_token"]
access_token: Final = _login_access_token(galileo_login_response)
self.headers = {
"accept": "application/json",
"Content-Type": "application/json",
@ -421,7 +450,7 @@ class GalileoObserve(CustomLogger):
try:
verbose_logger.debug(
"Galileo Logger HTTP error response json: %s",
response.json(),
_decoded_body(response),
)
except Exception:
pass

View file

@ -4,12 +4,80 @@ Now supports selecting a tag via `config["tag"]`; falls back to branch ("main").
"""
import base64
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Any, Final, Protocol, TypedDict
from urllib.parse import quote
from typing_extensions import ReadOnly
from litellm.llms.custom_httpx.http_handler import HTTPHandler
class GitLabFilePayload(TypedDict, total=False):
"""A repository-files API entry."""
content: ReadOnly[str]
encoding: ReadOnly[str]
class GitLabTreeEntry(TypedDict, total=False):
"""A repository-tree API entry."""
path: ReadOnly[str]
type: ReadOnly[str]
class GitLabBranch(TypedDict, total=False):
"""A repository-branches API entry."""
name: ReadOnly[str]
type: ReadOnly[str]
class GitLabFileMetadata(TypedDict):
"""The response headers a raw file request exposes as metadata."""
content_type: ReadOnly[str | None]
content_length: ReadOnly[str | None]
last_modified: ReadOnly[str | None]
class _FileJsonResponse(Protocol):
def json(self) -> GitLabFilePayload: ...
class _TreeJsonResponse(Protocol):
def json(self) -> Sequence[GitLabTreeEntry] | None: ...
class _ProjectJsonResponse(Protocol):
def json(self) -> Mapping[str, object]: ...
class _BranchesJsonResponse(Protocol):
def json(self) -> Sequence[GitLabBranch] | None: ...
def _file_payload(resp: _FileJsonResponse) -> GitLabFilePayload:
"""The JSON body of a repository-files response."""
return resp.json()
def _tree_entries(resp: _TreeJsonResponse) -> Sequence[GitLabTreeEntry]:
"""The entries of a repository-tree response."""
return resp.json() or []
def _project_info(resp: _ProjectJsonResponse) -> Mapping[str, object]:
"""The JSON body of a project response."""
return resp.json()
def _branch_entries(resp: _BranchesJsonResponse) -> Sequence[GitLabBranch] | None:
"""The JSON body of a repository-branches response."""
return resp.json()
class GitLabClient:
"""
Client for interacting with the GitLab API to fetch files.
@ -42,12 +110,12 @@ class GitLabClient:
self.project: str | int = project
self.access_token: str = str(access_token)
self.auth_method = config.get("auth_method", "token") # 'token' or 'oauth'
self.auth_method: str = config.get("auth_method", "token") # 'token' or 'oauth'
self.branch = config.get("branch", None)
if not self.branch:
self.branch = "main"
self.tag = config.get("tag")
self.base_url = config.get("base_url", "https://gitlab.com/api/v4")
self.base_url: str = config.get("base_url", "https://gitlab.com/api/v4")
if not all([self.project, self.access_token]):
raise ValueError("project and access_token are required")
@ -159,7 +227,7 @@ class GitLabClient:
if resp.status_code == 404:
return None
resp.raise_for_status()
data: Final = resp.json()
data: Final = _file_payload(resp)
content: Final = data.get("content")
encoding: Final = data.get("encoding", "")
if content and encoding == "base64":
@ -208,7 +276,7 @@ class GitLabClient:
return []
resp.raise_for_status()
data: Final = resp.json() or []
data: Final = _tree_entries(resp)
files: Final[list[str]] = []
for item in data:
if item.get("type") == "blob":
@ -229,13 +297,13 @@ class GitLabClient:
raise Exception("Authentication failed. Check your GitLab token and auth_method.")
raise Exception(f"Failed to list files in '{directory_path}': {e}")
def get_repository_info(self) -> dict[str, Any]:
def get_repository_info(self) -> Mapping[str, object]:
"""Get information about the project/repository."""
url: Final = f"{self.base_url}/projects/{self._project_enc}"
try:
resp: Final = self.http_handler.get(url, headers=self.headers)
resp.raise_for_status()
return resp.json()
return _project_info(resp)
except Exception as e:
raise Exception(f"Failed to get repository info: {e}")
@ -247,18 +315,18 @@ class GitLabClient:
except Exception:
return False
def get_branches(self) -> list[dict[str, Any]]:
def get_branches(self) -> list[GitLabBranch]:
"""Get list of branches in the repository."""
url: Final = f"{self.base_url}/projects/{self._project_enc}/repository/branches"
try:
resp: Final = self.http_handler.get(url, headers=self.headers)
resp.raise_for_status()
data: Final = resp.json()
data: Final = _branch_entries(resp)
return data if isinstance(data, list) else []
except Exception as e:
raise Exception(f"Failed to get branches: {e}")
def get_file_metadata(self, file_path: str, *, ref: str | None = None) -> dict[str, Any] | None:
def get_file_metadata(self, file_path: str, *, ref: str | None = None) -> GitLabFileMetadata | None:
"""
Get minimal metadata about a file via RAW endpoint headers at a given ref.

View file

@ -89,7 +89,7 @@ def _extract_cache_read_input_tokens(usage_obj) -> int:
# Check prompt_tokens_details.cached_tokens (used by Gemini and other providers)
if hasattr(usage_obj, "prompt_tokens_details"):
prompt_tokens_details: Final = getattr(usage_obj, "prompt_tokens_details", None)
prompt_tokens_details: Final[object] = getattr(usage_obj, "prompt_tokens_details", None)
if prompt_tokens_details is not None and hasattr(prompt_tokens_details, "cached_tokens"):
cached_tokens: Final = getattr(prompt_tokens_details, "cached_tokens", None)
if cached_tokens is not None and isinstance(cached_tokens, (int, float)) and cached_tokens > 0:
@ -623,9 +623,16 @@ class LangFuseLogger:
)
# Apply custom masking function if provided
if masking_function is not None and callable(masking_function):
input = self._apply_masking_function(input, masking_function)
output = self._apply_masking_function(output, masking_function)
masked_input: Final[object] = (
self._apply_masking_function(input, masking_function)
if masking_function is not None and callable(masking_function)
else input
)
masked_output: Final[object] = (
self._apply_masking_function(output, masking_function)
if masking_function is not None and callable(masking_function)
else output
)
clean_metadata = redact_user_api_key_info(metadata=clean_metadata)
@ -651,15 +658,15 @@ class LangFuseLogger:
# Special keys that are found in the function arguments and not the metadata
if "input" in update_trace_keys:
trace_params["input"] = input if not mask_input else "redacted-by-litellm"
trace_params["input"] = masked_input if not mask_input else "redacted-by-litellm"
if "output" in update_trace_keys:
trace_params["output"] = output if not mask_output else "redacted-by-litellm"
trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm"
else: # don't overwrite an existing trace
trace_params = {
"id": trace_id,
"name": trace_name,
"session_id": session_id,
"input": input if not mask_input else "redacted-by-litellm",
"input": masked_input if not mask_input else "redacted-by-litellm",
"version": clean_metadata.pop(
"trace_version", clean_metadata.get("version", None)
), # If provided just version, it will applied to the trace as well, if applied a trace version it will take precedence
@ -669,9 +676,9 @@ class LangFuseLogger:
trace_params[key.replace("trace_", "")] = clean_metadata.pop(key, None)
if level == "ERROR":
trace_params["status_message"] = output
trace_params["status_message"] = masked_output
else:
trace_params["output"] = output if not mask_output else "redacted-by-litellm"
trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm"
if debug is True or (isinstance(debug, str) and debug.lower() == "true"):
debug_metadata: Final = {
@ -708,7 +715,7 @@ class LangFuseLogger:
("aws_region_name", aws_region_name, bool(aws_region_name)),
("cache_hit", kwargs.get("cache_hit") or False, self._supports_tags() and "cache_hit" in kwargs),
)
enrichments: Final[Mapping[str, Any]] = {
enrichments: Final[Mapping[str, object]] = {
key: value for key, value, include in candidate_enrichments if include
}
@ -802,8 +809,8 @@ class LangFuseLogger:
"end_time": end_time,
"model": model_name,
"model_parameters": optional_params,
"input": input if not mask_input else "redacted-by-litellm",
"output": output if not mask_output else "redacted-by-litellm",
"input": masked_input if not mask_input else "redacted-by-litellm",
"output": masked_output if not mask_output else "redacted-by-litellm",
"usage": usage,
"usage_details": usage_details,
"metadata": {
@ -825,8 +832,8 @@ class LangFuseLogger:
prompt_management_metadata=prompt_management_metadata,
langfuse_client=self.Langfuse,
)
if output is not None and isinstance(output, str) and level == "ERROR":
generation_params["status_message"] = output
if masked_output is not None and isinstance(masked_output, str) and level == "ERROR":
generation_params["status_message"] = masked_output
if self._supports_completion_start_time():
generation_params["completion_start_time"] = kwargs.get("completion_start_time", None)
@ -935,7 +942,7 @@ class LangFuseLogger:
return Version(self.langfuse_sdk_version) >= Version("2.7.3")
@staticmethod
def _apply_masking_function(data: Any, masking_function: Callable[[Any], Any]) -> Any:
def _apply_masking_function(data: object, masking_function: Callable[[object], object]) -> object:
"""
Apply a masking function to data, handling different data types.
@ -1049,7 +1056,7 @@ def _add_prompt_to_generation_params(
generation_params: dict,
clean_metadata: dict,
prompt_management_metadata: StandardLoggingPromptManagementMetadata | None,
langfuse_client: Any,
langfuse_client: object,
) -> dict:
from langfuse import Langfuse
from langfuse.model import (

View file

@ -4,9 +4,12 @@ Opik Logger that logs LLM events to an Opik server
import asyncio
import traceback
from collections.abc import Mapping
from datetime import datetime
from typing import Any, Final
from typing_extensions import ReadOnly, TypedDict, Unpack
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.llms.custom_httpx.http_handler import (
@ -23,7 +26,7 @@ except Exception:
opik_client = None
def _should_skip_event(kwargs: dict[str, Any]) -> bool:
def _should_skip_event(kwargs: Mapping[str, object]) -> bool:
"""Check if event should be skipped due to missing standard_logging_object."""
if kwargs.get("standard_logging_object") is None:
verbose_logger.debug("OpikLogger skipping event; no standard_logging_object found")
@ -31,12 +34,24 @@ def _should_skip_event(kwargs: dict[str, Any]) -> bool:
return False
class _OpikLoggerKwargs(TypedDict, total=False):
"""Constructor options accepted by ``OpikLogger``."""
project_name: ReadOnly[str | None]
url: ReadOnly[str | None]
api_key: ReadOnly[str | None]
workspace: ReadOnly[str | None]
batch_size: ReadOnly[int | None]
flush_interval: ReadOnly[int | None]
max_queue_size: ReadOnly[int | None]
class OpikLogger(CustomBatchLogger):
"""
Opik Logger for logging events to an Opik Server
"""
def __init__(self, **kwargs: Any) -> None:
def __init__(self, **kwargs: Unpack[_OpikLoggerKwargs]) -> None:
self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
self.sync_httpx_client = _get_httpx_client()
@ -95,7 +110,7 @@ class OpikLogger(CustomBatchLogger):
async def async_log_success_event(
self,
kwargs: dict[str, Any],
kwargs: dict[str, object],
response_obj: Any,
start_time: datetime,
end_time: datetime,
@ -163,7 +178,7 @@ class OpikLogger(CustomBatchLogger):
except Exception as e:
verbose_logger.exception("OpikLogger failed to log success event - %s\n%s", e, traceback.format_exc())
def _sync_send(self, url: str, headers: dict[str, str], batch: dict[str, Any]) -> None:
def _sync_send(self, url: str, headers: dict[str, str], batch: dict[str, object]) -> None:
try:
response: Final = self.sync_httpx_client.post(
url=url,
@ -178,7 +193,7 @@ class OpikLogger(CustomBatchLogger):
def log_success_event(
self,
kwargs: dict[str, Any],
kwargs: dict[str, object],
response_obj: Any,
start_time: datetime,
end_time: datetime,
@ -247,7 +262,7 @@ class OpikLogger(CustomBatchLogger):
except Exception as e:
verbose_logger.exception("OpikLogger failed to log success event - %s\n%s", e, traceback.format_exc())
async def _submit_batch(self, url: str, headers: dict[str, str], batch: dict[str, Any]) -> None:
async def _submit_batch(self, url: str, headers: dict[str, str], batch: dict[str, object]) -> None:
try:
response: Final = await self.async_httpx_client.post(
url=url,

View file

@ -1,6 +1,7 @@
"""Data extraction functions for Opik payload building."""
import json
from collections.abc import Mapping
from typing import Any, Final
from litellm import _logging
@ -35,8 +36,8 @@ def normalize_provider_name(provider: str | None) -> str | None:
def extract_opik_metadata(
litellm_metadata: dict[str, Any],
standard_logging_metadata: dict[str, Any],
litellm_metadata: Mapping[str, Any],
standard_logging_metadata: Mapping[str, Any],
) -> dict[str, Any]:
"""
Merge Opik metadata from three sources in increasing priority order:
@ -97,7 +98,7 @@ def extract_span_identifiers(
def extract_tags(
opik_metadata: dict[str, Any],
opik_metadata: Mapping[str, Any],
custom_llm_provider: str | None,
) -> list[str]:
"""
@ -122,7 +123,7 @@ def apply_proxy_header_overrides(
project_name: str,
tags: list[str],
thread_id: str | None,
proxy_headers: dict[str, Any],
proxy_headers: Mapping[str, str],
) -> tuple[str, list[str], str | None]:
"""
Apply overrides from proxy request headers (opik_* prefix).
@ -148,7 +149,7 @@ def apply_proxy_header_overrides(
thread_id = value
elif param_key == "tags":
try:
parsed_tags = json.loads(value)
parsed_tags: object = json.loads(value)
if isinstance(parsed_tags, list):
tags.extend(parsed_tags)
except (json.JSONDecodeError, TypeError):
@ -158,11 +159,11 @@ def apply_proxy_header_overrides(
def extract_and_build_metadata(
opik_metadata: dict[str, Any],
standard_logging_metadata: dict[str, Any],
standard_logging_object: dict[str, Any],
litellm_kwargs: dict[str, Any],
) -> dict[str, Any]:
opik_metadata: Mapping[str, object],
standard_logging_metadata: Mapping[str, object],
standard_logging_object: Mapping[str, object],
litellm_kwargs: Mapping[str, object],
) -> dict[str, object]:
"""
Build the complete metadata dictionary from all available sources.

View file

@ -11,9 +11,10 @@ identical metrics. The attribute cardinality filter is reused from v1 by import
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime
from typing import Any, Final, TypeAlias
from typing import Any, Final, Literal, Protocol, TypeAlias
from opentelemetry.metrics import Histogram, Meter
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
@ -151,6 +152,29 @@ METRIC_ATTRIBUTE_CEILING: Final[frozenset[str]] = frozenset(
BOUNDED_HIDDEN_PARAM_KEYS: Final[tuple[str, ...]] = ("model_id",)
class _TokenUsage(TypedDict, total=False):
"""The token counts a response's ``usage`` carries, as the recorder reads them."""
prompt_tokens: ReadOnly[int]
completion_tokens: ReadOnly[int]
class _ResponseView(Protocol):
"""The one read the recorder makes on a litellm response object."""
def get(self, key: Literal["usage"], /) -> _TokenUsage | None: ...
class _MetricKwargs(TypedDict, total=False):
"""The logging kwargs the recorder reads directly."""
call_type: ReadOnly[str | None]
litellm_params: ReadOnly[Mapping[str, object] | None]
response_cost: ReadOnly[float | None]
completion_start_time: ReadOnly[datetime | float | str | None]
api_call_start_time: ReadOnly[datetime | float | str | None]
def resolve_error_type(kwargs: Mapping[str, Any]) -> str:
"""The ``error.type`` value for a failed request.
@ -192,8 +216,8 @@ class GenAIMetricRecorder:
def record(
self,
kwargs: Mapping[str, Any],
response_obj: Any,
kwargs: _MetricKwargs,
response_obj: _ResponseView | None,
start_time: datetime,
end_time: datetime,
) -> None:
@ -218,7 +242,7 @@ class GenAIMetricRecorder:
def record_failure(
self,
kwargs: Mapping[str, Any],
kwargs: _MetricKwargs,
start_time: datetime,
end_time: datetime,
) -> None:
@ -342,7 +366,7 @@ class GenAIMetricRecorder:
# Per-metric recording
# ------------------------------------------------------------------ #
def _record_token_usage(self, response_obj: Any, common_attrs: dict) -> None:
def _record_token_usage(self, response_obj: _ResponseView | None, common_attrs: dict) -> None:
if not response_obj:
return
usage: Final = response_obj.get("usage")
@ -353,7 +377,7 @@ class GenAIMetricRecorder:
self._metrics.token_usage.record(usage.get("prompt_tokens", 0), attributes=in_attrs)
self._metrics.token_usage.record(usage.get("completion_tokens", 0), attributes=out_attrs)
def _record_time_to_first_token(self, kwargs: Mapping[str, Any], common_attrs: dict) -> None:
def _record_time_to_first_token(self, kwargs: _MetricKwargs, common_attrs: dict) -> None:
time_to_first_chunk: Final = time_to_first_chunk_seconds(kwargs)
if time_to_first_chunk is None:
return
@ -361,15 +385,14 @@ class GenAIMetricRecorder:
def _record_time_per_output_token(
self,
kwargs: Mapping[str, Any],
response_obj: Any,
kwargs: _MetricKwargs,
response_obj: _ResponseView | None,
end_time: datetime,
duration_s: float,
common_attrs: dict,
) -> None:
completion_tokens = None
if response_obj and (usage := response_obj.get("usage")):
completion_tokens = usage.get("completion_tokens")
usage: Final = response_obj.get("usage") if response_obj else None
completion_tokens: Final = usage.get("completion_tokens") if usage else None
if completion_tokens is None or completion_tokens <= 0:
return

View file

@ -13,7 +13,7 @@ from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.utils import StandardCallbackDynamicParams
from litellm.types.utils import CallTypes, StandardCallbackDynamicParams
from litellm.types.vector_stores import (
LiteLLM_ManagedVectorStore,
VectorStoreResultContent,
@ -226,7 +226,7 @@ class VectorStorePreCallHook(CustomLogger):
self,
request_data: dict,
response: Any,
call_type: Any | None,
call_type: CallTypes | None,
) -> Any | None:
"""
Add search results to the response after successful LLM call.
@ -283,7 +283,7 @@ class VectorStorePreCallHook(CustomLogger):
self,
request_data: dict,
response_chunk: Any,
call_type: Any | None,
call_type: CallTypes | None,
) -> Any | None:
"""
Add search results to the final streaming chunk.

View file

@ -1500,6 +1500,6 @@ class RealTimeStreaming:
pass
def client_sent_openai_beta_realtime_header(websocket: Any) -> bool:
def client_sent_openai_beta_realtime_header(websocket: _ScopedWebSocket) -> bool:
"""True when the client WebSocket includes ``OpenAI-Beta: realtime=v1``."""
return RealTimeStreaming._detect_beta_header(websocket)

View file

@ -73,6 +73,18 @@ class _ContentChunk(TypedDict):
choices: Sequence[_ContentChoice]
class _FunctionCallDelta(TypedDict):
function_call: ReadOnly[FunctionCall]
class _FunctionCallChoice(TypedDict):
delta: ReadOnly[_FunctionCallDelta]
class _FunctionCallChunk(TypedDict):
choices: ReadOnly[Sequence[_FunctionCallChoice]]
class _AudioDelta(TypedDict, total=False):
audio: ChatCompletionAudioDelta | None
@ -588,7 +600,7 @@ class ChunkProcessor:
return tool_calls_list
def get_combined_function_call_content(self, function_call_chunks: list[dict[str, Any]]) -> FunctionCall:
def get_combined_function_call_content(self, function_call_chunks: Sequence["_FunctionCallChunk"]) -> FunctionCall:
argument_list: Final = []
delta = function_call_chunks[0]["choices"][0]["delta"]
function_call = delta.get("function_call", "")

View file

@ -11,8 +11,11 @@ A2A Protocol Format:
"""
import json
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final, Optional
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.types.utils import GenericGuardrailAPIInputs
@ -23,6 +26,13 @@ if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
class _A2ATextPart(TypedDict, total=False):
"""The subset of an A2A message part this handler reads text from."""
kind: ReadOnly[str]
text: ReadOnly[str]
class A2AGuardrailHandler(BaseTranslation):
"""
Handler for processing A2A Protocol messages with guardrails.
@ -41,7 +51,7 @@ class A2AGuardrailHandler(BaseTranslation):
data: dict,
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> Any:
) -> dict:
"""
Process A2A input messages by applying guardrails to text content.
@ -214,12 +224,12 @@ class A2AGuardrailHandler(BaseTranslation):
async def process_output_streaming_response(
self,
responses_so_far: list[Any],
responses_so_far: list[object],
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
request_data: dict | None = None,
) -> list[Any]:
) -> list[object]:
"""
Process A2A streaming output by applying guardrails to accumulated text.
@ -305,11 +315,12 @@ class A2AGuardrailHandler(BaseTranslation):
def _parse_streaming_responses(
self,
responses_so_far: list[Any],
) -> tuple[list[dict[str, Any] | None], list[tuple[int, dict[str, Any]]]]:
responses_so_far: list[object],
) -> tuple[list[dict[str, object] | None], list[tuple[int, dict[str, object]]]]:
"""Parse JSON-RPC items, returning aligned parsed list and valid entries."""
parsed: Final[list[dict[str, Any] | None]] = [None] * len(responses_so_far)
parsed: Final[list[dict[str, object] | None]] = [None] * len(responses_so_far)
for i, item in enumerate(responses_so_far):
obj: dict[str, object]
if isinstance(item, dict):
obj = item
elif isinstance(item, str):
@ -326,7 +337,7 @@ class A2AGuardrailHandler(BaseTranslation):
def _collect_text_from_parsed_chunks(
self,
valid_parsed: list[tuple[int, dict[str, Any]]],
valid_parsed: list[tuple[int, dict[str, object]]],
) -> tuple[str, list[int]]:
"""Collect text from parsed chunks, returning combined text and indices."""
from litellm.llms.a2a.common_utils import extract_text_from_a2a_response
@ -411,7 +422,7 @@ class A2AGuardrailHandler(BaseTranslation):
def _extract_texts_from_parts(
self,
parts: list[dict[str, Any]],
parts: Sequence[_A2ATextPart],
path: tuple[str, ...],
texts_to_check: list[str],
task_mappings: list[tuple[tuple[str, ...], int]],

View file

@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any, Final, NoReturn, cast
import httpx
from pydantic import ValidationError
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm.constants import (
@ -125,7 +126,25 @@ else:
_ANTHROPIC_TOOL_NAME_INVALID_CHARS: Final = re.compile(r"[^a-zA-Z0-9_-]")
_ANTHROPIC_TOOL_NAME_MAX_LEN: Final = 128
_ENUM_TYPE_CHECKS: Final[Mapping[str, Callable[[Any], bool]]] = MappingProxyType(
class _AnthropicUsageIteration(TypedDict, total=False):
"""One entry of the ``usage.iterations`` array on an Anthropic response."""
input_tokens: ReadOnly[int | None]
output_tokens: ReadOnly[int | None]
cache_creation_input_tokens: ReadOnly[int | None]
cache_read_input_tokens: ReadOnly[int | None]
class _AnthropicToolResultBlock(TypedDict, total=False):
"""A ``*_tool_result`` content block on an Anthropic response."""
type: ReadOnly[str]
tool_use_id: ReadOnly[str]
content: ReadOnly[object]
_ENUM_TYPE_CHECKS: Final[Mapping[str, Callable[[object], bool]]] = MappingProxyType(
{
"null": lambda v: v is None,
"boolean": lambda v: isinstance(v, bool),
@ -440,7 +459,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params.pop("speed", None)
@staticmethod
def _raise_invalid_reasoning_effort(model: str, value: Any, llm_provider: str) -> NoReturn:
def _raise_invalid_reasoning_effort(model: str, value: object, llm_provider: str) -> NoReturn:
"""Raise a ``BadRequestError`` for an unrecognised ``reasoning_effort``.
Args:
@ -2059,22 +2078,22 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
self, completion_response: dict
) -> tuple[
str,
list[Any] | None,
list[object] | None,
list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None,
str | None,
list[ChatCompletionToolCallChunk],
list[Any] | None,
list[Any] | None,
list[Any] | None,
list[object] | None,
list[_AnthropicToolResultBlock] | None,
list[object] | None,
]:
text_content = ""
citations: list[Any] | None = None
citations: list[object] | None = None
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
reasoning_content: str | None = None
tool_calls: Final[list[ChatCompletionToolCallChunk]] = []
web_search_results: list[Any] | None = None
tool_results: list[Any] | None = None
compaction_blocks: list[Any] | None = None
web_search_results: list[object] | None = None
tool_results: list[_AnthropicToolResultBlock] | None = None
compaction_blocks: list[object] | None = None
for idx, content in enumerate(completion_response["content"]):
if content["type"] == "text":
text_content += content["text"]
@ -2284,7 +2303,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
raw_speed: Final = _usage.get("speed")
resolved_speed: Final = raw_speed if isinstance(raw_speed, str) else speed
iterations: Final[list[Any] | None] = _usage.get("iterations")
iterations: Final[Sequence[_AnthropicUsageIteration] | None] = _usage.get("iterations")
if iterations:
prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations)
completion_tokens = sum(it.get("output_tokens", 0) or 0 for it in iterations)
@ -2377,7 +2396,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _build_code_interpreter_results(
self,
tool_results: list[Any],
tool_results: Sequence[_AnthropicToolResultBlock],
code_by_id: dict[str, str],
container_id: str | None,
) -> list[OutputCodeInterpreterCall]:
@ -2403,11 +2422,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def _build_provider_specific_fields(
self,
completion_response: dict,
citations: list[Any] | None,
citations: Sequence[object] | None,
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None,
web_search_results: list[Any] | None,
tool_results: list[Any] | None,
compaction_blocks: list[Any] | None,
web_search_results: Sequence[object] | None,
tool_results: Sequence[_AnthropicToolResultBlock] | None,
compaction_blocks: Sequence[object] | None,
tool_calls: list[ChatCompletionToolCallChunk],
) -> dict[str, Any]:
provider_specific_fields: Final[dict[str, Any]] = {

View file

@ -2,7 +2,7 @@ import asyncio
import json
import time
from collections.abc import Coroutine
from typing import Any, Final
from typing import Final
import httpx
@ -116,7 +116,7 @@ class AnthropicFilesHandler:
api_key: str | None = None,
timeout: float | httpx.Timeout = 600.0,
max_retries: int | None = None,
) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]:
) -> HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent]:
"""
Retrieve file content from Anthropic.

View file

@ -2,7 +2,7 @@ import asyncio
import json
import time
from collections.abc import Callable, Coroutine
from typing import Any, Final
from typing import Final
import httpx
from openai import (
@ -374,7 +374,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
except Exception as e:
status_code: Final = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_response: Final = getattr(e, "response", None)
error_response: Final[object] = getattr(e, "response", None)
error_body: Final = getattr(e, "body", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
@ -392,7 +392,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
model: str,
api_base: str,
data: dict,
timeout: Any,
timeout: float | httpx.Timeout,
dynamic_params: bool,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
@ -502,7 +502,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
dynamic_params: bool,
data: dict[str, object],
model: str,
timeout: Any,
timeout: float | httpx.Timeout,
max_retries: int,
azure_ad_token: str | None = None,
azure_ad_token_provider: Callable | None = None,
@ -578,7 +578,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
dynamic_params: bool,
data: dict,
model: str,
timeout: Any,
timeout: float | httpx.Timeout,
max_retries: int,
azure_ad_token: str | None = None,
azure_ad_token_provider: Callable | None = None,
@ -634,7 +634,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
except Exception as e:
status_code: Final = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
error_response: Final = getattr(e, "response", None)
error_response: Final[object] = getattr(e, "response", None)
message: Final = getattr(e, "message", str(e))
error_body: Final = getattr(e, "body", None)
if error_headers is None and error_response:
@ -754,7 +754,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
aembedding=None,
headers: dict | None = None,
litellm_params: dict | None = None,
) -> EmbeddingResponse | Coroutine[Any, Any, EmbeddingResponse]:
) -> EmbeddingResponse | Coroutine[object, object, EmbeddingResponse]:
if headers:
optional_params["extra_headers"] = headers
if self._client_session is None:
@ -1268,7 +1268,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
headers["Authorization"] = f"Bearer {azure_ad_token}"
# init AzureOpenAI Client
azure_client_params: Final[dict[str, Any]] = self.initialize_azure_sdk_client(
azure_client_params: Final[dict[str, object]] = self.initialize_azure_sdk_client(
litellm_params=litellm_params or {},
api_key=api_key,
model_name=model or "",

View file

@ -51,15 +51,13 @@ else:
AsyncHTTPHandler = Any
class _AzureRawAnnotation(TypedDict, total=False):
type: ReadOnly[str]
class _AzureRawAnnotation(ChatCompletionAnnotation, total=False):
text: ReadOnly[str]
start_index: ReadOnly[int]
end_index: ReadOnly[int]
url_citation: ReadOnly[ChatCompletionAnnotationURLCitation]
_TransformedAnnotation: TypeAlias = ChatCompletionAnnotation | _AzureRawAnnotation
_TransformedAnnotation: TypeAlias = ChatCompletionAnnotation
class _AzureText(TypedDict, total=False):
@ -223,18 +221,11 @@ class AzureAIAgentsHandler:
"""Build the ModelResponse from agent output."""
from litellm.types.utils import Choices, Message, Usage
message_kwargs: Final[dict[str, Any]] = {
"content": content,
"role": "assistant",
}
if annotations:
message_kwargs["annotations"] = annotations
model_response.choices = [
Choices(
finish_reason="stop",
index=0,
message=Message(**message_kwargs),
message=Message(content=content, role="assistant", annotations=annotations or None),
)
]
model_response.model = model
@ -655,9 +646,6 @@ class AzureAIAgentsHandler:
if data_str == "[DONE]":
# Send final chunk with finish_reason
final_delta_kwargs: dict[str, Any] = {"content": None}
if collected_annotations:
final_delta_kwargs["annotations"] = collected_annotations
final_chunk = ModelResponseStream(
id=response_id,
created=created,
@ -667,7 +655,7 @@ class AzureAIAgentsHandler:
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(**final_delta_kwargs),
delta=Delta(content=None, annotations=collected_annotations or None),
)
],
)

View file

@ -7,7 +7,7 @@ Transforms between OpenAI Realtime API format and Bedrock Nova Sonic format.
import base64
import json
import uuid as uuid_lib
from typing import Any, Final, cast
from typing import Final, cast
from pydantic import BaseModel
@ -633,7 +633,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
List of Bedrock format messages (JSON strings)
"""
try:
json_message: Final = json.loads(message)
json_message: Final[dict[str, object]] = json.loads(message)
except json.JSONDecodeError:
verbose_logger.warning("Invalid JSON message: %s", message[:200])
return []
@ -1182,7 +1182,7 @@ class BedrockRealtimeConfig(BaseRealtimeConfig):
# Create a function call arguments done event
# This is a custom event format that matches what clients expect
function_call_event: Final[dict[str, Any]] = {
function_call_event: Final[dict[str, object]] = {
"type": "response.function_call_arguments.done",
"event_id": f"event_{uuid.uuid4()}",
"response_id": current_response_id,

View file

@ -8,9 +8,11 @@ then we poll until the result is ready.
import asyncio
import time
from typing import Any, Final
from collections.abc import Coroutine, Mapping
from typing import Final, Protocol
import httpx
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
@ -33,6 +35,42 @@ from ..common_utils import (
from .transformation import BlackForestLabsImageEditConfig
class _BFLSubmitBody(TypedDict, total=False):
"""Decoded body of the BFL submit response, which hands back a polling URL."""
errors: ReadOnly[object]
polling_url: ReadOnly[str]
class _BFLPollBody(TypedDict, total=False):
"""Decoded body of a BFL polling response."""
status: ReadOnly[str]
class _BFLSubmitResponse(Protocol):
"""The submit call's HTTP response, read for its status, body text and decoded body."""
@property
def status_code(self) -> int: ...
@property
def text(self) -> str: ...
def json(self) -> _BFLSubmitBody: ...
class _BFLPollResponse(Protocol):
"""A polling call's HTTP response, read only for the task status it carries."""
def json(self) -> _BFLPollBody: ...
def _poll_status(response: _BFLPollResponse) -> str | None:
"""Read the task status out of a BFL polling response body."""
return response.json().get("status")
class BlackForestLabsImageEdit:
"""
Black Forest Labs Image Edit handler.
@ -53,10 +91,10 @@ class BlackForestLabsImageEdit:
litellm_params: GenericLiteLLMParams | dict,
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout | None,
extra_headers: dict[str, Any] | None = None,
extra_headers: Mapping[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
aimage_edit: bool = False,
) -> ImageResponse | Any:
) -> ImageResponse | Coroutine[object, object, ImageResponse]:
"""
Main entry point for image edit requests.
@ -185,7 +223,7 @@ class BlackForestLabsImageEdit:
litellm_params: GenericLiteLLMParams | dict,
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout | None,
extra_headers: dict[str, Any] | None = None,
extra_headers: Mapping[str, object] | None = None,
client: AsyncHTTPHandler | None = None,
) -> ImageResponse:
"""
@ -281,7 +319,7 @@ class BlackForestLabsImageEdit:
def _poll_for_result_sync(
self,
initial_response: httpx.Response,
initial_response: _BFLSubmitResponse,
headers: dict,
sync_client: HTTPHandler,
max_wait: float = DEFAULT_MAX_POLLING_TIME,
@ -356,8 +394,7 @@ class BlackForestLabsImageEdit:
message=f"Polling failed: {response.text}",
)
data = response.json()
status = data.get("status")
status = _poll_status(response)
verbose_logger.debug("BFL poll status: %s", status)
@ -383,7 +420,7 @@ class BlackForestLabsImageEdit:
async def _poll_for_result_async(
self,
initial_response: httpx.Response,
initial_response: _BFLSubmitResponse,
headers: dict,
async_client: AsyncHTTPHandler,
max_wait: float = DEFAULT_MAX_POLLING_TIME,
@ -447,8 +484,7 @@ class BlackForestLabsImageEdit:
message=f"Polling failed: {response.text}",
)
data = response.json()
status = data.get("status")
status = _poll_status(response)
verbose_logger.debug("BFL poll status: %s", status)

View file

@ -8,9 +8,11 @@ then we poll until the result is ready.
import asyncio
import time
from typing import Any, Final
from collections.abc import Coroutine, Mapping
from typing import Final, Protocol, TypedDict
import httpx
from typing_extensions import ReadOnly
import litellm
from litellm._logging import verbose_logger
@ -33,6 +35,23 @@ from ..common_utils import (
from .transformation import BlackForestLabsImageGenerationConfig
class _BFLTaskPayload(TypedDict, total=False):
"""The body BFL returns for a submitted or polled generation task."""
errors: ReadOnly[object]
polling_url: ReadOnly[str]
status: ReadOnly[str]
class _TaskJsonResponse(Protocol):
def json(self) -> _BFLTaskPayload: ...
def _task_payload(response: _TaskJsonResponse) -> _BFLTaskPayload:
"""The JSON body of a BFL task submission or poll response."""
return response.json()
class BlackForestLabsImageGeneration:
"""
Black Forest Labs Image Generation handler.
@ -53,10 +72,10 @@ class BlackForestLabsImageGeneration:
litellm_params: GenericLiteLLMParams | dict,
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout | None,
extra_headers: dict[str, Any] | None = None,
extra_headers: Mapping[str, str] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
aimg_generation: bool = False,
) -> ImageResponse | Any:
) -> ImageResponse | Coroutine[object, object, ImageResponse]:
"""
Main entry point for image generation requests.
@ -187,7 +206,7 @@ class BlackForestLabsImageGeneration:
litellm_params: GenericLiteLLMParams | dict,
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout | None,
extra_headers: dict[str, Any] | None = None,
extra_headers: Mapping[str, str] | None = None,
client: AsyncHTTPHandler | None = None,
) -> ImageResponse:
"""
@ -305,7 +324,7 @@ class BlackForestLabsImageGeneration:
# Parse initial response to get polling URL
try:
response_data: Final = initial_response.json()
response_data: Final = _task_payload(initial_response)
except Exception as e:
raise BlackForestLabsError(
status_code=initial_response.status_code,
@ -350,7 +369,7 @@ class BlackForestLabsImageGeneration:
message=f"Polling failed: {response.text}",
)
data = response.json()
data = _task_payload(response)
status = data.get("status")
verbose_logger.debug("BFL poll status: %s", status)
@ -396,7 +415,7 @@ class BlackForestLabsImageGeneration:
# Parse initial response to get polling URL
try:
response_data: Final = initial_response.json()
response_data: Final = _task_payload(initial_response)
except Exception as e:
raise BlackForestLabsError(
status_code=initial_response.status_code,
@ -441,7 +460,7 @@ class BlackForestLabsImageGeneration:
message=f"Polling failed: {response.text}",
)
data = response.json()
data = _task_payload(response)
status = data.get("status")
verbose_logger.debug("BFL poll status: %s", status)

View file

@ -4,9 +4,10 @@
import json
from collections.abc import Callable
from functools import partial
from typing import Final
from typing import Final, Protocol
import httpx
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
@ -23,6 +24,53 @@ from litellm.types.utils import TextChoices
from litellm.utils import CustomStreamWrapper, TextCompletionResponse
class _CodestralChoiceMessage(TypedDict):
"""`choices[].message` of a Codestral FIM completion."""
role: ReadOnly[NotRequired[str]]
content: ReadOnly[NotRequired[str | None]]
class _CodestralChoice(TypedDict):
"""One entry of `choices` in a Codestral FIM completion."""
index: ReadOnly[int]
message: ReadOnly[NotRequired[_CodestralChoiceMessage]]
finish_reason: ReadOnly[NotRequired[str | None]]
logprobs: ReadOnly[NotRequired[dict[str, object] | None]]
class _CodestralUsage(TypedDict):
"""Token accounting returned alongside a Codestral FIM completion."""
prompt_tokens: ReadOnly[NotRequired[int]]
completion_tokens: ReadOnly[NotRequired[int]]
total_tokens: ReadOnly[NotRequired[int]]
class _CodestralCompletionResponse(TypedDict):
"""Body returned by the Codestral `/v1/fim/completions` endpoint."""
id: ReadOnly[NotRequired[str]]
created: ReadOnly[NotRequired[int]]
model: ReadOnly[NotRequired[str]]
object: ReadOnly[NotRequired[str]]
usage: ReadOnly[NotRequired[_CodestralUsage]]
choices: ReadOnly[NotRequired[list[_CodestralChoice]]]
class _CodestralHTTPResponse(Protocol):
"""The Codestral completion response as this handler reads it."""
@property
def status_code(self) -> int: ...
@property
def text(self) -> str: ...
def json(self) -> _CodestralCompletionResponse: ...
class TextCompletionCodestralError(Exception):
def __init__(
self,
@ -115,7 +163,7 @@ class CodestralTextCompletion:
def process_text_completion_response(
self,
model: str,
response: httpx.Response,
response: _CodestralHTTPResponse,
model_response: TextCompletionResponse,
stream: bool,
logging_obj: LiteLLMLogging,

View file

@ -2,10 +2,11 @@
Translate between Cohere's `/rerank` format and Deepinfra's `/rerank` format.
"""
from collections.abc import Mapping
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Final, Protocol
import httpx
from typing_extensions import ReadOnly, TypedDict
from litellm._uuid import uuid
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -24,6 +25,36 @@ from litellm.types.rerank import (
)
class _DeepinfraInferenceStatus(TypedDict, total=False):
"""The ``inference_status`` block of a DeepInfra rerank response."""
status: ReadOnly[str]
runtime_ms: ReadOnly[float]
cost: ReadOnly[float]
tokens_generated: ReadOnly[int]
tokens_input: ReadOnly[int]
class _DeepinfraRerankResponse(TypedDict, total=False):
"""Body of a DeepInfra ``/rerank`` response."""
scores: ReadOnly[Sequence[float]]
input_tokens: ReadOnly[int]
request_id: ReadOnly[str | None]
inference_status: ReadOnly[_DeepinfraInferenceStatus]
class _DeepinfraRerankResponseSource(Protocol):
"""The DeepInfra ``/rerank`` HTTP response, read for the body it decodes to."""
def json(self) -> _DeepinfraRerankResponse: ...
def _deepinfra_rerank_body(response: _DeepinfraRerankResponseSource) -> _DeepinfraRerankResponse:
"""Decode the body of a DeepInfra ``/rerank`` response."""
return response.json()
class DeepinfraRerankConfig(BaseRerankConfig):
"""
Deepinfra Rerank - Follows the same Spec as Cohere Rerank
@ -95,7 +126,7 @@ class DeepinfraRerankConfig(BaseRerankConfig):
model: str,
drop_params: bool,
query: str,
documents: list[str | dict[str, Any]],
documents: list[str | dict[str, object]],
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: list[str] | None = None,
@ -150,7 +181,7 @@ class DeepinfraRerankConfig(BaseRerankConfig):
litellm_params: dict = {},
) -> RerankResponse:
try:
response_json: Final = raw_response.json()
response_json: Final = _deepinfra_rerank_body(raw_response)
logging_obj.post_call(original_response=raw_response.text)
# Extract the scores from the response

View file

@ -12,9 +12,10 @@ Schema versioning:
litellm.use_legacy_interactions_schema = True. Remove flag after June 8, 2026.
"""
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeAlias
import httpx
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
@ -41,6 +42,53 @@ else:
LiteLLMLoggingObj = Any
_JsonObject: TypeAlias = dict[str, object]
class _InteractionPayload(TypedDict, total=False):
"""JSON body of an Interactions API interaction, keyed as ``InteractionsAPIResponse`` fields."""
id: ReadOnly[str | None]
object: ReadOnly[str | None]
model: ReadOnly[str | None]
agent: ReadOnly[str | None]
status: ReadOnly[str | None]
created: ReadOnly[str | None]
updated: ReadOnly[str | None]
outputs: ReadOnly[list[_JsonObject] | None]
steps: ReadOnly[list[_JsonObject] | None]
usage: ReadOnly[_JsonObject | None]
class _CancelPayload(TypedDict, total=False):
"""JSON body of an Interactions API cancel response."""
id: ReadOnly[str | None]
status: ReadOnly[str | None]
class _InteractionPayloadSource(Protocol):
"""An Interactions API HTTP response, read for the interaction body it decodes to."""
def json(self) -> _InteractionPayload: ...
class _CancelPayloadSource(Protocol):
"""An Interactions API cancel HTTP response, read for the body it decodes to."""
def json(self) -> _CancelPayload: ...
def _interaction_body(response: _InteractionPayloadSource) -> _InteractionPayload:
"""Decode the body of an Interactions API interaction response."""
return response.json()
def _cancel_body(response: _CancelPayloadSource) -> _CancelPayload:
"""Decode the body of an Interactions API cancel response."""
return response.json()
class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
"""
Configuration for Google AI Studio Interactions API.
@ -143,7 +191,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
"""
use_legacy: Final[bool] = litellm.use_legacy_interactions_schema
request_body: Final[dict[str, Any]] = {}
request_body: Final[dict[str, object]] = {}
# Model or Agent (one required)
if model:
@ -189,7 +237,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
and (not isinstance(response_format, dict) or "mime_type" not in response_format)
):
# Wrap the legacy schema into the new polymorphic format.
new_rf: Final[dict[str, Any]] = {
new_rf: Final[dict[str, object]] = {
"type": "text",
"mime_type": response_mime_type,
}
@ -215,7 +263,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
if image_config is not None:
# Move image_config to response_format with type=image.
image_rf: Final[dict[str, Any]] = {"type": "image", **image_config}
image_rf: Final[_JsonObject] = {"type": "image", **image_config}
existing_rf: Final = request_body.get("response_format")
if existing_rf is None:
request_body["response_format"] = image_rf
@ -239,7 +287,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
original_response=raw_response.text,
additional_args={"complete_input_dict": {}},
)
raw_json: Final = raw_response.json()
raw_json: Final = _interaction_body(raw_response)
except Exception:
raise GeminiError(
message=raw_response.text,
@ -290,7 +338,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
logging_obj: LiteLLMLoggingObj,
) -> InteractionsAPIResponse:
try:
raw_json: Final = raw_response.json()
raw_json: Final = _interaction_body(raw_response)
except Exception:
raise GeminiError(
message=raw_response.text,
@ -355,7 +403,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
logging_obj: LiteLLMLoggingObj,
) -> CancelInteractionResult:
try:
raw_json: Final = raw_response.json()
raw_json: Final = _cancel_body(raw_response)
except Exception:
raise GeminiError(
message=raw_response.text,

View file

@ -1,4 +1,5 @@
import base64
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -54,8 +55,13 @@ def _convert_image_to_gemini_format(image_file) -> dict[str, str]:
return {"bytesBase64Encoded": base64_encoded, "mimeType": mime_type}
def _json_payload(raw_response: httpx.Response) -> object:
"""Read an HTTP response body as an opaque JSON payload."""
return raw_response.json()
def _usage_video_resolution_from_parameters(
parameters: dict[str, Any],
parameters: Mapping[str, object],
) -> str | None:
"""Normalize Veo ``parameters.resolution`` for usage and cost tracking."""
res: Final = parameters.get("resolution")
@ -97,7 +103,7 @@ class GeminiVideoConfig(BaseVideoConfig):
video_create_optional_params: VideoCreateOptionalRequestParams,
model: str,
drop_params: bool,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Map OpenAI-style parameters to Veo format.
@ -111,7 +117,7 @@ class GeminiVideoConfig(BaseVideoConfig):
All other params are passed through as-is to support Gemini-specific parameters.
"""
mapped_params: Final[dict[str, Any]] = {}
mapped_params: Final[dict[str, object]] = {}
# Get supported OpenAI params (exclude "model" and "prompt" which are handled separately)
supported_openai_params: Final = self.get_supported_openai_params(model)
@ -312,11 +318,11 @@ class GeminiVideoConfig(BaseVideoConfig):
- status: "processing"
- usage: includes duration_seconds and optional video_resolution for cost calculation
"""
response_data: Final = raw_response.json()
response_data: Final = _json_payload(raw_response)
# Parse response using Pydantic model for type safety
try:
operation_response: Final = GeminiLongRunningOperationResponse(**response_data)
operation_response: Final = GeminiLongRunningOperationResponse.model_validate(response_data)
except Exception as e:
raise ValueError(f"Failed to parse operation response: {e}")
@ -336,7 +342,7 @@ class GeminiVideoConfig(BaseVideoConfig):
model=model,
)
usage_data: Final[dict[str, Any]] = {}
usage_data: Final[dict[str, float | str]] = {}
if request_data:
parameters: Final = request_data.get("parameters", {})
duration: Final = parameters.get("durationSeconds") or DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS
@ -367,7 +373,7 @@ class GeminiVideoConfig(BaseVideoConfig):
"""
operation_name: Final = extract_original_video_id(video_id)
url: Final = f"{api_base.rstrip('/')}/v1beta/{operation_name}"
params: Final[dict[str, Any]] = {}
params: Final[dict[str, object]] = {}
return url, params
@ -403,9 +409,9 @@ class GeminiVideoConfig(BaseVideoConfig):
}
}
"""
response_data: Final = raw_response.json()
response_data: Final = _json_payload(raw_response)
# Parse response using Pydantic model for type safety
operation_response: Final = GeminiLongRunningOperationResponse(**response_data)
operation_response: Final = GeminiLongRunningOperationResponse.model_validate(response_data)
operation_name: Final = operation_response.name
is_done: Final = operation_response.done
@ -443,9 +449,9 @@ class GeminiVideoConfig(BaseVideoConfig):
client: Final = litellm.module_level_client
status_response: Final = client.get(url=status_url, headers=headers)
status_response.raise_for_status()
response_data: Final = status_response.json()
response_data: Final = _json_payload(status_response)
operation_response: Final = GeminiLongRunningOperationResponse(**response_data)
operation_response: Final = GeminiLongRunningOperationResponse.model_validate(response_data)
if not operation_response.done:
raise ValueError(
@ -458,7 +464,7 @@ class GeminiVideoConfig(BaseVideoConfig):
generated_samples: Final = operation_response.response.generateVideoResponse.generatedSamples
download_url: Final = generated_samples[0].video.uri
params: Final[dict[str, Any]] = {}
params: Final[dict[str, object]] = {}
return download_url, params
@ -480,7 +486,7 @@ class GeminiVideoConfig(BaseVideoConfig):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: dict[str, Any] | None = None,
extra_body: Mapping[str, object] | None = None,
) -> tuple[str, dict]:
"""
Video remix is not supported by Veo API.
@ -506,7 +512,7 @@ class GeminiVideoConfig(BaseVideoConfig):
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]:
"""
Video list is not supported by Veo API.
@ -547,7 +553,7 @@ class GeminiVideoConfig(BaseVideoConfig):
"""Video delete is not supported."""
raise NotImplementedError("Video delete is not supported by Google Veo.")
def transform_video_create_character_request(self, name, video, api_base, litellm_params, headers):
def transform_video_create_character_request(self, name, video: object, api_base, litellm_params, headers):
raise NotImplementedError("video create character is not supported for Gemini")
def transform_video_create_character_response(self, raw_response, logging_obj):

View file

@ -1,8 +1,9 @@
import json
import os
import time
from collections.abc import Sequence
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, Protocol
import httpx
@ -24,6 +25,8 @@ from litellm.utils import token_counter
from ..common_utils import HuggingFaceError, hf_task_list, hf_tasks, output_parser
if TYPE_CHECKING:
import tiktoken
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
LoggingClass = LiteLLMLoggingObj
@ -31,6 +34,12 @@ else:
LoggingClass = Any
class _TokenEncoding(Protocol):
"""Tokenizer handle the caller passes in; only `encode` is used, to count completion tokens."""
def encode(self, text: str, /) -> Sequence[object]: ...
tgi_models_cache = None
conv_models_cache = None
@ -369,7 +378,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
model_response: ModelResponse,
task: hf_tasks | None,
optional_params: dict,
encoding: Any,
encoding: "_TokenEncoding | None",
messages: list[AllMessageValues],
model: str,
):
@ -439,9 +448,10 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
if output_text is not None and len(output_text) > 0:
completion_tokens = 0
try:
completion_tokens = len(
encoding.encode(model_response["choices"][0]["message"].get("content", ""))
) ##[TODO] use the llama2 tokenizer here
if encoding is not None:
completion_tokens = len(
encoding.encode(model_response["choices"][0]["message"].get("content", ""))
) ##[TODO] use the llama2 tokenizer here
except Exception:
# this should remain non blocking we should not block a response returning if calculating usage fails
pass
@ -469,7 +479,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
encoding: "tiktoken.Encoding | None",
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:

View file

@ -325,7 +325,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
@overload
def _transform_messages(
self, messages: list[AllMessageValues], model: str, is_async: Literal[True]
) -> Coroutine[Any, Any, list[AllMessageValues]]:
) -> Coroutine[object, object, list[AllMessageValues]]:
...
@overload
@ -341,7 +341,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
def _transform_messages(
self, messages: list[AllMessageValues], model: str, is_async: bool = False
) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]:
) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]:
"""OpenAI no longer supports image_url as a string, so we need to convert it to a dict"""
stripped_messages: Final = drop_tool_reference_parts_from_tool_messages(messages)
hoisted_messages: Final = hoist_images_from_tool_messages(stripped_messages)
@ -497,8 +497,12 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
return None
tool_call_names: Final = get_tool_call_names(optional_params.get("tools", []))
try:
json_content: Final = json.loads(content)
if json_content.get("type") == "function" and json_content.get("name") in tool_call_names:
json_content: Final[object] = json.loads(content)
if (
isinstance(json_content, dict)
and json_content.get("type") == "function"
and json_content.get("name") in tool_call_names
):
return ChatCompletionMessageToolCall(
function=Function(
name=json_content.get("name"),
@ -622,7 +626,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
## RESPONSE OBJECT
try:
completion_response: Final = raw_response.json()
completion_response: Final[dict[str, object]] = raw_response.json()
except Exception as e:
response_headers: Final = getattr(raw_response, "headers", None)
raise OpenAIError(

View file

@ -51,6 +51,7 @@ if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
class OpenAIChatCompletionsHandler(BaseTranslation):
@ -80,7 +81,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
data: dict,
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
) -> Any:
) -> dict:
"""
Process input messages by applying guardrails to text content.
"""
@ -329,9 +330,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
response: "ModelResponse",
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
user_api_key_dict: Any | None = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
) -> Any:
) -> ModelResponse:
"""
Process output response by applying guardrails to text content.
@ -436,7 +437,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
responses_so_far: list["ModelResponseStream"],
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
user_api_key_dict: Any | None = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
stream_transform_sink: StreamTransformSink | None = None,
) -> list["ModelResponseStream"]:
@ -486,7 +487,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
responses_so_far: list["ModelResponseStream"],
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: "LiteLLMLoggingObj | None",
user_api_key_dict: Any | None,
user_api_key_dict: "UserAPIKeyAuth | None",
request_data: dict | None,
) -> list["ModelResponseStream"]:
"""Block-only streaming path: run the guardrail so an in-flight BLOCK can
@ -589,8 +590,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
def build_stream_error_items(
self,
exc: "HTTPException",
responses_so_far: Sequence[Any] | None = None,
) -> Sequence[Any] | None:
responses_so_far: Sequence[object] | None = None,
) -> Sequence[bytes] | None:
import json
from litellm.proxy.common_request_processing import sse_error_payload
@ -630,7 +631,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
responses_so_far: list["ModelResponseStream"],
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: "LiteLLMLoggingObj | None",
user_api_key_dict: Any | None,
user_api_key_dict: "UserAPIKeyAuth | None",
request_data: dict | None,
sink: StreamTransformSink,
) -> None:
@ -794,7 +795,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
# Determine content source and tool calls based on choice type
content = None
tool_calls: list[Any] | None = None
tool_calls: Sequence[object] | None = None
if isinstance(choice, litellm.Choices):
content = choice.message.content
tool_calls = choice.message.tool_calls

View file

@ -1,10 +1,11 @@
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, cast, get_type_hints
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_type_hints
import httpx
from openai.types.responses import ResponseReasoningItem
from pydantic import BaseModel, ValidationError
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
@ -37,6 +38,36 @@ _MODEL_FAMILIES_REJECTING_TOP_LEVEL_SCHEMA_COMBINATORS: Final = ("gpt-4", "gpt-3
_PROVIDERS_WITH_COMBINATOR_REJECTING_VALIDATOR: Final = frozenset({LlmProviders.AZURE, LlmProviders.OPENAI})
class _DeleteResponseBody(TypedDict):
"""Decoded body of the Responses API delete call."""
id: ReadOnly[str | None]
object: ReadOnly[str | None]
deleted: ReadOnly[bool | None]
class _DeleteResponse(Protocol):
"""The delete call's HTTP response, read for the decoded body it carries."""
def json(self) -> _DeleteResponseBody: ...
class _JsonObjectResponse(Protocol):
"""A Responses API HTTP response, read for the JSON object it decodes to."""
def json(self) -> dict[str, object]: ...
def _delete_response_body(response: _DeleteResponse) -> _DeleteResponseBody:
"""Decode a delete response body into the id, object and deleted fields it carries."""
return response.json()
def _json_object_body(response: _JsonObjectResponse) -> dict[str, object]:
"""Decode a Responses API response body into its JSON object form."""
return response.json()
class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
@property
def custom_llm_provider(self) -> LlmProviders:
@ -469,7 +500,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
return None
@staticmethod
def get_event_model_class(event_type: str) -> Any:
def get_event_model_class(event_type: str) -> type[BaseLiteLLMOpenAIResponseObject]:
"""
Returns the appropriate event model class based on the event type.
@ -583,7 +614,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
Transform the delete response API response into a DeleteResponseResult
"""
try:
raw_response_json: Final = raw_response.json()
raw_response_json: Final = _delete_response_body(raw_response)
except Exception:
raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code)
return DeleteResponseResult(**raw_response_json)
@ -618,7 +649,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
Transform the get response API response into a ResponsesAPIResponse
"""
try:
raw_response_json: Final = raw_response.json()
raw_response_json: Final = _json_object_body(raw_response)
except Exception:
raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code)
raw_response_headers: Final = dict(raw_response.headers)
@ -646,7 +677,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
) -> tuple[str, dict]:
encoded_response_id: Final = encode_url_path_segment(response_id, field_name="response_id")
url: Final = f"{api_base}/{encoded_response_id}/input_items"
params: Final[dict[str, Any]] = {}
params: Final[dict[str, object]] = {}
if after is not None:
params["after"] = after
if before is not None:
@ -665,7 +696,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
logging_obj: LiteLLMLoggingObj,
) -> dict:
try:
return raw_response.json()
return _json_object_body(raw_response)
except Exception:
raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code)
@ -699,7 +730,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
Transform the cancel response API response into a ResponsesAPIResponse
"""
try:
raw_response_json: Final = raw_response.json()
raw_response_json: Final = _json_object_body(raw_response)
except Exception:
raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code)
raw_response_headers: Final = dict(raw_response.headers)

View file

@ -5,10 +5,11 @@ For handling OpenAI-like chat completions, like IBM WatsonX, etc.
"""
import json
from collections.abc import Callable
from typing import Any, Final
from collections.abc import Callable, Mapping, Sequence
from typing import Final, TypedDict
import httpx
from typing_extensions import ReadOnly
import litellm
from litellm import LlmProviders
@ -25,6 +26,23 @@ from ..common_utils import OpenAILikeBase, OpenAILikeError
from .transformation import OpenAILikeChatConfig
class _OpenAILikeChatCompletion(TypedDict, total=False):
"""The chat-completion JSON body an OpenAI-like provider returns for a non-streamed call."""
id: ReadOnly[str]
choices: ReadOnly[Sequence[Mapping[str, object]]]
created: ReadOnly[int]
model: ReadOnly[str]
system_fingerprint: ReadOnly[str]
usage: ReadOnly[Mapping[str, object]]
object: ReadOnly[str]
def _fake_streamed_model_response(payload: _OpenAILikeChatCompletion) -> ModelResponse:
"""Build the single response a fake-streamed provider call replays as one chunk."""
return ModelResponse(**payload)
async def make_call(
client: AsyncHTTPHandler | None,
api_base: str,
@ -42,9 +60,9 @@ async def make_call(
response: Final = await client.post(api_base, headers=headers, data=data, stream=not fake_stream)
if streaming_decoder is not None:
completion_stream: Any = streaming_decoder.aiter_bytes(response.aiter_bytes(chunk_size=1024))
completion_stream = streaming_decoder.aiter_bytes(response.aiter_bytes(chunk_size=1024))
elif fake_stream:
model_response: Final = ModelResponse(**response.json())
model_response: Final = _fake_streamed_model_response(response.json())
completion_stream = MockResponseIterator(model_response=model_response)
else:
completion_stream = ModelResponseIterator(streaming_response=response.aiter_lines(), sync_stream=False)
@ -82,7 +100,7 @@ def make_sync_call(
if streaming_decoder is not None:
completion_stream = streaming_decoder.iter_bytes(response.iter_bytes(chunk_size=1024))
elif fake_stream:
model_response: Final = ModelResponse(**response.json())
model_response: Final = _fake_streamed_model_response(response.json())
completion_stream = MockResponseIterator(model_response=model_response)
else:
completion_stream = ModelResponseIterator(streaming_response=response.iter_lines(), sync_stream=True)

View file

@ -1,8 +1,10 @@
import asyncio
import time
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
import httpx
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_logger
from litellm.constants import (
@ -29,6 +31,16 @@ else:
LiteLLMLoggingObj = Any
class _RunwayMLTask(TypedDict, total=False):
"""The RunwayML task payload returned by POST /v1/text_to_image and GET /v1/tasks/{id}."""
id: ReadOnly[str]
status: ReadOnly[str]
output: ReadOnly[Sequence[str | Mapping[str, str]]]
failure: ReadOnly[str]
failureCode: ReadOnly[str]
class RunwayMLImageGenerationConfig(BaseImageGenerationConfig):
"""
Configuration for RunwayML image generation models.
@ -80,7 +92,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig):
@staticmethod
def _transform_runwayml_response_to_openai(
response_data: dict[str, Any],
response_data: _RunwayMLTask,
model_response: ImageResponse,
) -> ImageResponse:
"""
@ -155,7 +167,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig):
raise TimeoutError(f"RunwayML task polling timed out after {timeout_secs} seconds")
@staticmethod
def _check_task_status(response_data: dict[str, Any]) -> str:
def _check_task_status(response_data: _RunwayMLTask) -> str:
"""
Check RunwayML task status from response.
@ -227,7 +239,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig):
response = client.get(url=task_url, headers=headers)
response.raise_for_status()
response_data = response.json()
response_data: _RunwayMLTask = response.json()
# Check task status
status = self._check_task_status(response_data=response_data)
@ -276,7 +288,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig):
response = await client.get(url=task_url, headers=headers)
response.raise_for_status()
response_data = response.json()
response_data: _RunwayMLTask = response.json()
# Check task status
status = self._check_task_status(response_data=response_data)
@ -322,7 +334,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig):
}
"""
try:
response_data = raw_response.json()
response_data: _RunwayMLTask = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
@ -382,7 +394,7 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig):
We need to poll the task until it completes (status SUCCEEDED) using async polling.
"""
try:
response_data = raw_response.json()
response_data: _RunwayMLTask = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",

View file

@ -8,9 +8,10 @@ from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from pathlib import Path
from threading import Lock
from typing import Any, Final
from typing import Any, Final, Protocol
import httpx
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
@ -33,8 +34,8 @@ def _get_home() -> str:
return os.getenv(HOME_PATH_ENV_VAR, DEFAULT_HOME_PATH)
def _get_nested(d: dict[str, Any] | str, path: Sequence[str]) -> Any:
cur: Any = d
def _get_nested(d: object, path: Sequence[str]) -> object:
cur: object = d
if isinstance(cur, str):
# This shouldn't happen if service keys are pre-parsed correctly
try:
@ -54,7 +55,7 @@ def _get_nested(d: dict[str, Any] | str, path: Sequence[str]) -> Any:
return cur
def _load_json_env(var_name: str) -> dict[str, Any] | None:
def _load_json_env(var_name: str) -> dict[str, object] | None:
raw: Final = os.environ.get(var_name)
if not raw:
return None
@ -64,7 +65,7 @@ def _load_json_env(var_name: str) -> dict[str, Any] | None:
return None
def _str_or_none(value) -> str | None:
def _str_or_none(value: object) -> str | None:
try:
return str(value) if value is not None else None
except Exception:
@ -124,7 +125,7 @@ CREDENTIAL_VALUES: Final[list[CredentialsValue]] = [
]
def init_conf(profile: str | None = None) -> dict[str, Any]:
def init_conf(profile: str | None = None) -> dict[str, object]:
"""
Loads config JSON from:
1) $AICORE_CONFIG if set, otherwise
@ -191,7 +192,7 @@ def resolve_resource_group(sources: list[Source]) -> str | None:
def _parse_service_key_once(
service_key: str | dict | None,
) -> dict[str, Any] | None:
) -> dict[str, object] | None:
"""
Pre-parse service_key if it's a string to avoid repeated JSON parsing.
@ -348,8 +349,33 @@ def validate_credentials(
)
class _TokenBody(TypedDict):
"""Decoded body of the SAP AI Core OAuth2 token response."""
access_token: ReadOnly[str]
expires_in: ReadOnly[NotRequired[int]]
class _TokenResponse(Protocol):
"""The token endpoint's HTTP response, read for the decoded token body it carries."""
def json(self) -> _TokenBody: ...
def _bearer_token_and_expiry(response: _TokenResponse) -> tuple[str, datetime]:
"""Read a token response into the Authorization header value and the token's absolute expiry."""
payload: Final = response.json()
expires_in: Final = int(payload.get("expires_in", 3600))
access_token: Final = payload["access_token"]
return f"Bearer {access_token}", datetime.now(timezone.utc) + timedelta(seconds=expires_in)
def _request_token(
client_id: str, auth_url: str, timeout: float, cert_pair=None, client_secret=None
client_id: str,
auth_url: str,
timeout: float,
cert_pair: tuple[str, str] | None = None,
client_secret: str | None = None,
) -> tuple[str, datetime]:
data: Final = {"grant_type": "client_credentials", "client_id": client_id}
if client_secret:
@ -361,15 +387,10 @@ def _request_token(
with httpx.Client(cert=cert_pair) as raw_client:
handler = HTTPHandler(client=raw_client)
resp = handler.post(auth_url, data=data, timeout=timeout)
payload = resp.json()
else:
handler = _get_httpx_client()
resp = handler.post(auth_url, data=data, timeout=timeout)
payload = resp.json()
access_token: Final = payload["access_token"]
expires_in: Final = int(payload.get("expires_in", 3600))
expiry_date: Final = datetime.now(timezone.utc) + timedelta(seconds=expires_in)
return f"Bearer {access_token}", expiry_date
return _bearer_token_and_expiry(resp)
handler = _get_httpx_client()
resp = handler.post(auth_url, data=data, timeout=timeout)
return _bearer_token_and_expiry(resp)
except Exception as e:
msg: Final = resp.text if resp is not None else getattr(e, "text", str(e))
raise RuntimeError(f"Token request failed: {msg}") from e

View file

@ -12,7 +12,7 @@ from urllib.parse import quote, unquote
import httpx
from httpx import Headers, Response
from openai.types.file_deleted import FileDeleted
from typing_extensions import ReadOnly
from typing_extensions import ReadOnly, Required
import litellm
from litellm._uuid import uuid
@ -104,6 +104,27 @@ class _VertexBatchRow(TypedDict, total=False):
processed_time: ReadOnly[str]
class _VertexEmbeddingVector(TypedDict):
values: ReadOnly[list[float]]
class _VertexEmbeddingUsageMetadata(TypedDict, total=False):
promptTokenCount: ReadOnly[int]
class _VertexEmbeddingResponse(TypedDict, total=False):
embedding: ReadOnly[Required[_VertexEmbeddingVector]]
usageMetadata: ReadOnly[_VertexEmbeddingUsageMetadata]
tokenCount: ReadOnly[int]
class _VertexEmbeddingBatchRow(TypedDict, total=False):
key: ReadOnly[str]
request: ReadOnly[Mapping[str, object]]
status: ReadOnly[Required[str]]
response: ReadOnly[Required[_VertexEmbeddingResponse]]
class _OpenAIBatchOutputError(TypedDict):
code: ReadOnly[str]
message: ReadOnly[str]
@ -111,7 +132,7 @@ class _OpenAIBatchOutputError(TypedDict):
class _OpenAIBatchOutputResponse(TypedDict):
status_code: ReadOnly[int]
request_id: ReadOnly[str]
request_id: ReadOnly[object]
body: ReadOnly[Mapping[str, object]]
@ -218,7 +239,7 @@ def _get_litellm_batch_custom_id_from_labels(labels: Mapping[str, object] | None
return str(labels.get("litellm_custom_id", "unknown"))
def _is_vertex_embeddings_batch_output_row(vertex_output_row: Mapping[str, Any]) -> bool:
def _is_vertex_embeddings_batch_output_row(vertex_output_row: Mapping[str, object]) -> bool:
"""
Whether a Vertex batch output row came from an `EmbedContentRequest`.
@ -237,7 +258,7 @@ def _is_vertex_embeddings_batch_output_row(vertex_output_row: Mapping[str, Any])
def _openai_batch_output_row(
custom_id: str,
body: Mapping[str, Any] | None = None,
body: Mapping[str, object] | None = None,
error_code: str | None = None,
error_message: str = "",
) -> _OpenAIBatchOutputRow:
@ -259,7 +280,7 @@ def _openai_batch_output_row(
}
def _split_vertex_batch_key(vertex_output_row: Mapping[str, Any]) -> tuple[str, int, int]:
def _split_vertex_batch_key(vertex_output_row: Mapping[str, object]) -> tuple[str, int, int]:
"""
Resolve `(custom_id, index within that custom_id, group size)` for a Vertex batch
output row.
@ -278,7 +299,7 @@ def _split_vertex_batch_key(vertex_output_row: Mapping[str, Any]) -> tuple[str,
return unquote(match["custom_id"]), int(match["index"]), int(match["total"])
def _embedding_prompt_token_count(vertex_response: Mapping[str, Any]) -> int:
def _embedding_prompt_token_count(vertex_response: _VertexEmbeddingResponse) -> int:
"""
Prompt tokens billed for one Vertex Gemini Embedding batch row.
@ -293,7 +314,7 @@ def _embedding_prompt_token_count(vertex_response: Mapping[str, Any]) -> int:
def _vertex_embeddings_rows_to_openai_batch_output_row(
custom_id: str,
vertex_output_rows: tuple[Mapping[str, Any], ...],
vertex_output_rows: tuple[_VertexEmbeddingBatchRow, ...],
element_indices: tuple[int, ...],
element_count: int,
model: str | None,
@ -348,7 +369,7 @@ def _vertex_embeddings_rows_to_openai_batch_output_row(
def _transform_vertex_embeddings_batch_output_to_openai(
vertex_output_rows: Iterable[Mapping[str, Any]],
vertex_output_rows: Iterable[_VertexEmbeddingBatchRow],
model: str | None,
) -> tuple[_OpenAIBatchOutputRow, ...]:
"""
@ -388,7 +409,7 @@ def _model_from_managed_gcs_url(url: str) -> str | None:
return match.group(1) if match else None
def _is_embeddings_batch_entry(openai_entry: Mapping[str, Any]) -> bool:
def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool:
"""
Whether an OpenAI batch JSONL line targets the embeddings endpoint.
@ -431,7 +452,7 @@ def _vertex_batch_embeddings_key(custom_id: str, index: int, total: int) -> str:
return encoded_custom_id if total < 2 else f"{encoded_custom_id}#{index}/{total}"
def _vertex_embeddings_row(key: str | None, embed_content_request: Mapping[str, Any]) -> Mapping[str, Any]:
def _vertex_embeddings_row(key: str | None, embed_content_request: Mapping[str, object]) -> Mapping[str, object]:
"""
One Vertex Gemini Embedding batch input row.
@ -453,8 +474,8 @@ def _vertex_embeddings_row(key: str | None, embed_content_request: Mapping[str,
def _openai_batch_jsonl_entry_to_vertex_embeddings_rows(
openai_entry: Mapping[str, Any],
) -> tuple[Mapping[str, Any], ...]:
openai_entry: Mapping[str, object],
) -> tuple[Mapping[str, object], ...]:
"""
Transforms a single OpenAI `/v1/embeddings` batch entry into Vertex Gemini Embedding
batch rows, one per requested embedding.
@ -512,7 +533,7 @@ def _openai_batch_jsonl_entry_to_vertex_embeddings_rows(
def _openai_batch_jsonl_entry_to_vertex_rows(
openai_entry: dict[str, Any],
map_openai_to_vertex_params: Callable[[dict[str, Any]], dict[str, Any]],
) -> tuple[Mapping[str, Any], ...]:
) -> tuple[Mapping[str, object], ...]:
"""
Transforms a single OpenAI JSONL batch entry into the Vertex rows it maps to.
@ -533,7 +554,7 @@ def _openai_batch_jsonl_entry_to_vertex_rows(
cached_content=None,
)
custom_id: Final = openai_entry.get("custom_id")
custom_id: Final[object] = openai_entry.get("custom_id")
if custom_id is not None:
if "labels" not in vertex_request_body:
vertex_request_body["labels"] = {}

View file

@ -250,7 +250,7 @@ def _gs_uri_requires_content_type_metadata(url: str) -> bool:
def _image_url_payload_may_need_sync_gcs_metadata_fetch(
raw_image_url: Any,
raw_image_url: object,
) -> bool:
"""
True when this image_url value (content-part image_url or assistant ``images[]``
@ -326,7 +326,7 @@ def _openai_messages_may_need_sync_gcs_metadata_fetch(
def _get_gcs_object_content_type(
image_url: str,
vertex_project: str | None = None,
vertex_credentials: Any | None = None,
vertex_credentials: object = None,
) -> str | None:
"""
Resolve content type from GCS object metadata.
@ -479,7 +479,7 @@ def _process_gemini_media(
model: str | None = None,
video_metadata: dict[str, Any] | None = None,
vertex_project: str | None = None,
vertex_credentials: Any | None = None,
vertex_credentials: object = None,
) -> PartType:
"""
Given a media URL (image, audio, or video), return the appropriate PartType for Gemini
@ -1002,7 +1002,7 @@ def _gemini_convert_messages_with_history(
if isinstance(_ss_invocations, list):
for invocation in _ss_invocations:
# Re-inject toolCall part
tc_part: dict[str, Any] = {
tc_part: dict[str, object] = {
"toolCall": {
"toolType": invocation.get("tool_type"),
"id": invocation.get("id"),
@ -1015,13 +1015,13 @@ def _gemini_convert_messages_with_history(
# Re-inject toolResponse part if response is present
if "response" in invocation:
tr_dict: dict[str, Any] = {
tr_dict: dict[str, object] = {
"id": invocation.get("id"),
"response": invocation.get("response"),
}
if invocation.get("tool_type"):
tr_dict["toolType"] = invocation["tool_type"]
tr_part: dict[str, Any] = {"toolResponse": tr_dict}
tr_part: dict[str, object] = {"toolResponse": tr_dict}
if "response_thought_signature" in invocation:
tr_part["thoughtSignature"] = invocation["response_thought_signature"]
assistant_content.append(tr_part)
@ -1090,7 +1090,7 @@ def _pop_and_merge_extra_body(data: RequestBody, optional_params: dict) -> None:
data_dict[k] = v
def _has_google_maps_tool(tools: Any | None) -> bool:
def _has_google_maps_tool(tools: object) -> bool:
"""Return True if any tool object in the list has a 'googleMaps' key."""
if not isinstance(tools, list):
return False
@ -1127,7 +1127,7 @@ def _rewrite_mime_type_to_response_format(generation_config: GenerationConfig) -
schema = generation_config.pop("response_schema", None)
generation_config.pop("response_mime_type", None)
response_format: Final[dict[str, Any]] = {"text": {"mimeType": "APPLICATION_JSON"}}
response_format: Final[dict[str, dict[str, object]]] = {"text": {"mimeType": "APPLICATION_JSON"}}
if schema is not None:
response_format["text"]["schema"] = schema
generation_config["responseFormat"] = response_format
@ -1316,7 +1316,7 @@ async def async_transform_request_body(
timeout: float | httpx.Timeout | None,
extra_headers: dict | None,
optional_params: dict,
logging_obj: litellm.litellm_core_utils.litellm_logging.Logging,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
litellm_params: dict,
vertex_project: str | None,

View file

@ -9,7 +9,7 @@ import json
import os
import threading
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol
from urllib.parse import urlparse
import litellm
@ -47,6 +47,21 @@ else:
GoogleCredentialsObject = Any
class _VertexCredentialsObject(Protocol):
"""Structural view of the google-auth credentials handle that this class caches and refreshes."""
@property
def token(self) -> object: ...
@property
def quota_project_id(self) -> str | None: ...
@property
def expired(self) -> object: ...
def refresh(self, request: object) -> None: ...
class VertexBase:
def __init__(self) -> None:
super().__init__()
@ -55,7 +70,7 @@ class VertexBase:
self._credentials: GoogleCredentialsObject | None = None
self._credentials_project_mapping: dict[
tuple[VERTEX_CREDENTIALS_TYPES | None, str | None],
tuple[GoogleCredentialsObject, str | None],
tuple[_VertexCredentialsObject, str | None],
] = {}
self.project_id: str | None = None
self.async_handler: AsyncHTTPHandler | None = None
@ -109,7 +124,7 @@ class VertexBase:
self,
credentials: VERTEX_CREDENTIALS_TYPES | None,
project_id: str | None,
) -> tuple[Any, str]:
) -> tuple[_VertexCredentialsObject | None, str]:
if credentials is not None:
if isinstance(credentials, str):
_is_path: Final = os.path.exists(
@ -209,7 +224,7 @@ class VertexBase:
return creds, project_id
# Google Auth Helpers -- extracted for mocking purposes in tests
def _credentials_from_identity_pool(self, json_obj, scopes):
def _credentials_from_identity_pool(self, json_obj, scopes) -> _VertexCredentialsObject:
try:
from google.auth import identity_pool
except ImportError:
@ -220,7 +235,7 @@ class VertexBase:
creds = creds.with_scopes(scopes)
return creds
def _credentials_from_pluggable(self, json_obj, scopes):
def _credentials_from_pluggable(self, json_obj, scopes) -> _VertexCredentialsObject:
try:
from google.auth import pluggable
except ImportError:
@ -231,7 +246,7 @@ class VertexBase:
creds = creds.with_scopes(scopes)
return creds
def _credentials_from_identity_pool_with_aws(self, json_obj, scopes):
def _credentials_from_identity_pool_with_aws(self, json_obj, scopes) -> _VertexCredentialsObject:
try:
from google.auth import aws
except ImportError:
@ -242,7 +257,7 @@ class VertexBase:
creds = creds.with_scopes(scopes)
return creds
def _credentials_from_authorized_user(self, json_obj, scopes):
def _credentials_from_authorized_user(self, json_obj, scopes) -> _VertexCredentialsObject:
try:
import google.oauth2.credentials
except ImportError:
@ -250,7 +265,7 @@ class VertexBase:
return google.oauth2.credentials.Credentials.from_authorized_user_info(json_obj, scopes=scopes)
def _credentials_from_service_account(self, json_obj, scopes):
def _credentials_from_service_account(self, json_obj, scopes) -> _VertexCredentialsObject:
try:
import google.oauth2.service_account
except ImportError:
@ -258,7 +273,7 @@ class VertexBase:
return google.oauth2.service_account.Credentials.from_service_account_info(json_obj, scopes=scopes)
def _credentials_from_default_auth(self, scopes):
def _credentials_from_default_auth(self, scopes) -> tuple[_VertexCredentialsObject, str | None]:
try:
import google.auth as google_auth
except ImportError:
@ -350,7 +365,7 @@ class VertexBase:
)
return api_base
def refresh_auth(self, credentials: Any) -> None:
def refresh_auth(self, credentials: _VertexCredentialsObject) -> None:
try:
from google.auth.transport.requests import (
Request,
@ -426,7 +441,7 @@ class VertexBase:
self,
credential_cache_key: tuple,
project_id: str | None,
) -> tuple[str, str, "TokenState", Any, str | None] | None:
) -> tuple[str, str, "TokenState", _VertexCredentialsObject, str | None] | None:
"""
Look up cached credentials and return usable token info for FRESH or
STALE tokens (both are still valid for outbound requests). STALE
@ -449,7 +464,9 @@ class VertexBase:
return None
return creds.token, resolved_project, token_state, creds, cached_project_id
def _unpack_cached_credentials(self, credential_cache_key: tuple) -> tuple[Any, str | None]:
def _unpack_cached_credentials(
self, credential_cache_key: tuple
) -> tuple[_VertexCredentialsObject | None, str | None]:
"""
Return (credentials, project_id) from the cache, or (None, None) if
not cached. Handles both tuple and legacy cache formats.
@ -461,7 +478,7 @@ class VertexBase:
return cached_entry
return cached_entry, cached_entry.quota_project_id or getattr(cached_entry, "project_id", None)
def _get_token_state(self, credentials: Any) -> "TokenState":
def _get_token_state(self, credentials: _VertexCredentialsObject) -> "TokenState":
"""
Return the token state using google-auth's TokenState enum.
@ -485,7 +502,7 @@ class VertexBase:
credentials: VERTEX_CREDENTIALS_TYPES | None,
project_id: str | None,
credential_cache_key: tuple,
) -> tuple[Any, str | None]:
) -> tuple[_VertexCredentialsObject, str | None]:
"""Load credentials via load_auth (in thread) and cache the result."""
try:
_credentials, credential_project_id = await asyncify(self.load_auth)(
@ -505,7 +522,7 @@ class VertexBase:
async def _background_refresh_credentials(
self,
credentials: Any,
credentials: _VertexCredentialsObject,
credential_cache_key: tuple,
credential_project_id: str | None,
) -> None:
@ -557,7 +574,7 @@ class VertexBase:
def _schedule_background_refresh(
self,
credentials: Any,
credentials: _VertexCredentialsObject,
credential_cache_key: tuple,
credential_project_id: str | None,
) -> None:
@ -575,7 +592,7 @@ class VertexBase:
self._background_refresh_credentials(credentials, credential_cache_key, credential_project_id)
)
def _drop_background_refresh_task(_fut: asyncio.Future[Any]) -> None:
def _drop_background_refresh_task(_fut: asyncio.Future[None]) -> None:
if self._background_refresh_tasks.get(credential_cache_key) is _fut:
self._background_refresh_tasks.pop(credential_cache_key, None)
@ -888,7 +905,7 @@ class VertexBase:
# Convert dict credentials to string for caching
cache_credentials: Final = json.dumps(credentials) if isinstance(credentials, dict) else credentials
credential_cache_key: Final = (cache_credentials, project_id)
_credentials: GoogleCredentialsObject | None = None
_credentials: _VertexCredentialsObject | None = None
verbose_logger.debug("Checking cached credentials for project_id: %s", project_id)

View file

@ -6,7 +6,7 @@ from __future__ import annotations
import asyncio
import contextvars
from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Generator, Iterator
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Coroutine, Generator, Iterator
from functools import partial
from types import TracebackType
from typing import Any, Final, cast
@ -27,19 +27,19 @@ base_llm_http_handler = BaseLLMHTTPHandler()
from .utils import BasePassthroughUtils
async def _as_async_generator(iterable: AsyncIterator[bytes]) -> AsyncGenerator[bytes, Any]:
async def _as_async_generator(iterable: AsyncIterator[bytes]) -> AsyncGenerator[bytes, bytes]:
async for chunk in iterable:
yield chunk
def _as_generator(iterable: Iterator[bytes]) -> Generator[bytes, Any, Any]:
def _as_generator(iterable: Iterator[bytes]) -> Generator[bytes, bytes, None]:
yield from iterable
class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]):
def __init__(
self,
response: Coroutine[Any, Any, httpx.Response],
response: Awaitable[httpx.Response],
litellm_logging_obj: LiteLLMLoggingObj,
provider_config: BasePassthroughConfig,
) -> None:
@ -48,7 +48,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
self._headers = httpx.Headers()
self._response_coro = response
self._response: httpx.Response
self._iterator: AsyncGenerator[bytes, Any]
self._iterator: AsyncGenerator[bytes, bytes]
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
@ -172,7 +172,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
pass
class PassthroughStreamingResponse(Generator[Any, Any, Any]):
class PassthroughStreamingResponse(Generator[bytes, bytes, None]):
def __init__(
self,
response: httpx.Response,
@ -184,7 +184,7 @@ class PassthroughStreamingResponse(Generator[Any, Any, Any]):
self.status_code = response.status_code
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._iterator: Generator[bytes, Any, Any] = _as_generator(response.iter_bytes())
self._iterator: Generator[bytes, bytes, None] = _as_generator(response.iter_bytes())
self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks
self._flush_scheduled = False
@ -263,7 +263,7 @@ async def allm_passthrough_route(
cookies: CookieTypes | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
**kwargs,
) -> httpx.Response | AsyncGenerator[Any, Any]:
) -> httpx.Response | AsyncGenerator[bytes, bytes]:
"""
Async: Reranks a list of documents based on their relevance to the query
"""
@ -390,10 +390,10 @@ def llm_passthrough_route(
**kwargs,
) -> (
httpx.Response
| Coroutine[Any, Any, httpx.Response]
| Coroutine[Any, Any, httpx.Response | AsyncGenerator[Any, Any]]
| Generator[Any, Any, Any]
| AsyncGenerator[Any, Any]
| Coroutine[object, object, httpx.Response]
| Coroutine[object, object, httpx.Response | AsyncGenerator[bytes, bytes]]
| Generator[bytes, bytes, None]
| AsyncGenerator[bytes, bytes]
):
"""
Pass through requests to the LLM APIs.
@ -592,7 +592,7 @@ async def _async_passthrough_request(
is_streaming_request: bool,
litellm_logging_obj: LiteLLMLoggingObj,
provider_config: BasePassthroughConfig,
) -> httpx.Response | AsyncGenerator[Any, Any]:
) -> httpx.Response | AsyncGenerator[bytes, bytes]:
"""
Handle async passthrough requests.
Uses async client to send request and properly handles streaming.

View file

@ -5,6 +5,7 @@ Filters MCP tools semantically for /chat/completions and /responses endpoints.
"""
import asyncio
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_logger
@ -74,7 +75,7 @@ class SemanticMCPToolFilter:
self.router_instance = litellm_router_instance
self.tool_router: SemanticRouter | None = None
self.context_window_error: str | None = None
self._tool_map: dict[str, Any] = {} # MCPTool objects or OpenAI function dicts
self._tool_map: dict[str, object] = {} # MCPTool objects or OpenAI function dicts
self._index_sync_lock = asyncio.Lock()
async def build_router_from_mcp_registry(self) -> None:
@ -182,11 +183,11 @@ class SemanticMCPToolFilter:
return
raise
def _has_tools_missing_from_index(self, tools: list[Any]) -> bool:
def _has_tools_missing_from_index(self, tools: Sequence[object]) -> bool:
"""Allocation-free check for any named tool not yet in the semantic index."""
return any(name and name not in self._tool_map for name in (self._extract_tool_info(t)[0] for t in tools))
def _tools_missing_from_index(self, tools: list[Any]) -> dict[str, Any]:
def _tools_missing_from_index(self, tools: Sequence[object]) -> Mapping[str, object]:
"""Map name -> tool for every named tool not yet in the semantic index."""
return {
name: tool
@ -194,7 +195,7 @@ class SemanticMCPToolFilter:
if name and name not in self._tool_map
}
async def _ensure_tools_indexed(self, available_tools: list[Any]) -> None:
async def _ensure_tools_indexed(self, available_tools: Sequence[object]) -> None:
"""
Index request-time tools the startup build never saw.
@ -385,7 +386,7 @@ class SemanticMCPToolFilter:
separator: Final = client_name[-len(canonical) - 1]
return separator in ("_", "-")
def _get_tools_by_names(self, tool_names: list[str], available_tools: list[Any]) -> list[Any]:
def _get_tools_by_names(self, tool_names: Sequence[str], available_tools: Sequence[object]) -> list[object]:
"""
Get tools from available_tools by their names, preserving the
semantic router's ordering.
@ -401,14 +402,14 @@ class SemanticMCPToolFilter:
# Exact matches win over suffix matches when both are present, and
# each incoming tool is returned at most once even if two canonical
# names happen to be tail-compatible with the same incoming name.
available_by_name: Final[dict[str, Any]] = {}
available_by_name: Final[dict[str, object]] = {}
for tool in available_tools:
client_name, _ = self._extract_tool_info(tool)
if client_name and client_name not in available_by_name:
available_by_name[client_name] = tool
matched: Final[list[Any]] = []
used_ids: Final[set] = set()
matched: Final[list[object]] = []
used_ids: Final[set[int]] = set()
for canonical in tool_names:
tool = available_by_name.get(canonical)
if tool is None:
@ -430,7 +431,7 @@ class SemanticMCPToolFilter:
used_ids.add(id(tool))
return matched
def extract_user_query(self, messages: list[dict[str, Any]]) -> str:
def extract_user_query(self, messages: Sequence[Mapping[str, object]]) -> str:
"""
Extract user query from messages for /chat/completions or /responses.

View file

@ -14,7 +14,7 @@ import json
from collections.abc import AsyncGenerator, Mapping
from copy import deepcopy
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, Protocol
from urllib.parse import urlparse
from fastapi import APIRouter, Depends, HTTPException, Request, Response
@ -215,11 +215,20 @@ def _enforce_inbound_trace_id(agent: "AgentResponse", request: Request) -> None:
)
class _JsonRpcResponse(Protocol):
def json(self) -> dict[str, object]: ...
def _jsonrpc_body(response: _JsonRpcResponse) -> dict[str, object]:
"""The decoded JSON-RPC body of ``response``."""
return response.json()
async def _forward_jsonrpc(
agent_url: str,
body: dict[str, object],
extra_headers: Mapping[str, str] | None = None,
) -> dict[str, Any]:
) -> dict[str, object]:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -230,7 +239,7 @@ async def _forward_jsonrpc(
)
resp: Final = await handler.post(agent_url, json=body, headers=headers)
try:
result: Final = resp.json()
result: Final = _jsonrpc_body(resp)
except Exception:
resp.raise_for_status()
raise
@ -940,8 +949,8 @@ async def invoke_agent_a2a(
)
result = await _forward_jsonrpc(agent_url, forward_body, extra_headers=caller_headers)
if method == "agent/getAuthenticatedExtendedCard":
if isinstance(result.get("result"), dict):
card: Final = result["result"]
card: Final = result.get("result")
if isinstance(card, dict):
proxy_url: Final = get_custom_url(str(request.base_url), route=f"a2a/{agent_id}")
# Rewrite the upstream agent URL in both 0.3 (top-level `url`)
# and 1.0 (`supportedInterfaces[0].url`) wire formats so that

View file

@ -14,8 +14,8 @@ import hashlib
import os
import re
import time
from collections.abc import Awaitable, Callable
from typing import Any, Final, Literal, NoReturn, TypeVar, cast
from collections.abc import Awaitable, Callable, Sequence
from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast
import httpx
import jwt
@ -24,6 +24,7 @@ from cryptography.hazmat.backends import default_backend
from cryptography.hazmat.primitives import serialization
from fastapi import HTTPException, status
from jwt.api_jwk import PyJWK
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
@ -93,6 +94,47 @@ UNREACHABLE_CACHE_KEY_PREFIX: Final = "litellm_jwks_unreachable_"
_CachedValueT = TypeVar("_CachedValueT", bound=JWKKeyValue | str)
class _JWTAuthSettings(Protocol):
"""The JWT auth settings block this handler reads back through ``getattr``, when one is configured."""
@property
def issuers(self) -> Sequence[JWTIssuerConfig] | None: ...
@property
def public_key_ttl(self) -> float: ...
@property
def public_key_stale_ttl(self) -> float: ...
class _OIDCDiscoveryBody(TypedDict, total=False):
"""Decoded OIDC discovery document, read for the JWKS endpoint it advertises."""
jwks_uri: ReadOnly[str]
class _OIDCDiscoveryResponse(Protocol):
"""The discovery endpoint's HTTP response, read for the decoded document it carries."""
def json(self) -> _OIDCDiscoveryBody: ...
class _UserInfoResponse(Protocol):
"""The OIDC UserInfo endpoint's HTTP response, read for the identity document it carries."""
def json(self) -> dict[str, object]: ...
def _discovery_document(response: _OIDCDiscoveryResponse) -> _OIDCDiscoveryBody:
"""Decode an OIDC discovery response body."""
return response.json()
def _userinfo_document(response: _UserInfoResponse) -> dict[str, object]:
"""Decode an OIDC UserInfo response body into its JSON object form."""
return response.json()
def jwks_unavailable_exception(error: JWKSUnreachableError) -> ProxyException:
return ProxyException(
message=(
@ -794,7 +836,7 @@ class JWTHandler:
f"JWT Auth: OIDC discovery endpoint {url} returned status {response.status_code}: {response.text}"
)
try:
discovery: Final = response.json()
discovery: Final = _discovery_document(response)
except Exception as e:
raise Exception(f"JWT Auth: Failed to parse OIDC discovery document at {url}: {e}")
@ -806,13 +848,13 @@ class JWTHandler:
return jwks_uri
def _get_public_key_cache_ttl(self) -> float:
litellm_jwtauth: Final = getattr(self, "litellm_jwtauth", None)
litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
return 600
return litellm_jwtauth.public_key_ttl
def _get_public_key_stale_ttl(self) -> float:
litellm_jwtauth: Final = getattr(self, "litellm_jwtauth", None)
litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
return DEFAULT_JWKS_STALE_TTL
return litellm_jwtauth.public_key_stale_ttl
@ -938,7 +980,7 @@ class JWTHandler:
if response.status_code != 200:
raise Exception(f"OIDC UserInfo endpoint returned status {response.status_code}: {response.text}")
userinfo: Final = response.json()
userinfo: Final = _userinfo_document(response)
verbose_proxy_logger.debug("Received OIDC UserInfo: %s", userinfo)
# Cache the userinfo response
@ -996,7 +1038,7 @@ class JWTHandler:
}
def _get_configured_issuer(self, token: str) -> JWTIssuerConfig | None:
litellm_jwtauth: Final = getattr(self, "litellm_jwtauth", None)
litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
return None

View file

@ -6,9 +6,11 @@ import os
import sys
import tracemalloc
from collections import Counter
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Any, Final, NamedTuple, Protocol, TypedDict
from fastapi import APIRouter, Depends, HTTPException, Query
from typing_extensions import ReadOnly
from litellm import get_secret_str
from litellm._logging import verbose_proxy_logger
@ -194,6 +196,42 @@ async def memory_usage_in_mem_cache_items(
}
class _ProcessMemoryInfo(Protocol):
"""The resident and virtual sizes psutil reports for a process."""
@property
def rss(self) -> int: ...
@property
def vms(self) -> int: ...
class _ProcessHandle(Protocol):
"""The psutil process handle members this module reads."""
def memory_info(self) -> _ProcessMemoryInfo: ...
def memory_percent(self) -> float: ...
class _ProcessMemoryUsage(NamedTuple):
"""Memory usage of a single worker process."""
resident_megabytes: float
virtual_megabytes: float
percent: float
def _process_memory_usage(process: _ProcessHandle) -> _ProcessMemoryUsage:
"""Read resident/virtual megabytes and system memory share for ``process``."""
memory_info: Final = process.memory_info()
return _ProcessMemoryUsage(
resident_megabytes=memory_info.rss / (1024 * 1024),
virtual_megabytes=memory_info.vms / (1024 * 1024),
percent=process.memory_percent(),
)
@router.get("/debug/memory/summary", include_in_schema=False)
async def get_memory_summary(
_: UserAPIKeyAuth = Depends(user_api_key_auth),
@ -227,10 +265,9 @@ async def get_memory_summary(
try:
import psutil
process: Final = psutil.Process()
memory_info: Final = process.memory_info()
memory_mb: Final = memory_info.rss / (1024 * 1024)
memory_percent: Final = process.memory_percent()
usage: Final = _process_memory_usage(psutil.Process())
memory_mb: Final = usage.resident_megabytes
memory_percent: Final = usage.percent
process_memory = {
"summary": f"{memory_mb:.1f} MB ({memory_percent:.1f}% of system memory)",
@ -252,7 +289,7 @@ async def get_memory_summary(
process_memory["error"] = str(e)
# Get cache information
caches: Final[dict[str, Any]] = {}
caches: Final[dict[str, object]] = {}
total_cache_items = 0
try:
@ -313,7 +350,7 @@ async def get_memory_summary(
}
def _get_gc_statistics() -> dict[str, Any]:
def _get_gc_statistics() -> Mapping[str, object]:
"""Get garbage collector statistics."""
return {
"enabled": gc.isenabled(),
@ -341,30 +378,42 @@ def _get_gc_statistics() -> dict[str, Any]:
}
def _get_object_type_counts(top_n: int) -> tuple[int, list[dict[str, Any]]]:
class _ObjectTypeCount(TypedDict):
"""One row of the tracked-object histogram."""
type: ReadOnly[str]
count: ReadOnly[int]
count_readable: ReadOnly[str]
def _type_name_counts(objects: Sequence[object]) -> Counter[str]:
"""Count ``objects`` by the name of their type."""
return Counter(type(obj).__name__ for obj in objects)
def _get_object_type_counts(top_n: int) -> tuple[int, list[_ObjectTypeCount]]:
"""Count objects by type and return total count and top N types."""
type_counts: Final[Counter] = Counter()
total_objects = 0
type_counts: Final = _type_name_counts(gc.get_objects())
for obj in gc.get_objects():
total_objects += 1
obj_type = type(obj).__name__
type_counts[obj_type] += 1
top_object_types: Final = [
top_object_types: Final[list[_ObjectTypeCount]] = [
{"type": obj_type, "count": count, "count_readable": f"{count:,}"}
for obj_type, count in type_counts.most_common(top_n)
]
return total_objects, top_object_types
return sum(type_counts.values()), top_object_types
def _get_uncollectable_objects_info() -> dict[str, Any]:
def _type_names(objects: Sequence[object]) -> Sequence[str]:
"""The type name of each object in ``objects``."""
return [type(obj).__name__ for obj in objects]
def _get_uncollectable_objects_info() -> Mapping[str, object]:
"""Get information about uncollectable objects (potential memory leaks)."""
uncollectable: Final = gc.garbage
return {
"count": len(uncollectable),
"sample_types": [type(obj).__name__ for obj in uncollectable[:10]],
"sample_types": _type_names(uncollectable[:10]),
"warning": (
"If count > 0, you may have reference cycles preventing garbage collection"
if len(uncollectable) > 0
@ -373,9 +422,11 @@ def _get_uncollectable_objects_info() -> dict[str, Any]:
}
def _get_cache_memory_stats(user_api_key_cache, llm_router, proxy_logging_obj, redis_usage_cache) -> dict[str, Any]:
def _get_cache_memory_stats(
user_api_key_cache, llm_router, proxy_logging_obj, redis_usage_cache
) -> Mapping[str, object]:
"""Calculate memory usage for all caches."""
cache_stats: Final[dict[str, Any]] = {}
cache_stats: Final[dict[str, object]] = {}
try:
# User API key cache
user_cache_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.cache_dict)
@ -439,9 +490,9 @@ def _get_cache_memory_stats(user_api_key_cache, llm_router, proxy_logging_obj, r
return cache_stats
def _get_router_memory_stats(llm_router) -> dict[str, Any]:
def _get_router_memory_stats(llm_router) -> Mapping[str, object]:
"""Get memory usage statistics for LiteLLM router."""
litellm_router_memory: dict[str, Any] = {}
litellm_router_memory: dict[str, object] = {}
try:
if llm_router is not None:
# Model list memory size
@ -505,7 +556,7 @@ def _get_router_memory_stats(llm_router) -> dict[str, Any]:
return litellm_router_memory
def _get_process_memory_info(worker_pid: int, include_process_info: bool) -> dict[str, Any] | None:
def _get_process_memory_info(worker_pid: int, include_process_info: bool) -> Mapping[str, object] | None:
"""Get process-level memory information using psutil."""
if not include_process_info:
return None
@ -514,10 +565,10 @@ def _get_process_memory_info(worker_pid: int, include_process_info: bool) -> dic
import psutil
process: Final = psutil.Process()
memory_info: Final = process.memory_info()
ram_usage_mb: Final = round(memory_info.rss / (1024 * 1024), 2)
virtual_memory_mb: Final = round(memory_info.vms / (1024 * 1024), 2)
memory_percent: Final = round(process.memory_percent(), 2)
usage: Final = _process_memory_usage(process)
ram_usage_mb: Final = round(usage.resident_megabytes, 2)
virtual_memory_mb: Final = round(usage.virtual_megabytes, 2)
memory_percent: Final = round(usage.percent, 2)
return {
"pid": worker_pid,

View file

@ -211,7 +211,7 @@ class DBSpendUpdateWriter:
org_id: str | None,
# Completion object fields
kwargs: dict | None,
completion_response: litellm.ModelResponse | Any | Exception | None,
completion_response: object,
start_time: datetime | None,
end_time: datetime | None,
response_cost: float | None,
@ -323,7 +323,7 @@ class DBSpendUpdateWriter:
async def _enqueue_tool_usage_transaction(
self,
payload: SpendLogsPayload,
completion_response: "litellm.ModelResponse | Any | Exception | None",
completion_response: object,
prisma_client: "PrismaClient | None",
kwargs: "dict | None" = None,
) -> None:
@ -396,7 +396,7 @@ class DBSpendUpdateWriter:
def _enqueue_tool_registry_upsert(
self,
kwargs: dict | None,
completion_response: Any | None,
completion_response: object,
hashed_token: str | None = None,
team_id: str | None = None,
) -> None:
@ -849,7 +849,7 @@ class DBSpendUpdateWriter:
return
# Parse tags from JSON string
tags = []
tags: Sequence[object] = []
if isinstance(request_tags, str):
tags = safe_json_loads(request_tags, default=[])
if not tags:
@ -2260,7 +2260,7 @@ class DBSpendUpdateWriter:
verbose_proxy_logger.debug("request_tags is None for request. Skipping incrementing tag spend.")
return
request_tags = []
request_tags: Sequence[str] = []
if isinstance(payload["request_tags"], str):
request_tags = json.loads(payload["request_tags"])
elif isinstance(payload["request_tags"], list):

View file

@ -162,10 +162,10 @@ class AktoGuardrail(CustomGuardrail):
def build_request_body(
inputs: GenericGuardrailAPIInputs,
request_data: dict | None = None,
) -> dict[str, Any]:
) -> dict[str, object]:
"""Build the LLM request body from guardrail inputs (messages, model, tools)."""
model: Final = inputs.get("model", "") or ""
body: Final[dict[str, Any]] = {"model": model}
body: Final[dict[str, object]] = {"model": model}
structured: Final = inputs.get("structured_messages")
if structured:
@ -194,7 +194,7 @@ class AktoGuardrail(CustomGuardrail):
def build_response_body(
inputs: GenericGuardrailAPIInputs,
request_data: dict | None = None,
) -> dict[str, Any]:
) -> dict[str, object]:
"""Build the LLM response body, preferring the actual model response if available."""
model_response: Final = request_data.get("response") if request_data else None
if model_response is not None and hasattr(model_response, "model_dump"):
@ -224,7 +224,7 @@ class AktoGuardrail(CustomGuardrail):
*,
status_code: int = 200,
include_response: bool = False,
) -> dict[str, Any]:
) -> dict[str, object]:
"""Build the flat MIRRORING payload sent to Akto's HTTP proxy endpoint.
All body fields use double-encoding: json.dumps({"body": json.dumps(actual_body)})

View file

@ -2,9 +2,10 @@
import os
import time
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol
from fastapi import HTTPException
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
@ -27,6 +28,33 @@ if TYPE_CHECKING:
GRAYSWAN_BLOCK_ERROR_MSG: Final = "Blocked by Gray Swan Guardrail"
class _GraySwanMonitorResponse(TypedDict):
"""Body returned by Gray Swan's `/cygnal/monitor` endpoint."""
violation: ReadOnly[NotRequired[float | None]]
violated_rules: ReadOnly[NotRequired[list[object]]]
violated_rule_descriptions: ReadOnly[NotRequired[list[object]]]
mutation: ReadOnly[NotRequired[bool | None]]
ipi: ReadOnly[NotRequired[bool | None]]
class _GraySwanMonitorHTTPResponse(Protocol):
def raise_for_status(self) -> object: ...
def json(self) -> _GraySwanMonitorResponse: ...
class _GraySwanMonitorHTTPClient(Protocol):
async def post(
self,
*,
url: str,
headers: dict[str, str],
json: dict[str, object],
timeout: float,
) -> _GraySwanMonitorHTTPResponse: ...
class GraySwanGuardrailMissingSecrets(Exception):
"""Raised when the Gray Swan API key is missing."""
@ -77,7 +105,9 @@ class GraySwanGuardrail(CustomGuardrail):
guardrail_timeout: float | None = 30.0,
**kwargs: Any,
) -> None:
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.async_handler: _GraySwanMonitorHTTPClient = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
api_key_value: Final = api_key or os.getenv("GRAYSWAN_API_KEY")
if not api_key_value:
@ -266,7 +296,7 @@ class GraySwanGuardrail(CustomGuardrail):
# Legacy Test Interface (for backward compatibility)
# ------------------------------------------------------------------
async def run_grayswan_guardrail(self, payload: dict) -> dict[str, Any]:
async def run_grayswan_guardrail(self, payload: dict[str, object]) -> _GraySwanMonitorResponse:
"""
Run the GraySwan guardrail on a payload.
@ -285,7 +315,7 @@ class GraySwanGuardrail(CustomGuardrail):
def _process_grayswan_response(
self,
response_json: dict,
response_json: _GraySwanMonitorResponse,
data: dict | None = None,
hook_type: GuardrailEventHooks | None = None,
) -> None:
@ -385,7 +415,7 @@ class GraySwanGuardrail(CustomGuardrail):
# Core GraySwan API interaction
# ------------------------------------------------------------------
async def _call_grayswan_api(self, payload: dict) -> dict[str, Any]:
async def _call_grayswan_api(self, payload: dict[str, object]) -> _GraySwanMonitorResponse:
"""Call the GraySwan monitoring API."""
headers: Final = self._prepare_headers()
@ -406,7 +436,7 @@ class GraySwanGuardrail(CustomGuardrail):
def _process_response_internal(
self,
response_json: dict[str, Any],
response_json: _GraySwanMonitorResponse,
request_data: dict,
inputs: GenericGuardrailAPIInputs,
is_output: bool,
@ -534,8 +564,8 @@ class GraySwanGuardrail(CustomGuardrail):
dynamic_body: dict,
request_data: dict,
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> dict[str, Any] | None:
payload: Final[dict[str, Any]] = {"messages": messages}
) -> dict[str, object] | None:
payload: Final[dict[str, object]] = {"messages": messages}
categories: Final = dynamic_body.get("categories") or self.categories
if categories:
@ -563,13 +593,13 @@ class GraySwanGuardrail(CustomGuardrail):
{**existing_headers, **inbound_headers} if isinstance(existing_headers, dict) else inbound_headers
)
if cleaned_litellm_metadata:
sanitized: Final = safe_json_loads(safe_dumps(cleaned_litellm_metadata), default={})
sanitized: Final[object] = safe_json_loads(safe_dumps(cleaned_litellm_metadata), default={})
if isinstance(sanitized, dict) and sanitized:
payload["litellm_metadata"] = sanitized
return payload
def _format_violation_message(self, detection_info: Any, is_output: bool = False) -> str:
def _format_violation_message(self, detection_info: object, is_output: bool = False) -> str:
"""
Format detection info into a user-friendly violation message.

View file

@ -8,6 +8,7 @@
import json
import os
import uuid
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
try:
@ -128,7 +129,7 @@ class LassoGuardrail(CustomGuardrail):
@staticmethod
def _extract_tool_call_fields(
call: Any,
call: object,
) -> tuple[str | None, str | None, dict[str, object] | None]:
"""Extract (call_id, name, parsed_input) from a tool call.
@ -476,7 +477,7 @@ class LassoGuardrail(CustomGuardrail):
def _map_masked_messages_back(
self,
original_messages: list[dict[str, Any]],
masked_messages: list[dict[str, Any]],
masked_messages: Sequence[Mapping[str, object]],
) -> list[dict[str, object]]:
"""Map Lasso-format masked messages back onto the original OpenAI-format messages.
@ -638,7 +639,7 @@ class LassoGuardrail(CustomGuardrail):
},
)
def _expand_messages_for_classification(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
def _expand_messages_for_classification(self, messages: list[dict[str, Any]]) -> list[dict[str, object]]:
"""
Convert raw OpenAI-format messages to Lasso API format with content blocks.
@ -646,7 +647,7 @@ class LassoGuardrail(CustomGuardrail):
- role=tool messages developer role + tool_result block
- plain text messages pass through unchanged
"""
expanded: Final[list[dict[str, Any]]] = []
expanded: Final[list[dict[str, object]]] = []
for msg in messages:
role = msg.get("role", "")
content = msg.get("content")
@ -917,7 +918,7 @@ class LassoGuardrail(CustomGuardrail):
def _apply_masking_to_model_response(
self,
model_response: litellm.ModelResponse,
masked_messages: list[dict[str, Any]],
masked_messages: Sequence[Mapping[str, object]],
) -> None:
"""Apply masking to the actual model response when mask=True and masked content is available."""
# Index masked tool_use blocks by id for O(1) lookup.

View file

@ -8,11 +8,12 @@
# Standard library imports
import json
import os
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol
from urllib.parse import quote
# Third-party imports
from fastapi import HTTPException
from typing_extensions import NotRequired, ReadOnly, TypedDict
# LiteLLM imports
from litellm import DualCache
@ -42,7 +43,34 @@ if TYPE_CHECKING:
MAX_PILLAR_HEADER_VALUE_BYTES: Final = 8 * 1024
def _encode_json_for_header(data: Any) -> str:
class _PillarProtectResponse(TypedDict):
"""Body returned by Pillar's `/api/v1/protect` endpoint."""
flagged: ReadOnly[NotRequired[bool]]
session_id: ReadOnly[NotRequired[str]]
scanners: ReadOnly[NotRequired[dict[str, object]]]
evidence: ReadOnly[NotRequired[list[object]]]
masked_session_messages: ReadOnly[NotRequired[list[object]]]
class _PillarProtectHTTPResponse(Protocol):
def raise_for_status(self) -> object: ...
def json(self) -> _PillarProtectResponse: ...
class _PillarProtectHTTPClient(Protocol):
async def post(
self,
*,
url: str,
headers: dict[str, str],
json: dict[str, object],
timeout: float,
) -> _PillarProtectHTTPResponse: ...
def _encode_json_for_header(data: object) -> str:
"""
JSON-serialize and URL-encode data for safe header transmission.
"""
@ -50,7 +78,9 @@ def _encode_json_for_header(data: Any) -> str:
return quote(json_payload, safe="")
def _truncate_evidence_payload(evidence: Any, max_bytes: int = MAX_PILLAR_HEADER_VALUE_BYTES) -> tuple[Any, str, bool]:
def _truncate_evidence_payload(
evidence: object, max_bytes: int = MAX_PILLAR_HEADER_VALUE_BYTES
) -> tuple[object, str, bool]:
"""
Truncate evidence payload so the encoded header value stays within max_bytes.
@ -66,12 +96,12 @@ def _truncate_evidence_payload(evidence: Any, max_bytes: int = MAX_PILLAR_HEADER
truncated_value: Final = "[truncated]"
return truncated_value, _encode_json_for_header(truncated_value), True
truncated: Final[list[Any]] = []
truncated: Final[list[object]] = []
encoded = _encode_json_for_header(truncated)
truncated_flag = False
for entry in evidence:
working_entry: Any
working_entry: object
if isinstance(entry, dict):
working_entry = dict(entry)
else:
@ -105,7 +135,7 @@ def _truncate_evidence_payload(evidence: Any, max_bytes: int = MAX_PILLAR_HEADER
return truncated, encoded, truncated_flag
def build_pillar_response_headers(metadata_store: dict[str, Any]) -> dict[str, str]:
def build_pillar_response_headers(metadata_store: dict[str, object]) -> dict[str, str]:
"""
Create URL-safe Pillar response headers and apply truncation metadata.
"""
@ -191,7 +221,9 @@ class PillarGuardrail(CustomGuardrail):
LiteLLM virtual key context (user_id, team_id, key_alias, etc.) is always
automatically passed as X-LiteLLM-* headers to enable application/user tracking.
"""
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.async_handler: _PillarProtectHTTPClient = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.api_key = api_key or os.environ.get("PILLAR_API_KEY")
if self.api_key is None:
@ -686,7 +718,7 @@ class PillarGuardrail(CustomGuardrail):
)
return payload
async def _call_pillar_api(self, headers: dict[str, str], payload: dict[str, Any]) -> dict[str, Any]:
async def _call_pillar_api(self, headers: dict[str, str], payload: dict[str, Any]) -> _PillarProtectResponse:
"""
Call the Pillar API and return the response.
@ -714,7 +746,7 @@ class PillarGuardrail(CustomGuardrail):
verbose_proxy_logger.debug("Pillar Guardrail: Analysis complete - flagged=%s, session=%s", flagged, session_id)
return res
def _process_pillar_response(self, pillar_response: dict[str, Any], original_data: dict) -> None:
def _process_pillar_response(self, pillar_response: _PillarProtectResponse, original_data: dict) -> None:
"""
Process the Pillar API response and handle detections based on configuration.
@ -774,7 +806,7 @@ class PillarGuardrail(CustomGuardrail):
build_pillar_response_headers(metadata_store)
def _raise_pillar_detection_exception(self, pillar_response: dict[str, Any]) -> None:
def _raise_pillar_detection_exception(self, pillar_response: _PillarProtectResponse) -> None:
"""
Raise an HTTPException for Pillar security detections.
@ -784,7 +816,7 @@ class PillarGuardrail(CustomGuardrail):
Raises:
HTTPException: Always raises with security detection details
"""
pillar_response_dict: Final = {
pillar_response_dict: Final[dict[str, object]] = {
"session_id": pillar_response.get("session_id"),
}

View file

@ -6,7 +6,7 @@ via embedding similarity. Smarter than regex (understands intent), lighter
than an LLM call (~20-50ms per request for embedding).
"""
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, Protocol
from litellm._logging import verbose_logger
from litellm.integrations.custom_guardrail import (
@ -50,7 +50,7 @@ class SemanticGuardrail(CustomGuardrail):
similarity_threshold: float,
route_templates: list[str] | None = None,
custom_routes_file: str | None = None,
custom_routes: list[dict[str, Any]] | None = None,
custom_routes: list[dict[str, object]] | None = None,
on_flagged_action: str = "block",
event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None,
default_on: bool = False,
@ -157,7 +157,14 @@ class SemanticGuardrail(CustomGuardrail):
return response
def _get_top_route_choice(result: Any) -> Any:
class _RouteChoice(Protocol):
"""The semantic-router match this guardrail reads: the route that fired, if any."""
@property
def name(self) -> str | None: ...
def _get_top_route_choice(result: _RouteChoice | list[_RouteChoice] | None) -> _RouteChoice | None:
"""Extract the top RouteChoice from SemanticRouter result.
SemanticRouter.__call__ can return RouteChoice or List[RouteChoice].
@ -194,7 +201,7 @@ def _extract_response_text(response: Any) -> str:
return ""
def _content_to_text(content: Any) -> str:
def _content_to_text(content: object) -> str:
if isinstance(content, str):
return content
if isinstance(content, list):

View file

@ -1,9 +1,10 @@
import json
import re
from collections.abc import AsyncGenerator, AsyncIterable, Mapping, Sequence
from typing import Any, Final, Literal
from typing import Any, Final, Literal, TypedDict
from fastapi import HTTPException
from typing_extensions import ReadOnly, Required
from litellm import ChatCompletionToolParam
from litellm._logging import verbose_proxy_logger
@ -51,6 +52,27 @@ def _object_list(value: object) -> Sequence[object] | None:
return value if isinstance(value, list) else None
class _ToolPermissionRuleFields(TypedDict, total=False):
"""The config-file shape a :class:`ToolPermissionRule` is built from."""
id: ReadOnly[Required[str]]
tool_name: ReadOnly[str | None]
tool_type: ReadOnly[str | None]
decision: ReadOnly[Required[Literal["allow", "deny"]]]
allowed_param_patterns: ReadOnly[dict[str, str] | None]
def _rule_from_fields(fields: _ToolPermissionRuleFields) -> ToolPermissionRule:
"""Validate one config-file rule entry into a :class:`ToolPermissionRule`."""
return ToolPermissionRule(**fields)
def _is_tool_use_block(block: object) -> bool:
"""Whether ``block`` is an Anthropic ``tool_use`` content block."""
fields: Final = _object_mapping(block)
return fields is not None and fields.get("type") == "tool_use"
class ToolPermissionGuardrail(CustomGuardrail):
def __init__(
self,
@ -101,7 +123,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
compiled_patterns: Final[dict[str, dict[str, re.Pattern]]] = {}
for rule_item in rules or []:
rule = rule_item if isinstance(rule_item, ToolPermissionRule) else ToolPermissionRule(**rule_item)
rule = rule_item if isinstance(rule_item, ToolPermissionRule) else _rule_from_fields(rule_item)
target_patterns: dict[str, re.Pattern | None] = {
"tool_name": None,
@ -440,7 +462,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
return is_allowed, None, message
@staticmethod
def _get_mapping_value(item: Any, key: str) -> Any:
def _get_mapping_value(item: object, key: str) -> Any:
if isinstance(item, dict):
return item.get(key)
return getattr(item, key, None)
@ -450,7 +472,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
return f"legacy_function_call_{choice_index}"
def _legacy_function_call_to_tool_call(
self, function_call: Any, choice_index: int
self, function_call: object, choice_index: int
) -> ChatCompletionMessageToolCall | None:
if function_call is None:
return None
@ -549,7 +571,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
def _modify_anthropic_content_with_permission_errors(
self,
response: object,
content: tuple[Any, ...],
content: tuple[object, ...],
denied_tools: tuple[tuple[ChatCompletionMessageToolCall, PermissionError], ...],
) -> None:
if not denied_tools or not isinstance(response, dict):
@ -557,27 +579,33 @@ class ToolPermissionGuardrail(CustomGuardrail):
verbose_proxy_logger.info("Blocking %s unauthorized tool uses", len(denied_tools))
error_by_tool_use_id: Final = { # mutable-ok: read-only lookup, never mutated after construction
error_by_tool_use_id: Final[
Mapping[object, str]
] = { # mutable-ok: read-only lookup, never mutated after construction
tool_call.id: self._create_permission_error_result(tool_call, error).content
for tool_call, error in denied_tools
}
denied_block_ids: Final = frozenset(error_by_tool_use_id)
def _is_denied(block: object) -> bool:
return isinstance(block, dict) and block.get("type") == "tool_use" and block.get("id") in denied_block_ids
def _denied_message(block: object) -> str | None:
fields: Final = _object_mapping(block)
if fields is None or fields.get("type") != "tool_use":
return None
return error_by_tool_use_id.get(fields.get("id"))
error_messages: Final = tuple(error_by_tool_use_id[block["id"]] for block in content if _is_denied(block))
kept_blocks: Final = tuple(block for block in content if not _is_denied(block))
error_messages: Final = tuple(
message for message in (_denied_message(block) for block in content) if message is not None
)
kept_blocks: Final = tuple(block for block in content if _denied_message(block) is None)
new_content: Final = [ # mutable-ok: response content is a JSON array on the wire
*kept_blocks,
{"type": "text", "text": "\n".join(error_messages)}, # mutable-ok: content block is a JSON object
]
response["content"] = new_content # rebind-ok: the guardrail rewrites the provider response in place
if not any(isinstance(block, dict) and block.get("type") == "tool_use" for block in kept_blocks):
if not any(_is_tool_use_block(block) for block in kept_blocks):
response["stop_reason"] = "end_turn" # rebind-ok: dropping every tool_use ends the turn
def _get_request_tool_name(self, tool: Any) -> tuple[str | None, str | None]:
def _get_request_tool_name(self, tool: object) -> tuple[str | None, str | None]:
tool_type: Final = self._get_mapping_value(tool, "type")
if tool_type != "function":
return None, tool_type
@ -586,7 +614,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
tool_name: Final = self._get_mapping_value(function, "name")
return tool_name, tool_type
def _get_legacy_function_name(self, function: Any) -> str | None:
def _get_legacy_function_name(self, function: object) -> str | None:
return self._get_mapping_value(function, "name")
def _get_named_tool_choice(self, data: dict) -> str | None:

View file

@ -433,7 +433,7 @@ class VigilGuardGuardrail(CustomGuardrail):
return collected
@staticmethod
def _clamp_metadata_value(value: Any) -> _MetadataValue | None:
def _clamp_metadata_value(value: object) -> _MetadataValue | None:
if isinstance(value, bool):
return None
if isinstance(value, str):

View file

@ -67,6 +67,24 @@ class _ChatMessage(Protocol):
def tool_calls(self) -> Sequence[_ChatToolCall] | None: ...
class _ChatChoice(Protocol):
@property
def message(self) -> _ChatMessage: ...
@property
def finish_reason(self) -> str | None: ...
class _ChatCompletion(Protocol):
@property
def choices(self) -> Sequence[_ChatChoice]: ...
def _first_choice(response: _ChatCompletion) -> _ChatChoice:
"""The first choice of an OpenAI shaped completion response."""
return response.choices[0]
class SkillsInjectionHook(CustomLogger):
"""
Pre/Post-call hook that processes skills from container.skills parameter.
@ -738,8 +756,9 @@ print('No executable skill module found')
for iteration in range(self.max_iterations):
# OpenAI format response has choices[0].message
assistant_message: _ChatMessage = current_response.choices[0].message
stop_reason: str | None = current_response.choices[0].finish_reason
choice: _ChatChoice = _first_choice(current_response)
assistant_message: _ChatMessage = choice.message
stop_reason: str | None = choice.finish_reason
# Build assistant message for conversation history
assistant_msg_dict: dict[str, object] = {

View file

@ -8,7 +8,7 @@ import asyncio
import binascii
import os
import uuid
from collections.abc import Callable, Mapping, Sequence, Set
from collections.abc import Awaitable, Callable, Mapping, Sequence, Set
from contextvars import ContextVar
from dataclasses import dataclass, field
from datetime import datetime
@ -386,6 +386,12 @@ CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None]
ParallelGaugeCacheValue: TypeAlias = dict[str, object] | int | float | str | bytes
class _AsyncLuaScript(Protocol):
"""A Lua script registered against the async Redis client, called with KEYS and ARGV."""
def __call__(self, *, keys: Sequence[str], args: Sequence[object]) -> Awaitable[list[CacheCounterValue]]: ...
class RateLimitDescriptorRateLimitObject(TypedDict, total=False):
requests_per_unit: int | None
tokens_per_unit: int | None
@ -577,6 +583,14 @@ def _parse_output_cap_value(raw_value: object) -> int | None:
class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
batch_rate_limiter_script: _AsyncLuaScript | None
token_increment_script: _AsyncLuaScript | None
check_and_increment_by_n_script: _AsyncLuaScript | None
window_guarded_token_increment_script: _AsyncLuaScript | None
parallel_acquire_script: _AsyncLuaScript | None
parallel_release_script: _AsyncLuaScript | None
parallel_count_script: _AsyncLuaScript | None
def __init__(
self,
internal_usage_cache: InternalUsageCache,
@ -3855,7 +3869,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
expected_window_start = operation.get("expected_window_start")
if window_key is None or expected_window_start is None:
continue
active_window_start = await self.internal_usage_cache.async_get_cache(
active_window_start: CacheCounterValue | None = await self.internal_usage_cache.async_get_cache(
key=window_key,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
@ -4144,7 +4158,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
def _collect_tpm_scope_targets(
self,
standard_logging_metadata: dict[str, Any],
kwargs: Any,
kwargs: object,
model_group: str | None,
) -> list[tuple[str, str]]:
"""
@ -4301,8 +4315,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
def _build_success_event_pipeline_operations(
self,
kwargs: Any,
response_obj: Any,
kwargs: dict[str, Any],
response_obj: object,
rate_limit_type: Literal["output", "input", "total"],
) -> list[RedisPipelineIncrementOperation]:
"""Build Redis pipeline increment ops for TPM / parallel-request counters."""

View file

@ -543,7 +543,7 @@ def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> Pr
_raise_if_ptu_cost_attribution_disabled(updated_patch.model_info.model_dump(exclude_none=True))
merged_model_name: Final = updated_patch.model_name or db_model.model_name
merged_litellm_params: Final = db_model.litellm_params.model_dump(exclude_none=True)
merged_model_info: Final = db_model.model_info.model_dump(exclude_none=True)
merged_model_info: Final[dict[str, object]] = db_model.model_info.model_dump(exclude_none=True)
# update litellm params
if updated_patch.litellm_params:
@ -1982,7 +1982,7 @@ async def update_model(
### MERGE WITH EXISTING DATA ###
merged_dictionary: Final = {}
_mp: Final = model_params.litellm_params.dict()
_mp: Final[dict[str, object]] = model_params.litellm_params.dict()
for key, value in _mp.items():
if value is not None:

View file

@ -487,12 +487,11 @@ async def new_organization(
for m in data.models:
await can_user_call_model(m, llm_router=llm_router, user_object=user_object_correct_type)
organization_row: Final = LiteLLM_OrganizationTable(
**data.json(exclude_none=True),
object_permission_id=object_permission_id,
created_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
)
organization_payload: Final = _STR_OBJECT_DICT_ADAPTER.validate_python(data.json(exclude_none=True))
organization_payload["object_permission_id"] = object_permission_id
organization_payload["created_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
organization_payload["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
organization_row: Final = LiteLLM_OrganizationTable.model_validate(organization_payload)
for field in LiteLLM_ManagementEndpoint_MetadataFields:
if getattr(data, field, None) is not None:
@ -644,7 +643,7 @@ async def update_organization(
)
# Transform UI payload to expected format
raw_data: Final = await request.json()
raw_data: Final[dict[str, object]] = await request.json()
raw_data_with_flat_budget_fields: Final = handle_nested_budget_structure_in_organization_update_request(raw_data)
# Create validated data model
@ -691,7 +690,7 @@ async def update_organization(
# Merge metadata from existing organization with updated metadata
if updated_organization_row_json.get("metadata") is not None:
existing_metadata: Final = existing_organization_row.metadata or {}
updated_metadata: Final = updated_organization_row_json.get("metadata", {})
updated_metadata: Final[dict[str, object]] = updated_organization_row_json.get("metadata", {})
merged_metadata: Final[Mapping[str, object]] = _update_dictionary(
existing_dict=cast( # cast-ok: prisma de-serializes a Json column to the plain python dict it stores
"dict[str, object]", existing_metadata

View file

@ -502,7 +502,7 @@ def _set_nested_metadata_value(metadata: dict[str, object], key_path: str, value
placeholder: Final = "\x00"
parts = key_path.replace("\\.", placeholder).split(".")
parts = [p.replace(placeholder, ".") for p in parts]
current: Any = metadata
current: dict[str, object] = metadata
for part in parts[:-1]:
existing = current.get(part)
if not isinstance(existing, dict):
@ -4076,7 +4076,7 @@ class SSOAuthenticationHandler:
)
if resp.status_code == 200:
try:
userinfo_raw: Final = resp.json()
userinfo_raw: Final[dict[str, object] | None] = resp.json()
if not userinfo_raw:
# JSON null (None) or empty dict ({}) — no identity claims.
# Treat as failure so id_token fallback can be attempted.
@ -4406,7 +4406,7 @@ class MicrosoftSSOHandler:
) -> tuple[list[str], str | None]:
"""Helper function to fetch and parse group data from a URL"""
response: Final = await async_client.get(url, headers=headers)
response_json: Final = response.json()
response_json: Final[dict[str, object]] = response.json()
response_typed: Final = await MicrosoftSSOHandler._cast_graph_api_response_dict(response=response_json)
group_ids: Final = MicrosoftSSOHandler._get_group_ids_from_graph_api_response(response=response_typed)
return group_ids, response_typed.get("odata_nextLink")

View file

@ -267,7 +267,7 @@ class VertexPassthroughLoggingHandler:
model: Final = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
_json_response: Final = httpx_response.json()
_json_response: Final[dict[str, object]] = httpx_response.json()
litellm_prediction_response: ModelResponse | EmbeddingResponse | ImageResponse = ModelResponse()
if vertex_image_generation_class.is_image_generation_response(_json_response):
@ -422,7 +422,7 @@ class VertexPassthroughLoggingHandler:
- Creates standard logging object
- Logs in litellm callbacks
"""
kwargs: dict[str, Any] = {}
kwargs: dict[str, object] = {}
vertex_location: Final = get_vertex_location_from_url(url_route)
if vertex_location is not None:
litellm_logging_obj.optional_params["vertex_location"] = vertex_location

View file

@ -52,7 +52,7 @@ _TOOL_PAYLOAD_KEYS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
"function": ("name", "description", "parameters", "strict"),
}
)
_EMPTY_TOOL_PAYLOAD: Final[Mapping[str, Any]] = MappingProxyType({})
_EMPTY_TOOL_PAYLOAD: Final[Mapping[str, object]] = MappingProxyType({})
def _convert_tool_payload_value(key: str, value: object, *, to_chat: bool) -> object:
@ -105,7 +105,7 @@ def _normalize_tool_dialect(
return {**data, **{key: value for key, value in replaceable if key in data}} # mutable-ok: plain body dict
def _is_chat_completions_body(data: Mapping[str, Any]) -> bool:
def _is_chat_completions_body(data: Mapping[str, object]) -> bool:
messages: Final = data.get("messages")
if isinstance(messages, list) and messages:
return True
@ -1373,7 +1373,7 @@ async def _enforce_responses_ws_first_frame_model_auth(
request: Request,
model: str,
user_api_key_dict: UserAPIKeyAuth,
llm_router: Any | None,
llm_router: "Router | None",
) -> None:
from litellm.proxy.auth.user_api_key_auth import (
_enforce_key_and_fallback_model_access,
@ -1417,7 +1417,7 @@ async def _enforce_responses_ws_first_frame_model_auth(
async def responses_websocket_endpoint(
websocket: WebSocket,
model: str | None = fastapi.Query(None, description="The model to use for the responses WebSocket session."),
user_api_key_dict=Depends(user_api_key_auth_websocket),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket),
):
"""
Responses API WebSocket mode endpoint.
@ -1462,7 +1462,7 @@ async def responses_websocket_endpoint(
return
model, first_message = result
data: dict[str, Any] = {
data: dict[str, object] = {
"model": model,
"websocket": websocket,
}
@ -1471,7 +1471,7 @@ async def responses_websocket_endpoint(
# Construct a synthetic Request for pre-call processing
headers_list: Final = list(websocket.scope.get("headers") or [])
scope: Final[dict[str, Any]] = {
scope: Final[dict[str, object]] = {
"type": "http",
"method": "POST",
"path": "/v1/responses",

View file

@ -50,12 +50,12 @@ def _route_user_config_request(data: dict, route_type: str):
return ret_val
def _is_a2a_agent_model(model_name: Any) -> bool:
def _is_a2a_agent_model(model_name: object) -> bool:
"""Check if the model name is for an A2A agent (a2a/ prefix)."""
return isinstance(model_name, str) and model_name.startswith("a2a/")
def _raise_if_model_fully_blocked(llm_router: LitellmRouter, model_name: Any, team_id: str | None) -> None:
def _raise_if_model_fully_blocked(llm_router: LitellmRouter, model_name: object, team_id: str | None) -> None:
if not isinstance(model_name, str) or not model_name:
return
if not isinstance(llm_router, litellm.Router):

View file

@ -1,6 +1,6 @@
#### Video Endpoints #####
from typing import Any, Final
from typing import Final
from fastapi import APIRouter, Depends, File, Form, Request, Response, UploadFile
from fastapi.responses import ORJSONResponse
@ -161,7 +161,7 @@ async def video_list(
# Read query parameters
query_params: Final = dict(request.query_params)
data: Final[dict[str, Any]] = {"query_params": query_params}
data: Final[dict[str, object]] = {"query_params": query_params}
# Extract custom_llm_provider from headers, query params, or body
custom_llm_provider: Final = (
@ -246,7 +246,7 @@ async def video_status(
)
# Create data with video_id
data: Final[dict[str, Any]] = {"video_id": video_id}
data: Final[dict[str, object]] = {"video_id": video_id}
decoded: Final = decode_video_id_with_provider(video_id)
provider_from_id: Final = decoded.get("custom_llm_provider")
@ -345,7 +345,7 @@ async def video_content(
)
# Create data with video_id
data: Final[dict[str, Any]] = {"video_id": video_id}
data: Final[dict[str, object]] = {"video_id": video_id}
decoded: Final = decode_video_id_with_provider(video_id)
provider_from_id: Final = decoded.get("custom_llm_provider")
@ -653,7 +653,7 @@ async def video_get_character(
)
original_requested_character_id: Final = character_id
data: Final[dict[str, Any]] = {"character_id": character_id}
data: Final[dict[str, object]] = {"character_id": character_id}
decoded: Final = decode_character_id_with_provider(character_id)
provider_from_id: Final = decoded.get("custom_llm_provider")

View file

@ -29,6 +29,7 @@ from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion
from litellm.rag.ingestion.s3_vectors_ingestion import S3VectorsRAGIngestion
from litellm.rag.ingestion.vertex_ai_ingestion import VertexAIRAGIngestion
from litellm.rag.rag_query import RAGQuery
from litellm.types.llms.openai import AllMessageValues
from litellm.types.rag import (
RAGIngestOptions,
RAGIngestResponse,
@ -204,7 +205,7 @@ def _suppressed_sub_call_billing() -> Iterator[None]:
async def _execute_query_pipeline(
model: str,
messages: list[Any],
messages: list[AllMessageValues],
retrieval_config: dict[str, Any],
rerank: dict[str, Any] | None = None,
stream: bool = False,
@ -311,7 +312,7 @@ async def _execute_query_pipeline(
@client
async def aquery(
model: str,
messages: list[Any],
messages: list[AllMessageValues],
retrieval_config: dict[str, Any],
rerank: dict[str, Any] | None = None,
stream: bool = False,
@ -358,12 +359,12 @@ async def aquery(
@client
def query(
model: str,
messages: list[Any],
messages: list[AllMessageValues],
retrieval_config: dict[str, Any],
rerank: dict[str, Any] | None = None,
stream: bool = False,
**kwargs,
) -> ModelResponse | Coroutine[Any, Any, ModelResponse]:
) -> ModelResponse | Coroutine[None, None, ModelResponse]:
"""
Query a RAG pipeline.
"""
@ -410,7 +411,7 @@ def ingest(
file_id: str | None = None,
timeout: float | httpx.Timeout | None = None,
**kwargs,
) -> RAGIngestResponse | Coroutine[Any, Any, RAGIngestResponse]:
) -> RAGIngestResponse | Coroutine[None, None, RAGIngestResponse]:
"""
Ingest a document into a vector store.

View file

@ -399,7 +399,7 @@ class LiteLLM_Proxy_MCP_Handler:
@staticmethod
async def _process_mcp_tools_without_openai_transform(
user_api_key_auth: Any,
user_api_key_auth: "UserAPIKeyAuth | None",
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]],
litellm_trace_id: str | None = None,
mcp_auth_header: str | None = None,
@ -636,7 +636,7 @@ class LiteLLM_Proxy_MCP_Handler:
async def _execute_tool_calls(
tool_server_map: dict[str, str],
tool_calls: Sequence[object],
user_api_key_auth: Any,
user_api_key_auth: "UserAPIKeyAuth | None",
mcp_auth_header: str | None = None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
oauth2_headers: dict[str, str] | None = None,

View file

@ -20,6 +20,7 @@ anthropic:
import asyncio
import builtins
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
from typing import Any, Final
@ -54,19 +55,19 @@ class _LiteLLMParamsDictView:
__slots__ = ("_params",)
def __init__(self, params: dict[str, Any]):
def __init__(self, params: Mapping[str, object]):
self._params = params
def __getattr__(self, key: str) -> Any:
def __getattr__(self, key: str) -> object:
return self._params.get(key)
def __getitem__(self, key: str) -> Any:
def __getitem__(self, key: str) -> object:
return self._params.get(key)
def __contains__(self, key: str) -> bool:
return key in self._params
def get(self, key: str, default: Any = None) -> Any:
def get(self, key: str, default: object = None) -> object:
return self._params.get(key, default)
def keys(self):
@ -84,10 +85,10 @@ class _LiteLLMParamsDictView:
def __len__(self) -> int:
return len(self._params)
def dict(self) -> dict[str, Any]:
def dict(self) -> builtins.dict[str, object]:
return dict(self._params)
def model_dump(self) -> builtins.dict[str, Any]:
def model_dump(self) -> builtins.dict[str, object]:
return dict(self._params)

View file

@ -282,7 +282,7 @@ def _response_cost_or_none(response: ModelResponse) -> float | None:
return float(cost)
def _effective_turn_off_message_logging(request_kwargs: Mapping[str, Any] | None) -> bool | None:
def _effective_turn_off_message_logging(request_kwargs: Mapping[str, object] | None) -> bool | None:
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
initialize_standard_callback_dynamic_params,
)
@ -1925,7 +1925,7 @@ class ComplexityRouter(CustomLogger):
ceiling_severity: Final = self._active_tier_severity(hard_ceiling) if hard_ceiling is not None else None
best_model: str | None = None
best_score = float("-inf")
candidate_scores: Final[list[dict[str, Any]]] = []
candidate_scores: Final[list[dict[str, object]]] = []
for model in candidates:
if floor_severity is not None and all(
self._active_tier_severity(model_tier) < floor_severity

View file

@ -1,7 +1,9 @@
import os
from typing import Any, Final
from collections.abc import Mapping
from typing import Final, Protocol
import httpx
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
@ -17,6 +19,72 @@ from litellm.proxy._types import KeyManagementSystem
from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name
class _VaultAuthData(TypedDict):
"""The ``auth`` block Vault returns from a login endpoint."""
client_token: ReadOnly[str]
lease_duration: ReadOnly[int]
class _VaultLoginResponse(TypedDict):
"""Body of a Vault ``/v1/auth/.../login`` response."""
auth: ReadOnly[_VaultAuthData]
class _VaultSecretTarget(TypedDict):
"""Resolved coordinates of one Vault KV v2 secret."""
url: ReadOnly[str]
data_key: ReadOnly[str]
secret_name: ReadOnly[str]
class _VaultSecretDataBlock(TypedDict, total=False):
"""The inner ``data`` block of a Vault KV v2 read body."""
data: ReadOnly[Mapping[str, object]]
class _VaultSecretReadResponse(TypedDict, total=False):
"""Body of a Vault KV v2 secret read, narrowed to the nesting this module walks."""
data: ReadOnly[_VaultSecretDataBlock]
class _VaultLoginResponseSource(Protocol):
"""A Vault login call's HTTP response, read for the auth block it carries."""
def json(self) -> _VaultLoginResponse: ...
class _VaultSecretReadSource(Protocol):
"""A Vault KV v2 read response, read for the nested secret data it carries."""
def json(self) -> _VaultSecretReadResponse: ...
class _JsonObjectSource(Protocol):
"""A Vault response whose body is a JSON object nothing further is assumed about."""
def json(self) -> dict[str, object]: ...
def _vault_login_body(response: _VaultLoginResponseSource) -> _VaultLoginResponse:
"""Decode the body of a Vault login response."""
return response.json()
def _vault_secret_read_body(response: _VaultSecretReadSource) -> _VaultSecretReadResponse:
"""Decode the body of a Vault KV v2 secret read response."""
return response.json()
def _json_object_body(response: _JsonObjectSource) -> dict[str, object]:
"""Decode a Vault response body as a plain JSON object."""
return response.json()
class HashicorpSecretManager(BaseSecretManager):
def __init__(self):
from litellm.proxy.proxy_server import CommonProxyErrors, premium_user
@ -130,7 +198,8 @@ class HashicorpSecretManager(BaseSecretManager):
)
resp.raise_for_status()
auth_data: Final = resp.json()["auth"]
login_response: Final = _vault_login_body(resp)
auth_data: Final = login_response["auth"]
token: Final = auth_data["client_token"]
_lease_duration: Final = auth_data["lease_duration"]
@ -191,8 +260,10 @@ class HashicorpSecretManager(BaseSecretManager):
json=self._get_tls_cert_auth_body(),
)
resp.raise_for_status()
token: Final = resp.json()["auth"]["client_token"]
_lease_duration: Final = resp.json()["auth"]["lease_duration"]
token_response: Final = _vault_login_body(resp)
token: Final = token_response["auth"]["client_token"]
lease_response: Final = _vault_login_body(resp)
_lease_duration: Final = lease_response["auth"]["lease_duration"]
verbose_logger.debug("Successfully obtained Vault token via TLS cert auth.")
self.cache.set_cache(key="hcp_vault_token", value=token, ttl=_lease_duration)
return token
@ -205,9 +276,9 @@ class HashicorpSecretManager(BaseSecretManager):
def get_url(
self,
secret_name: str,
namespace: str | None = None,
mount_name: str | None = None,
path_prefix: str | None = None,
namespace: object = None,
mount_name: object = None,
path_prefix: object = None,
) -> str:
"""
Constructs the Vault URL for KV v2 secrets.
@ -238,7 +309,7 @@ class HashicorpSecretManager(BaseSecretManager):
_url += secret_name
return _url
def _sanitize_plain_value(self, value: str | int | None) -> str | None:
def _sanitize_plain_value(self, value: object) -> str | None:
if value is None:
return None
value_str: Final = str(value).strip()
@ -246,23 +317,23 @@ class HashicorpSecretManager(BaseSecretManager):
return None
return value_str
def _sanitize_path_component(self, value: str | int | None) -> str | None:
def _sanitize_path_component(self, value: object) -> str | None:
sanitized_value = self._sanitize_plain_value(value)
if sanitized_value is None:
return None
sanitized_value = sanitized_value.strip("/")
return sanitized_value or None
def _extract_secret_manager_settings(self, optional_params: dict | None) -> dict[str, Any]:
def _extract_secret_manager_settings(self, optional_params: dict | None) -> dict[str, object]:
if not isinstance(optional_params, dict):
return {}
candidate: Final = optional_params.get("secret_manager_settings")
source: Final = candidate if isinstance(candidate, dict) else optional_params
source: Final[Mapping[str, object]] = candidate if isinstance(candidate, dict) else optional_params
allowed_keys: Final = {"namespace", "mount", "path_prefix", "data"}
return {k: source[k] for k in allowed_keys if k in source}
def _build_secret_target(self, secret_name: str, optional_params: dict | None) -> dict[str, Any]:
def _build_secret_target(self, secret_name: str, optional_params: dict | None) -> _VaultSecretTarget:
settings: Final = self._extract_secret_manager_settings(optional_params)
namespace: Final = settings.get("namespace", self.vault_namespace)
@ -331,7 +402,7 @@ class HashicorpSecretManager(BaseSecretManager):
response.raise_for_status()
# For KV v2, the secret is in response.json()["data"]["data"]
json_resp: Final = response.json()
json_resp: Final = _json_object_body(response)
_value: Final = self._get_secret_value_from_json_response(json_resp)
self.cache.set_cache(secret_name, _value)
return _value
@ -362,7 +433,7 @@ class HashicorpSecretManager(BaseSecretManager):
response.raise_for_status()
# For KV v2, the secret is in response.json()["data"]["data"]
json_resp: Final = response.json()
json_resp: Final = _json_object_body(response)
_value: Final = self._get_secret_value_from_json_response(json_resp)
self.cache.set_cache(secret_name, _value)
return _value
@ -379,7 +450,7 @@ class HashicorpSecretManager(BaseSecretManager):
optional_params: dict | None = None,
timeout: float | httpx.Timeout | None = None,
tags: dict | list | None = None,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Writes a secret to Vault KV v2 using an async HTTPX client.
@ -413,7 +484,7 @@ class HashicorpSecretManager(BaseSecretManager):
json=data,
)
response.raise_for_status()
return response.json()
return _json_object_body(response)
except Exception as e:
verbose_logger.exception("Error writing secret to Hashicorp Vault: %s", e)
return {"status": "error", "message": str(e)}
@ -500,7 +571,7 @@ class HashicorpSecretManager(BaseSecretManager):
headers=self._get_request_headers(),
)
response.raise_for_status()
json_resp: Final = response.json()
json_resp: Final = _vault_secret_read_body(response)
# Use data_key from target to get the correct value
data_key: Final = new_target["data_key"]
new_secret_value_from_vault: Final = json_resp.get("data", {}).get("data", {}).get(data_key, None)

View file

@ -327,6 +327,10 @@ class BatchGuardrailReport(BaseModel):
"""Every record that was redacted or dropped, in file order."""
_JsonValue: TypeAlias = object
"""Alias for ``object``, usable inside model bodies that declare a field named ``object``."""
BATCH_GUARDRAIL_RESPONSE_FIELD: Final = "litellm_batch_guardrail"
@ -1191,7 +1195,7 @@ class ShellToolParam(TypedDict, total=False):
type: Required[Literal["shell"] | str]
"""The type of tool. Use ``\"shell\"``."""
environment: Required[dict[str, Any]]
environment: Required[dict[str, object]]
"""Environment config: ``type`` (e.g. ``\"container_auto\"``, ``\"container_reference\"``, ``\"local\"``), optional ``container_id``, ``network_policy``, ``domain_secrets``, ``skills``."""
@ -1308,7 +1312,7 @@ class ResponseAPIUsage(BaseLiteLLMOpenAIResponseObject):
@field_validator("cost", mode="before")
@classmethod
def parse_cost(cls, v: Any) -> float | None:
def parse_cost(cls, v: object) -> object:
"""Normalise cost: accept either a float or a dict with a ``total_cost`` key."""
if isinstance(v, dict):
return v.get("total_cost")
@ -1805,7 +1809,7 @@ class ErrorEventError(BaseLiteLLMOpenAIResponseObject):
type: str # e.g., 'invalid_request_error'
code: str # e.g., 'context_length_exceeded'
message: str
param: str | dict[str, Any] | None = None
param: str | dict[str, object] | None = None
class ErrorEvent(BaseLiteLLMOpenAIResponseObject):
@ -2418,7 +2422,7 @@ class OpenAIVideoObject(BaseModel):
expires_at: int | None = None
"""Unix timestamp (seconds) for when the downloadable assets expire, if set."""
error: dict[str, Any] | None = None
error: dict[str, _JsonValue] | None = None
"""Error payload that explains why generation failed, if applicable."""
progress: int | None = None
@ -2436,15 +2440,15 @@ class OpenAIVideoObject(BaseModel):
model: str | None = None
"""The video generation model that produced the job."""
_hidden_params: dict[str, Any] = {}
_hidden_params: dict[str, _JsonValue] = {}
def __contains__(self, key) -> bool:
return hasattr(self, key)
def get(self, key, default=None):
def get(self, key, default=None) -> _JsonValue:
return getattr(self, key, default)
def __getitem__(self, key):
def __getitem__(self, key) -> _JsonValue:
return getattr(self, key)
def json(self, **kwargs):

View file

@ -369,7 +369,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
@model_validator(mode="before")
@classmethod
def preprocess_input_data(cls, data: Any) -> Any:
def preprocess_input_data(cls, data: object) -> object:
"""
Pre-process input data before validation:
1. Filter out reserved Python keywords ('self', 'params', '__class__') to prevent
@ -627,6 +627,11 @@ class AlertingConfig(BaseModel):
alerting_threshold: float | None = 300
def _resolved_annotations(model_class: type[object]) -> Mapping[str, object]:
"""Resolve a class's annotations, keeping each resolved annotation opaque."""
return get_type_hints(model_class)
class ModelGroupInfo(BaseModel):
model_group: str
providers: list[str]
@ -655,7 +660,7 @@ class ModelGroupInfo(BaseModel):
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None
def __init__(self, **data) -> None:
for field_name, field_type in get_type_hints(self.__class__).items():
for field_name, field_type in _resolved_annotations(self.__class__).items():
if field_type is bool and data.get(field_name) is None:
data[field_name] = False
super().__init__(**data)

View file

@ -112,7 +112,9 @@ class VectorStoreRegistry:
Dynamically extracts all parameters defined in VECTOR_STORE_OPENAI_PARAMS.
"""
# Get the list of supported param names from the Literal type
supported_params: Final = get_args(VECTOR_STORE_OPENAI_PARAMS)
supported_params: Final = tuple(
param for param in get_args(VECTOR_STORE_OPENAI_PARAMS) if isinstance(param, str)
)
# Extract only the params that exist in the tool
kwargs: Final = {param: tool.get(param) for param in supported_params if param in tool}
@ -503,7 +505,7 @@ class VectorStoreRegistry:
vector_stores_from_db.append(_litellm_managed_vector_store)
return vector_stores_from_db
def get_credentials_for_vector_store(self, vector_store_id: str) -> dict[str, Any]:
def get_credentials_for_vector_store(self, vector_store_id: str) -> dict[str, object]:
"""
Get the credentials for a vector store

View file

@ -1,6 +1,6 @@
{
"ANN001": {
"limit": 2991
"limit": 2985
},
"ANN002": {
"limit": 71
@ -9,13 +9,13 @@
"limit": 809
},
"ANN201": {
"limit": 2002
"limit": 2001
},
"ANN202": {
"limit": 841
"limit": 835
},
"ANN204": {
"limit": 694
"limit": 693
},
"ANN205": {
"limit": 112
@ -24,7 +24,7 @@
"limit": 133
},
"ANN401": {
"limit": 387
"limit": 307
},
"ASYNC230": {
"limit": 11
@ -231,7 +231,7 @@
"limit": 5
},
"TID251": {
"limit": 1084
"limit": 1073
},
"TRY002": {
"limit": 524
@ -246,7 +246,7 @@
"limit": 113
},
"TRY300": {
"limit": 855
"limit": 854
},
"UP028": {
"limit": 2

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22403
"limit": 22367
},
"LIT002": {
"limit": 26780
"limit": 26777
},
"LIT003": {
"limit": 269
@ -27,10 +27,10 @@
"limit": 0
},
"LIT010": {
"limit": 16512
"limit": 16507
},
"LIT011": {
"limit": 5537
"limit": 5535
},
"LIT012": {
"limit": 4495