Merge pull request #40251 from BerriAI/litellm_decrease_anys_opus5_r5

refactor(types): remove 1,173 Any errors across 169 backend files
This commit is contained in:
Mateo Wang 2026-09-21 15:11:44 -07:00 • committed by GitHub
commit e7bff277a6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
170 changed files with 1413 additions and 830 deletions

View file

@ -2,10 +2,11 @@
Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if the cost has been tracked.
"""
from collections.abc import Sequence
from dataclasses import replace as dataclasses_replace
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, cast
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Protocol, Tuple, cast
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -18,8 +19,8 @@ if TYPE_CHECKING:
from prisma import models as prisma_models
from litellm.integrations.prometheus import PrometheusLogger
from litellm.proxy._types import LiteLLM_ManagedObjectTable
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.prisma_protocols import TableActions
from litellm.router import Router
from litellm.types.router import Deployment
from litellm.types.utils import LiteLLMBatch
@ -41,6 +42,42 @@ TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
)
class _ManagedObjectRow(Protocol):
@property
def id(self) -> str: ...
@property
def unified_object_id(self) -> str: ...
@property
def created_by(self) -> str | None: ...
@property
def file_object(self) -> object: ...
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
return table
def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]":
table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable
return table
def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]":
table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = (
prisma_client.db.litellm_verificationtoken
)
return table
def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]":
table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable
return table
class CheckBatchCost:
def __init__(
self,
@ -73,7 +110,7 @@ class CheckBatchCost:
inline for a batch the first poll cycle then accounts again.
"""
try:
await self.prisma_client.db.litellm_managedobjecttable.find_first(
await _managed_object_table(self.prisma_client).find_first(
where={"file_purpose": "batch", "batch_processed": False}
)
except Exception as probe_err:
@ -97,10 +134,8 @@ class CheckBatchCost:
if not user_id:
return {}
try:
user_row: prisma_models.LiteLLM_UserTable | None = (
await self.prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_id}
)
user_row: prisma_models.LiteLLM_UserTable | None = await _user_table(self.prisma_client).find_unique(
where={"user_id": user_id}
)
if user_row is None:
return {}
@ -117,11 +152,9 @@ class CheckBatchCost:
if not api_key:
return None
try:
key_row: prisma_models.LiteLLM_VerificationToken | None = (
await self.prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": api_key}
)
)
key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table(
self.prisma_client
).find_unique(where={"token": api_key})
return getattr(key_row, "key_alias", None) if key_row is not None else None
except Exception as e:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up key alias for batch {batch_id}: {e}")
@ -132,17 +165,15 @@ class CheckBatchCost:
if not team_id:
return None
try:
team_row: prisma_models.LiteLLM_TeamTable | None = (
await self.prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
)
team_row: prisma_models.LiteLLM_TeamTable | None = await _team_table(self.prisma_client).find_unique(
where={"team_id": team_id}
)
return getattr(team_row, "team_alias", None) if team_row is not None else None
except Exception as e:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up team alias for team {team_id}: {e}")
return None
async def _get_org_id(self, job: "LiteLLM_ManagedObjectTable", batch_id: str) -> str | None:
async def _get_org_id(self, job: "_ManagedObjectRow", batch_id: str) -> str | None:
org_id = getattr(job, "org_id", None)
if org_id:
return org_id
@ -150,11 +181,9 @@ class CheckBatchCost:
team_id = getattr(job, "team_id", None)
if api_key:
try:
key_row: prisma_models.LiteLLM_VerificationToken | None = (
await self.prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": api_key}
)
)
key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table(
self.prisma_client
).find_unique(where={"token": api_key})
key_org_id = getattr(key_row, "organization_id", None) if key_row is not None else None
if key_org_id:
return key_org_id
@ -166,10 +195,8 @@ class CheckBatchCost:
if not team_id:
return None
try:
team_row: prisma_models.LiteLLM_TeamTable | None = (
await self.prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
)
team_row: prisma_models.LiteLLM_TeamTable | None = await _team_table(self.prisma_client).find_unique(
where={"team_id": team_id}
)
return getattr(team_row, "organization_id", None) if team_row is not None else None
except Exception as e:
@ -177,7 +204,7 @@ class CheckBatchCost:
return None
async def _build_creator_attribution_metadata(
self, job: "LiteLLM_ManagedObjectTable", batch_id: str
self, job: "_ManagedObjectRow", batch_id: str
) -> dict[str, object]:
"""
Rebuild the spend-tracking metadata for the key, team, and tags that created the
@ -225,7 +252,7 @@ class CheckBatchCost:
should not be polled.
"""
cutoff: Final = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
result: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many(
result: Final = await _managed_object_table(self.prisma_client).update_many(
where={
"file_purpose": "batch",
"status": {"not_in": list(TERMINAL_MANAGED_OBJECT_STATUSES)},
@ -244,7 +271,7 @@ class CheckBatchCost:
# A row already in a terminal status is never rewritten by the sweep above, so
# without this it keeps a poll-page slot forever and starves newer batches.
retired: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many(
retired: Final = await _managed_object_table(self.prisma_client).update_many(
where={
"file_purpose": "batch",
"batch_processed": False,
@ -259,9 +286,9 @@ class CheckBatchCost:
f"{MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days that were never costed"
)
async def _fallback_find_jobs(self) -> list:
async def _fallback_find_jobs(self) -> "Sequence[_ManagedObjectRow]":
"""Query batch jobs without the batch_processed filter (for older schemas)."""
return await self.prisma_client.db.litellm_managedobjecttable.find_many(
return await _managed_object_table(self.prisma_client).find_many(
where={
"file_purpose": "batch",
"status": {
@ -279,7 +306,7 @@ class CheckBatchCost:
order={"created_at": "asc"},
)
async def _retire_job(self, job: "LiteLLM_ManagedObjectTable", reason: str) -> None:
async def _retire_job(self, job: "_ManagedObjectRow", reason: str) -> None:
"""
Take a row that can never be costed out of the poll page. Leaving it selectable
would burn one of the MAX_OBJECTS_PER_POLL_CYCLE slots on every future cycle, and
@ -292,7 +319,7 @@ class CheckBatchCost:
else {"status": "stale_expired"}
)
try:
await self.prisma_client.db.litellm_managedobjecttable.update(
await _managed_object_table(self.prisma_client).update(
where={"id": job.id},
data=data,
)
@ -306,7 +333,7 @@ class CheckBatchCost:
"so it will no longer be polled"
)
async def _claim_job_for_costing(self, job: "LiteLLM_ManagedObjectTable") -> bool:
async def _claim_job_for_costing(self, job: "_ManagedObjectRow") -> bool:
"""
Atomically flip batch_processed from false to true, returning whether this pod won
the row. Every pod and uvicorn worker schedules its own poller against the shared
@ -321,7 +348,7 @@ class CheckBatchCost:
if not self._has_batch_processed_column:
return True
try:
claimed: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many(
claimed: Final = await _managed_object_table(self.prisma_client).update_many(
where={"id": job.id, "batch_processed": False},
data={"batch_processed": True},
)
@ -332,7 +359,7 @@ class CheckBatchCost:
return False
return claimed > 0
async def _release_job_claim(self, job: "LiteLLM_ManagedObjectTable") -> None:
async def _release_job_claim(self, job: "_ManagedObjectRow") -> None:
"""Give a claimed row back once billing it failed, so a later poll cycle retries it.
Safe to match on batch_processed=True: while this poller is active the retrieve
@ -342,7 +369,7 @@ class CheckBatchCost:
if not self._has_batch_processed_column:
return
try:
await self.prisma_client.db.litellm_managedobjecttable.update_many(
await _managed_object_table(self.prisma_client).update_many(
where={"id": job.id, "batch_processed": True},
data={"batch_processed": False},
)
@ -353,7 +380,7 @@ class CheckBatchCost:
)
@staticmethod
def _has_unified_id_without_model(job: "LiteLLM_ManagedObjectTable") -> bool:
def _has_unified_id_without_model(job: "_ManagedObjectRow") -> bool:
"""A unified id that decodes but carries no model_id can never be routed."""
from litellm.proxy.openai_files_endpoints.common_utils import (
convert_b64_uid_to_unified_uid,
@ -402,7 +429,7 @@ class CheckBatchCost:
return isinstance(error, (NotFoundError, openai.NotFoundError)) and output_file_id in str(error)
async def _finalize_unbilled_terminal_job(
self, job: "prisma_models.LiteLLM_ManagedObjectTable", response: "LiteLLMBatch"
self, job: "_ManagedObjectRow", response: "LiteLLMBatch"
) -> None:
"""Persist a terminal batch that has nothing billable, converting any raw
provider file ids to managed ids, and take it out of the poll page."""
@ -426,7 +453,7 @@ class CheckBatchCost:
"file_object": response.model_dump_json(),
**({"batch_processed": True} if self._has_batch_processed_column else {}),
}
await self.prisma_client.db.litellm_managedobjecttable.update(
await _managed_object_table(self.prisma_client).update(
where={"id": job.id},
data=update_data,
)
@ -447,7 +474,7 @@ class CheckBatchCost:
def _resolve_job_routing(
self,
job: "LiteLLM_ManagedObjectTable",
job: "_ManagedObjectRow",
prom_logger: Optional["PrometheusLogger"],
) -> Optional[Tuple[str, str]]:
"""
@ -524,7 +551,7 @@ class CheckBatchCost:
def _resolve_unmanaged_provider_routing(
self,
job: "LiteLLM_ManagedObjectTable",
job: "_ManagedObjectRow",
prom_logger: Optional["PrometheusLogger"],
llm_provider: str,
bare_model_name: str,
@ -620,7 +647,7 @@ class CheckBatchCost:
@classmethod
def _get_managed_file_model_name(
cls,
job: "LiteLLM_ManagedObjectTable",
job: "_ManagedObjectRow",
deployment_info: "Deployment",
) -> Optional[str]:
"""
@ -640,7 +667,7 @@ class CheckBatchCost:
)
@staticmethod
def _get_input_file_id(job: "LiteLLM_ManagedObjectTable") -> Optional[str]:
def _get_input_file_id(job: "_ManagedObjectRow") -> Optional[str]:
import json
from litellm.types.utils import LiteLLMBatch
@ -660,7 +687,7 @@ class CheckBatchCost:
async def _track_completed_batch_cost(
self,
job: "LiteLLM_ManagedObjectTable",
job: "_ManagedObjectRow",
response: "LiteLLMBatch",
model_id: str,
batch_id: str,
@ -936,7 +963,7 @@ class CheckBatchCost:
# endpoint may transition a batch to "complete" before
# CheckBatchCost runs. The batch_processed=False filter
# already prevents reprocessing finished batches.
jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many(
jobs = await _managed_object_table(self.prisma_client).find_many(
where={
"file_purpose": "batch",
"batch_processed": False,
@ -1038,7 +1065,7 @@ class CheckBatchCost:
}
if self._has_batch_processed_column:
update_data["batch_processed"] = True
await self.prisma_client.db.litellm_managedobjecttable.update(
await _managed_object_table(self.prisma_client).update(
where={"id": job.id},
data=update_data,
)

View file

@ -6,7 +6,7 @@ same route are non-inference and free.
"""
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Dict, Optional, cast
from typing import TYPE_CHECKING, Dict, Final, Optional, Protocol, cast
import litellm
from litellm._logging import verbose_proxy_logger
@ -22,11 +22,31 @@ from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.prisma_protocols import TableActions
from litellm.router import Router
TERMINAL_RESPONSE_STATUSES = frozenset({"completed", "failed", "cancelled", "incomplete"})
class _ManagedObjectRow(Protocol):
@property
def id(self) -> str: ...
@property
def unified_object_id(self) -> str: ...
@property
def created_by(self) -> str | None: ...
@property
def file_object(self) -> object: ...
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
return table
class CheckResponsesCost:
def __init__(
self,
@ -128,7 +148,7 @@ class CheckResponsesCost:
f"CheckResponsesCost: stale cleanup failed (poll will continue): {cleanup_err}"
)
jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many(
jobs = await _managed_object_table(self.prisma_client).find_many(
where={
"status": {"in": ["queued", "in_progress"]},
"file_purpose": "response",
@ -138,7 +158,7 @@ class CheckResponsesCost:
)
verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check")
completed_jobs = []
completed_jobs: Final[list[_ManagedObjectRow]] = []
for job in jobs:
unified_object_id = job.unified_object_id
@ -189,7 +209,7 @@ class CheckResponsesCost:
# Mark completed jobs in the database
if len(completed_jobs) > 0:
await self.prisma_client.db.litellm_managedobjecttable.update_many(
await _managed_object_table(self.prisma_client).update_many(
where={"id": {"in": [job.id for job in completed_jobs]}},
data={"status": "completed"},
)

View file

@ -481,10 +481,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"""
if self.prisma_client is None:
return
managed_object = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"OR": [{"unified_object_id": object_id}, {"model_object_id": object_id}]}
)
managed_object = await _managed_object_table(self.prisma_client).find_first(
where={"OR": [{"unified_object_id": object_id}, {"model_object_id": object_id}]}
)
if managed_object is None:
return
@ -509,10 +507,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"""
if self.prisma_client is None:
return
managed_file = (
await self.prisma_client.db.litellm_managedfiletable.find_first(
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
)
managed_file = await _managed_file_table(self.prisma_client).find_first(
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
)
if managed_file is None:
return
@ -535,8 +531,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
provider_file_ids = tuple(
file_id
for file_id in (
getattr(response, "output_file_id", None),
getattr(response, "error_file_id", None),
response.output_file_id,
response.error_file_id,
)
if file_id and not _is_base64_encoded_unified_file_id(file_id)
)
@ -544,10 +540,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
return
if self.prisma_client is None:
return
batch_row = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"unified_object_id": response.id}
)
batch_row = await _managed_object_table(self.prisma_client).find_first(
where={"unified_object_id": response.id}
)
if batch_row is None or (
batch_row.created_by is None and batch_row.team_id is None

View file

@ -23,7 +23,7 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig):
params: dict[str, Any],
api_base: str | None = None,
**kwargs: Any,
) -> dict[str, Any]:
) -> dict[str, object]:
"""Handle non-streaming request to Pydantic AI agent."""
if api_base is None:
raise ValueError("api_base is required for PydanticAIProviderConfig")
@ -41,7 +41,7 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig):
params: dict[str, Any],
api_base: str | None = None,
**kwargs,
) -> AsyncIterator[dict[str, Any]]:
) -> AsyncIterator[dict[str, object]]:
"""Handle streaming request with fake streaming."""
if not api_base:
raise ValueError("api_base is required for Pydantic AI agents")

View file

@ -81,7 +81,7 @@ class Cache:
s3_aws_access_key_id: str | None = None,
s3_aws_secret_access_key: str | None = None,
s3_aws_session_token: str | None = None,
s3_config: Any | None = None,
s3_config: object | None = None,
s3_path: str | None = None,
gcs_bucket_name: str | None = None,
gcs_path_service_account: str | None = None,

View file

@ -74,7 +74,7 @@ class CachingHandlerResponse(BaseModel):
For embeddings there can be a cache hit for some of the inputs in the list and a cache miss for others
"""
cached_result: Any | None = None
cached_result: object | None = None
final_embedding_cached_response: EmbeddingResponse | None = None
embedding_all_elements_cache_hit: bool = False # this is set to True when all elements in the list have a cache hit in the embedding cache, if true return the final_embedding_cached_response no need to make an API call
@ -722,7 +722,7 @@ class LLMCachingHandler:
async def _retrieve_from_cache(
self, call_type: str, kwargs: dict[str, object], args: tuple[object, ...]
) -> Any | None:
) -> object | None:
"""
Internal method to
- get cache key
@ -968,7 +968,7 @@ class LLMCachingHandler:
def _convert_cached_stream_response(
self,
cached_result: Any,
cached_result: dict[str, object],
call_type: str,
logging_obj: LiteLLMLoggingObj,
model: str,
@ -997,7 +997,7 @@ class LLMCachingHandler:
async def async_set_cache(
self,
result: Any,
result: object,
original_function: Callable,
kwargs: dict[str, Any],
args: tuple[object, ...] | None = None,
@ -1065,7 +1065,7 @@ class LLMCachingHandler:
def sync_set_cache(
self,
result: Any,
result: object,
kwargs: dict[str, object],
args: tuple[object, ...] | None = None,
):

View file

@ -981,7 +981,7 @@ class RedisCache(BaseCache):
client: object = None,
) -> object:
async def execute() -> object:
executor: Callable[..., Awaitable[Any]] | None = litellm.in_memory_llm_clients_cache.get_cache(
executor: Callable[..., Awaitable[object]] | None = litellm.in_memory_llm_clients_cache.get_cache(
key=script_cache_key
)
if executor is None:
@ -993,7 +993,7 @@ class RedisCache(BaseCache):
return run_script
def _register_script_for_current_loop(self, script: str) -> Callable[..., Awaitable[Any]]:
def _register_script_for_current_loop(self, script: str) -> Callable[..., Awaitable[object]]:
"""
Register the script against the current event loop's Redis client.

View file

@ -2,7 +2,7 @@
Handler for transforming /chat/completions api requests to litellm.responses requests
"""
from collections.abc import Coroutine
from collections.abc import AsyncIterable, Coroutine, Iterable
from typing import TYPE_CHECKING, Any, Final, Union
from typing_extensions import TypedDict
@ -74,7 +74,7 @@ class ResponsesToCompletionBridgeHandler:
existing.setdefault(key, value)
return response
def _collect_response_from_stream(self, stream_iter: Any) -> "ResponsesAPIResponse":
def _collect_response_from_stream(self, stream_iter: Iterable[object]) -> "ResponsesAPIResponse":
for _ in stream_iter:
pass
@ -89,7 +89,7 @@ class ResponsesToCompletionBridgeHandler:
raise ValueError("Stream completed response is invalid")
return response
async def _collect_response_from_stream_async(self, stream_iter: Any) -> "ResponsesAPIResponse":
async def _collect_response_from_stream_async(self, stream_iter: AsyncIterable[object]) -> "ResponsesAPIResponse":
async for _ in stream_iter:
pass

View file

@ -6,7 +6,7 @@ import json
import os
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast, get_args
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, Union, cast, get_args
from openai.types.chat import ChatCompletion
from openai.types.responses import Response
@ -52,7 +52,7 @@ from litellm.types.llms.openai import (
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
if TYPE_CHECKING:
from openai.types.responses import ResponseInputImageParam
from openai.types.responses import ResponseInputImageParam, ResponseOutputItem
from openai.types.responses.response_text_config_param import (
ResponseTextConfigParam as ResponseText,
)
@ -197,6 +197,9 @@ def _as_chat_reasoning_items(
return cast(list[ChatCompletionReasoningItem], list(reasoning_items))
_ToolChoiceT = TypeVar("_ToolChoiceT")
def _map_incomplete_reason_to_finish_reason(incomplete_reason: str | None) -> Literal["length", "content_filter"]:
if incomplete_reason == "content_filter":
return "content_filter"
@ -291,7 +294,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
def __init__(self):
pass
def _normalize_tool_choice_for_responses_api(self, tool_choice: Any) -> Any:
def _normalize_tool_choice_for_responses_api(
self, tool_choice: _ToolChoiceT
) -> _ToolChoiceT | ToolChoiceFunctionParam | ToolChoiceCustomParam | Literal["auto", "none", "required"]:
"""Chat tool_choice nests the name under function/custom; Responses API expects top-level name."""
if not isinstance(tool_choice, dict):
return tool_choice
@ -497,7 +502,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
responses_api_request["max_output_tokens"] = value
elif key == "tools" and value is not None:
responses_api_request["tools"] = self._convert_tools_to_responses_format(
cast(list[dict[str, Any]], value)
cast(list[dict[str, object]], value)
)
elif key == "response_format":
text_format = self._transform_response_format_to_text_format(value)
@ -828,7 +833,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
response_output: Final = response_payload.get("output")
if not isinstance(response_output, list) or len(response_output) == 0:
return None
return cast(list[dict[str, Any]], response_output)
return cast(list[dict[str, object]], response_output)
@classmethod
def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, object]]:
@ -911,10 +916,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
output_items = raw_response.output
if len(output_items) == 0:
recovered_output_items: Final = self._recover_output_items_from_logging(logging_obj)
recovered_output_items: Final[list[ResponseOutputItem | dict[str, object]]] = [
*self._recover_output_items_from_logging(logging_obj)
]
if recovered_output_items:
output_items = cast(Any, recovered_output_items)
raw_response.output = cast(Any, recovered_output_items)
output_items = recovered_output_items
raw_response.output = recovered_output_items
verbose_logger.warning(
"Recovered empty Responses API output from raw SSE for model=%s",
model,
@ -1110,7 +1117,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
verbose_logger.debug("Chat provider: Other content type -> %s", result)
return result
def _convert_tools_to_responses_format(self, tools: list[dict[str, Any]]) -> list["ALL_RESPONSES_API_TOOL_PARAMS"]:
def _convert_tools_to_responses_format(
self, tools: list[dict[str, object]]
) -> list["ALL_RESPONSES_API_TOOL_PARAMS"]:
"""Convert chat completion tools to responses API tools format"""
responses_tools: Final[list[ALL_RESPONSES_API_TOOL_PARAMS]] = []
for tool in tools:
@ -1126,12 +1135,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
description=function_tool.get("description"),
)
)
elif tool.get("type") == "custom" and isinstance(tool.get("custom"), dict):
elif tool.get("type") == "custom" and isinstance(custom_payload := tool.get("custom"), dict):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_custom_tool_format_to_responses_shape,
)
custom_payload = tool["custom"]
flat_custom = CustomToolParam(type="custom", name=custom_payload.get("name", ""))
if custom_payload.get("description") is not None:
flat_custom["description"] = custom_payload["description"]

View file

@ -353,7 +353,7 @@ def cost_per_token(
data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
### VERTEX LOCATION ###
vertex_location: str | None = None, # for Vertex AI regional-endpoint uplift (e.g. "us-east5", "global")
response: Any | None = None,
response: object | None = None,
### REQUEST MODEL ###
request_model: str | None = None, # original request model for router detection
custom_model_info: OCRPricing | None = None,
@ -609,7 +609,7 @@ def cost_per_token(
model=model,
custom_llm_provider=custom_llm_provider,
number_of_queries=number_of_queries or 1,
optional_params=(response._hidden_params if response and hasattr(response, "_hidden_params") else None),
optional_params=(getattr(response, "_hidden_params", None) if response else None),
)
elif custom_llm_provider == "vertex_ai":
cost_router: Final = google_cost_router(
@ -999,7 +999,7 @@ def _is_known_usage_objects(usage_obj):
)
def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: Any) -> CallTypesLiteral | None:
def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: object) -> CallTypesLiteral | None:
if call_type is not None:
return call_type

View file

@ -7,7 +7,8 @@ import json
import os
import random
import types
from typing import Any, Final
from collections.abc import Mapping
from typing import Final
import httpx
from pydantic import BaseModel
@ -69,7 +70,7 @@ class ArgillaLogger(CustomBatchLogger):
self.flush_lock = asyncio.Lock()
super().__init__(**kwargs, flush_lock=self.flush_lock)
def validate_argilla_transformation_object(self, argilla_transformation_object: dict[str, Any]):
def validate_argilla_transformation_object(self, argilla_transformation_object: Mapping[str, object]):
if not isinstance(argilla_transformation_object, dict):
raise Exception("'argilla_transformation_object' must be a dictionary, to log your payload to Argilla.")
@ -115,7 +116,7 @@ class ArgillaLogger(CustomBatchLogger):
ARGILLA_DATASET_NAME=_credentials_dataset_name,
)
def get_chat_messages(self, payload: StandardLoggingPayload) -> list[dict[str, Any]]:
def get_chat_messages(self, payload: StandardLoggingPayload) -> list[dict[str, object]]:
payload_messages: Final = payload.get("messages", None)
if payload_messages is None:

View file

@ -139,13 +139,13 @@ class BraintrustLogger(CustomLogger):
):
output = None
elif response_obj is not None and isinstance(response_obj, litellm.ModelResponse):
output = response_obj["choices"][0]["message"].json()
output = response_obj.choices[0].message.json()
choices = response_obj["choices"]
elif response_obj is not None and isinstance(response_obj, litellm.TextCompletionResponse):
output = response_obj.choices[0].text
choices = response_obj.choices
elif response_obj is not None and isinstance(response_obj, litellm.ImageResponse):
output = response_obj["data"]
output = response_obj.data
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
dynamic_metadata: Final = litellm_params.get("metadata", {}) or {}
@ -264,13 +264,13 @@ class BraintrustLogger(CustomLogger):
):
output = None
elif response_obj is not None and isinstance(response_obj, litellm.ModelResponse):
output = response_obj["choices"][0]["message"].json()
output = response_obj.choices[0].message.json()
choices = response_obj["choices"]
elif response_obj is not None and isinstance(response_obj, litellm.TextCompletionResponse):
output = response_obj.choices[0].text
choices = response_obj.choices
elif response_obj is not None and isinstance(response_obj, litellm.ImageResponse):
output = response_obj["data"]
output = response_obj.data
litellm_params: Final = kwargs.get("litellm_params", {})
dynamic_metadata: Final = litellm_params.get("metadata", {}) or {}

View file

@ -150,7 +150,7 @@ class CustomGuardrail(CustomLogger):
def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks
super().__init_subclass__(**kwargs)
own_apply_guardrail: Final = cls.__dict__.get("apply_guardrail")
own_apply_guardrail: Final[object] = cls.__dict__.get("apply_guardrail")
if own_apply_guardrail is None or LOGS_GUARDRAIL_INFORMATION_MARKER in vars(own_apply_guardrail):
return
cls.apply_guardrail = log_guardrail_information(own_apply_guardrail)

View file

@ -54,7 +54,7 @@ from litellm.types.utils import (
StandardLoggingPayloadErrorInformation,
)
_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({})
_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
_EMPTY_MESSAGE: Final[Message] = {"role": "", "content": ""}
_MAX_PARSED_TOOL_ARGUMENT_CHARS: Final = 256 * 1024
_SAFE_REDACTED_MESSAGE_ROLES: Final = frozenset(
@ -154,7 +154,7 @@ def _guardrail_information_without_prompt_carriers(
return tuple(_guardrail_entry_without_prompt_carriers(entry) for entry in _guardrail_entries(guardrail_information))
def _metadata_without_prompt_carriers(standard_logging_metadata: Mapping[str, Any]) -> Mapping[str, Any]:
def _metadata_without_prompt_carriers(standard_logging_metadata: Mapping[str, object]) -> Mapping[str, object]:
"""The metadata minus the records that quote prompts, tool arguments, tool results, or retrieved text."""
return MappingProxyType(
{
@ -237,7 +237,7 @@ def _declared_cost_tags(span_tags: Sequence[str]) -> tuple[str, ...]:
return tuple(dimension for dimension in _COST_DIMENSIONS if dimension in present)
def _reasoning_output_tokens(usage_object: Mapping[str, Any] | None) -> float:
def _reasoning_output_tokens(usage_object: Mapping[str, object] | None) -> float:
"""The provider's reasoning-token count, from either the chat or the responses spelling."""
if usage_object is None:
return 0.0
@ -254,20 +254,24 @@ def _reasoning_output_tokens(usage_object: Mapping[str, Any] | None) -> float:
)
def _mapping_field(source: Mapping[str, Any], key: str) -> Mapping[str, Any]:
def _mapping_field(source: Mapping[str, object], key: str) -> Mapping[str, object]:
"""The value at `key` when it is a mapping, else an empty one."""
value: Final = source.get(key)
return value if isinstance(value, dict) else _EMPTY_MAPPING
def _content_blocks(message: Mapping[str, Any]) -> tuple[Mapping[str, Any], ...]:
def _text_field(source: Mapping[str, object], key: str, default: str = "") -> str:
return _safe_identifier(source.get(key, default))
def _content_blocks(message: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
content: Final = message.get("content")
if not isinstance(content, list):
return ()
return tuple(block for block in content if isinstance(block, dict))
def _to_dd_arguments(raw_arguments: object) -> dict[str, Any] | str:
def _to_dd_arguments(raw_arguments: object) -> dict[str, object] | str:
"""
Arguments as the object LLM Obs types them as, or the raw string when they are not one.
@ -282,7 +286,7 @@ def _to_dd_arguments(raw_arguments: object) -> dict[str, Any] | str:
return parsed if isinstance(parsed, dict) else raw_arguments
def _to_dd_tool_calls(message: Mapping[str, Any]) -> tuple[ToolCall, ...]:
def _to_dd_tool_calls(message: Mapping[str, object]) -> tuple[ToolCall, ...]:
"""
The tool calls a message carries, in LLM Obs' ToolCall schema, from either dialect.
@ -293,10 +297,10 @@ def _to_dd_tool_calls(message: Mapping[str, Any]) -> tuple[ToolCall, ...]:
raw_tool_calls: Final = message.get("tool_calls")
openai_calls: Final = tuple(
ToolCall(
name=function.get("name", ""),
name=_text_field(function, "name"),
arguments=_to_dd_arguments(function.get("arguments", "")),
tool_id=tool_call.get("id", ""),
type=tool_call.get("type", "function"),
tool_id=_text_field(tool_call, "id"),
type=_text_field(tool_call, "type", "function"),
)
for tool_call in (raw_tool_calls if isinstance(raw_tool_calls, list) else ())
if isinstance(tool_call, dict)
@ -304,9 +308,9 @@ def _to_dd_tool_calls(message: Mapping[str, Any]) -> tuple[ToolCall, ...]:
)
anthropic_calls: Final = tuple(
ToolCall(
name=block.get("name", ""),
name=_text_field(block, "name"),
arguments=_to_dd_arguments(block.get("input") or {}),
tool_id=block.get("id", ""),
tool_id=_text_field(block, "id"),
type="tool_use",
)
for block in _content_blocks(message)
@ -315,7 +319,7 @@ def _to_dd_tool_calls(message: Mapping[str, Any]) -> tuple[ToolCall, ...]:
return openai_calls + anthropic_calls
def _to_dd_tool_results(message: Mapping[str, Any], tool_call_names: Mapping[str, str]) -> tuple[ToolResult, ...]:
def _to_dd_tool_results(message: Mapping[str, object], tool_call_names: Mapping[str, str]) -> tuple[ToolResult, ...]:
"""
The tool results a message carries, linked back to the call each answers.
@ -400,14 +404,14 @@ def _to_dd_messages(messages: object) -> tuple[Message, ...]:
return tuple(_to_dd_message(message, tool_call_names) for message in messages)
def _to_dd_tool_definition(entry: Mapping[str, Any]) -> ToolDefinition | None:
def _to_dd_tool_definition(entry: Mapping[str, object]) -> ToolDefinition | None:
function: Final = entry.get("function")
declared: Final[Mapping[str, Any]] = function if isinstance(function, dict) else entry
name: Final = declared.get("name")
declared: Final[Mapping[str, object]] = function if isinstance(function, dict) else entry
name: Final = _text_field(declared, "name")
if not name:
return None
schema: Final = declared.get("parameters") or declared.get("input_schema")
description: Final = declared.get("description", "")
description: Final = _text_field(declared, "description")
if not isinstance(schema, dict):
return ToolDefinition(name=name, description=description)
return ToolDefinition(name=name, description=description, schema=schema)
@ -683,7 +687,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
if callable(current_span_fn):
current_span: Final = current_span_fn()
if current_span is not None:
trace_id: Final = getattr(current_span, "trace_id", None)
trace_id: Final[object] = getattr(current_span, "trace_id", None)
if trace_id is not None:
return str(trace_id)
except Exception:
@ -716,7 +720,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
def redacts_messages_itself(self) -> bool:
return True
def _payload_logging_is_off(self, kwargs: Mapping[str, Any]) -> bool:
def _payload_logging_is_off(self, kwargs: Mapping[str, object]) -> bool:
return (
bool(self.turn_off_message_logging)
or self.message_logging is not True

View file

@ -3,12 +3,21 @@
import os
import traceback
from typing import Any, Final
from collections.abc import Mapping
from typing import Final, Protocol
import litellm
from litellm._uuid import uuid
class _DynamoTable(Protocol):
def put_item(self, *, Item: Mapping[str, object]) -> object: ...
class _DynamoResource(Protocol):
def Table(self, name: str) -> _DynamoTable: ...
class DyanmoDBLogger:
# Class variables or attributes
@ -16,7 +25,7 @@ class DyanmoDBLogger:
# Instance variables
import boto3
self.dynamodb: Any = boto3.resource("dynamodb", region_name=os.environ["AWS_REGION_NAME"])
self.dynamodb: Final[_DynamoResource] = boto3.resource("dynamodb", region_name=os.environ["AWS_REGION_NAME"])
if litellm.dynamodb_table_name is None:
raise ValueError(
"LiteLLM Error, trying to use DynamoDB but not table name passed. Create a table and set `litellm.dynamodb_table_name=<your-table>`"
@ -41,7 +50,7 @@ class DyanmoDBLogger:
id: Final = response_obj.get("id", str(uuid.uuid4()))
# Build the initial payload
payload: Final = {
payload: Final[dict[str, object]] = {
"id": id,
"call_type": call_type,
"startTime": start_time,

View file

@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Any, Final
from typing import Final
import polars as pl
@ -32,7 +32,7 @@ class FocusLiteLLMDatabase:
client: Final = self._ensure_prisma_client()
where_clauses: Final[list[str]] = []
query_params: Final[list[Any]] = []
query_params: Final[list[datetime | int]] = []
placeholder_index = 1
if start_time_utc:
where_clauses.append(f"dus.updated_at >= ${placeholder_index}::timestamptz")
@ -112,7 +112,7 @@ class FocusLiteLLMDatabase:
except Exception as exc:
raise RuntimeError(f"Error retrieving usage data: {exc}") from exc
async def get_table_info(self) -> dict[str, Any]:
async def get_table_info(self) -> dict[str, object]:
"""Return metadata about the spend table for diagnostics."""
client: Final = self._ensure_prisma_client()

View file

@ -4,7 +4,8 @@ from __future__ import annotations
import csv
import io
from typing import Any, Final
from collections.abc import Mapping
from typing import Final
import httpx # noqa: F401 - used at runtime (AsyncClient, HTTPStatusError)
@ -94,7 +95,7 @@ class FocusVantageDestination(FocusDestination):
self,
*,
prefix: str,
config: dict[str, Any] | None = None,
config: Mapping[str, object] | None = None,
) -> None:
config = config or {}
api_key: Final = config.get("api_key")

View file

@ -396,12 +396,13 @@ class GalileoObserve(CustomLogger):
)
@staticmethod
def _log_v2_payload_validation(payload: dict[str, Any]) -> None:
def _log_v2_payload_validation(payload: dict[str, object]) -> None:
missing_fields: Final[list[str]] = []
traces: Final[Sequence[object]] = payload.get("traces", [])
if not traces:
traces_value: Final = payload.get("traces", [])
if not traces_value:
missing_fields.append("traces")
traces: Final[Sequence[object]] = traces_value if isinstance(traces_value, list) else []
for trace_index, trace in enumerate(traces):
if not isinstance(trace, dict):
continue
@ -425,8 +426,8 @@ class GalileoObserve(CustomLogger):
missing_fields,
)
def _log_flush_payload(self, url: str, payload: dict[str, Any]) -> None:
traces: Final[Sequence[object]] = payload.get("traces", [])
def _log_flush_payload(self, url: str, payload: dict[str, object]) -> None:
traces: Final = payload.get("traces")
verbose_logger.debug(
"Galileo Logger flush URL: %s trace_count=%s",
url,

View file

@ -4,7 +4,7 @@ import inspect
import os
import re
import traceback
from collections.abc import Callable, Iterable, Mapping
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import datetime
from functools import lru_cache
from types import MappingProxyType
@ -432,7 +432,7 @@ class LangFuseLogger:
prompt: dict,
level: str,
status_message: str | None,
) -> tuple[dict | None, str | dict | list | None]:
) -> tuple[dict | None, str | dict | Sequence[object] | None]:
"""
Get the input and output content for Langfuse logging
@ -448,7 +448,7 @@ class LangFuseLogger:
output: The output content for Langfuse logging
"""
input = None
output: str | dict | list[Any] | None = None
output: str | dict | Sequence[object] | None = None
if level == "ERROR" and status_message is not None and isinstance(status_message, str):
input = prompt
output = status_message
@ -508,7 +508,7 @@ class LangFuseLogger:
user_id: str | None,
metadata: dict[str, object],
litellm_params: dict,
output: str | dict | list | None,
output: str | dict | Sequence[object] | None,
start_time: datetime | None,
end_time: datetime | None,
kwargs: dict,

View file

@ -5,6 +5,7 @@ Relevant Issue: https://github.com/BerriAI/litellm/issues/13764
"""
import json
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
from pydantic import BaseModel
@ -40,7 +41,7 @@ def get_output_content_by_type(
| HttpxBinaryResponseContent
| ResponsesAPIResponse
| list,
kwargs: dict[str, Any] | None = None,
kwargs: Mapping[str, object] | None = None,
) -> str:
"""
Extract output content from response objects based on their type.

View file

@ -77,9 +77,9 @@ class LangsmithLogger(CustomBatchLogger):
if _batch_size:
self.batch_size = int(_batch_size)
self.log_queue: list[LangsmithQueueObject] = []
self._flush_task: asyncio.Task[Any] | None = self._start_periodic_flush_task()
self._flush_task: asyncio.Task[None] | None = self._start_periodic_flush_task()
def _start_periodic_flush_task(self) -> asyncio.Task[Any] | None:
def _start_periodic_flush_task(self) -> asyncio.Task[None] | None:
"""Start the periodic flush task only when an event loop is already running."""
try:
loop: Final = asyncio.get_running_loop()
@ -154,9 +154,9 @@ class LangsmithLogger(CustomBatchLogger):
return self._redact_metadata(extra_metadata)
def _build_outputs_with_usage(self, payload: StandardLoggingPayload) -> dict[str, Any]:
def _build_outputs_with_usage(self, payload: StandardLoggingPayload) -> dict[str, object]:
response: Final = payload["response"]
outputs: dict[str, Any]
outputs: dict[str, object]
if isinstance(response, dict):
outputs = {**response}
else:

View file

@ -36,7 +36,7 @@ model. They coincide on the SDK path, which is correct.
from __future__ import annotations
from collections.abc import Iterator, Mapping
from collections.abc import Callable, Iterator, Mapping
from dataclasses import dataclass, field
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, cast
@ -61,7 +61,7 @@ class RequestIdentity:
# The team's free-form metadata, carried raw (empty/missing -> None) and
# filtered to an operator allowlist only at Baggage-promotion time, so an
# unconfigured deployment never promotes any of it.
team_metadata: Mapping[str, Any] | None = None
team_metadata: Mapping[str, object] | None = None
key_hash: str | None = None
end_user: str | None = None
# The model litellm dispatched to the provider. Only known once the call
@ -111,7 +111,7 @@ class RequestIdentity:
snapshot) is flattened to dotted keys so ``requester_metadata.<key>``
resolves too.
"""
get: Final = lambda name: getattr(auth, name, None) # noqa: E731
get: Final[Callable[[str], object]] = lambda name: getattr(auth, name, None) # noqa: E731
auth_meta: Final = tuple(
(meta_key, str(value))
for meta_key, attr in (
@ -228,7 +228,7 @@ class LLMCallEvent:
trace: TraceControls
@classmethod
def from_dict(cls, kwargs: Mapping[str, Any]) -> LLMCallEvent:
def from_dict(cls, kwargs: Mapping[str, object]) -> LLMCallEvent:
raw_payload: Final = kwargs.get("standard_logging_object")
payload: Final = cast("StandardLoggingPayload", raw_payload) if raw_payload else None
operation: Final = resolve_operation(as_str(kwargs.get("call_type")))
@ -251,7 +251,7 @@ def time_to_first_chunk_seconds(kwargs: Mapping[str, Any]) -> float | None:
to the first streamed chunk (``completion_start_time``); ``None`` for
non-streaming calls, where ``completion_start_time`` is backfilled with the
end time and would not measure first-chunk latency."""
optional_params: Final = cast(Mapping[str, Any], kwargs.get("optional_params") or {})
optional_params: Final = cast(Mapping[str, object], kwargs.get("optional_params") or {})
if not optional_params.get("stream"):
return None
api_call_start: Final = to_seconds(kwargs.get("api_call_start_time"))
@ -312,7 +312,7 @@ def _metadata_dicts(
)
def _call_id(payload: StandardLoggingPayload | None, kwargs: Mapping[str, Any]) -> str | None:
def _call_id(payload: StandardLoggingPayload | None, kwargs: Mapping[str, object]) -> str | None:
"""The call id from the payload (when closed) or the bare kwargs (at pre_call)."""
if payload is not None:
call_id: Final = as_str(payload.get("litellm_call_id")) or as_str(payload.get("id"))
@ -385,7 +385,7 @@ def _model_info_id(model_info: object) -> str | None:
return None
def _team_metadata_dict(value: object) -> Mapping[str, Any] | None:
def _team_metadata_dict(value: object) -> Mapping[str, object] | None:
"""The team's free-form metadata as a raw mapping, or ``None`` when missing
or empty.

View file

@ -12,11 +12,14 @@ when the feature gate is off.
"""
import os
from typing import Any, Final
from typing import TYPE_CHECKING, Final, Protocol
from litellm._logging import verbose_logger
from litellm.integrations.otel.model.config import is_otel_v2_enabled
if TYPE_CHECKING:
from fastapi import FastAPI
# Routes excluded from server-span tracing by default: high-frequency pollers and
# static UI/docs assets, none of which are LLM traffic. Entries are substring-matched
# against the request path (unanchored, so they survive a ``server_root_path`` prefix
@ -65,7 +68,15 @@ PASSTHROUGH_PREFIXES: Final = frozenset(
)
def _passthrough_span_name_hook(span: Any, scope: dict) -> None:
class _RenameableSpan(Protocol):
def is_recording(self) -> bool: ...
def update_name(self, name: str) -> None: ...
def set_attribute(self, key: str, value: str) -> None: ...
def _passthrough_span_name_hook(span: "_RenameableSpan | None", scope: dict) -> None:
"""FastAPI ``server_request_hook``: give passthrough server spans a useful name.
The instrumentation matches the route at span creation, so both the span name
@ -88,7 +99,7 @@ def _passthrough_span_name_hook(span: Any, scope: dict) -> None:
pass
def instrument_fastapi_app(app: Any) -> None:
def instrument_fastapi_app(app: "FastAPI") -> None:
"""Attach OTel server-span instrumentation to the proxy FastAPI app.
Safe no-op when the V2 gate is off or ``opentelemetry-instrumentation-fastapi``

View file

@ -16,7 +16,7 @@ class CoroutineChecker:
"""
def __init__(self):
self._cache = WeakKeyDictionary()
self._cache: WeakKeyDictionary[object, bool] = WeakKeyDictionary()
self._max_size = COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY
def is_async_callable(self, callback: Any) -> bool:
@ -33,10 +33,10 @@ class CoroutineChecker:
pass
# Determine target - optimized path for common cases
target = callback
target: object = callback
if not inspect.isfunction(target) and not inspect.ismethod(target):
try:
call_attr: Final = getattr(target, "__call__", None) # noqa: B004 # value unwrap so iscoroutinefunction sees through functors
call_attr: Final[object] = getattr(target, "__call__", None) # noqa: B004 # value unwrap so iscoroutinefunction sees through functors
if call_attr is not None:
target = call_attr
except Exception:

View file

@ -4,7 +4,7 @@ import re
import traceback
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final, Protocol, cast
from typing import Final, Protocol, cast
import httpx
@ -194,7 +194,7 @@ def _get_response_headers(original_exception: Exception) -> httpx.Headers | None
_response_headers: httpx.Headers | None = None
try:
_response_headers = getattr(original_exception, "headers", None)
error_response: Final = getattr(original_exception, "response", None)
error_response: Final[object] = getattr(original_exception, "response", None)
if not _response_headers and error_response:
_response_headers = getattr(error_response, "headers", None)
if not _response_headers:
@ -211,7 +211,7 @@ def _accepted_init_kwargs(exception_class: type[Exception], candidates: Mapping[
def extract_and_raise_litellm_exception(
response: Any | None,
response: object | None,
error_str: str,
model: str,
custom_llm_provider: str,

View file

@ -6,7 +6,7 @@ from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone, tzinfo
from types import MappingProxyType
from typing import Any, Final, Literal, TypedDict, cast
from typing import Final, Literal, TypedDict, cast
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from typing_extensions import ReadOnly
@ -100,7 +100,7 @@ def _requested_image_size(optional_params: Mapping[str, object] | None) -> str |
return value if value is not None and _IMAGE_SIZE_PATTERN.fullmatch(value) else None
def get_web_search_requests(server_tool_use: Any) -> int | None:
def get_web_search_requests(server_tool_use: object) -> int | None:
"""
Tolerantly read ``web_search_requests`` from a ``server_tool_use`` value
that may be ``None``, a ``dict``, a ``ServerToolUse`` pydantic instance,
@ -1653,7 +1653,7 @@ def calculate_image_response_cost_from_usage(
if prompt_tokens == 0 and completion_tokens == 0 and total_tokens == 0:
return None
input_tokens_details: Final = getattr(usage, "input_tokens_details", None)
input_tokens_details: Final[object] = getattr(usage, "input_tokens_details", None)
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
if input_tokens_details is not None:
# input_tokens_details may be a dict (e.g. OpenAI image edit responses)
@ -1666,9 +1666,12 @@ def calculate_image_response_cost_from_usage(
cached_tokens=0,
)
output_tokens_details = getattr(usage, "completion_tokens_details", None)
if output_tokens_details is None:
output_tokens_details = getattr(usage, "output_tokens_details", None)
completion_tokens_details_attr: Final[object] = getattr(usage, "completion_tokens_details", None)
output_tokens_details: Final[object] = (
getattr(usage, "output_tokens_details", None)
if completion_tokens_details_attr is None
else completion_tokens_details_attr
)
if output_tokens_details is None:
completion_tokens_details = CompletionTokensDetailsWrapper(

View file

@ -1,7 +1,7 @@
import datetime
from collections.abc import Mapping
from functools import reduce
from typing import Any, Final
from typing import Final
import httpx
@ -106,7 +106,7 @@ class ResponseMetadata:
Handles setting and managing `_hidden_params`, `response_time_ms`, and `litellm_overhead_time_ms` for LiteLLM responses
"""
def __init__(self, result: Any):
def __init__(self, result: object):
self.result = result
self._hidden_params: HiddenParams | dict = getattr(result, "_hidden_params", {}) or {}

View file

@ -13,14 +13,6 @@ from pathlib import Path
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast
from openai.types.chat.chat_completion_custom_tool_param import (
CustomFormatGrammar,
CustomFormatGrammarGrammar,
)
from openai.types.shared_params.custom_tool_input_format import (
Grammar as ResponsesGrammarFormat,
)
import litellm
from litellm import verbose_logger
from litellm.router_utils.batch_utils import InMemoryFile
@ -59,7 +51,7 @@ if TYPE_CHECKING:
def handle_any_messages_to_chat_completion_str_messages_conversion(
messages: Any,
messages: object,
) -> list[dict[str, str]]:
"""
Handles any messages to chat completion str messages conversion
@ -804,7 +796,7 @@ def extract_file_metadata(file_data: FileTypes) -> tuple[str | None, str | None]
"""
filename: str | None = None
content_type: str | None = None
file_content: Any = None
file_content: object = None
if isinstance(file_data, tuple):
if len(file_data) == 2:
@ -1002,7 +994,7 @@ def unpack_defs(
# Use iterative approach with queue to avoid recursion
# Each item in queue is (node, parent_container, key/index, active_defs, ref_chain)
queue: Final[deque[tuple[Any, dict | list | None, str | int | None, dict, set]]] = deque(
queue: Final[deque[tuple[object, dict | list | None, str | int | None, dict, set]]] = deque(
[(schema, None, None, root_defs, set())]
)
inlined_bytes = 0
@ -1624,7 +1616,10 @@ def is_function_call(optional_params: dict) -> bool:
return False
def convert_custom_tool_format_to_chat_shape(format_obj: Mapping[str, Any]) -> Mapping[str, Any]:
_CUSTOM_GRAMMAR_FIELDS: Final = ("definition", "syntax")
def convert_custom_tool_format_to_chat_shape(format_obj: Mapping[str, object]) -> Mapping[str, object]:
"""
Responses API grammar formats are flat ({"type": "grammar", "definition", "syntax"});
Chat Completions wraps the same fields in a "grammar" object. Text formats are
@ -1632,15 +1627,11 @@ def convert_custom_tool_format_to_chat_shape(format_obj: Mapping[str, Any]) -> M
"""
if format_obj.get("type") != "grammar" or "grammar" in format_obj:
return format_obj
grammar: Final = CustomFormatGrammarGrammar()
if "definition" in format_obj:
grammar["definition"] = format_obj["definition"]
if "syntax" in format_obj:
grammar["syntax"] = format_obj["syntax"]
return CustomFormatGrammar(type="grammar", grammar=grammar)
grammar: Final[Mapping[str, object]] = {key: format_obj[key] for key in _CUSTOM_GRAMMAR_FIELDS if key in format_obj}
return {"type": "grammar", "grammar": grammar}
def convert_custom_tool_format_to_responses_shape(format_obj: Mapping[str, Any]) -> Mapping[str, Any]:
def convert_custom_tool_format_to_responses_shape(format_obj: Mapping[str, object]) -> Mapping[str, object]:
"""
Inverse of convert_custom_tool_format_to_chat_shape: unwrap the Chat Completions
"grammar" object into the flat Responses API grammar shape.
@ -1648,12 +1639,10 @@ def convert_custom_tool_format_to_responses_shape(format_obj: Mapping[str, Any])
grammar: Final = format_obj.get("grammar")
if format_obj.get("type") != "grammar" or not isinstance(grammar, dict):
return format_obj
flat: Final = ResponsesGrammarFormat(type="grammar")
if "definition" in grammar:
flat["definition"] = grammar["definition"]
if "syntax" in grammar:
flat["syntax"] = grammar["syntax"]
return flat
return {
"type": "grammar",
**{key: grammar[key] for key in _CUSTOM_GRAMMAR_FIELDS if key in grammar},
}
def get_file_ids_from_messages(messages: list[AllMessageValues]) -> list[str]:

View file

@ -1,6 +1,8 @@
import json
from datetime import datetime
from typing import Any, Final
from typing import Any, Final, Literal
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
@ -9,6 +11,20 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.types.llms.custom_http import httpxSpecialProvider
class _TokenizerConfigResult(TypedDict):
"""Outcome of a tokenizer_config.json fetch, carrying the parsed document when the fetch succeeded."""
status: ReadOnly[Literal["success", "failure"]]
tokenizer: NotRequired[ReadOnly[object]]
class _ChatTemplateFileResult(TypedDict):
"""Outcome of a chat template file fetch, carrying the template body when the fetch succeeded."""
status: ReadOnly[Literal["success", "failure"]]
chat_template: NotRequired[ReadOnly[str]]
def strftime_now(fmt: str) -> str:
"""
Custom function for templates that need current date/time formatting (e.g., gpt-oss)
@ -22,7 +38,7 @@ def strftime_now(fmt: str) -> str:
return datetime.now().strftime(fmt)
def _get_tokenizer_config(hf_model_name: str) -> dict[str, Any]:
def _get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult:
"""
Fetch tokenizer_config.json from HuggingFace (sync)
@ -45,7 +61,7 @@ def _get_tokenizer_config(hf_model_name: str) -> dict[str, Any]:
return {"status": "failure"}
async def _aget_tokenizer_config(hf_model_name: str) -> dict[str, Any]:
async def _aget_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult:
"""
Fetch tokenizer_config.json from HuggingFace (async)
@ -70,7 +86,7 @@ async def _aget_tokenizer_config(hf_model_name: str) -> dict[str, Any]:
return {"status": "failure"}
def _get_chat_template_file(hf_model_name: str) -> dict[str, Any]:
def _get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult:
"""
Fetch chat template from separate .jinja file (sync)
@ -98,7 +114,7 @@ def _get_chat_template_file(hf_model_name: str) -> dict[str, Any]:
return {"status": "failure"}
async def _aget_chat_template_file(hf_model_name: str) -> dict[str, Any]:
async def _aget_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult:
"""
Fetch chat template from separate .jinja file (async)

View file

@ -93,13 +93,13 @@ class SensitiveDataMasker:
def _mask_sequence(
self,
values: list[Any],
values: Sequence[object],
depth: int,
max_depth: int,
excluded_keys: set[str] | None,
key_is_sensitive: bool,
) -> list[Any]:
masked_items: Final[list[Any]] = []
) -> Sequence[object]:
masked_items: Final[list[object]] = []
if depth >= max_depth:
return values
@ -222,7 +222,7 @@ class _PayloadWalker:
return [self.walk(item, key_is_sensitive, depth + 1) for item in node]
def mask_sensitive_keys(data: dict[str, Any], sensitive_fields: set[str]) -> dict[str, Any]:
def mask_sensitive_keys(data: Mapping[str, object], sensitive_fields: set[str]) -> dict[str, object]:
"""Return a new dict with values masked for keys listed in ``sensitive_fields``.
Unlike :meth:`SensitiveDataMasker.mask_dict`, this does exact key-name
@ -234,7 +234,7 @@ def mask_sensitive_keys(data: dict[str, Any], sensitive_fields: set[str]) -> dic
range and are replaced with a fixed-length all-mask string, so a short
credential is never returned verbatim.
"""
masked: Final[dict[str, Any]] = {}
masked: Final[dict[str, object]] = {}
mask_char: Final = _default_masker.mask_char
min_visible: Final = _default_masker.visible_prefix + _default_masker.visible_suffix
for key, value in data.items():

View file

@ -125,7 +125,7 @@ class A2AGuardrailHandler(BaseTranslation):
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
request_data: dict | None = None,
) -> Any:
) -> object:
"""
Process A2A output response by applying guardrails to text content.

View file

@ -6,7 +6,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NoReturn, cast
import httpx
from pydantic import ValidationError
from pydantic import BaseModel, ValidationError
from typing_extensions import ReadOnly, TypedDict
import litellm
@ -151,7 +151,7 @@ class _AnthropicToolResultBlock(TypedDict, total=False):
content: ReadOnly[object]
_ENUM_TYPE_CHECKS: Final[Mapping[str, Callable[[object], bool]]] = MappingProxyType(
_ENUM_TYPE_CHECKS: Final[Mapping[object, Callable[[object], bool]]] = MappingProxyType(
{
"null": lambda v: v is None,
"boolean": lambda v: isinstance(v, bool),
@ -164,7 +164,7 @@ _ENUM_TYPE_CHECKS: Final[Mapping[str, Callable[[object], bool]]] = MappingProxyT
)
def _enum_conflicts_with_declared_type(schema: Mapping[str, Any]) -> bool:
def _enum_conflicts_with_declared_type(schema: Mapping[str, object]) -> bool:
"""Whether ``schema``'s ``enum`` cannot match its declared ``type``."""
enum_values: Final = schema.get("enum")
declared_type: Final = schema.get("type")
@ -659,7 +659,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return result
def get_json_schema_from_pydantic_object(self, response_format: Any | dict | None) -> dict | None:
def get_json_schema_from_pydantic_object(self, response_format: type[BaseModel] | dict | None) -> dict | None:
return type_to_response_format_param(
response_format, ref_template="/$defs/{model}"
) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755
@ -1072,7 +1072,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _sanitize_tool_names_in_request(
optional_params: dict[str, Any],
optional_params: dict[str, object],
) -> tuple[dict[str, str], dict[str, str]]:
"""Sanitize ``optional_params['tools']`` and ``optional_params['tool_choice']``
in place so every name matches Anthropic's ``^[a-zA-Z0-9_-]{1,128}$``.
@ -1119,7 +1119,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
# so a caller reusing the same tool list/dicts across requests
# doesn't see its inputs permanently rewritten (which would also
# drop the original key from `forward` on the next request).
new_tools: Final[list[Any]] = []
new_tools: Final[list[object]] = []
for t in tools:
if (
isinstance(t, dict)
@ -1442,7 +1442,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
entry_type = entry.get("type")
if entry_type == "compaction":
anthropic_edit: dict[str, Any] = {"type": "compact_20260112"}
anthropic_edit: dict[str, object] = {"type": "compact_20260112"}
compact_threshold = entry.get("compact_threshold")
# Rewrite to 'trigger' with correct nesting if threshold exists
if compact_threshold is not None and isinstance(compact_threshold, (int, float)):
@ -2442,9 +2442,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
code_by_id: Final[dict[str, str]] = {}
for tc in tool_calls:
try:
args = json.loads(tc.get("function", {}).get("arguments", "{}"))
args: object = json.loads(tc.get("function", {}).get("arguments", "{}"))
if not isinstance(args, Mapping):
continue
call_id = tc.get("id")
command = args.get("command", "")
command: object = args.get("command", "")
if isinstance(call_id, str):
code_by_id[call_id] = command if isinstance(command, str) else ""
except Exception:
@ -2514,8 +2516,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
tool_results: Sequence[_AnthropicToolResultBlock] | None,
compaction_blocks: Sequence[object] | None,
tool_calls: list[ChatCompletionToolCallChunk],
) -> dict[str, Any]:
provider_specific_fields: Final[dict[str, Any]] = {
) -> dict[str, object]:
provider_specific_fields: Final[dict[str, object]] = {
"citations": citations,
"thinking_blocks": thinking_blocks,
}

View file

@ -7,7 +7,7 @@ import re
from collections.abc import Mapping, MutableMapping, Sequence
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Any, Final, Literal
from typing import Any, Final, Literal, TypeVar
import httpx
from pydantic import BaseModel, ConfigDict, StrictBool, TypeAdapter, ValidationError
@ -40,6 +40,8 @@ from litellm.types.llms.anthropic import (
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.model_listing import ModelInfoResponse
_MessageT = TypeVar("_MessageT")
DROP_FORCED_TOOL_CHOICE_WARNING: Final = (
"Downgrading forced tool_choice to 'auto' for model=%s (drop_params=True): this model rejects tool_choice type "
"'any'/'tool' with a 400 because thinking is always on and a forced call would skip it."
@ -1121,7 +1123,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
return AnthropicTokenCounter()
def strip_advisor_blocks_from_messages(messages: list[Any], replace_with_text: bool = False) -> list[Any]:
def strip_advisor_blocks_from_messages(messages: list[_MessageT], replace_with_text: bool = False) -> list[_MessageT]:
"""
Remove (or replace) server_tool_use (name='advisor') and advisor_tool_result blocks
from assistant message content.
@ -1228,7 +1230,7 @@ def is_anthropic_invalid_thinking_block_error(error_text: str) -> bool:
return "must contain thinking" in lower
def strip_thinking_blocks_from_anthropic_messages(messages: list[Any]) -> list[Any]:
def strip_thinking_blocks_from_anthropic_messages(messages: Sequence[object]) -> list[object]:
"""
Return a new message list with thinking / redacted_thinking content blocks removed
from each message. Used to recover from invalid thinking signatures on retry.
@ -1236,7 +1238,7 @@ def strip_thinking_blocks_from_anthropic_messages(messages: list[Any]) -> list[A
Messages whose content is a list and becomes empty after stripping are omitted,
since Anthropic rejects empty content arrays.
"""
out: Final[list[Any]] = []
out: Final[list[object]] = []
for m in messages:
if not isinstance(m, dict):
out.append(m)

View file

@ -25,6 +25,9 @@ from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.anthropic_messages.transformation import (
BaseAnthropicMessagesConfig,
)
HOLD_BACK_PING_INTERVAL_SECONDS: Final = 15.0
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES: Final = (
@ -182,7 +185,7 @@ class AgenticAnthropicStreamingIterator:
http_handler: Any,
model: str,
messages: list[dict],
anthropic_messages_provider_config: Any,
anthropic_messages_provider_config: "BaseAnthropicMessagesConfig",
anthropic_messages_optional_request_params: dict,
logging_obj: "LiteLLMLoggingObj",
custom_llm_provider: str,
@ -402,7 +405,7 @@ class AgenticAnthropicStreamingIterator:
@staticmethod
def _rebuild_anthropic_response_from_sse(
raw_bytes: list[bytes],
) -> dict[str, Any] | None:
) -> dict[str, object] | None:
"""
Parse collected SSE bytes into an Anthropic Messages response dict.
@ -416,17 +419,18 @@ class AgenticAnthropicStreamingIterator:
"""
events: Final = _parse_sse_events(b"".join(raw_bytes))
response: Final[dict[str, Any]] = {
content: Final[list[dict[str, object]]] = []
response: Final[dict[str, object]] = {
"id": "",
"type": "message",
"role": "assistant",
"model": "",
"content": [],
"content": content,
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": 0, "output_tokens": 0},
}
content_blocks: Final[dict[int, dict[str, Any]]] = {}
content_blocks: Final[dict[int, dict[str, object]]] = {}
saw_message_start = False
for event_type, data in events:
@ -448,6 +452,6 @@ class AgenticAnthropicStreamingIterator:
for idx in sorted(content_blocks.keys()):
block = content_blocks[idx]
block.pop("_partial_json", None)
response["content"].append(block)
content.append(block)
return response

View file

@ -185,7 +185,11 @@ class AnthropicFilesHandler:
if not line.strip():
continue
anthropic_result = json.loads(line)
anthropic_result: object = json.loads(line)
if not isinstance(anthropic_result, dict):
raise TypeError(
f"Anthropic batch result line is not a JSON object: {type(anthropic_result).__name__}"
)
custom_id = anthropic_result.get("custom_id", "")
result = anthropic_result.get("result", {})
result_type = result.get("type", "")

View file

@ -1,5 +1,5 @@
from collections.abc import Coroutine, Iterable
from typing import Any, Final, Literal, TypedDict
from typing import Final, Literal, TypedDict
import httpx
from openai import AsyncAzureOpenAI, AzureOpenAI
@ -715,7 +715,8 @@ class AzureAssistantsAPI(BaseAzureLLM):
event_handler: AssistantEventHandler | None,
litellm_params: dict | None = None,
) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]:
data: Final[dict[str, Any]] = {
stream_fn: Final = client.beta.threads.runs.stream
base_data: Final[_RunThreadStreamData] = {
"thread_id": thread_id,
"assistant_id": assistant_id,
"additional_instructions": additional_instructions,
@ -725,8 +726,8 @@ class AzureAssistantsAPI(BaseAzureLLM):
"tools": tools,
}
if event_handler is not None:
data["event_handler"] = event_handler
return client.beta.threads.runs.stream(**data)
return stream_fn(**base_data, event_handler=event_handler)
return stream_fn(**base_data)
def run_thread_stream(
self,

View file

@ -19,6 +19,7 @@ from litellm.secret_managers.main import get_secret_str
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.llms.openai import HttpxBinaryResponseContent
else:
LiteLLMLoggingObj = Any
@ -67,15 +68,15 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig):
litellm_params_dict: dict,
logging_obj: "LiteLLMLoggingObj",
timeout: float | httpx.Timeout,
extra_headers: dict[str, Any] | None,
base_llm_http_handler: Any,
extra_headers: dict[str, object] | None,
base_llm_http_handler: "BaseLLMHTTPHandler",
aspeech: bool,
api_base: str | None,
api_key: str | None,
**kwargs: Any,
**kwargs: object,
) -> Union[
"HttpxBinaryResponseContent",
Coroutine[Any, Any, "HttpxBinaryResponseContent"],
Coroutine[object, object, "HttpxBinaryResponseContent"],
]:
"""
Dispatch method to handle Azure AVA TTS requests

View file

@ -33,7 +33,7 @@ class AzureAIAnthropicCountTokensHandler(AzureAIAnthropicCountTokensConfig):
litellm_params: dict[str, Any] | None = None,
timeout: float | httpx.Timeout | None = None,
tools: list[dict[str, Any]] | None = None,
system: Any | None = None,
system: object = None,
) -> dict[str, Any]:
"""
Handle a CountTokens request using httpx with Azure authentication.

View file

@ -180,7 +180,7 @@ class BaseVideoConfig(ABC):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: dict[str, Any] | None = None,
extra_body: dict[str, object] | None = None,
) -> tuple[str, dict]:
"""
Transform the video remix request into a URL and data
@ -207,7 +207,7 @@ class BaseVideoConfig(ABC):
after: str | None = None,
limit: int | None = None,
order: str | None = None,
extra_query: dict[str, Any] | None = None,
extra_query: dict[str, object] | None = None,
) -> tuple[str, dict]:
"""
Transform the video list request into a URL and params
@ -355,8 +355,8 @@ class BaseVideoConfig(ABC):
litellm_params: GenericLiteLLMParams,
headers: dict,
video_file: FileContent | None = None,
extra_body: dict[str, Any] | None = None,
prefetched_source_data: dict[str, Any] | None = None,
extra_body: dict[str, object] | None = None,
prefetched_source_data: dict[str, object] | None = None,
) -> tuple[str, Mapping[str, object], RequestFiles | None]:
"""
Transform the video edit request into a URL plus either JSON data or
@ -386,7 +386,7 @@ class BaseVideoConfig(ABC):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: dict[str, Any] | None = None,
extra_body: dict[str, object] | None = None,
) -> tuple[str, dict]:
"""
Transform the video extension request into a URL and JSON data.

View file

@ -1126,7 +1126,7 @@ class AmazonConverseConfig(BaseConfig):
return optional_params
def _map_request_metadata_param(self, value: Any, optional_params: dict) -> None:
def _map_request_metadata_param(self, value: object, optional_params: dict) -> None:
if value is not None and isinstance(value, dict):
self._validate_request_metadata(value)
optional_params["requestMetadata"] = value

View file

@ -2,6 +2,7 @@
Bedrock Token Counter implementation using the CountTokens API.
"""
from collections.abc import Mapping, Sequence
from typing import Any, Final
from litellm._logging import verbose_logger
@ -26,12 +27,12 @@ class BedrockTokenCounter(BaseTokenCounter):
async def count_tokens(
self,
model_to_use: str,
messages: list[dict[str, Any]] | None,
contents: list[dict[str, Any]] | None,
messages: Sequence[Mapping[str, object]] | None,
contents: Sequence[Mapping[str, object]] | None,
deployment: dict[str, Any] | None = None,
request_model: str = "",
tools: list[dict[str, Any]] | None = None,
system: Any | None = None,
tools: Sequence[Mapping[str, object]] | None = None,
system: object | None = None,
) -> TokenCountResponse | None:
"""
Count tokens using AWS Bedrock's CountTokens API.
@ -56,7 +57,7 @@ class BedrockTokenCounter(BaseTokenCounter):
litellm_params: Final = deployment.get("litellm_params", {})
# Build request data in the format expected by BedrockCountTokensHandler
request_data: Final[dict[str, Any]] = {
request_data: Final[dict[str, object]] = {
"model": model_to_use,
"messages": messages,
}

View file

@ -375,7 +375,7 @@ def _listed_managed_file(
)
def _uploaded_object_size(litellm_params: Mapping[str, object], raw_response: Response) -> int:
def _uploaded_object_size(litellm_params: Mapping[str, object], response_headers: Mapping[str, str]) -> int:
"""
S3 answers PutObject with an empty body, so the stored object size comes from the
signed request recorded by `transform_create_file_request`, not the response headers.
@ -383,7 +383,7 @@ def _uploaded_object_size(litellm_params: Mapping[str, object], raw_response: Re
uploaded_size: Final = litellm_params.get(UPLOAD_CONTENT_LENGTH_PARAM)
if isinstance(uploaded_size, int):
return uploaded_size
response_content_length: Final = raw_response.headers.get("Content-Length", "0")
response_content_length: Final = response_headers.get("Content-Length", "0")
return int(response_content_length) if response_content_length.isdigit() else 0
@ -1277,7 +1277,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
filename=filename,
created_at=int(time.time()), # Current timestamp
status="uploaded",
bytes=_uploaded_object_size(litellm_params=litellm_params, raw_response=raw_response),
bytes=_uploaded_object_size(litellm_params=litellm_params, response_headers=raw_response.headers),
object="file",
)

View file

@ -125,7 +125,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig):
}
# Create a copy to not mutate original - convert TypedDict to regular dict
mapped_params: Final[dict[str, Any]] = dict(image_edit_optional_params)
mapped_params: Final[dict[str, object]] = dict(image_edit_optional_params)
for k, v in image_edit_optional_params.items():
if k in param_mapping:
@ -172,7 +172,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig):
Returns the request body dict that will be JSON-encoded by the handler.
"""
# Build Bedrock Stability request
data: Final[dict[str, Any]] = {
data: Final[dict[str, object]] = {
"output_format": "png", # Default to PNG
}

View file

@ -14,6 +14,9 @@ from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.pass_through.guardrail_translation.handler import (
PassThroughEndpointHandler,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
@ -27,7 +30,7 @@ def _is_converse_endpoint(endpoint: str) -> bool:
return bool(parts) and parts[-1] in _CONVERSE_ACTIONS
def _generic_passthrough_handler() -> BaseTranslation:
def _generic_passthrough_handler() -> "PassThroughEndpointHandler":
"""
Fallback for non-Converse Bedrock routes (e.g. invoke). The generic
handler scans the full request/response payload so blocking guardrails

View file

@ -16,8 +16,8 @@ BaseAWSLLM._sign_request after the request body is finalized.
"""
import json
from collections.abc import Mapping
from typing import Any, Final, cast # noqa: TID251 # map_openai_params returns the filtered params as a bare dict
from collections.abc import Mapping, Sequence
from typing import Final, cast # noqa: TID251 # map_openai_params returns the filtered params as a bare dict
import httpx
from typing_extensions import ReadOnly, TypedDict
@ -142,9 +142,9 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
return False
@staticmethod
def _filter_unsupported_tools(tools: list[Any]) -> list[Any]:
def _filter_unsupported_tools(tools: "Sequence[object]") -> "list[object]":
"""Keep only tool types Mantle's Responses API accepts."""
kept: Final[list[Any]] = []
kept: Final[list[object]] = []
dropped_types: Final[list[str]] = []
for tool in tools:
if not isinstance(tool, dict):

View file

@ -268,7 +268,7 @@ def _normalize_litellm_params(litellm_params: Any | None) -> dict:
return {}
def get_chatgpt_session_id(litellm_params: Any | None) -> str | None:
def get_chatgpt_session_id(litellm_params: object) -> str | None:
params: Final = _normalize_litellm_params(litellm_params)
for key in ("litellm_session_id", "session_id"):
value = params.get(key)
@ -286,5 +286,5 @@ def get_chatgpt_session_id(litellm_params: Any | None) -> str | None:
return None
def ensure_chatgpt_session_id(litellm_params: Any | None) -> str:
def ensure_chatgpt_session_id(litellm_params: object) -> str:
return get_chatgpt_session_id(litellm_params) or str(uuid4())

View file

@ -1,5 +1,8 @@
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import httpx
from litellm.exceptions import AuthenticationError
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
@ -13,6 +16,7 @@ from litellm.responses.sse_output_recovery import (
record_output_text_chunk,
)
from litellm.types.llms.openai import (
ResponseInputParam,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
@ -64,7 +68,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
def transform_responses_api_request(
self,
model: str,
input: Any,
input: str | ResponseInputParam,
response_api_optional_request_params: dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
@ -109,9 +113,9 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
def transform_response_api_response(
self,
model: str,
raw_response: Any,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
):
) -> ResponsesAPIResponse:
body_text: Final = raw_response.text or ""
if not self._should_parse_as_sse(raw_response=raw_response, body_text=body_text):
return super().transform_response_api_response(
@ -135,7 +139,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
self._attach_response_headers(completed_response=completed_response, raw_response=raw_response)
return completed_response
def _should_parse_as_sse(self, raw_response: Any, body_text: str) -> bool:
def _should_parse_as_sse(self, raw_response: httpx.Response, body_text: str) -> bool:
content_type: Final = (raw_response.headers or {}).get("content-type", "")
if "text/event-stream" in content_type.lower():
return True
@ -150,8 +154,8 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
def _extract_completed_response_from_sse(self, body_text: str) -> tuple[ResponsesAPIResponse | None, str | None]:
completed_response = None
error_message = None
streamed_output_items: Final[dict[int, dict]] = {}
text_only_output_items: Final[dict[int, dict]] = {}
streamed_output_items: Final[dict[int, dict[str, object]]] = {}
text_only_output_items: Final[dict[int, dict[str, object]]] = {}
for chunk in body_text.splitlines():
parsed_chunk = parse_sse_json_chunk(chunk)
if parsed_chunk is None:
@ -178,7 +182,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
# output_index, but text-only items at indices without a
# matching OUTPUT_ITEM_DONE must still be preserved (e.g.
# providers that emit only OUTPUT_TEXT_DONE for some indices).
merged_items: dict[int, dict] = {**text_only_output_items}
merged_items: dict[int, dict[str, object]] = {**text_only_output_items}
merged_items.update(streamed_output_items)
completed_response = self._build_completed_response_from_chunk(
parsed_chunk=parsed_chunk,
@ -197,7 +201,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
return completed_response, error_message
def _build_completed_response_from_chunk(
self, parsed_chunk: dict[str, Any], streamed_output_items: dict[int, dict]
self, parsed_chunk: Mapping[str, object], streamed_output_items: Mapping[int, dict[str, object]]
) -> ResponsesAPIResponse | None:
response_payload = parsed_chunk.get("response")
if not isinstance(response_payload, dict):
@ -223,7 +227,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
def _attach_response_headers(
self,
completed_response: ResponsesAPIResponse,
raw_response: Any,
raw_response: httpx.Response,
) -> None:
raw_headers: Final = dict(raw_response.headers)
processed_headers: Final = process_response_headers(raw_headers)

View file

@ -110,7 +110,7 @@ class CohereChatConfig(BaseConfig):
tool_results: list | None = None,
seed: int | None = None,
) -> None:
locals_: Final = locals().copy()
locals_: Final[dict[str, object]] = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)

View file

@ -2,7 +2,8 @@
Legacy /v1/embedding transformation logic for Bedrock Cohere.
"""
from typing import Any, Final
from collections.abc import Sized
from typing import Final, Protocol
import httpx
@ -16,6 +17,12 @@ from litellm.types.utils import EmbeddingResponse, PromptTokensDetailsWrapper, U
from litellm.utils import is_base64_encoded
class _SupportsEncode(Protocol):
"""Tokenizer handle: the embedding usage path only encodes text to measure its token length."""
def encode(self, text: str, /) -> Sized: ...
class CohereEmbeddingConfig:
"""
Reference: https://docs.cohere.com/v2/reference/embed
@ -61,7 +68,7 @@ class CohereEmbeddingConfig:
return transformed_request
def _calculate_usage(self, input: list[str], encoding: Any, meta: dict) -> Usage:
def _calculate_usage(self, input: list[str], encoding: _SupportsEncode, meta: dict) -> Usage:
input_tokens = 0
text_tokens: Final[int | None] = meta.get("billed_units", {}).get("input_tokens")
@ -97,7 +104,7 @@ class CohereEmbeddingConfig:
data: dict | CohereEmbeddingRequest,
model_response: EmbeddingResponse,
model: str,
encoding: Any,
encoding: _SupportsEncode,
input: list,
) -> EmbeddingResponse:
response_json: Final = response.json()
@ -121,7 +128,7 @@ class CohereEmbeddingConfig:
response_json: dict,
model_response: EmbeddingResponse,
model: str,
encoding: Any,
encoding: _SupportsEncode,
input: list,
) -> EmbeddingResponse:
"""

View file

@ -479,6 +479,11 @@ def _safe_get_response_text(response: httpx.Response) -> str:
return ""
def header_value(headers: Mapping[str, str], name: str) -> str | None:
"""Read one header as ``str | None``; ``httpx.Headers.get`` itself is typed ``Any``."""
return headers.get(name)
async def _safe_aread_response(response: httpx.Response, timeout: float | None = None) -> bytes:
"""Safely read async response body, falling back to empty bytes on errors."""
try:

View file

@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to Databricks' `/chat/completion
"""
import os
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, cast, overload
import httpx
@ -67,7 +67,7 @@ def _is_bare_assistant_message(message_dict: Mapping[str, object]) -> bool:
)
def _sanitize_empty_content(message_dict: dict[str, Any]) -> None:
def _sanitize_empty_content(message_dict: dict[str, object]) -> None:
"""
Remove or filter content so empty text blocks are not sent.
Databricks Model Serving uses Anthropic Messages API spec and rejects empty text blocks.
@ -430,7 +430,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
@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
def _transform_messages(
@ -442,7 +442,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
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]]:
"""
Databricks does not support:
- 'name' in user message.
@ -564,7 +564,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
@staticmethod
def extract_citations(
content: AllDatabricksContentValues | None,
) -> list[Any] | None:
) -> Sequence[Sequence[Mapping[str, object]]] | None:
if content is None:
return None
citations: Final = []

View file

@ -1,6 +1,6 @@
import json
from collections.abc import AsyncIterator, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from typing import TYPE_CHECKING, Final, Literal, cast
import httpx
@ -759,7 +759,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse,
sync_stream: bool,
json_mode: bool | None = False,
) -> Any:
) -> "FireworksAIChatCompletionStreamingHandler":
return FireworksAIChatCompletionStreamingHandler(
streaming_response=streaming_response,
sync_stream=sync_stream,

View file

@ -7,14 +7,32 @@ import os
import re
import threading
from collections.abc import Callable
from typing import Any, Final, Protocol
from typing import Final, Protocol
from urllib.parse import urlsplit
from typing_extensions import ReadOnly, TypedDict, Unpack
import litellm
from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig
from litellm.types.llms.openai import AllMessageValues
class _OpenAIGPTConfigOptions(TypedDict, total=False):
"""The sampling defaults ``OpenAIGPTConfig.__init__`` accepts and stashes on the class."""
frequency_penalty: ReadOnly[int | None]
function_call: ReadOnly[str | dict[str, object] | None]
functions: ReadOnly[list[object] | None]
logit_bias: ReadOnly[dict[str, object] | None]
max_tokens: ReadOnly[int | None]
n: ReadOnly[int | None]
presence_penalty: ReadOnly[int | None]
stop: ReadOnly[str | list[object] | None]
temperature: ReadOnly[int | None]
top_p: ReadOnly[int | None]
response_format: ReadOnly[dict[str, object] | None]
class _GDCHAudienceCredentials(Protocol):
"""A GDCH service account credential already bound to an audience, ready to mint a bearer token."""
@ -32,7 +50,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
_GDCH_CREDENTIAL_TYPE: Final[str] = "gdch_service_account"
_PATH_ID_PATTERN: Final[re.Pattern[str]] = re.compile(r"^[a-zA-Z0-9_-]+$")
def __init__(self, **kwargs: Any) -> None:
def __init__(self, **kwargs: Unpack[_OpenAIGPTConfigOptions]) -> None:
super().__init__(**kwargs)
self._creds_lock = threading.Lock()
self._gdch_creds_cache: dict[tuple[str, str], _GDCHAudienceCredentials] = {}

View file

@ -84,7 +84,7 @@ class GoogleAIStudioTokenCounter:
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
**kwargs,
**kwargs: object,
) -> dict[str, Any]:
"""
Count tokens using Google Gen AI Studio countTokens endpoint.

View file

@ -1,4 +1,5 @@
import base64
from collections.abc import Mapping
from io import BufferedReader, BytesIO
from typing import TYPE_CHECKING, Any, Final, cast
@ -44,7 +45,7 @@ class GeminiImageEditConfig(BaseImageEditConfig):
image_edit_optional_params: ImageEditOptionalRequestParams,
model: str,
drop_params: bool,
) -> dict[str, Any]:
) -> dict[str, object]:
return map_openai_image_params_to_gemini(
params=image_edit_optional_params,
model=model,
@ -87,10 +88,10 @@ class GeminiImageEditConfig(BaseImageEditConfig):
model: str,
prompt: str | None,
image: FileTypes | None,
image_edit_optional_request_params: dict[str, Any],
image_edit_optional_request_params: Mapping[str, object],
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> tuple[dict[str, Any], RequestFiles | None]:
) -> tuple[dict[str, object], RequestFiles | None]:
inline_parts: Final = self._prepare_inline_image_parts(image) if image else []
if not inline_parts:
raise ValueError("Gemini image edit requires at least one image.")
@ -106,7 +107,7 @@ class GeminiImageEditConfig(BaseImageEditConfig):
}
]
request_body: Final[dict[str, Any]] = {"contents": contents}
request_body: Final[dict[str, object]] = {"contents": contents}
request_body["generationConfig"] = get_gemini_image_generation_config(
model=model,
@ -153,14 +154,14 @@ class GeminiImageEditConfig(BaseImageEditConfig):
model_response.usage = transform_gemini_image_usage(response_json["usageMetadata"])
return model_response
def _prepare_inline_image_parts(self, image: FileTypes | list[FileTypes]) -> list[dict[str, Any]]:
def _prepare_inline_image_parts(self, image: FileTypes | list[FileTypes]) -> list[dict[str, object]]:
images: list[FileTypes]
if isinstance(image, list):
images = image
else:
images = [image]
inline_parts: Final[list[dict[str, Any]]] = []
inline_parts: Final[list[dict[str, object]]] = []
for img in images:
if img is None:
continue

View file

@ -81,9 +81,17 @@ class GigaChatConfig(BaseConfig):
repetition_penalty: float | None = None,
profanity_check: bool | None = None,
) -> None:
locals_: Final = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
config_params: Final[Mapping[str, float | int | bool | None]] = MappingProxyType(
{
"temperature": temperature,
"top_p": top_p,
"max_tokens": max_tokens,
"repetition_penalty": repetition_penalty,
"profanity_check": profanity_check,
}
)
for key, value in config_params.items():
if value is not None:
setattr(self.__class__, key, value)
# Instance variables for current request context
self._current_credentials: str | None = None

View file

@ -19,6 +19,7 @@ from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfi
from litellm.types.llms.openai import (
ResponseInputParam,
ResponsesAPIOptionalRequestParams,
ResponsesAPIStreamingResponse,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
@ -129,7 +130,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
model: str,
parsed_chunk: dict,
logging_obj: LiteLLMLoggingObj,
) -> Any:
) -> ResponsesAPIStreamingResponse:
parsed_chunk = self._normalize_stream_item_id(parsed_chunk)
return super().transform_streaming_response(
model=model,
@ -262,7 +263,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
# Return the responses endpoint
return f"{effective_api_base}/responses"
def _handle_reasoning_item(self, item: dict[str, Any]) -> dict[str, Any]:
def _handle_reasoning_item(self, item: dict[str, object]) -> dict[str, object]:
"""
Handle reasoning items for GitHub Copilot, preserving encrypted_content.
@ -280,7 +281,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
# Filter out None values for known problematic fields,
# but preserve encrypted_content even if it exists
filtered_item: Final[dict[str, Any]] = {}
filtered_item: Final[dict[str, object]] = {}
for k, v in item.items():
# Always include encrypted_content if present (even if None)
if k == "encrypted_content":

View file

@ -4,7 +4,7 @@ Translate from OpenAI's `/v1/chat/completions` to VLLM's `/v1/chat/completions`
import json
from collections.abc import Coroutine
from typing import Any, Final, Literal, cast, overload
from typing import Final, Literal, cast, overload
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_get_image_mime_type_from_url,
@ -28,12 +28,12 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig
class HostedVLLMChatConfig(OpenAIGPTConfig):
def _convert_custom_tools_to_function_tools(self, tools: list[dict[str, Any]]) -> list[dict[str, Any]]:
def _convert_custom_tools_to_function_tools(self, tools: list[dict[str, object]]) -> list[dict[str, object]]:
"""
vLLM chat completions currently accepts only OpenAI function tools.
Convert custom tools into function tools so request validation does not fail.
"""
converted_tools: Final[list[dict[str, Any]]] = []
converted_tools: Final[list[dict[str, object]]] = []
for idx, tool in enumerate(tools):
if not isinstance(tool, dict):
converted_tools.append(tool)
@ -63,17 +63,14 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
"required": ["input"],
}
function_tool: dict[str, Any] = {
"type": "function",
"function": {
"name": str(tool_name),
"parameters": tool_parameters,
},
function_definition: dict[str, object] = {
"name": str(tool_name),
"parameters": tool_parameters,
}
if isinstance(tool_description, str):
function_tool["function"]["description"] = tool_description
function_definition["description"] = tool_description
converted_tools.append(function_tool)
converted_tools.append({"type": "function", "function": function_definition})
return converted_tools
@ -148,7 +145,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
@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
def _transform_messages(
@ -160,7 +157,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
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]]:
"""
Support translating:
- video files from file_id or file_data to video_url

View file

@ -84,13 +84,13 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
typical_p: float | None = None,
watermark: bool | None = None,
) -> None:
locals_: Final = locals().copy()
locals_: Final[dict[str, object]] = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
@classmethod
def get_config(cls):
def get_config(cls) -> dict[str, object]:
return super().get_config()
def get_special_options_params(self):
@ -352,17 +352,17 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
model: str,
data: dict,
api_key: str | None = None,
) -> list[dict[str, Any]]:
) -> list[dict[str, str]]:
streamed_response: Final = CustomStreamWrapper(
completion_stream=response.iter_lines(),
model=model,
custom_llm_provider="huggingface",
logging_obj=logging_obj,
)
content = ""
content: str = ""
for chunk in streamed_response:
content += chunk["choices"][0]["delta"]["content"]
completion_response: Final[list[dict[str, Any]]] = [{"generated_text": content}]
completion_response: Final[list[dict[str, str]]] = [{"generated_text": content}]
## LOGGING
logging_obj.post_call(
input=data,

View file

@ -27,7 +27,6 @@ without the optional STT extras installed.
import asyncio
import inspect
from collections.abc import Callable, Iterable
from types import ModuleType
from typing import TYPE_CHECKING, Any, Final, Protocol
from litellm.litellm_core_utils.audio_utils.utils import (
@ -95,11 +94,37 @@ class _AudioEncoding(Protocol):
def LINEAR_PCM(self) -> object: ...
def _auth_factory(riva_module: ModuleType) -> Callable[..., _RivaAuth]:
class _RivaClientModule(Protocol):
"""The ``riva.client`` entry points this handler calls."""
@property
def Auth(self) -> Callable[..., _RivaAuth]: ...
@property
def ASRService(self) -> Callable[[_RivaAuth], _AsrService]: ...
class _RivaAsrModule(Protocol):
"""The protobuf constructors this handler calls, from whichever module exposes them."""
@property
def AudioEncoding(self) -> _AudioEncoding: ...
@property
def RecognitionConfig(self) -> Callable[..., _RecognitionConfig]: ...
@property
def StreamingRecognitionConfig(self) -> Callable[..., _StreamingRecognitionConfig]: ...
@property
def EndpointingConfig(self) -> Callable[..., _EndpointingConfig]: ...
def _auth_factory(riva_module: _RivaClientModule) -> Callable[..., _RivaAuth]:
return riva_module.Auth
def _audio_encoding(riva_asr_module: ModuleType) -> _AudioEncoding:
def _audio_encoding(riva_asr_module: _RivaAsrModule) -> _AudioEncoding:
return riva_asr_module.AudioEncoding
@ -317,7 +342,7 @@ class NvidiaRivaAudioTranscription:
def _construct_auth(
self,
riva_module: ModuleType,
riva_module: _RivaClientModule,
api_base: str,
api_key: str | None,
optional_params: dict,
@ -349,7 +374,7 @@ class NvidiaRivaAudioTranscription:
return _auth_factory(riva_module)(None, use_ssl, api_base, metadata)
def _build_recognition_config_proto(
self, riva_asr_module: ModuleType, recognition_config_dict: dict[str, Any]
self, riva_asr_module: _RivaAsrModule, recognition_config_dict: dict[str, Any]
) -> _RecognitionConfig:
encoding_name: Final = (recognition_config_dict.get("encoding") or "LINEAR_PCM").upper()
encoding_enum: Final[object] = getattr(
@ -436,7 +461,7 @@ class NvidiaRivaAudioTranscription:
return final_results
def _import_riva() -> tuple[ModuleType, ModuleType]:
def _import_riva() -> tuple[_RivaClientModule, _RivaAsrModule]:
"""
Lazy import of ``riva.client`` and ``riva.client.proto.riva_asr_pb2``.

View file

@ -124,7 +124,7 @@ class OllamaChatConfig(BaseConfig):
setattr(self.__class__, key, value)
@classmethod
def get_config(cls):
def get_config(cls) -> dict[str, object]:
return super().get_config()
def get_supported_openai_params(self, model: str):

View file

@ -231,7 +231,7 @@ class OllamaConfig(BaseConfig):
model: str,
api_base: str | None = None,
api_key: str | None = None,
) -> Any:
) -> dict[str, object] | None:
"""
curl http://localhost:11434/api/show -d '{
"name": "mistral"

View file

@ -17,7 +17,7 @@ This pattern can be replicated for other message formats (e.g., Anthropic).
import json
import time
import uuid
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Union, cast
@ -269,7 +269,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
def _extract_inputs(
self,
message: dict[str, Any],
message: Mapping[str, object],
msg_idx: int,
texts_to_check: list[str],
images_to_check: list[str],
@ -330,7 +330,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
async def _apply_guardrail_responses_to_input_texts(
self,
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
responses: list[str],
task_mappings: list[tuple[int, int | None]],
) -> None:
@ -355,12 +355,12 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
elif isinstance(content, list) and content_idx_optional is not None:
# Replace specific text item in list content
messages[msg_idx]["content"][content_idx_optional]["text"] = guardrail_response
content[content_idx_optional]["text"] = guardrail_response
async def _apply_guardrail_responses_to_input_tool_calls(
self,
messages: list[dict[str, Any]],
tool_calls: list[dict[str, Any]],
messages: Sequence[Mapping[str, object]],
tool_calls: Sequence[Mapping[str, object]],
task_mappings: list[tuple[int, int]],
) -> None:
"""
@ -412,7 +412,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
texts_to_check: Final[list[str]] = []
images_to_check: Final[list[str]] = []
tool_calls_to_check: Final[list[dict[str, Any]]] = []
tool_calls_to_check: Final[list[dict[str, object]]] = []
text_task_mappings: Final[list[tuple[int, int | None]]] = []
tool_call_task_mappings: Final[list[tuple[int, int]]] = []
# text_task_mappings: Track (choice_index, content_index) for each text
@ -461,8 +461,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
guardrailed_texts: Final = guardrailed_inputs.get("texts", [])
returned_tool_calls: Final = guardrailed_inputs.get("tool_calls")
guardrailed_tool_calls: Final[list[dict[str, Any]]] = (
cast(list[dict[str, Any]], returned_tool_calls)
guardrailed_tool_calls: Final[list[dict[str, object]]] = (
cast(list[dict[str, object]], returned_tool_calls)
if isinstance(returned_tool_calls, list) and len(returned_tool_calls) == len(tool_calls_to_check)
else tool_calls_to_check
)
@ -939,7 +939,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
choice_idx: int,
texts_to_check: list[str],
images_to_check: list[str],
tool_calls_to_check: list[dict[str, Any]],
tool_calls_to_check: list[dict[str, object]],
text_task_mappings: list[tuple[int, int | None]],
tool_call_task_mappings: list[tuple[int, int]],
) -> None:

View file

@ -10,7 +10,7 @@ import ssl
import time
import uuid
from collections.abc import AsyncIterator, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Optional
from typing import TYPE_CHECKING, Final, Literal, NamedTuple, Optional
from urllib.parse import urlsplit
import httpx
@ -88,8 +88,8 @@ class OpenAIError(BaseLLMException):
###################################################################
def drop_params_from_unprocessable_entity_error(
e: openai.UnprocessableEntityError | httpx.HTTPStatusError,
data: dict[str, Any],
) -> dict[str, Any]:
data: Mapping[str, object],
) -> dict[str, object]:
"""
Helper function to read OpenAI UnprocessableEntityError and drop the params that raised an error from the error message.

View file

@ -1,7 +1,7 @@
import time
import types
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
from typing import TYPE_CHECKING, Final, Literal, Optional, cast
import httpx
@ -2756,7 +2756,12 @@ class OpenAIAssistantsAPI(BaseLLM):
message_thread: Final = await openai_client.beta.threads.create(**data)
return Thread(**message_thread.dict())
return Thread(
id=message_thread.id,
created_at=message_thread.created_at,
metadata=message_thread.metadata,
object=message_thread.object,
)
# fmt: off
@ -2842,7 +2847,12 @@ class OpenAIAssistantsAPI(BaseLLM):
message_thread: Final = openai_client.beta.threads.create(**data)
return Thread(**message_thread.dict())
return Thread(
id=message_thread.id,
created_at=message_thread.created_at,
metadata=message_thread.metadata,
object=message_thread.object,
)
async def async_get_thread(
self,
@ -2865,7 +2875,12 @@ class OpenAIAssistantsAPI(BaseLLM):
response: Final = await openai_client.beta.threads.retrieve(thread_id=thread_id)
return Thread(**response.dict())
return Thread(
id=response.id,
created_at=response.created_at,
metadata=response.metadata,
object=response.object,
)
# fmt: off
@ -2931,7 +2946,12 @@ class OpenAIAssistantsAPI(BaseLLM):
response: Final = openai_client.beta.threads.retrieve(thread_id=thread_id)
return Thread(**response.dict())
return Thread(
id=response.id,
created_at=response.created_at,
metadata=response.metadata,
object=response.object,
)
def delete_thread(self):
pass
@ -2988,18 +3008,27 @@ class OpenAIAssistantsAPI(BaseLLM):
tools: Iterable[AssistantToolParam] | None,
event_handler: AssistantEventHandler | None,
) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]:
data: Final[dict[str, Any]] = {
"thread_id": thread_id,
"assistant_id": assistant_id,
"additional_instructions": additional_instructions,
"instructions": instructions,
"metadata": metadata,
"model": model,
"tools": tools,
}
runs_stream: Final = client.beta.threads.runs.stream
if event_handler is not None:
data["event_handler"] = event_handler
return client.beta.threads.runs.stream(**data)
return runs_stream(
thread_id=thread_id,
assistant_id=assistant_id,
additional_instructions=additional_instructions,
instructions=instructions,
metadata=metadata,
model=model,
tools=tools,
event_handler=event_handler,
)
return runs_stream(
thread_id=thread_id,
assistant_id=assistant_id,
additional_instructions=additional_instructions,
instructions=instructions,
metadata=metadata,
model=model,
tools=tools,
)
def run_thread_stream(
self,

View file

@ -237,7 +237,7 @@ class OpenAIVideoConfig(BaseVideoConfig):
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: dict[str, Any] | None = None,
extra_body: dict[str, object] | None = None,
) -> tuple[str, dict]:
"""
Transform the video remix request for OpenAI API.
@ -252,7 +252,7 @@ class OpenAIVideoConfig(BaseVideoConfig):
url: Final = f"{api_base.rstrip('/')}/{encoded_video_id}/remix"
# Prepare the request data
data: Final = {"prompt": prompt}
data: Final[dict[str, object]] = {"prompt": prompt}
# Add any extra body parameters
if extra_body:
@ -305,7 +305,7 @@ class OpenAIVideoConfig(BaseVideoConfig):
after: str | None = None,
limit: int | None = None,
order: str | None = None,
extra_query: dict[str, Any] | None = None,
extra_query: dict[str, object] | None = None,
) -> tuple[str, dict]:
"""
Transform the video list request for OpenAI API.

View file

@ -90,20 +90,21 @@ class OpenRouterImageEditConfig(BaseImageEditConfig):
drop_params: bool,
) -> dict:
supported_params: Final = self.get_supported_openai_params(model)
mapped_params: Final[dict[str, Any]] = {}
mapped_params: Final[dict[str, object]] = {}
image_config: Final[dict[str, str]] = {}
for key, value in image_edit_optional_params.items():
if key in supported_params:
if key == "size":
if "image_config" not in mapped_params:
mapped_params["image_config"] = {}
mapped_params["image_config"]["aspect_ratio"] = self._map_size_to_aspect_ratio(cast(str, value))
mapped_params["image_config"] = image_config
image_config["aspect_ratio"] = self._map_size_to_aspect_ratio(cast(str, value))
elif key == "quality":
image_size = self._map_quality_to_image_size(cast(str, value))
if image_size:
if "image_config" not in mapped_params:
mapped_params["image_config"] = {}
mapped_params["image_config"]["image_size"] = image_size
mapped_params["image_config"] = image_config
image_config["image_size"] = image_size
else:
mapped_params[key] = value

View file

@ -130,7 +130,7 @@ class PerplexityEmbeddingConfig(BaseEmbeddingConfig):
if isinstance(embedding_value, str):
raw_bytes: Final = base64.b64decode(embedding_value)
count: Final = len(raw_bytes)
int8_values: Final = struct.unpack(f"{count}b", raw_bytes)
int8_values: Final[tuple[int, ...]] = struct.unpack(f"{count}b", raw_bytes)
return [float(v) / 127.0 for v in int8_values]
return embedding_value

View file

@ -315,7 +315,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
for msg in messages:
if isinstance(msg, dict):
role = msg.get("role", "")
content: Any = msg.get("content", "")
content: object = msg.get("content", "")
msg_cache_control: object = msg.get("cache_control")
else:
role = getattr(msg, "role", "")
@ -463,7 +463,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
return body
def _transform_tool_choice_to_anthropic(self, tool_choice: Any) -> dict[str, Any]:
def _transform_tool_choice_to_anthropic(self, tool_choice: object) -> Mapping[str, object]:
"""
Convert tool_choice from OpenAI format to Anthropic format.

View file

@ -74,7 +74,7 @@ class StabilityImageEditConfig(BaseImageEditConfig):
}
# Create a copy to not mutate original - convert TypedDict to regular dict
mapped_params: Final[dict[str, Any]] = dict(image_edit_optional_params)
mapped_params: Final[dict[str, object]] = dict(image_edit_optional_params)
for k, v in image_edit_optional_params.items():
if k in param_mapping:
@ -182,7 +182,7 @@ class StabilityImageEditConfig(BaseImageEditConfig):
# Build Stability request
# Populate multipart form-data as separate text fields (data) and files.
# Stability expects prompt/output_format/etc. as normal form fields, not file parts.
data: Final[dict[str, Any]] = {
data: Final[dict[str, object]] = {
"output_format": "png", # Default to PNG
}

View file

@ -4,7 +4,7 @@ Translates from OpenAI's `/v1/chat/completions` endpoint to Triton's `/generate`
import json
from collections.abc import AsyncIterator, Iterator
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Final, Literal
from httpx import Headers, Response
@ -172,7 +172,7 @@ class TritonConfig(BaseConfig):
streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse,
sync_stream: bool,
json_mode: bool | None = False,
) -> Any:
) -> "TritonResponseIterator":
return TritonResponseIterator(
streaming_response=streaming_response,
sync_stream=sync_stream,
@ -195,14 +195,14 @@ class TritonGenerateConfig(TritonConfig):
) -> dict:
inference_params: Final = optional_params.copy()
stream: Final = inference_params.pop("stream", False)
data_for_triton: Final[dict[str, Any]] = {
data_for_triton: Final[dict[str, object]] = {
"text_input": prompt_factory(model=model, messages=messages),
"parameters": {
"max_tokens": int(optional_params.get("max_tokens", DEFAULT_MAX_TOKENS_FOR_TRITON)),
**inference_params,
},
"stream": bool(stream),
}
data_for_triton["parameters"].update(inference_params)
return data_for_triton
def transform_response(

View file

@ -280,7 +280,7 @@ class VertexFineTuningAPI(VertexLLM):
vertex_location: str,
vertex_credentials: str,
request_route: str,
):
) -> object:
_auth_header, vertex_project = await self._ensure_access_token_async(
credentials=vertex_credentials,
project_id=vertex_project,
@ -341,5 +341,4 @@ class VertexFineTuningAPI(VertexLLM):
f"Error creating fine tuning job. Status code: {response.status_code}. Response: {response.text}"
)
response_json: Final = response.json()
return response_json
return response.json()

View file

@ -179,7 +179,7 @@ def _apply_gemini_metadata(
part: PartType,
model: str | None,
media_resolution_enum: dict[str, str] | None,
video_metadata: dict[str, Any] | None,
video_metadata: Mapping[str, object] | None,
) -> PartType:
"""
Apply media_resolution and video_metadata parameters to a Gemini part.
@ -480,7 +480,7 @@ def _process_gemini_media(
format: str | None = None,
media_resolution_enum: dict[str, str] | None = None,
model: str | None = None,
video_metadata: dict[str, Any] | None = None,
video_metadata: Mapping[str, object] | None = None,
vertex_project: str | None = None,
vertex_credentials: object = None,
) -> PartType:

View file

@ -1,6 +1,7 @@
import base64
import json
import os
from collections.abc import Mapping
from io import BufferedRandom, BufferedReader, BytesIO
from pathlib import Path
from typing import TYPE_CHECKING, Any, Final, cast
@ -47,11 +48,11 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
image_edit_optional_params: ImageEditOptionalRequestParams,
model: str,
drop_params: bool,
) -> dict[str, Any]:
) -> dict[str, object]:
supported_params: Final = self.get_supported_openai_params(model)
filtered_params = {key: value for key, value in image_edit_optional_params.items() if key in supported_params}
mapped_params: Final[dict[str, Any]] = {}
mapped_params: Final[dict[str, object]] = {}
# Map OpenAI parameters to Imagen format
if "n" in filtered_params:
@ -148,10 +149,10 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
model: str,
prompt: str | None,
image: FileTypes | None,
image_edit_optional_request_params: dict[str, Any],
image_edit_optional_request_params: Mapping[str, object],
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> tuple[dict[str, Any], RequestFiles | None]:
) -> tuple[dict[str, object], RequestFiles | None]:
# Prepare reference images in the correct Imagen format
if image is None:
raise ValueError("Vertex AI Imagen image edit requires at least one reference image.")
@ -182,14 +183,14 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
parameters["guidanceScale"] = 7.5 # Default guidance scale
parameters["seed"] = None # Let Vertex AI choose random seed
request_body: Final[dict[str, Any]] = {
request_body: Final[dict[str, object]] = {
"instances": instances,
"parameters": parameters,
}
payload: Final[Any] = json.dumps(request_body)
payload: Final = json.dumps(request_body)
empty_files: Final = cast(RequestFiles, [])
return cast(tuple[dict[str, Any], RequestFiles | None], (payload, empty_files))
return cast(tuple[dict[str, object], RequestFiles | None], (payload, empty_files))
def transform_image_edit_response(
self,
@ -237,8 +238,8 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
def _prepare_reference_images(
self,
image: FileTypes | list[FileTypes],
image_edit_optional_request_params: dict[str, Any],
) -> list[dict[str, Any]]:
image_edit_optional_request_params: Mapping[str, object],
) -> list[dict[str, object]]:
"""
Prepare reference images in the correct Imagen API format
"""
@ -248,7 +249,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
else:
images = [image]
reference_images: Final[list[dict[str, Any]]] = []
reference_images: Final[list[dict[str, object]]] = []
for idx, img in enumerate(images):
if img is None:
@ -258,7 +259,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
base64_data = base64.b64encode(image_bytes).decode("utf-8")
# Create reference image structure
reference_image = {
reference_image: dict[str, object] = {
"referenceType": "REFERENCE_TYPE_RAW",
"referenceId": idx + 1,
"referenceImage": {"bytesBase64Encoded": base64_data},
@ -272,7 +273,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
mask_bytes: Final = self._read_all_bytes(mask_image)
mask_base64: Final = base64.b64encode(mask_bytes).decode("utf-8")
mask_reference: Final = {
mask_reference: Final[dict[str, object]] = {
"referenceType": "REFERENCE_TYPE_MASK",
"referenceId": len(reference_images) + 1,
"referenceImage": {"bytesBase64Encoded": mask_base64},

View file

@ -218,10 +218,10 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
contents: Final = [{"role": "user", "parts": [{"text": prompt}]}]
# Prepare generation config
generation_config: Final[dict[str, Any]] = {"responseModalities": ["IMAGE"]}
generation_config: Final[dict[str, object]] = {"responseModalities": ["IMAGE"]}
# Seed from user-supplied imageConfig dict; flat params are overlaid for backward compat.
image_config: Final[dict[str, Any]] = dict(optional_params.get("imageConfig") or {})
image_config: Final[dict[str, object]] = dict(optional_params.get("imageConfig") or {})
if "aspectRatio" in optional_params:
image_config["aspectRatio"] = optional_params["aspectRatio"]
@ -242,7 +242,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
elif "n" in optional_params:
generation_config["candidateCount"] = optional_params["n"]
request_body: Final[dict[str, Any]] = {
request_body: Final[dict[str, object]] = {
"contents": contents,
"generationConfig": generation_config,
}

View file

@ -7,7 +7,7 @@ Why separate file? Make it easy to see how transformation works
import math
import uuid
from collections.abc import Mapping
from typing import Any, Final
from typing import Final
import httpx
@ -232,7 +232,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
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,

View file

@ -1,6 +1,6 @@
import types
from collections.abc import AsyncIterator, Iterator
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Final
import httpx
@ -95,7 +95,7 @@ class VertexAILlama3Config(OpenAIGPTConfig):
streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse,
sync_stream: bool,
json_mode: bool | None = False,
) -> Any:
) -> "VertexAILlama3StreamingHandler":
return VertexAILlama3StreamingHandler(
streaming_response=streaming_response,
sync_stream=sync_stream,

View file

@ -27,6 +27,7 @@ if TYPE_CHECKING:
import tiktoken
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
class VertexGemmaConfig(OpenAIGPTConfig):
@ -56,7 +57,7 @@ class VertexGemmaConfig(OpenAIGPTConfig):
self,
model_response: ModelResponse,
stream: bool,
) -> ModelResponse | Any:
) -> "ModelResponse | MockResponseIterator":
"""
Helper method to return fake stream iterator if streaming is requested.
@ -138,7 +139,7 @@ class VertexGemmaConfig(OpenAIGPTConfig):
client: HTTPHandler | httpx.Client | None,
api_base: str,
headers: dict[str, str], # mutable-ok: forwarded to post(headers: dict | None)
request_data: dict[str, Any], # mutable-ok: forwarded to post(json: dict | ...)
request_data: dict[str, object], # mutable-ok: forwarded to post(json: dict | ...)
timeout: float | httpx.Timeout | None,
) -> httpx.Response:
if isinstance(client, HTTPHandler):
@ -173,7 +174,7 @@ class VertexGemmaConfig(OpenAIGPTConfig):
client: AsyncHTTPHandler | httpx.AsyncClient | None,
api_base: str,
headers: dict[str, str], # mutable-ok: forwarded to post(headers: dict | None)
request_data: dict[str, Any], # mutable-ok: forwarded to post(json: dict | ...)
request_data: dict[str, object], # mutable-ok: forwarded to post(json: dict | ...)
timeout: float | httpx.Timeout | None,
) -> httpx.Response:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client

View file

@ -3,7 +3,8 @@ Volcengine Embedding Transformation
Transforms OpenAI embedding requests to Volcengine format
"""
from typing import Any, Final
from collections.abc import Mapping
from typing import Final
import httpx
@ -83,11 +84,11 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig):
def map_openai_params(
self,
non_default_params: dict[str, Any],
optional_params: dict[str, Any],
non_default_params: Mapping[str, object],
optional_params: dict[str, object],
model: str,
drop_params: bool,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Map OpenAI embedding parameters to Volcengine format.

View file

@ -4,8 +4,8 @@ Transformation logic for Voyage AI's /v1/rerank endpoint.
Docs - https://docs.voyageai.com/docs/reranker
"""
from collections.abc import Mapping
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Final
import httpx
@ -34,7 +34,7 @@ class VoyageRerankConfig(BaseRerankConfig):
model: str,
drop_params: bool,
query: str,
documents: list[str | dict[str, Any]],
documents: Sequence[str | Mapping[str, object]],
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: list[str] | None = None,

View file

@ -5,8 +5,8 @@ Docs - https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank
"""
import uuid
from collections.abc import Mapping
from typing import Any, Final, cast
from collections.abc import Mapping, Sequence
from typing import Final, cast
import httpx
@ -96,7 +96,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
model: str,
drop_params: bool,
query: str,
documents: list[str | dict[str, Any]],
documents: Sequence[str | Mapping[str, object]],
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: list[str] | None = None,
@ -178,7 +178,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
transformed_results: Final = []
for result in _results:
transformed_result: dict[str, Any] = {
transformed_result: dict[str, object] = {
"index": result["index"],
"relevance_score": result["score"],
}

View file

@ -16,7 +16,7 @@ construction time (see ``handler.py``) so all normalization is isolated here
and ``RealTimeStreaming`` stays provider-agnostic.
"""
from typing import Any, Final
from typing import Final
class XAIRealtimeNormalizer:
@ -58,7 +58,7 @@ class XAIRealtimeNormalizer:
# Cache content-part objects keyed by (response_id, item_id, content_index)
# so that ``response.content_part.done`` events missing ``part`` can be
# back-filled from earlier ``content_part.added`` / delta-done events.
self._content_part_by_key: dict[tuple, dict[str, Any]] = {}
self._content_part_by_key: dict[tuple, dict[str, object]] = {}
# ---------------------------------------------------------------------------
# Public interface consumed by RealTimeStreaming
@ -140,7 +140,7 @@ class XAIRealtimeNormalizer:
}
self._content_part_by_key[key] = updated
def _resolve_content_part(self, event: dict) -> dict[str, Any]:
def _resolve_content_part(self, event: dict) -> dict[str, object]:
part: Final = event.get("part")
if isinstance(part, dict):
return part
@ -214,7 +214,7 @@ class XAIRealtimeNormalizer:
needs_content: Final = event_type in self._EVENTS_NEEDING_CONTENT_INDEX
if not needs_output and not needs_content:
return event
patch: Final[dict[str, Any]] = {}
patch: Final[dict[str, object]] = {}
if needs_output and "output_index" not in event:
patch["output_index"] = 0
if needs_content and "content_index" not in event:
@ -228,8 +228,8 @@ class XAIRealtimeNormalizer:
# ---------------------------------------------------------------------------
@staticmethod
def _default_ga_usage() -> dict[str, Any]:
default_details: Final[dict[str, Any]] = {
def _default_ga_usage() -> dict[str, object]:
default_details: Final[dict[str, int]] = {
"cached_tokens": 0,
"text_tokens": 0,
"audio_tokens": 0,
@ -243,7 +243,7 @@ class XAIRealtimeNormalizer:
}
@staticmethod
def _normalize_usage(usage: object, *, empty_as_null: bool) -> dict[str, Any] | None:
def _normalize_usage(usage: object, *, empty_as_null: bool) -> dict[str, object] | None:
"""Coerce a usage object into the full OpenAI GA shape.
``empty_as_null=True`` for ``response.created`` (usage optional).
@ -253,12 +253,12 @@ class XAIRealtimeNormalizer:
return None
if not usage:
return None if empty_as_null else XAIRealtimeNormalizer._default_ga_usage()
default_details: Final[dict[str, Any]] = {
default_details: Final[dict[str, int]] = {
"cached_tokens": 0,
"text_tokens": 0,
"audio_tokens": 0,
}
normalized: Final[dict[str, Any]] = {
normalized: Final[dict[str, object]] = {
"total_tokens": usage.get("total_tokens", 0),
"input_tokens": usage.get("input_tokens", 0),
"output_tokens": usage.get("output_tokens", 0),

View file

@ -46,6 +46,7 @@ from litellm.repositories.table_repositories import (
MCPServerOAuthClientRepository,
MCPServerRepository,
MCPUserCredentialsRepository,
PrismaTableRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.verification_token_repository import (
@ -535,11 +536,14 @@ def _user_credential_actions(
return table
class _MCPUserEnvVarsRepository(PrismaTableRepository["prisma_db_models.LiteLLM_MCPUserEnvVars"]):
table_name = "litellm_mcpuserenvvars"
def _user_env_var_actions(
prisma_client: PrismaClient,
) -> "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]":
table: Final[TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]] = prisma_client.db.litellm_mcpuserenvvars
return table
return _MCPUserEnvVarsRepository(prisma_client).table
async def _db_find_user_credential_row(

View file

@ -5612,7 +5612,7 @@ class MCPServerManager:
async def pre_call_tool_check(
self,
name: str,
arguments: dict[str, Any],
arguments: _ToolArguments,
server_name: str,
user_api_key_auth: UserAPIKeyAuth | None,
proxy_logging_obj: ProxyLogging | None,

View file

@ -49,6 +49,7 @@ from litellm.litellm_core_utils.url_utils import async_safe_get
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
header_value,
httpxSpecialProvider,
)
from litellm.proxy._experimental.mcp_server.tool_registry import (
@ -457,7 +458,7 @@ def _raise_for_upstream_failure(
if response.status_code == 401 and relays_upstream_auth:
raise MCPUpstreamAuthError(
status_code=response.status_code,
www_authenticate=response.headers.get("www-authenticate"),
www_authenticate=header_value(response.headers, "www-authenticate"),
server_name=upstream,
)
raise MCPOpenApiUpstreamError(response.status_code, upstream)

View file

@ -18,6 +18,7 @@ TTL ``MCP_SSO_ASSERTION_CACHE_TTL_SECONDS``; invalidation also guards against st
from __future__ import annotations
import json
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Final, Protocol
@ -29,6 +30,8 @@ from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, MCP_SSO_ASSERTION_CACHE_TTL_SECONDS
if TYPE_CHECKING:
from prisma.models import LiteLLM_SSOIdentityAssertion
from litellm.proxy.utils import PrismaClient
_ASSERTION_DECRYPT_LOG_KEY: Final = "sso_identity_assertion"
@ -36,6 +39,34 @@ _STR_ADAPTER: Final[TypeAdapter[str]] = TypeAdapter(str)
_MAYBE_STR_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
class _SSOAssertionTable(Protocol):
"""The ``LiteLLM_SSOIdentityAssertion`` table operations this store calls."""
async def find_unique(self, *, where: Mapping[str, str]) -> LiteLLM_SSOIdentityAssertion | None: ...
async def find_many(self) -> Sequence[LiteLLM_SSOIdentityAssertion]: ...
async def upsert(self, *, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> object: ...
async def update(self, *, where: Mapping[str, str], data: Mapping[str, str]) -> object: ...
class _MCPServerTable(Protocol):
"""The ``LiteLLM_MCPServerTable`` lookup the retention gate calls."""
async def find_first(self, *, where: Mapping[str, str]) -> object | None: ...
def _assertion_table(prisma_client: PrismaClient) -> _SSOAssertionTable:
"""The SSO assertion table, typed so the untyped prisma client surface stops here."""
return prisma_client.db.litellm_ssoidentityassertion
def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable:
"""The MCP server table, typed so the untyped prisma client surface stops here."""
return prisma_client.db.litellm_mcpservertable
class SSOIdentityAssertion(BaseModel):
"""The IdP material an EMA exchange needs: ``id_token`` is the RFC 8693 subject token,
``expires_at`` bounds its usefulness, and the refresh token renews it without re-login."""
@ -163,9 +194,7 @@ async def ema_assertion_retention_enabled() -> bool:
return True
if prisma_client is None:
return False
row: Final = await prisma_client.db.litellm_mcpservertable.find_first(
where={"auth_type": MCPAuth.oauth2_id_jag.value}
)
row: Final = await _mcp_server_table(prisma_client).find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value})
return row is not None
@ -184,7 +213,7 @@ async def persist_sso_identity_assertion(
**({"expires_at": assertion.expires_at.isoformat()} if assertion.expires_at else {}),
}
encoded: Final = _STR_ADAPTER.validate_python(encrypt_value_helper(json.dumps(payload)))
await prisma_client.db.litellm_ssoidentityassertion.upsert(
await _assertion_table(prisma_client).upsert(
where={"user_id": user_id},
data={
"create": {"user_id": user_id, "assertion_b64": encoded},
@ -200,7 +229,7 @@ async def _read_assertion_from_db(user_id: str) -> SSOIdentityAssertion | None:
if prisma_client is None:
return None
row: Final = await prisma_client.db.litellm_ssoidentityassertion.find_unique(where={"user_id": user_id})
row: Final = await _assertion_table(prisma_client).find_unique(where={"user_id": user_id})
if row is None:
return None
raw: Final = _MAYBE_STR_ADAPTER.validate_python(
@ -310,13 +339,13 @@ async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient,
re_encrypted: Final = _STR_ADAPTER.validate_python(
encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
)
await prisma_client.db.litellm_ssoidentityassertion.update(
await _assertion_table(prisma_client).update(
where={"user_id": row.user_id},
data={"assertion_b64": re_encrypted},
)
return True
rows: Final = await prisma_client.db.litellm_ssoidentityassertion.find_many()
rows: Final = await _assertion_table(prisma_client).find_many()
outcomes: Final = [await _rotate_row(row) for row in rows]
verbose_proxy_logger.info(
"rotate_sso_identity_assertions_master_key: rotated %d row(s), skipped %d",

View file

@ -15,7 +15,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from starlette.routing import BaseRoute, Match
from starlette.types import Receive, Scope, Send
from starlette.types import ASGIApp, Receive, Scope, Send
from litellm._logging import verbose_proxy_logger
from litellm.proxy.route_priority import hot_routes_first
@ -304,7 +304,7 @@ class LazyFeatureMiddleware:
def __init__(
self,
app,
app: ASGIApp,
fastapi_app: "FastAPI",
features: tuple[LazyFeature, ...] = LAZY_FEATURES,
):

View file

@ -2,19 +2,29 @@ import asyncio
import json
from collections.abc import Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Final
from typing import Final, Protocol
from pydantic import BaseModel, TypeAdapter
from litellm.proxy._types import LiteLLMRoutes
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
UNKNOWN_CALL_TYPE: Final = "Unknown"
INFO_ROUTES_JSON: Final = json.dumps(LiteLLMRoutes.info_routes.value)
class _SupportsQueryRaw(Protocol):
"""The single database operation the cache-activity queries issue."""
async def query_raw(self, query: str, *args: object) -> Sequence[object]: ...
class _SupportsRawQueryDb(Protocol):
"""A prisma client handle, narrowed to the raw-query surface used here."""
@property
def db(self) -> _SupportsQueryRaw: ...
class CacheActivityGroup(BaseModel):
call_type: str
api_requests: int
@ -150,7 +160,7 @@ def compute_totals(groups: Sequence[CacheActivityGroup]) -> CacheActivityTotals:
async def get_cache_activity(
prisma_client: "PrismaClient",
prisma_client: _SupportsRawQueryDb,
start_date: datetime,
end_date: datetime,
key_aliases: Sequence[str],

View file

@ -167,7 +167,7 @@ def check_regex_or_str_match(request_body_value: Any, regex_str: str) -> bool:
def _is_param_allowed(
param: str,
request_body_value: Any,
request_body_value: object,
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS,
) -> bool:
"""
@ -190,7 +190,7 @@ def _is_param_allowed(
def _allow_model_level_clientside_configurable_parameters(
model: str, param: str, request_body_value: Any, llm_router: Router | None
model: str, param: str, request_body_value: object, llm_router: Router | None
) -> bool:
"""
Check if model is allowed to use configurable client-side params
@ -533,7 +533,7 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router:
return True
def _coerce_metadata_to_dict(value: Any) -> dict[str, Any] | None:
def _coerce_metadata_to_dict(value: object) -> dict[str, object] | None:
"""Return ``value`` as a dict, parsing it from JSON if delivered as a string.
Multipart/form-data and ``extra_body`` callers send ``litellm_metadata``
@ -892,7 +892,7 @@ async def check_if_request_size_is_safe(request: Request) -> bool:
return True
async def check_response_size_is_safe(response: Any) -> bool:
async def check_response_size_is_safe(response: object) -> bool:
"""
Enterprise Only:
- Checks if the response size is within the limit
@ -1525,7 +1525,7 @@ def get_customer_user_header_from_mapping(user_id_mapping) -> list | None:
def _get_customer_id_from_standard_headers(
request_headers: dict | None,
request_headers: Mapping[str, object] | None,
) -> str | None:
"""
Check standard customer ID headers for a customer/end-user ID.
@ -1551,7 +1551,7 @@ def _get_customer_id_from_standard_headers(
return None
def _coerce_user_id_to_str(value: Any) -> str | None:
def _coerce_user_id_to_str(value: object) -> str | None:
"""Return a usable end-user identifier string, or None if the value isn't one.
Always drops non-string structured values (dict/list/tuple/set) because
@ -1578,7 +1578,7 @@ def _coerce_user_id_to_str(value: Any) -> str | None:
# behind the flag preserves backwards compatibility for deployments
# that intentionally pass JSON-encoded user identifiers.
if litellm.validate_end_user_id_in_db and stripped[:1] in ("{", "["):
parsed: Final = safe_json_loads(stripped)
parsed: Final[object] = safe_json_loads(stripped)
if isinstance(parsed, (dict, list)):
return None
return stripped
@ -1586,7 +1586,9 @@ def _coerce_user_id_to_str(value: Any) -> str | None:
return None
def get_end_user_id_from_request_body(request_body: dict, request_headers: dict | None = None) -> str | None:
def get_end_user_id_from_request_body(
request_body: Mapping[str, object], request_headers: Mapping[str, object] | None = None
) -> str | None:
# Import general_settings here to avoid potential circular import issues at module level
# and to ensure it's fetched at runtime.
from litellm.proxy.proxy_server import general_settings
@ -1635,7 +1637,7 @@ def get_end_user_id_from_request_body(request_body: dict, request_headers: dict
if user_id_str:
return user_id_str
def _as_dict(value: Any) -> dict:
def _as_dict(value: object) -> dict:
# metadata / litellm_metadata can arrive as JSON strings from
# multipart/form-data or extra_body; coerce so string-encoded
# payloads can't evade end-user attribution.
@ -1720,11 +1722,11 @@ _MODEL_ROUTING_ID_FIELDS: Final = (
)
def _append_model_candidates(candidates: list[str], value: Any) -> None:
def _append_model_candidates(candidates: list[str], value: object) -> None:
if value is None:
return
values: Final = value if isinstance(value, (list, tuple, set)) else [value]
values: Final[tuple[object, ...]] = tuple(value) if isinstance(value, (list, tuple, set)) else (value,)
for item in values:
if item is None:
continue
@ -1765,7 +1767,7 @@ def _route_uses_model_routing_sources(route: str) -> bool:
def _extract_models_from_managed_resource_id(
resource_id: Any,
resource_id: object,
resource_id_field: str | None = None,
llm_router: Router | None = None,
) -> list[str]:

View file

@ -98,13 +98,15 @@ def _preflight(target: str) -> None:
raise click.ClickException(str(e)) from e
def _start(ctx: click.Context, api_key: str | None, target: str = _CLAUDE_TARGET) -> tuple[StaticToken, _Listing]:
def _start(
ctx: click.Context, base_url: str, api_key: str | None, target: str = _CLAUDE_TARGET
) -> tuple[StaticToken, _Listing]:
_preflight(target)
try:
credential: Final = resolve_credential(ctx, api_key)
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
return credential, _listed_models(ctx.obj["base_url"], credential.token, target)
return credential, _listed_models(base_url, credential.token, target)
def _listing_error(base_url: str, error: PiSyncError, target: str) -> str:
@ -147,9 +149,7 @@ def _validated_model(model: str | None, listing: _Listing, base_url: str) -> str
return starting
def _apply_claude(ctx: click.Context, credential: StaticToken, listing: _Listing, model: str | None) -> None:
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
def _apply_claude(base_url: str, credential: StaticToken, listing: _Listing, model: str | None) -> None:
listed: Final = listing.ids
starting: Final = _validated_model(model, listing, base_url)
settings_path: Final = claude_settings_path(os.environ)
@ -214,8 +214,7 @@ def _pick_codex_model(listed: Sequence[str]) -> str:
return str(inquirer.fuzzy(message="Model Codex starts on (type to filter):", choices=choices).execute())
def _apply_codex(ctx: click.Context, credential: StaticToken, listing: _Listing, model: str) -> None:
base_url: Final[str] = ctx.obj["base_url"]
def _apply_codex(base_url: str, credential: StaticToken, listing: _Listing, model: str) -> None:
_validated_model(model, listing, base_url)
settings_path: Final = codex_config_path(os.environ)
try:
@ -237,13 +236,12 @@ class _Setup:
def _choose_setup(
ctx: click.Context,
base_url: str,
target: str,
credential: StaticToken,
pick_model: Callable[[Sequence[str]], str | None],
pick_codex_model: Callable[[Sequence[str]], str],
) -> _Setup:
base_url: Final[str] = ctx.obj["base_url"]
listing: Final = _listed_models(base_url, credential.token, target)
model: Final = (
pick_model(tuple(item.source_model or item.id for item in listing.models))
@ -270,12 +268,15 @@ def interactive_configure(
credential: Final = resolve_credential(ctx, None)
except ClaudeSettingsError as e:
raise click.ClickException(str(e)) from e
setups: Final = tuple(_choose_setup(ctx, target, credential, pick_model, pick_codex_model) for target in targets)
base_url: Final[str] = ctx.obj["base_url"]
setups: Final = tuple(
_choose_setup(base_url, target, credential, pick_model, pick_codex_model) for target in targets
)
for setup in setups:
if setup.target == _CLAUDE_TARGET:
_apply_claude(ctx, credential, setup.listing, setup.model)
_apply_claude(base_url, credential, setup.listing, setup.model)
elif setup.model is not None:
_apply_codex(ctx, credential, setup.listing, setup.model)
_apply_codex(base_url, credential, setup.listing, setup.model)
class _ConnectionOptions(BaseModel):
@ -283,7 +284,8 @@ class _ConnectionOptions(BaseModel):
gateway_url: str | None = None
def _connection_context(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> click.Context:
def _connection_settings(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> CliContextObj:
"""The context object a subcommand runs with: its own --api-key / --gateway-url over the group's, over `lite`'s."""
ctx_obj: Final[CliContextObj] = ctx.obj
group: Final = (
_ConnectionOptions.model_validate(ctx.parent.params)
@ -300,7 +302,11 @@ def _connection_context(ctx: click.Context, api_key: str | None, gateway_url: st
"api_key": key if key is not None else ctx_obj.get("api_key"),
"api_key_from_token_file": False if key is not None else ctx_obj.get("api_key_from_token_file", False),
}
return click.Context(ctx.command, parent=ctx.parent, obj=connection)
return connection
def _connection_context(ctx: click.Context, settings: CliContextObj) -> click.Context:
return click.Context(ctx.command, parent=ctx.parent, obj=settings)
@click.group(name="configure", invoke_without_command=True)
@ -316,19 +322,19 @@ def configure_group(ctx: click.Context, api_key: str | None, gateway_url: str |
"""
if ctx.invoked_subcommand is not None:
return
connection: Final = _connection_context(ctx, api_key, gateway_url)
settings: Final = _connection_settings(ctx, api_key, gateway_url)
connection: Final = _connection_context(ctx, settings)
if not sys.stdin.isatty():
raise click.ClickException(
"`lite configure` asks questions, so it needs a terminal. Non-interactively, run "
"`lite configure claude --api-key <key> --model <model>` or "
"`lite configure codex --api-key <key> --model <model>`."
)
prompted: Final = (
connection
if connection.obj.get("base_url_explicit")
else _connection_context(connection, None, click.prompt("Gateway URL", default=connection.obj["base_url"]))
)
interactive_configure(prompted)
if settings.get("base_url_explicit"):
interactive_configure(connection)
return
prompted: Final = _connection_settings(connection, None, click.prompt("Gateway URL", default=settings["base_url"]))
interactive_configure(_connection_context(connection, prompted))
@click.group(name="unconfigure")
@ -356,9 +362,9 @@ def configure_claude(ctx: click.Context, api_key: str | None, model: str | None,
setting is kept, and what changed is recorded so `lite unconfigure claude` can put it back.
Assumes the proxy is already running.
"""
connection: Final = _connection_context(ctx, api_key, gateway_url)
credential, listing = _start(connection, api_key)
_apply_claude(connection, credential, listing, model)
settings: Final = _connection_settings(ctx, api_key, gateway_url)
credential, listing = _start(_connection_context(ctx, settings), settings["base_url"], api_key)
_apply_claude(settings["base_url"], credential, listing, model)
@configure_group.command(name="codex")
@ -368,9 +374,9 @@ def configure_claude(ctx: click.Context, api_key: str | None, model: str | None,
@click.pass_context
def configure_codex(ctx: click.Context, api_key: str | None, gateway_url: str | None, model: str) -> None:
"""Route plain `codex` through the gateway until `lite unconfigure codex`."""
connection: Final = _connection_context(ctx, api_key, gateway_url)
credential, listing = _start(connection, api_key, _CODEX_TARGET)
_apply_codex(connection, credential, listing, model)
settings: Final = _connection_settings(ctx, api_key, gateway_url)
credential, listing = _start(_connection_context(ctx, settings), settings["base_url"], api_key, _CODEX_TARGET)
_apply_codex(settings["base_url"], credential, listing, model)
@unconfigure_group.command(name="codex")

View file

@ -1,3 +1,4 @@
from collections.abc import Mapping
from typing import Final, Literal
import click
@ -5,10 +6,17 @@ import rich
import rich.table
from ... import Client
from ._cli_context import cli_context_values
def create_client(ctx: click.Context) -> Client:
return Client(base_url=ctx.obj["base_url"], api_key=ctx.obj["api_key"])
context: Final = cli_context_values(ctx)
return Client(base_url=context["base_url"], api_key=context["api_key"])
def _rendered_field(group: Mapping[str, object], key: str, default: str) -> str:
"""The rendered value of one model group field, or ``default`` when the group omits it."""
return str(group.get(key, default))
@click.group(name="model-groups")
@ -46,10 +54,10 @@ def list_model_groups(ctx: click.Context, output_format: Literal["table", "json"
for group in groups:
table.add_row(
str(group.get("model_group", "")),
str(group.get("mode", "chat")),
str(group.get("input_cost_per_token", "")),
str(group.get("output_cost_per_token", "")),
_rendered_field(group, "model_group", ""),
_rendered_field(group, "mode", "chat"),
_rendered_field(group, "input_cost_per_token", ""),
_rendered_field(group, "output_cost_per_token", ""),
)
rich.print(table)

View file

@ -166,7 +166,8 @@ def up(ctx: click.Context) -> None:
is already running (this does not start one for you). Cursor is not
supported: it has no equivalent file-based config to patch.
"""
base_url: Final = ctx.obj["base_url"]
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
try:
ensure_fresh_login(ctx)

View file

@ -1,3 +1,4 @@
from collections.abc import Mapping
from typing import Any, Final
import requests
@ -50,7 +51,7 @@ class UsersManagementClient:
response.raise_for_status()
return response.json()
def create_user(self, user_data: dict[str, Any]) -> dict[str, Any]:
def create_user(self, user_data: Mapping[str, object]) -> dict[str, Any]:
"""Create a new user (POST /user/new)"""
url: Final = f"{self.base_url}/user/new"
response: Final = requests.post(url, headers=self._get_headers(), json=user_data, timeout=self.timeout)

View file

@ -37,7 +37,7 @@ class CacheCodec:
"""
@staticmethod
def serialize(value: Any, model_type: type[T] | None = None) -> Any:
def serialize(value: object, model_type: type[T] | None = None) -> object:
"""
Encode a value for DualCache / Redis (``json.dumps``-safe).

View file

@ -714,7 +714,7 @@ def strip_callback_config(metadata: dict[str, object] | None) -> dict[str, objec
return {k: v for k, v in metadata.items() if k not in _CALLBACK_CONFIG_SLOTS}
def encrypt_callback_vars(metadata: Any) -> Any:
def encrypt_callback_vars(metadata: object) -> Any:
"""Return a deep copy of metadata with callback_vars values encrypted at rest.
Idempotent: a value that already decrypts cleanly is left unchanged so
@ -723,7 +723,7 @@ def encrypt_callback_vars(metadata: Any) -> Any:
return _transform_callback_vars(metadata, _encrypt_if_plaintext)
def decrypt_callback_vars(metadata: Any) -> Any:
def decrypt_callback_vars(metadata: object) -> Any:
"""Return a deep copy of metadata with callback_vars values decrypted.
Legacy plaintext rows pass through unchanged (decrypt failure → original).
@ -731,7 +731,7 @@ def decrypt_callback_vars(metadata: Any) -> Any:
return _transform_callback_vars(metadata, _decrypt_or_passthrough)
def _transform_callback_vars(metadata: object, transform: Callable[[str, Any], Any]) -> object:
def _transform_callback_vars(metadata: object, transform: Callable[[str, object], object]) -> object:
if not isinstance(metadata, dict):
return metadata
out: Final = copy.deepcopy(metadata)

View file

@ -56,14 +56,18 @@ def _unqualified(annotation: object) -> object:
return _unqualified(qualified[0])
def _union_members(annotation: object) -> tuple[object, ...]:
"""The non-``None`` members of a union annotation, or the annotation itself when it is not a union."""
if get_origin(annotation) not in (Union, UnionType):
return (annotation,)
members: Final[tuple[object, ...]] = get_args(annotation)
return tuple(arg for arg in members if arg is not type(None))
def _numeric_form_type(annotation: object) -> type[int] | type[float] | None:
"""The scalar to parse an ``int``/``float``-typed field as, else ``None``."""
unwrapped: Final = _unqualified(annotation)
candidates: Final = (
tuple(arg for arg in get_args(unwrapped) if arg is not type(None))
if get_origin(unwrapped) in (Union, UnionType)
else (unwrapped,)
)
candidates: Final = _union_members(unwrapped)
if len(candidates) != 1:
return None
if candidates[0] is int:

View file

@ -66,7 +66,7 @@ def map_v3_rate_limit_type(
return None
def _coerce_message(detail: Any) -> str:
def _coerce_message(detail: object) -> str:
"""Best-effort, JSON-friendly stringification of an HTTPException-style detail."""
if detail is None:
return ""
@ -144,7 +144,7 @@ class ProxyRateLimitError(HTTPException, RateLimitError):
def __init__(
self,
detail: Any,
headers: Mapping[str, Any] | None = None,
headers: Mapping[str, object] | None = None,
category: str | RateLimitErrorCategory = RateLimitErrorCategory.LITELLM_RATE_LIMIT,
rate_limit_type: str | RateLimitType | None = None,
model: str | None = None,

View file

@ -49,7 +49,7 @@ from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManage
from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.prisma_protocols import SpendLinkedTable
from litellm.repositories.prisma_protocols import PrismaBatch, SpendLinkedTable
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import (
EndUserRepository,
@ -478,6 +478,11 @@ class ResetBudgetJob:
self.reset_settings: BudgetResetSettings = reset_settings or get_budget_reset_settings()
self.pod_lock_manager: PodLockManager | None = pod_lock_manager
@property
def _new_batch(self) -> Callable[[], PrismaBatch]:
new_batch: Final[Callable[[], PrismaBatch]] = self.prisma_client.db.batch_
return new_batch
async def _lease_is_held(self, lock_manager: PodLockManager) -> bool:
"""True only when the lease is readable and someone holds it.
@ -837,7 +842,7 @@ class ResetBudgetJob:
)
async def _commit_budget_cascade_once(self, cascade: _BudgetCascade) -> None:
async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow:
async with budget_cascade_unit_of_work(self._new_batch) as uow:
_queue_budget_linked_resets(uow.team_memberships, cascade)
_queue_budget_linked_resets(uow.keys, cascade, extra=_LINKED_KEYS_WHERE)
_queue_budget_linked_resets(uow.organizations, cascade, extra=_SPENT_ROWS_WHERE)
@ -959,7 +964,7 @@ class ResetBudgetJob:
)
async def _write_key_reset_updates_once(self, updated_keys: Sequence[_RowReset[LiteLLM_VerificationToken]]) -> None:
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
async with spend_reset_unit_of_work(self._new_batch) as uow:
for k in updated_keys:
if k.row.token is None:
continue
@ -983,7 +988,7 @@ class ResetBudgetJob:
)
async def _write_user_reset_updates_once(self, updated_users: Sequence[_RowReset[LiteLLM_UserTable]]) -> None:
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
async with spend_reset_unit_of_work(self._new_batch) as uow:
for u in updated_users:
uow.users.queue_spend_reset(
user_id=u.row.user_id,
@ -1005,7 +1010,7 @@ class ResetBudgetJob:
)
async def _write_team_reset_updates_once(self, updated_teams: Sequence[_RowReset[LiteLLM_TeamTable]]) -> None:
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
async with spend_reset_unit_of_work(self._new_batch) as uow:
for t in updated_teams:
uow.teams.queue_spend_reset(
team_id=t.row.team_id,

View file

@ -106,7 +106,7 @@ async def create_container(
# Process request using ProxyBaseLLMRequestProcessing
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
response: Final = await processor.base_process_llm_request(
response: Final[object] = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -216,7 +216,7 @@ async def list_containers(
or get_custom_llm_provider_from_request_query(request=request)
or "openai"
)
data: Final[dict[str, Any]] = {
data: Final[dict[str, object]] = {
"query_params": query_params,
"model": query_params.get("model"),
"order": order,
@ -341,7 +341,7 @@ async def retrieve_container(
# Process request using ProxyBaseLLMRequestProcessing
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
container: Final[object] = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -366,6 +366,7 @@ async def retrieve_container(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
return container
@router.delete(
@ -446,7 +447,7 @@ async def delete_container(
# Process request using ProxyBaseLLMRequestProcessing
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
deleted_container: Final[object] = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -471,6 +472,7 @@ async def delete_container(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
return deleted_container
# Register JSON-configured container file endpoints

View file

@ -15,7 +15,7 @@ import traceback
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, cast, overload
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, TypeVar, cast, overload
from urllib.parse import quote, unquote
from pydantic import TypeAdapter
@ -136,6 +136,25 @@ class _SpendBatch(Protocol):
litellm_modelaccessgroupbudgettable: BatchTable
_EntitySpendTable: TypeAlias = Literal[
"litellm_tagtable", "litellm_agentstable", "litellm_modelaccessgroupbudgettable", "litellm_projecttable"
]
_ENTITY_SPEND_TABLES: Final[Mapping[_EntitySpendTable, Callable[[_SpendBatch], BatchTable]]] = MappingProxyType(
{
"litellm_tagtable": lambda batcher: batcher.litellm_tagtable,
"litellm_agentstable": lambda batcher: batcher.litellm_agentstable,
"litellm_modelaccessgroupbudgettable": lambda batcher: batcher.litellm_modelaccessgroupbudgettable,
"litellm_projecttable": lambda batcher: batcher.litellm_projecttable,
}
)
def _entity_spend_table(batcher: _SpendBatch, table_accessor: _EntitySpendTable) -> BatchTable:
return _ENTITY_SPEND_TABLES[table_accessor](batcher)
class _SpendBatchManager(Protocol):
async def __aenter__(self) -> _SpendBatch: ...
@ -2159,9 +2178,7 @@ class DBSpendUpdateWriter:
async def _update_entity_spend_in_db(
entity_name: str,
transactions: dict[str, float] | None,
table_accessor: Literal[
"litellm_tagtable", "litellm_agentstable", "litellm_modelaccessgroupbudgettable", "litellm_projecttable"
],
table_accessor: _EntitySpendTable,
where_field: str,
n_retry_times: int,
prisma_client: PrismaClient,
@ -2195,7 +2212,7 @@ class DBSpendUpdateWriter:
entity_id,
response_cost,
)
getattr(batcher, table_accessor).update_many(
_entity_spend_table(batcher, table_accessor).update_many(
where={where_field: entity_id},
data={"spend": {"increment": response_cost}},
)

Some files were not shown because too many files have changed in this diff Show more