mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
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:
commit
e7bff277a6
170 changed files with 1413 additions and 830 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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``
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", "")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue