Merge litellm_internal_staging into fix/config-yaml-utf8-encoding

Append-append conflict in tests/test_litellm/proxy/test_proxy_server.py:
both sides added tests at the end of the file. Kept both.
This commit is contained in:
Chirag 2026-08-18 08:13:54 +05:30
commit cd5d414c13
430 changed files with 18593 additions and 3092 deletions

View file

@ -43,6 +43,7 @@ jobs:
tests/test_litellm/proxy/video_endpoints
tests/test_litellm/proxy/response_api_endpoints
tests/test_litellm/proxy/image_endpoints
tests/test_litellm/proxy/ocr_endpoints
tests/test_litellm/proxy/vector_store_endpoints
tests/test_litellm/proxy/agent_endpoints
tests/test_litellm/proxy/a2a

View file

@ -35,6 +35,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
# Models & routing config
"/model/",
"/v1/model/info",
"/v1/model/deprecations",
"/v2/model/",
"/model_group",
"/model_access_group/",

View file

@ -1,9 +1,9 @@
{
"reportAny": {
"limit": 22945
"limit": 22343
},
"reportArgumentType": {
"limit": 2579
"limit": 2578
},
"reportAssignmentType": {
"limit": 323
@ -24,13 +24,13 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 7311
"limit": 6991
},
"reportFunctionMemberAccess": {
"limit": 7
},
"reportGeneralTypeIssues": {
"limit": 157
"limit": 154
},
"reportIncompatibleMethodOverride": {
"limit": 56
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5707
"limit": 5681
},
"reportMissingTypeArgument": {
"limit": 15640
"limit": 15608
},
"reportMissingTypeStubs": {
"limit": 40
@ -72,7 +72,7 @@
"limit": 0
},
"reportOptionalMemberAccess": {
"limit": 1069
"limit": 1061
},
"reportOptionalOperand": {
"limit": 0
@ -84,7 +84,7 @@
"limit": 56
},
"reportPrivateUsage": {
"limit": 1824
"limit": 1823
},
"reportRedeclaration": {
"limit": 8
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44776
"limit": 44709
},
"reportUnknownLambdaType": {
"limit": 113
"limit": 112
},
"reportUnknownMemberType": {
"limit": 39237
"limit": 39154
},
"reportUnknownParameterType": {
"limit": 19967
"limit": 19947
},
"reportUnknownVariableType": {
"limit": 30881
"limit": 30772
},
"reportUnnecessaryCast": {
"limit": 117
@ -123,7 +123,7 @@
"limit": 5
},
"reportUnnecessaryIsInstance": {
"limit": 853
"limit": 851
},
"reportUntypedBaseClass": {
"limit": 0

View file

@ -24,12 +24,16 @@ if TYPE_CHECKING:
CHECK_BATCH_COST_USER_AGENT = "LiteLLM Proxy/CheckBatchCost"
TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
PROVIDER_TERMINAL_BATCH_STATUSES: Final[Tuple[str, ...]] = (
"completed",
"complete",
"failed",
"expired",
"cancelled",
)
TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
*PROVIDER_TERMINAL_BATCH_STATUSES,
"stale_expired",
)
@ -286,6 +290,57 @@ class CheckBatchCost:
404 must not retire the row; the staleness sweep bounds it instead."""
return self.llm_router.get_deployment(model_id=model_id) is not None
@staticmethod
def _is_output_file_gone_at_provider(error: Exception, output_file_id: Optional[str]) -> bool:
"""A 404 naming the output file means there is nothing to fetch on this or any
later poll: providers like Vertex AI advertise an output path for every batch,
including terminal ones that never wrote it. Any other failure may be
transient, so it keeps retrying until the staleness sweep bounds it."""
import openai
from litellm.exceptions import NotFoundError
if not output_file_id:
return False
return isinstance(error, (NotFoundError, openai.NotFoundError)) and output_file_id in str(error)
async def _finalize_unbilled_terminal_job(
self, job: "LiteLLM_ManagedObjectTable", 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."""
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
ensure_batch_response_managed_file_ids,
)
response.id = job.unified_object_id
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=self.proxy_logging_obj.get_proxy_hook("managed_files"),
prisma_client=self.prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
db_batch_object=job,
unified_batch_id=_is_base64_encoded_unified_file_id(job.unified_object_id),
)
update_data: Final[dict] = {
"status": response.status,
"file_object": response.model_dump_json(),
**({"batch_processed": True} if self._has_batch_processed_column else {}),
}
await self.prisma_client.db.litellm_managedobjecttable.update(
where={"id": job.id},
data=update_data,
)
verbose_proxy_logger.info(
f"CheckBatchCost: marked job {job.id} as {response.status} in DB"
)
except Exception as db_err:
verbose_proxy_logger.error(
f"CheckBatchCost: failed to mark job {job.id} as {response.status} in DB: {db_err}"
)
@staticmethod
def _record_error(
prom_logger: Optional["PrometheusLogger"], error_type: str
@ -528,6 +583,7 @@ class CheckBatchCost:
from litellm.files.main import afile_content
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
@ -648,15 +704,20 @@ class CheckBatchCost:
f"{_file_attr}={_raw_file_id!r}: {_e}"
)
# Pass deployment model_info so custom batch pricing
# (input_cost_per_token_batches etc.) is used for cost calc
deployment_model_info = deployment_info.model_info.model_dump() if deployment_info.model_info else {}
# Pass the deployment's router-registered pricing (litellm_params custom
# rates merged with the model's published rates) so custom batch pricing
# (input_cost_per_token_batches etc.) is used for cost calc, exactly as
# the inline retrieve path does.
deployment_model_info = deployment_pricing_model_info(
model_id=model_id,
deployment_model=litellm_model_name,
)
batch_cost, batch_usage, batch_models = (
await calculate_batch_cost_and_usage(
file_content_dictionary=file_content_as_dict,
custom_llm_provider=llm_provider, # type: ignore
model_name=model_name,
model_info=deployment_model_info, # type: ignore[arg-type]
model_info=deployment_model_info,
)
)
logging_obj = LiteLLMLogging(
@ -796,7 +857,7 @@ class CheckBatchCost:
## RETRIEVE THE BATCH JOB OUTPUT FILE
if (
response.status in ("completed", "complete", "expired")
response.status in PROVIDER_TERMINAL_BATCH_STATUSES
and response.output_file_id is not None
):
try:
@ -808,6 +869,15 @@ class CheckBatchCost:
prom_logger=prom_logger,
)
except Exception as tracking_err:
if self._is_output_file_gone_at_provider(
tracking_err, response.output_file_id
) and self._batch_deployment_exists(model_id):
verbose_proxy_logger.warning(
f"CheckBatchCost: output file {response.output_file_id} of batch {batch_id} "
f"does not exist at the provider; retiring job {job.id} unbilled"
)
await self._finalize_unbilled_terminal_job(job, response)
continue
verbose_proxy_logger.error(
f"CheckBatchCost: failed to track cost for batch {batch_id} "
f"(job {job.id}); leaving it unprocessed so the next poll retries: {tracking_err}"
@ -837,45 +907,8 @@ class CheckBatchCost:
f"CheckBatchCost: failed to mark job {job.id} complete in DB: {db_err}"
)
elif response.status in (
"completed",
"complete",
"failed",
"expired",
"cancelled",
):
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
ensure_batch_response_managed_file_ids,
)
response.id = job.unified_object_id
await ensure_batch_response_managed_file_ids(
response=response,
managed_files_obj=self.proxy_logging_obj.get_proxy_hook("managed_files"),
prisma_client=self.prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
db_batch_object=job,
unified_batch_id=_is_base64_encoded_unified_file_id(job.unified_object_id),
)
update_data = {
"status": response.status,
"file_object": response.model_dump_json(),
}
if self._has_batch_processed_column:
update_data["batch_processed"] = True
await self.prisma_client.db.litellm_managedobjecttable.update(
where={"id": job.id},
data=update_data,
)
verbose_proxy_logger.info(
f"CheckBatchCost: marked job {job.id} as {response.status} in DB"
)
except Exception as db_err:
verbose_proxy_logger.error(
f"CheckBatchCost: failed to mark job {job.id} as {response.status} in DB: {db_err}"
)
elif response.status in PROVIDER_TERMINAL_BATCH_STATUSES:
await self._finalize_unbilled_terminal_job(job, response)
# Record polling run metrics (always, even if nothing was processed)
if prom_logger:

View file

@ -41,6 +41,7 @@ from litellm.proxy._types import (
CallTypes,
LiteLLM_ManagedFileTable,
LiteLLM_ManagedObjectTable,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.openai_files_endpoints.common_utils import (
@ -423,13 +424,26 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# This is because the encoded object ids stored in the managed objects table do not contain the provider information
# To support provider filtering, we would need to store the provider information in the encoded object ids
if provider:
raise Exception("Filtering by 'provider' is not supported when using managed batches.")
raise ProxyException(
message="Filtering by 'provider' is not supported when using managed batches.",
type="invalid_request_error",
param="provider",
code=400,
)
# Model name filtering is not supported for managed batches
# This is because the encoded object ids stored in the managed objects table do not contain the model name
# A hash of the model name + litellm_params for the model name is encoded as the model id. This is not sufficient to reliably map the target model names to the model ids.
if target_model_names:
raise Exception("Filtering by 'target_model_names' is not supported when using managed batches.")
raise ProxyException(
message="Filtering by 'target_model_names' is not supported when using managed batches.",
type="invalid_request_error",
param="target_model_names",
code=400,
)
if limit == 0:
return build_list_page([])
owner_filter = build_owner_filter(user_api_key_dict)
if owner_filter is None:

View file

@ -11,7 +11,7 @@ Endpoints for /project operations
#### PROJECT MANAGEMENT ####
import json
from collections.abc import Mapping, Sequence
from collections.abc import Sequence
from typing import TYPE_CHECKING
from fastapi import APIRouter, Depends, HTTPException, Request
@ -29,7 +29,11 @@ from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
if TYPE_CHECKING:
from prisma import models as prisma_models
from prisma.actions import LiteLLM_TeamTableActions
from prisma.actions import (
LiteLLM_ProjectTableActions,
LiteLLM_TeamTableActions,
LiteLLM_VerificationTokenActions,
)
router = APIRouter()
@ -39,6 +43,27 @@ def _team_table(prisma_client: PrismaClient) -> "LiteLLM_TeamTableActions[prisma
return team_table
def _project_table(prisma_client: PrismaClient) -> "LiteLLM_ProjectTableActions[prisma_models.LiteLLM_ProjectTable]":
project_table: LiteLLM_ProjectTableActions[prisma_models.LiteLLM_ProjectTable] = (
prisma_client.db.litellm_projecttable
)
return project_table
def _verification_token_table(
prisma_client: PrismaClient,
) -> "LiteLLM_VerificationTokenActions[prisma_models.LiteLLM_VerificationToken]":
verification_token_table: LiteLLM_VerificationTokenActions[prisma_models.LiteLLM_VerificationToken] = (
prisma_client.db.litellm_verificationtoken
)
return verification_token_table
def _jsonified(prisma_client: PrismaClient, payload: dict[str, object]) -> dict[str, object]:
jsonified: dict[str, object] = prisma_client.jsonify_object(payload)
return jsonified
async def _check_user_permission_for_project(
user_api_key_dict: UserAPIKeyAuth,
team_id: str | None,
@ -137,7 +162,7 @@ def _check_team_project_limits(
# --- Validate project models are a subset of team models ---
project_models = data.models
team_models = team_object.models or []
team_models: list[str] = team_object.models or []
if project_models and len(team_models) > 0:
# If team has 'all-proxy-models', skip validation as it allows all models
if SpecialModelNames.all_proxy_models.value not in team_models:
@ -188,11 +213,11 @@ async def _create_budget_for_project(
) -> str:
"""Create a budget for the project and return budget_id."""
budget_params = LiteLLM_BudgetTable.model_fields.keys()
_json_data: Mapping[str, object] = data.json(exclude_none=True)
_json_data: dict[str, object] = data.model_dump(exclude_none=True)
_budget_data = {k: v for k, v in _json_data.items() if k in budget_params}
budget_row = LiteLLM_BudgetTable.model_validate(_budget_data)
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
new_budget = _jsonified(prisma_client, budget_row.model_dump(exclude_none=True))
_budget: prisma_models.LiteLLM_BudgetTable = await prisma_client.db.litellm_budgettable.create(
data={
@ -227,7 +252,7 @@ async def _set_project_object_permission(
return None
def _remove_budget_fields_from_project_data(project_data: dict) -> dict:
def _remove_budget_fields_from_project_data(project_data: dict[str, object]) -> dict[str, object]:
"""
Remove budget fields from project data.
Budget fields belong to LiteLLM_BudgetTable, not LiteLLM_ProjectTable.
@ -396,9 +421,7 @@ async def new_project(
data.project_id = str(uuid.uuid4())
else:
# Check if project_id already exists
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
where={"project_id": data.project_id}
)
existing_project = await _project_table(prisma_client).find_unique(where={"project_id": data.project_id})
if existing_project is not None:
raise ProxyException(
message=f"Project id = {data.project_id} already exists. Please use a different project id.",
@ -423,11 +446,14 @@ async def new_project(
)
# Create project row (following organization_endpoints.py pattern)
project_row = LiteLLM_ProjectTable(
**data.json(exclude_none=True),
object_permission_id=object_permission_id,
created_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name,
project_row_payload: dict[str, object] = data.model_dump(exclude_none=True)
project_row = LiteLLM_ProjectTable.model_validate(
{
**project_row_payload,
"object_permission_id": object_permission_id,
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
}
)
for field in LiteLLM_ManagementEndpoint_MetadataFields:
@ -438,7 +464,7 @@ async def new_project(
value=getattr(data, field),
)
new_project_row = prisma_client.jsonify_object(project_row.json(exclude_none=True))
new_project_row = _jsonified(prisma_client, project_row.model_dump(exclude_none=True))
# Remove budget fields (following organization_endpoints.py pattern)
new_project_row = _remove_budget_fields_from_project_data(new_project_row)
@ -560,7 +586,7 @@ async def update_project(
# Fetch existing project
existing_project: (
prisma_models.LiteLLM_ProjectTable | None
) = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": data.project_id})
) = await _project_table(prisma_client).find_unique(where={"project_id": data.project_id})
if existing_project is None:
raise ProxyException(
@ -617,8 +643,7 @@ async def update_project(
)
# Prepare update data
update_data = data.json(exclude_none=True, exclude={"project_id"})
update_data = prisma_client.jsonify_object(update_data)
update_data = _jsonified(prisma_client, data.model_dump(exclude_none=True, exclude={"project_id"}))
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
# Handle budget updates
@ -660,9 +685,10 @@ async def update_project(
# Handle metadata fields
for field in LiteLLM_ManagementEndpoint_MetadataFields:
if field in update_data:
if update_data.get("metadata") is None:
update_data["metadata"] = {}
update_data["metadata"][field] = update_data.pop(field)
existing_metadata = update_data.get("metadata")
metadata_dict: dict[str, object] = existing_metadata if isinstance(existing_metadata, dict) else {}
metadata_dict[field] = update_data.pop(field)
update_data["metadata"] = metadata_dict
# Remove budget fields (following organization_endpoints.py pattern)
update_data = _remove_budget_fields_from_project_data(update_data)
@ -748,11 +774,11 @@ async def delete_project(
detail={"error": "Only admins can delete projects"},
)
deleted_projects = []
deleted_projects: list[prisma_models.LiteLLM_ProjectTable | None] = []
for project_id in data.project_ids:
# Check if project exists
existing_project = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": project_id})
existing_project = await _project_table(prisma_client).find_unique(where={"project_id": project_id})
if existing_project is None:
raise ProxyException(
@ -765,7 +791,7 @@ async def delete_project(
# Check if there are any keys associated with this project
associated_keys: Sequence[
prisma_models.LiteLLM_VerificationToken
] = await prisma_client.db.litellm_verificationtoken.find_many(where={"project_id": project_id})
] = await _verification_token_table(prisma_client).find_many(where={"project_id": project_id})
if len(associated_keys) > 0:
raise ProxyException(
@ -778,7 +804,7 @@ async def delete_project(
# Delete the project
deleted_project: (
prisma_models.LiteLLM_ProjectTable | None
) = await prisma_client.db.litellm_projecttable.delete(where={"project_id": project_id})
) = await _project_table(prisma_client).delete(where={"project_id": project_id})
await delete_cached_project_object(
project_id=project_id,
@ -829,7 +855,7 @@ async def project_info(
)
# Fetch project
project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.find_unique(
project: prisma_models.LiteLLM_ProjectTable | None = await _project_table(prisma_client).find_unique(
where={"project_id": project_id},
include={"litellm_budget_table": True, "object_permission": True},
)
@ -901,7 +927,7 @@ async def list_projects(
if user_api_key_has_admin_view(user_api_key_dict):
projects: Sequence[
prisma_models.LiteLLM_ProjectTable
] = await prisma_client.db.litellm_projecttable.find_many(
] = await _project_table(prisma_client).find_many(
include={"litellm_budget_table": True, "object_permission": True}
)
else:
@ -911,9 +937,9 @@ async def list_projects(
user_record: prisma_models.LiteLLM_UserTable | None = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_api_key_dict.user_id},
)
user_team_ids: Sequence[str] = user_record.teams if user_record is not None and user_record.teams else []
user_team_ids: list[str] = user_record.teams if user_record is not None and user_record.teams else []
projects = await prisma_client.db.litellm_projecttable.find_many(
projects = await _project_table(prisma_client).find_many(
where={"team_id": {"in": user_team_ids}},
include={"litellm_budget_table": True, "object_permission": True},
)

View file

@ -83,6 +83,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/azure_ai/",
"/aws/",
"/bedrock/",
"/comprehendmedical",
"/cohere/",
"/gemini/",
"/google/",

View file

@ -119,4 +119,7 @@ spec:
{{- end }}
ttlSecondsAfterFinished: {{ .Values.migrationJob.ttlSecondsAfterFinished }}
backoffLimit: {{ .Values.migrationJob.backoffLimit }}
{{- with .Values.migrationJob.activeDeadlineSeconds }}
activeDeadlineSeconds: {{ . }}
{{- end }}
{{- end }}

View file

@ -314,3 +314,31 @@ tests:
operator: Equal
value: litellm-e2e
effect: NoSchedule
- it: bounds the Job with a deadline by default, so a blocked migration cannot stall the release forever
set:
migrationJob:
enabled: true
asserts:
- equal:
path: spec.activeDeadlineSeconds
value: 1800
- it: honours an operator-supplied deadline
set:
migrationJob:
enabled: true
activeDeadlineSeconds: 600
asserts:
- equal:
path: spec.activeDeadlineSeconds
value: 600
- it: omits the deadline entirely when it is nulled out, restoring the unbounded behaviour
set:
migrationJob:
enabled: true
activeDeadlineSeconds: null
asserts:
- notExists:
path: spec.activeDeadlineSeconds

View file

@ -427,6 +427,13 @@ migrationJob:
enabled: true # Enable or disable the schema migration Job
retries: 3 # Number of retries for the Job in case of failure
backoffLimit: 4 # Backoff limit for Job restarts
# Wall-clock budget for the whole Job, shared across every `backoffLimit`
# retry rather than granted per attempt. Without it a migration that blocks
# on the database never fails, and when the Helm hook is enabled the release
# waits on it forever: `helm upgrade` and any GitOps controller driving it
# stop reconciling the whole chart until someone deletes the Job by hand.
# Set to null to opt out and restore the unbounded behaviour.
activeDeadlineSeconds: 1800
disableSchemaUpdate: false # Skip schema migrations for specific environments. When True, the job will exit with code 0.
# Optional service account for the migration job.
# Only used when migrationJob.hooks.helm.enabled=true and serviceAccount.create=true.

View file

@ -24,7 +24,7 @@
"/v1/video" "/v1/videos" "/video" "/videos" "/v1/search" "/search"
"/v1/containers" "/containers" "/v1/evals" "/v1/memory" "/queue/chat"
"/v1beta" "/interactions"
"/anthropic" "/azure" "/azure_ai" "/aws" "/bedrock" "/cohere" "/gemini" "/google"
"/anthropic" "/azure" "/azure_ai" "/aws" "/bedrock" "/comprehendmedical" "/cohere" "/gemini" "/google"
"/vertex_ai" "/vertex-ai" "/assemblyai" "/eu.assemblyai" "/langfuse" "/vllm"
"/mistral" "/groq" "/voyage" "/cursor" "/milvus" "/openai_passthrough"
"/toolset"

View file

@ -21,6 +21,9 @@ metadata:
spec:
backoffLimit: {{ .Values.migrationJob.backoffLimit }}
ttlSecondsAfterFinished: {{ .Values.migrationJob.ttlSecondsAfterFinished }}
{{- with .Values.migrationJob.activeDeadlineSeconds }}
activeDeadlineSeconds: {{ . }}
{{- end }}
template:
metadata:
{{- /* The Job's selector is generated by the controller rather than

View file

@ -167,3 +167,24 @@ tests:
- equal:
path: spec.template.metadata.labels['app.kubernetes.io/component']
value: batch-migrations
- it: bounds the Job with a deadline by default, so a blocked migration cannot stall the release forever
asserts:
- equal:
path: spec.activeDeadlineSeconds
value: 1800
- it: honours an operator-supplied deadline
set:
migrationJob.activeDeadlineSeconds: 600
asserts:
- equal:
path: spec.activeDeadlineSeconds
value: 600
- it: omits the deadline entirely when it is nulled out, restoring the unbounded behaviour
set:
migrationJob.activeDeadlineSeconds: null
asserts:
- notExists:
path: spec.activeDeadlineSeconds

View file

@ -56,6 +56,15 @@ migrationJob:
enabled: true
backoffLimit: 4
ttlSecondsAfterFinished: 120
# Wall-clock budget for the whole Job, shared across every `backoffLimit`
# retry rather than granted per attempt. Without it a migration that blocks
# on the database never fails, and because this is a pre-upgrade hook the
# release waits on it forever: `helm upgrade` and any GitOps controller
# driving it stop reconciling the whole chart until someone deletes the Job
# by hand. A migration that has exhausted its retries is not going to
# succeed on the next one, so failing is strictly better than hanging.
# Set to null to opt out and restore the unbounded behaviour.
activeDeadlineSeconds: 1800
resources: {}
# ServiceAccount for the Job pod only.
#

View file

@ -0,0 +1,16 @@
-- CreateTable
CREATE TABLE "LiteLLM_DailyGuardrailUsageUnits" (
"guardrail_id" TEXT NOT NULL,
"date" TEXT NOT NULL,
"team_id" TEXT NOT NULL,
"api_key" TEXT NOT NULL,
"usage_unit" TEXT NOT NULL,
"units" BIGINT NOT NULL DEFAULT 0,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
CONSTRAINT "LiteLLM_DailyGuardrailUsageUnits_pkey" PRIMARY KEY ("guardrail_id","date","team_id","api_key","usage_unit")
);
-- CreateIndex
CREATE INDEX "LiteLLM_DailyGuardrailUsageUnits_date_idx" ON "LiteLLM_DailyGuardrailUsageUnits"("date");

View file

@ -1069,6 +1069,21 @@ model LiteLLM_DailyGuardrailMetrics {
@@index([guardrail_id])
}
// Daily guardrail billable usage units (one row per guardrail/day/team/key/unit type)
model LiteLLM_DailyGuardrailUsageUnits {
guardrail_id String
date String // YYYY-MM-DD
team_id String // empty string when the request had no team
api_key String // hashed virtual key; empty string when unknown
usage_unit String // provider counter name, e.g. Bedrock's contentPolicyUnits
units BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@id([guardrail_id, date, team_id, api_key, usage_unit])
@@index([date])
}
// Daily policy metrics for usage dashboard (one row per policy per day)
model LiteLLM_DailyPolicyMetrics {
policy_id String

View file

@ -26,6 +26,7 @@ def _dev_env_hot_reload_enabled() -> bool:
if os.getenv("LITELLM_MODE", "DEV") == "DEV":
_dotenv.load_dotenv(override=_dev_env_hot_reload_enabled())
from collections.abc import Sequence
from typing import (
Any,
Callable,
@ -217,6 +218,9 @@ add_user_information_to_llm_headers: Optional[bool] = (
overwrite_user_with_key_hash: bool = (
False # force the outgoing `user` param to the hashed api key, so providers see a stable, tamper-proof id
)
bedrock_request_metadata_fields: Optional[Sequence[str]] = (
None # allow-list of `user_api_key_*` fields (+ `spend_logs_metadata`) sent as Bedrock `requestMetadata`
)
store_audit_logs = False # Enterprise feature, allow users to see audit logs
skip_system_message_in_guardrail: bool = False
skip_tool_message_in_guardrail: bool = False

View file

@ -48,6 +48,7 @@ async def _handle_completed_batch(
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
model_name: str | None = None,
litellm_params: dict | None = None,
model_info: ModelInfo | None = None,
) -> tuple[float, Usage, list[str]]:
"""Fetch a completed batch's output file and aggregate its cost, usage, and
models in a single pass over the JSONL lines, so the parsed file content is
@ -58,7 +59,21 @@ async def _handle_completed_batch(
custom_llm_provider: The LLM provider
model_name: Optional model name
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
model_info: Optional deployment-level model info with custom pricing,
threaded through so a deployment's configured rates win over the
global cost map.
"""
# A completed batch whose request lines all failed has no output file - the
# results are written to a separate error_file_id and output_file_id is None.
# There is nothing to price or measure, so report an empty result set instead
# of calling _fetch_batch_output_file_content, which raises on a missing
# output file. Without this guard the logging worker crashes on every
# aretrieve_batch poll and the completed batch's zero-cost accounting is lost.
# The generic retrieval helper keeps raising for callers that explicitly ask
# for a missing output file.
if batch.output_file_id is None:
return 0.0, Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0), []
file_content = await _fetch_batch_output_file_content(batch, custom_llm_provider, litellm_params=litellm_params)
if (
@ -75,6 +90,7 @@ async def _handle_completed_batch(
entries=_iter_batch_input_entries(file_content),
custom_llm_provider=custom_llm_provider,
model_name=model_name,
model_info=model_info,
)
@ -430,11 +446,23 @@ def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_prov
"""
if custom_llm_provider in ("anthropic", "bedrock"):
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
return AnthropicConfig().calculate_usage(
usage_object=response_body.get("usage", None) or {},
usage_object: Final = response_body.get("usage", None) or {}
if custom_llm_provider == "bedrock" and AmazonConverseConfig.is_converse_usage_shape(usage_object):
return AmazonConverseConfig().usage_from_batch_output(usage_object)
anthropic_usage: Final = AnthropicConfig().calculate_usage(
usage_object=usage_object,
reasoning_content=None,
)
if usage_object and anthropic_usage.total_tokens == 0:
verbose_logger.warning(
"batch output line reported usage this parser does not understand, so it will be billed at $0. "
"provider=%s usage_keys=%s",
custom_llm_provider,
sorted(usage_object.keys()),
)
return anthropic_usage
from litellm.responses.utils import ResponseAPILoggingUtils
_usage_dict: Final = response_body.get("usage", None) or {}

View file

@ -826,7 +826,7 @@ def list_batches(
async def acancel_batch(
batch_id: str,
model: str | None = None,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock"] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -872,7 +872,7 @@ async def acancel_batch(
def cancel_batch(
batch_id: str,
model: str | None = None,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] | str = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock"] | str = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -993,9 +993,14 @@ def cancel_batch(
timeout=timeout,
max_retries=optional_params.max_retries,
)
elif custom_llm_provider == "bedrock":
response = BedrockBatchesHandler.cancel_batch(
batch_id=batch_id,
**kwargs,
)
else:
raise litellm.exceptions.BadRequestError(
message=f"LiteLLM doesn't support {custom_llm_provider} for 'cancel_batch'. Only 'openai', 'azure', and 'vertex_ai' are supported.",
message=f"LiteLLM doesn't support {custom_llm_provider} for 'cancel_batch'. Only 'openai', 'azure', 'vertex_ai', and 'bedrock' are supported.",
model="n/a",
llm_provider=custom_llm_provider,
response=httpx.Response(

View file

@ -49,7 +49,7 @@ if TYPE_CHECKING:
cluster_pipeline = ClusterPipeline
async_redis_client = Redis
async_redis_cluster_client = RedisCluster
Span = _Span | Any
Span = _Span
else:
pipeline = Any
cluster_pipeline = Any
@ -625,7 +625,11 @@ class RedisCache(BaseCache):
f"{self.namespace}-{hashlib.sha256(script.encode()).hexdigest()[:16]}"
)
async def run_script(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
async def run_script(
keys: Sequence[str],
args: Sequence[str | bytes | int | float],
client: object = None,
) -> object:
async def execute() -> object:
executor: Callable[..., Awaitable[Any]] | None = litellm.in_memory_llm_clients_cache.get_cache(
key=script_cache_key
@ -650,7 +654,11 @@ class RedisCache(BaseCache):
if hasattr(_redis_client, "register_script"):
registered_script: Final = _redis_client.register_script(script)
async def standalone_executor(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
async def standalone_executor(
keys: Sequence[str],
args: Sequence[str | bytes | int | float],
client: object = None,
) -> object:
namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys)
return await registered_script(keys=namespaced_keys, args=args, client=client)
@ -659,7 +667,11 @@ class RedisCache(BaseCache):
if hasattr(_redis_client, "script_load"):
script_sha: Final = _redis_client.script_load(script)
async def cluster_executor(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
async def cluster_executor(
keys: Sequence[str],
args: Sequence[str | bytes | int | float],
client: object = None,
) -> object:
namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys)
return await _redis_client.evalsha(script_sha, len(namespaced_keys), *namespaced_keys, *args)
@ -757,7 +769,7 @@ class RedisCache(BaseCache):
async def _pipeline_helper(
self,
pipe: pipeline | cluster_pipeline,
cache_list: list[tuple[Any, Any]],
cache_list: Sequence[tuple[str, object]],
ttl: float | None,
) -> list:
"""
@ -783,7 +795,9 @@ class RedisCache(BaseCache):
return results
@_redis_circuit_breaker_guard
async def async_set_cache_pipeline(self, cache_list: list[tuple[Any, Any]], ttl: float | None = None, **kwargs):
async def async_set_cache_pipeline(
self, cache_list: Sequence[tuple[str, object]], ttl: float | None = None, **kwargs
):
"""
Use Redis Pipelines for bulk write operations
"""
@ -795,7 +809,7 @@ class RedisCache(BaseCache):
start_time: Final = time.time()
print_verbose(f"Set Async Redis Cache: key list: {cache_list}\nttl={ttl}, redis_version={self.redis_version}")
cache_value: Final[Any] = None
cache_value: Final = None
try:
async with _redis_client.pipeline(transaction=False) as pipe:
results: Final = await self._pipeline_helper(pipe, cache_list, ttl)
@ -1074,7 +1088,7 @@ class RedisCache(BaseCache):
# NON blocking - notify users Redis is throwing an exception
verbose_logger.error("litellm.caching.caching: get() - Got exception from REDIS: ", e)
def _run_redis_mget_operation(self, keys: list[str]) -> list[Any]:
def _run_redis_mget_operation(self, keys: list[str]) -> Sequence[bytes | str | None]:
"""
Wrapper to call `mget` on the redis client
@ -1082,7 +1096,7 @@ class RedisCache(BaseCache):
"""
return self.redis_client.mget(keys=keys)
async def _async_run_redis_mget_operation(self, keys: list[str]) -> list[Any]:
async def _async_run_redis_mget_operation(self, keys: list[str]) -> Sequence[bytes | str | None]:
"""
Wrapper to call `mget` on the redis client
@ -1115,7 +1129,7 @@ class RedisCache(BaseCache):
cache_key = self.check_and_fix_namespace(key=cache_key or "")
_keys.append(cache_key)
start_time: Final = time.time()
results: Final[list] = self._run_redis_mget_operation(keys=_keys)
results: Final = self._run_redis_mget_operation(keys=_keys)
end_time: Final = time.time()
_duration: Final = end_time - start_time
self.service_logger_obj.service_success_hook(
@ -1522,7 +1536,7 @@ class RedisCache(BaseCache):
async def async_rpush(
self,
key: str,
values: list[Any],
values: Sequence[str | bytes | int | float],
parent_otel_span: Span | None = None,
**kwargs,
) -> int:

View file

@ -185,6 +185,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
if not isinstance(tool_choice, dict):
return tool_choice
choice_type: Final = tool_choice.get("type")
if isinstance(choice_type, str) and choice_type in ("auto", "none", "required"):
return choice_type
if choice_type not in ("function", "custom"):
return tool_choice
if isinstance(tool_choice.get("name"), str) and tool_choice.get("name"):

View file

@ -1487,6 +1487,7 @@ WEEKLY_SPEND_REPORT_JOB_ID: Final = "weekly_spend_report_job"
MONTHLY_SPEND_REPORT_JOB_ID: Final = "monthly_spend_report_job"
PROMETHEUS_FALLBACK_STATS_JOB_ID: Final = "prometheus_fallback_stats_job"
SLACK_DAILY_REPORT_LOCK_ID: Final = "slack_daily_report"
SLACK_MODEL_DEPRECATION_LOCK_ID: Final = "slack_model_deprecation_warning"
SPEND_LOG_RUN_LOOPS: Final = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500))
SPEND_LOG_CLEANUP_BATCH_SIZE: Final = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000))
SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3))
@ -1593,6 +1594,13 @@ DEFAULT_MCP_ACCESS_GROUP_NEGATIVE_CACHE_TTL: Final = 10
# in a single ``/{name1,name2,...}/mcp`` URL. Bounds the per-request DB / cache
# fan-out an authenticated caller can trigger by stuffing the path with tokens.
DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS: Final = 16
# Ceilings on the cached auth registries; larger tables fall back to per-row lookups
# instead of holding an unbounded id set in every worker.
TAG_REGISTRY_MAX_SIZE: Final = 5000
END_USER_RESTRICTED_REGISTRY_MAX_SIZE: Final = 5000
# How long a failed registry load is remembered as "unusable", so a degraded Postgres
# is not re-scanned on every request on top of the per-id lookups it falls back to.
REGISTRY_ERROR_NEGATIVE_CACHE_TTL: Final = 30
# Sentry Scrubbing Configuration
SENTRY_DENYLIST: Final = [

View file

@ -1,9 +1,11 @@
import asyncio
import contextvars
import json
from collections.abc import Coroutine
from collections.abc import Coroutine, Mapping
from functools import partial
from typing import Any, Final, Literal, overload
from typing import Final, Literal, overload
import httpx
import litellm
from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT
@ -48,16 +50,16 @@ __all__ = [
@client
async def acreate_container(
name: str,
expires_after: dict[str, Any] | None = None,
expires_after: Mapping[str, object] | None = None,
file_ids: list[str] | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
# LiteLLM specific params,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerObject:
"""Asynchronously calls the `create_container` function with the given arguments and keyword arguments.
@ -120,9 +122,9 @@ async def acreate_container(
@overload
def create_container(
name: str,
expires_after: dict[str, Any] | None = None,
expires_after: Mapping[str, object] | None = None,
file_ids: list[str] | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -130,16 +132,16 @@ def create_container(
*,
acreate_container: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerObject]:
) -> Coroutine[object, object, ContainerObject]:
...
@overload
def create_container(
name: str,
expires_after: dict[str, Any] | None = None,
expires_after: Mapping[str, object] | None = None,
file_ids: list[str] | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -156,20 +158,20 @@ def create_container(
@client
def create_container(
name: str,
expires_after: dict[str, Any] | None = None,
expires_after: Mapping[str, object] | None = None,
file_ids: list[str] | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerObject | Coroutine[Any, Any, ContainerObject]:
) -> ContainerObject | Coroutine[object, object, ContainerObject]:
"""Create a container using the OpenAI Container API.
Currently supports OpenAI
@ -281,13 +283,13 @@ async def alist_containers(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerListResponse:
"""Asynchronously list containers.
@ -351,7 +353,7 @@ def list_containers(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -359,7 +361,7 @@ def list_containers(
*,
alist_containers: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerListResponse]:
) -> Coroutine[object, object, ContainerListResponse]:
...
@ -368,7 +370,7 @@ def list_containers(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -387,18 +389,18 @@ def list_containers(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerListResponse | Coroutine[Any, Any, ContainerListResponse]:
) -> ContainerListResponse | Coroutine[object, object, ContainerListResponse]:
"""List containers using the OpenAI Container API.
Currently supports OpenAI
@ -481,13 +483,13 @@ def list_containers(
@client
async def aretrieve_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerObject:
"""Asynchronously retrieve a container.
@ -545,7 +547,7 @@ async def aretrieve_container(
@overload
def retrieve_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -553,14 +555,14 @@ def retrieve_container(
*,
aretrieve_container: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerObject]:
) -> Coroutine[object, object, ContainerObject]:
...
@overload
def retrieve_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -577,18 +579,18 @@ def retrieve_container(
@client
def retrieve_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerObject | Coroutine[Any, Any, ContainerObject]:
) -> ContainerObject | Coroutine[object, object, ContainerObject]:
"""Retrieve a container using the OpenAI Container API.
Currently supports OpenAI
@ -696,13 +698,13 @@ def retrieve_container(
@client
async def adelete_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> DeleteContainerResult:
"""Asynchronously delete a container.
@ -760,7 +762,7 @@ async def adelete_container(
@overload
def delete_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -768,14 +770,14 @@ def delete_container(
*,
adelete_container: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, DeleteContainerResult]:
) -> Coroutine[object, object, DeleteContainerResult]:
...
@overload
def delete_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -792,18 +794,18 @@ def delete_container(
@client
def delete_container(
container_id: str,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> DeleteContainerResult | Coroutine[Any, Any, DeleteContainerResult]:
) -> DeleteContainerResult | Coroutine[object, object, DeleteContainerResult]:
"""Delete a container using the OpenAI Container API.
Currently supports OpenAI
@ -914,11 +916,11 @@ async def alist_container_files(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerFileListResponse:
"""Asynchronously list files in a container.
@ -985,7 +987,7 @@ def list_container_files(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600,
timeout: float | httpx.Timeout = 600,
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -993,7 +995,7 @@ def list_container_files(
*,
alist_container_files: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerFileListResponse]:
) -> Coroutine[object, object, ContainerFileListResponse]:
...
@ -1003,7 +1005,7 @@ def list_container_files(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600,
timeout: float | httpx.Timeout = 600,
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -1023,16 +1025,16 @@ def list_container_files(
after: str | None = None,
limit: int | None = None,
order: str | None = None,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerFileListResponse | Coroutine[Any, Any, ContainerFileListResponse]:
) -> ContainerFileListResponse | Coroutine[object, object, ContainerFileListResponse]:
"""List files in a container using the OpenAI Container API.
Currently supports OpenAI
@ -1125,11 +1127,11 @@ def list_container_files(
async def aupload_container_file(
container_id: str,
file: FileTypes,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerFileObject:
"""Asynchronously upload a file to a container.
@ -1211,7 +1213,7 @@ async def aupload_container_file(
def upload_container_file(
container_id: str,
file: FileTypes,
timeout=600,
timeout: float | httpx.Timeout = 600,
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -1219,7 +1221,7 @@ def upload_container_file(
*,
aupload_container_file: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, ContainerFileObject]:
) -> Coroutine[object, object, ContainerFileObject]:
...
@ -1227,7 +1229,7 @@ def upload_container_file(
def upload_container_file(
container_id: str,
file: FileTypes,
timeout=600,
timeout: float | httpx.Timeout = 600,
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
@ -1245,16 +1247,16 @@ def upload_container_file(
def upload_container_file(
container_id: str,
file: FileTypes,
timeout=600, # default to 10 minutes
timeout: float | httpx.Timeout = 600, # default to 10 minutes
api_key: str | None = None,
api_base: str | None = None,
api_version: str | None = None,
custom_llm_provider: Literal["openai", "azure", "azure_text"] = "openai",
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
extra_body: dict[str, object] | None = None,
**kwargs,
) -> ContainerFileObject | Coroutine[Any, Any, ContainerFileObject]:
) -> ContainerFileObject | Coroutine[object, object, ContainerFileObject]:
"""Upload a file to a container using the OpenAI Container API.
This endpoint allows uploading files directly to a container session,

View file

@ -2160,7 +2160,7 @@ def batch_cost_calculator(
output_cost_per_token: Final = model_info.get("output_cost_per_token")
total_prompt_cost = 0.0
total_completion_cost = 0.0
if input_cost_per_token_batches:
if input_cost_per_token_batches is not None:
total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches
elif input_cost_per_token:
details: Final = parse_prompt_tokens_details(usage)
@ -2180,7 +2180,7 @@ def batch_cost_calculator(
cache_creation_cost: Final = model_info.get("cache_creation_input_token_cost") or input_cost_per_token
total_prompt_cost += cache_creation_tokens * cache_creation_cost / 2
if output_cost_per_token_batches:
if output_cost_per_token_batches is not None:
total_completion_cost = usage.completion_tokens * output_cost_per_token_batches
elif output_cost_per_token:
total_completion_cost = (

View file

@ -5,6 +5,7 @@ import datetime
import os
import random
import time
from collections.abc import Callable
from datetime import timedelta
from typing import TYPE_CHECKING, Any, Final, Literal
@ -17,7 +18,11 @@ import litellm.litellm_core_utils.litellm_logging
import litellm.types
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.constants import HOURS_IN_A_DAY, SLACK_DAILY_REPORT_LOCK_ID
from litellm.constants import (
HOURS_IN_A_DAY,
SLACK_DAILY_REPORT_LOCK_ID,
SLACK_MODEL_DEPRECATION_LOCK_ID,
)
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_type
from litellm.integrations.SlackAlerting.hanging_request_check import (
@ -45,6 +50,10 @@ from litellm.repositories.table_repositories import InvitationLinkRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
from litellm.types.integrations.slack_alerting import *
from litellm.types.proxy.model_deprecation import (
DEFAULT_DEPRECATION_CHECK_INTERVAL_SECONDS,
DEPRECATION_IDLE_POLL_SECONDS,
)
from ..email_templates.templates import *
from .batching_handler import send_to_webhook, squash_payloads
@ -59,6 +68,12 @@ else:
Router = Any
def _proxy_llm_router() -> Router | None:
from litellm.proxy.proxy_server import llm_router
return llm_router
class SlackAlerting(CustomBatchLogger):
"""
Class for sending Slack Alerts
@ -1044,6 +1059,99 @@ Model Info:
async def model_removed_alert(self, model_name: str):
pass
def _deprecation_alerts_enabled(self) -> bool:
return self.alerting is not None and AlertType.model_deprecation_warnings in self.alert_types
async def send_model_deprecation_alert(
self,
llm_router: Router | None = None,
pod_lock_manager: "PodLockManager | None" = None,
) -> bool:
"""Alert on the router's deprecated and imminent models, True when one was sent
The daily lock is claimed only once there is something to say, so an empty pass never blocks a
later real one, and a sent alert is stamped in the shared cache for a day so sibling pods stop asking
"""
if not self._deprecation_alerts_enabled():
return False
from litellm.proxy.common_utils.model_deprecation import (
collect_model_deprecations,
format_deprecation_alert_message,
)
snapshot: Final = collect_model_deprecations(llm_router=llm_router)
message: Final = format_deprecation_alert_message(snapshot)
if message is None:
return False
if not await self._claimed_deprecation_alert_window(pod_lock_manager):
return False
level: Final[Literal["Low", "Medium", "High"]] = "High" if snapshot.deprecated else "Medium"
await self.send_alert(
message=message,
level=level,
alert_type=AlertType.model_deprecation_warnings,
alerting_metadata={ # mutable-ok: send_alert takes a dict payload
"deprecated_count": len(snapshot.deprecated),
"imminent_count": len(snapshot.imminent),
"upcoming_count": len(snapshot.upcoming),
},
)
await self.internal_usage_cache.async_set_cache(
key=SlackAlertingCacheKeys.deprecation_alert_sent_key.value,
value=time.time(),
ttl=DEFAULT_DEPRECATION_CHECK_INTERVAL_SECONDS,
)
return True
async def _claimed_deprecation_alert_window(self, pod_lock_manager: "PodLockManager | None") -> bool:
"""Without a redis backed lock there is no fleet to coordinate, so a lone pod always alerts"""
if pod_lock_manager is None:
return True
return (
await pod_lock_manager.acquire_lock(
cronjob_id=SLACK_MODEL_DEPRECATION_LOCK_ID,
ttl=DEFAULT_DEPRECATION_CHECK_INTERVAL_SECONDS,
allow_reentrant=False,
)
) is not False
async def _deprecation_alert_sent_within_a_day(self) -> bool:
return (
await self.internal_usage_cache.async_get_cache(key=SlackAlertingCacheKeys.deprecation_alert_sent_key.value)
) is not None
async def _run_deprecation_alert_pass(
self, llm_router: Router | None, pod_lock_manager: "PodLockManager | None"
) -> bool:
if llm_router is None or not self._deprecation_alerts_enabled():
return False
if await self._deprecation_alert_sent_within_a_day():
return False
return await self.send_model_deprecation_alert(llm_router=llm_router, pod_lock_manager=pod_lock_manager)
async def run_scheduled_deprecation_check(
self,
get_llm_router: Callable[[], Router | None] = _proxy_llm_router,
pod_lock_manager: "PodLockManager | None" = None,
) -> None:
"""Poll every pass for a loaded router, the alert being on, and no alert in the last day, then alert
A pass that could not alert (no router yet, alert type off, a sibling pod holds the daily lock, or a
redis blip at claim time) is retried on the next poll instead of costing a day, while a pass that
raised (a missing webhook, say) backs off a full day so a misconfiguration logs once, not every poll
"""
while True:
try:
await self._run_deprecation_alert_pass(get_llm_router(), pod_lock_manager)
except Exception as e: # noqa: BLE001 # a failed alert must not kill the loop
verbose_proxy_logger.exception("Error in model deprecation alert loop: %s", e)
await asyncio.sleep(DEFAULT_DEPRECATION_CHECK_INTERVAL_SECONDS)
continue
await asyncio.sleep(DEPRECATION_IDLE_POLL_SECONDS)
async def send_webhook_alert(self, webhook_event: WebhookEvent) -> bool:
"""
Sends structured alert to webhook, if set.

View file

@ -102,6 +102,9 @@ class CustomGuardrail(CustomLogger):
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
use_native_during_call_hook: ClassVar[bool] = False
# If True, every proxy lifecycle event runs this guardrail's own hooks, not apply_guardrail.
use_native_lifecycle_hooks: ClassVar[bool] = False
records_own_guardrail_information: ClassVar[bool] = False
def __init__(
@ -632,7 +635,7 @@ class CustomGuardrail(CustomLogger):
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
def _deployment_pre_call_target(self) -> "CustomLogger":
if not self.uses_apply_guardrail_interface():
if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks:
return self
try:
from litellm.proxy.utils import unified_guardrail

View file

@ -9,13 +9,14 @@ across pods or stop races; the hook reads active jobs through a short-TTL cache.
import asyncio
import hashlib
import random
import traceback
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from itertools import groupby
from operator import itemgetter
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, Literal
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError, field_validator, model_validator
@ -32,6 +33,7 @@ from litellm.litellm_core_utils.llm_judge import (
parse_json_verdict,
)
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.management_endpoints.auto_router_endpoints import ShadowEvalDirection
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN
@ -55,7 +57,7 @@ _MAX_JUDGE_PROMPT_CHARS: Final = 24_000
# The judge answers with a small JSON object; a tighter budget truncates the JSON
# mid-object and the attempt is lost to an error row.
JUDGE_MAX_OUTPUT_TOKENS: Final = 500
JUDGE_MAX_OUTPUT_TOKENS: Final = 1500
_MAX_ERROR_CHARS: Final = 500
@ -305,16 +307,21 @@ Criteria: correctness, completeness, clarity, conciseness.
Return ONLY valid JSON in this exact format, no other text:
{
"preference": "A" | "B" | "tie",
"confidence": <0.0 to 1.0>,
"reasoning": "<one sentence>"
"confidence": <0.0 to 1.0>
}"""
class PairwiseVerdict(BaseModel):
"""The judge's blind A/B verdict, validated at the parse boundary."""
"""The judge's blind A/B verdict: the response_format schema sent with the judge call
and the validation contract on its reply. Both fields are required and preference is
closed over the prompt's labels, so a malformed or truncated reply is an
unparseable-verdict error row, never a defaulted or fabricated verdict."""
preference: str = "tie"
confidence: float = 0.0
preference: Literal["A", "B", "tie"]
confidence: float
PAIRWISE_JUDGE_RESPONSE_FORMAT: Final = type_to_response_format_param(PairwiseVerdict)
def _sample_hits(request_id: str, job_id: str, percentage: float) -> bool:
@ -325,6 +332,14 @@ def _sample_hits(request_id: str, job_id: str, percentage: float) -> bool:
return bucket * 100.0 < percentage
def _failure_detail(e: BaseException) -> str:
"""Exception class, message, and the raising frame, so an attempt's error row names
the faulty code path without needing debug logs on the pod."""
frames: Final = traceback.extract_tb(e.__traceback__)
location: Final = f" at {frames[-1].filename.rsplit('/', 1)[-1]}:{frames[-1].lineno}" if frames else ""
return f"{type(e).__name__}{location}: {e}"
def _judge_call_cost(response: object) -> float:
"""Price a judge call, treating an unmapped judge model as free rather than fatal."""
import litellm
@ -764,7 +779,9 @@ class ShadowEvalLogger(CustomLogger):
try:
response: Final = await router.acompletion(
model=target_model,
messages=messages, # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts
messages=[ # mutable-ok: provider transforms rewrite messages in place, so the router gets its own copy
dict(m) for m in messages
], # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts
metadata=shadow_metadata,
num_retries=0,
fallbacks=[], # mutable-ok: SDK kwarg; a failed shadow is a recorded error, never a spend multiplier
@ -772,7 +789,7 @@ class ShadowEvalLogger(CustomLogger):
)
except Exception as e: # noqa: BLE001 # provider errors become error rows, not crashes
verbose_logger.debug("shadow_eval: router call failed: %s", e)
return _CallFailure(f"shadow router call failed: {e}")
return _CallFailure(f"shadow router call failed: {_failure_detail(e)}")
text: Final = _chat_final_text(response)
if not text:
return _CallFailure("shadow router returned an empty response")
@ -815,6 +832,7 @@ class ShadowEvalLogger(CustomLogger):
judge_messages, # pyright: ignore[reportArgumentType] # plain SDK message dicts
temperature=0,
max_tokens=JUDGE_MAX_OUTPUT_TOKENS,
response_format=PAIRWISE_JUDGE_RESPONSE_FORMAT,
metadata=judge_metadata,
)
except Exception as e: # noqa: BLE001 # judge outages become error rows, not crashes

View file

@ -13,6 +13,7 @@ import traceback
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime as dt_object
from functools import lru_cache
from types import TracebackType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
from httpx import Response
@ -107,6 +108,7 @@ from litellm.types.utils import (
LiteLLMBatch,
LiteLLMLoggingBaseClass,
LiteLLMRealtimeStreamLoggingObject,
ModelInfo,
ModelResponse,
ModelResponseStream,
RawRequestTypedDict,
@ -306,6 +308,66 @@ def _get_cached_prometheus_logger():
return _PrometheusLogger
_DEPLOYMENT_PRICING_KEYS: Final = (
"input_cost_per_token",
"output_cost_per_token",
"input_cost_per_token_batches",
"output_cost_per_token_batches",
)
def deployment_pricing_model_info(model_id: str | None, deployment_model: str | None) -> ModelInfo | None:
"""Pricing the router registered under this deployment's model_info.id.
Returns None when the deployment declares no pricing of its own, so the
caller falls back to the global cost map. The raw registration is what
decides that: the router registers an entry for every deployment, and
get_model_info fills absent costs with 0, so asking it directly cannot
tell "configured as free" apart from "no pricing configured". A deployment
may declare only one side of its pricing, so the side it leaves out keeps
the model's published rates instead of billing as zero. Ownership is per
token direction: declaring either rate for a direction takes that whole
direction, so a published batch rate can never displace a standard rate
the deployment configured itself.
"""
if model_id is None:
return None
registered: Final = litellm.model_cost.get(model_id)
if not isinstance(registered, dict) or not any(registered.get(key) is not None for key in _DEPLOYMENT_PRICING_KEYS):
return None
try:
merged: Final = litellm.get_model_info(model=model_id).copy()
except Exception: # noqa: BLE001 # get_model_info raises for ids it cannot resolve a provider for
return None
published: Final = _published_pricing(deployment_model)
if published is None:
return merged
declares_input: Final = (
registered.get("input_cost_per_token") is not None or registered.get("input_cost_per_token_batches") is not None
)
declares_output: Final = (
registered.get("output_cost_per_token") is not None
or registered.get("output_cost_per_token_batches") is not None
)
if not declares_input:
merged["input_cost_per_token"] = published.get("input_cost_per_token")
merged["input_cost_per_token_batches"] = published.get("input_cost_per_token_batches")
if not declares_output:
merged["output_cost_per_token"] = published.get("output_cost_per_token")
merged["output_cost_per_token_batches"] = published.get("output_cost_per_token_batches")
return merged
def _published_pricing(deployment_model: str | None) -> ModelInfo | None:
"""The cost map's own entry for the deployment's model, when it resolves."""
if deployment_model is None:
return None
try:
return litellm.get_model_info(model=deployment_model)
except Exception: # noqa: BLE001 # no published entry to layer the declared rates over
return None
class Logging(LiteLLMLoggingBaseClass):
global \
supabaseClient, \
@ -578,6 +640,28 @@ class Logging(LiteLLMLoggingBaseClass):
return model_id
return None
def get_deployment_model_for_cost(self) -> str | None:
"""The provider-qualified model to price against.
On a batch retrieve both self.model and litellm_params["model"] can be
unset, and self.model can otherwise carry the router's model_group alias,
which no cost map resolves. model_call_details holds the deployment's own
provider-qualified model, so it is preferred.
"""
candidates: Final = (
(self.model_call_details or {}).get("model") if hasattr(self, "model_call_details") else None,
self.litellm_params.get("model") if hasattr(self, "litellm_params") else None,
self.model,
)
return next((candidate for candidate in candidates if isinstance(candidate, str) and candidate), None)
def get_router_deployment_model_info(self) -> ModelInfo | None:
"""See deployment_pricing_model_info; None means fall back to the global cost map."""
return deployment_pricing_model_info(
model_id=self.get_router_model_id(),
deployment_model=self.get_deployment_model_for_cost(),
)
def update_environment_variables(
self,
litellm_params: dict,
@ -1189,6 +1273,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["additional_args"] = additional_args
self.model_call_details["log_event_type"] = "post_api_call"
attr: Literal["warning", "debug"]
if self.litellm_request_debug:
attr = "warning"
else:
@ -1802,7 +1887,7 @@ class Logging(LiteLLMLoggingBaseClass):
if self.model_call_details.get("litellm_params") is None:
return
metadata_hidden_params: Final = hidden_params.copy()
response_cost: Final = self.model_call_details.get("response_cost")
response_cost: Final[object] = self.model_call_details.get("response_cost")
if metadata_hidden_params.get("response_cost") is None and response_cost is not None:
metadata_hidden_params["response_cost"] = response_cost
@ -1844,7 +1929,10 @@ class Logging(LiteLLMLoggingBaseClass):
logging_result, start_time, end_time
)
if (standard_logging_payload := self.model_call_details.get("standard_logging_object")) is not None:
standard_logging_payload: Final[StandardLoggingPayload | None] = self.model_call_details.get(
"standard_logging_object"
)
if standard_logging_payload is not None:
emit_standard_logging_payload(standard_logging_payload)
def _build_standard_logging_payload(
@ -2109,7 +2197,7 @@ class Logging(LiteLLMLoggingBaseClass):
def _success_handler_body(
self,
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
result: object = None,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
cache_hit: bool | None = None,
@ -2150,7 +2238,10 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["standard_logging_object"] = self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
if (standard_logging_payload := self.model_call_details.get("standard_logging_object")) is not None:
standard_logging_payload: Final[StandardLoggingPayload | None] = self.model_call_details.get(
"standard_logging_object"
)
if standard_logging_payload is not None:
# Only emit for sync requests (async_success_handler handles async)
if is_sync_request:
emit_standard_logging_payload(standard_logging_payload)
@ -2592,7 +2683,9 @@ class Logging(LiteLLMLoggingBaseClass):
) = await _handle_completed_batch(
batch=result,
custom_llm_provider=self.custom_llm_provider,
model_name=self.get_deployment_model_for_cost(),
litellm_params=self.litellm_params,
model_info=self.get_router_deployment_model_info(),
)
result._hidden_params["response_cost"] = response_cost
@ -2981,7 +3074,7 @@ class Logging(LiteLLMLoggingBaseClass):
global_callbacks=litellm.failure_callback,
)
result = None # result sent to all loggers, init this to None incase it's not created
result: object = None # result sent to all loggers, init this to None incase it's not created
result = redact_message_input_output_from_logging(
model_call_details=(self.model_call_details if hasattr(self, "model_call_details") else {}),
@ -3395,11 +3488,11 @@ class Logging(LiteLLMLoggingBaseClass):
def _get_assembled_streaming_response(
self,
result: ModelResponse | TextCompletionResponse | ModelResponseStream | ResponseCompletedEvent | Any,
result: ModelResponse | TextCompletionResponse | ModelResponseStream | ResponseCompletedEvent | object,
start_time: datetime.datetime,
end_time: datetime.datetime,
is_async: bool,
streaming_chunks: list[Any],
streaming_chunks: list[object],
) -> ModelResponse | TextCompletionResponse | ResponsesAPIResponse | None:
if self.stream is not True:
return None
@ -3677,9 +3770,7 @@ def set_callbacks(callback_list, function_id=None):
from sentry_sdk.scrubber import EventScrubber
sentry_sdk_instance = sentry_sdk
sentry_trace_rate = (
os.environ.get("SENTRY_API_TRACE_RATE") if "SENTRY_API_TRACE_RATE" in os.environ else "1.0"
)
sentry_trace_rate = os.environ.get("SENTRY_API_TRACE_RATE", "1.0")
sentry_sample_rate = (
os.environ.get("SENTRY_API_SAMPLE_RATE") if "SENTRY_API_SAMPLE_RATE" in os.environ else "1.0"
)
@ -5150,13 +5241,13 @@ class StandardLoggingPayloadSetup:
# ProxyException uses .code, LiteLLM exceptions use .status_code,
# httpx.HTTPStatusError exposes status only as .response.status_code.
# Stringified for Prisma JSON compatibility.
error_code_attr: Final = getattr(original_exception, "code", None)
error_code_attr: Final[object] = getattr(original_exception, "code", None)
if error_code_attr is not None and str(error_code_attr) not in ("", "None"):
error_status: str = str(error_code_attr)
else:
status_code_attr = getattr(original_exception, "status_code", None)
status_code_attr: object = getattr(original_exception, "status_code", None)
if status_code_attr is None:
response_attr: Final = getattr(original_exception, "response", None)
response_attr: Final[object] = getattr(original_exception, "response", None)
status_code_attr = getattr(response_attr, "status_code", None)
error_status = str(status_code_attr) if status_code_attr is not None else ""
error_class: Final[str] = str(original_exception.__class__.__name__) if original_exception else ""
@ -5165,7 +5256,7 @@ class StandardLoggingPayloadSetup:
# Get traceback information (first 100 lines)
traceback_info = traceback_str or ""
if original_exception:
tb: Final = getattr(original_exception, "__traceback__", None)
tb: Final[TracebackType | None] = getattr(original_exception, "__traceback__", None)
if tb:
tb_lines: Final = traceback.format_tb(tb)
traceback_info += "".join(tb_lines[:MAXIMUM_TRACEBACK_LINES_TO_LOG]) # Limit to first 100 lines
@ -5276,11 +5367,11 @@ class StandardLoggingPayloadSetup:
"""
dynamic_litellm_session_id: Final = litellm_params.get("litellm_session_id")
dynamic_litellm_trace_id: Final = litellm_params.get("litellm_trace_id")
metadata: Final = litellm_params.get("metadata")
metadata: Final[Mapping[str, object] | None] = litellm_params.get("metadata")
metadata_session_id: Final = metadata.get("session_id") if metadata else None
metadata_trace_id: Final = metadata.get("trace_id") if metadata else None
ordered_candidates: Final[tuple[Any, Any, Any, Any]] = (
ordered_candidates: Final[tuple[object, object, object, object]] = (
(dynamic_litellm_trace_id, dynamic_litellm_session_id, metadata_trace_id, metadata_session_id)
if litellm.request_correlation_in_logs
else (dynamic_litellm_session_id, dynamic_litellm_trace_id, metadata_session_id, metadata_trace_id)
@ -5305,10 +5396,10 @@ class StandardLoggingPayloadSetup:
"""
if not litellm.request_correlation_in_logs:
return ""
dynamic_litellm_session_id: Final = litellm_params.get("litellm_session_id")
dynamic_litellm_session_id: Final[object] = litellm_params.get("litellm_session_id")
if dynamic_litellm_session_id:
return str(dynamic_litellm_session_id)
metadata: Final = litellm_params.get("metadata")
metadata: Final[Mapping[str, object] | None] = litellm_params.get("metadata")
metadata_session_id: Final = metadata.get("session_id") if metadata else None
if metadata_session_id:
return str(metadata_session_id)

View file

@ -745,6 +745,8 @@ class RealTimeStreaming:
for callback in litellm.callbacks:
if not isinstance(callback, CustomGuardrail):
continue
if callback.use_native_lifecycle_hooks:
continue
if id(callback) in _already_run:
continue
if not any(callback.should_run_guardrail(data=_check_data, event_type=et) for et in _realtime_event_types):

View file

@ -258,6 +258,12 @@ def perform_redaction(model_call_details: dict, result, redact_streaming_respons
# For async objects, return a simple redacted response without deepcopy
return {"text": "redacted-by-litellm"}
if not (
isinstance(result, (litellm.ModelResponse, litellm.ResponsesAPIResponse, litellm.EmbeddingResponse))
or (isinstance(result, dict) and ("choices" in result or "output" in result))
):
return {"text": "redacted-by-litellm"}
_result: Final = copy.deepcopy(result)
if isinstance(_result, litellm.ModelResponse):
if hasattr(_result, "choices") and _result.choices is not None:

View file

@ -3,7 +3,9 @@ import time
from collections.abc import Iterator, Mapping, Sequence
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypedDict, Union, cast
from typing import TYPE_CHECKING, Any, Final, TypeAlias, TypedDict, Union, cast
from typing_extensions import ReadOnly, Required
from litellm._logging import verbose_logger
from litellm.types.llms.openai import (
@ -14,6 +16,9 @@ from litellm.types.utils import (
CacheCreationTokenDetails,
ChatCompletionAudioResponse,
ChatCompletionCustomToolCallPayload,
ChatCompletionDeltaCustomToolCall,
ChatCompletionDeltaCustomToolCallPayload,
ChatCompletionDeltaToolCall,
ChatCompletionMessageCustomToolCall,
ChatCompletionMessageToolCall,
Choices,
@ -25,6 +30,7 @@ from litellm.types.utils import (
ModelResponseStream,
PromptTokensDetailsWrapper,
ServerToolUse,
StreamingChoices,
Usage,
)
from litellm.utils import print_verbose, token_counter
@ -79,6 +85,51 @@ class _AudioChunk(TypedDict):
choices: Sequence[_AudioChoice]
_ChunkHiddenParams: TypeAlias = dict[str, object]
class _BaseChunk(TypedDict, total=False):
id: ReadOnly[str]
object: ReadOnly[str]
created: ReadOnly[int]
model: ReadOnly[str]
system_fingerprint: ReadOnly[str | None]
choices: ReadOnly[Required[Sequence[StreamingChoices]]]
_hidden_params: ReadOnly[_ChunkHiddenParams]
class _ToolCallFunctionFragment(TypedDict, total=False):
name: ReadOnly[str]
arguments: ReadOnly[str]
provider_specific_fields: ReadOnly[dict[str, object]]
class _ToolCallCustomFragment(TypedDict, total=False):
name: ReadOnly[str]
input: ReadOnly[str]
class _ToolCallFragment(TypedDict, total=False):
index: ReadOnly[int]
id: ReadOnly[str | None]
type: ReadOnly[str | None]
function: ReadOnly[_ToolCallFunctionFragment | Function | None]
custom: ReadOnly[_ToolCallCustomFragment | None]
provider_specific_fields: ReadOnly[dict[str, object] | None]
class _ToolCallDelta(TypedDict, total=False):
tool_calls: ReadOnly[Sequence[_ToolCallFragment | ChatCompletionDeltaToolCall | ChatCompletionDeltaCustomToolCall]]
class _ToolCallChoice(TypedDict, total=False):
delta: ReadOnly[_ToolCallDelta]
class _ToolCallChunk(TypedDict):
choices: ReadOnly[Sequence[_ToolCallChoice]]
class _UsageBearingChunk(TypedDict, total=False):
usage: Usage | None
_hidden_params: Mapping[str, str]
@ -158,7 +209,7 @@ class ChunkProcessor:
return chunks
def update_model_response_with_hidden_params(
self, model_response: ModelResponse, chunk: Mapping[str, dict[str, object]] | None = None
self, model_response: ModelResponse, chunk: "_BaseChunk | None" = None
) -> ModelResponse:
if chunk is None:
return model_response
@ -214,18 +265,18 @@ class ChunkProcessor:
)
@staticmethod
def _get_chunk_id(chunks: Sequence[Mapping[str, str]]) -> str:
def _get_chunk_id(chunks: Sequence["_BaseChunk"]) -> str:
"""
Chunks:
[{"id": ""}, {"id": "1"}, {"id": "1"}]
"""
for chunk in chunks:
if chunk.get("id"):
return chunk["id"]
if chunk_id := chunk.get("id"):
return chunk_id
return ""
@staticmethod
def _get_model_from_chunks(chunks: Sequence[Mapping[str, str]], first_chunk_model: str) -> str:
def _get_model_from_chunks(chunks: Sequence["_BaseChunk"], first_chunk_model: str) -> str:
"""
Get the actual model from chunks, preferring a model that differs from the first chunk.
@ -241,7 +292,7 @@ class ChunkProcessor:
# Fall back to first chunk's model if no different model found
return first_chunk_model
def build_base_response(self, chunks: list[dict[str, Any]]) -> ModelResponse:
def build_base_response(self, chunks: Sequence["_BaseChunk"]) -> ModelResponse:
chunk = self.first_chunk
id: Final = ChunkProcessor._get_chunk_id(chunks)
object: Final = chunk["object"]
@ -292,7 +343,7 @@ class ChunkProcessor:
@staticmethod
def _iter_tool_call_fragments(
tool_call_chunks: Sequence[Mapping[str, Any]],
tool_call_chunks: Sequence["_ToolCallChunk"],
) -> Iterator[tuple[int, str, str]]:
for chunk in tool_call_chunks:
for choice in chunk["choices"]:
@ -306,21 +357,21 @@ class ChunkProcessor:
index = tool_call.get("index", 0)
function = tool_call.get("function")
if isinstance(function, dict):
if function.get("arguments"):
yield index, "arguments", function["arguments"]
elif getattr(function, "arguments", None):
yield index, "arguments", function.arguments
if fragment_arguments := function.get("arguments"):
yield index, "arguments", fragment_arguments
elif function_arguments := getattr(function, "arguments", None):
yield index, "arguments", function_arguments
custom = tool_call.get("custom")
if isinstance(custom, dict) and custom.get("input"):
yield index, "custom_input", custom["input"]
if isinstance(custom, dict) and (custom_input := custom.get("input")):
yield index, "custom_input", custom_input
else:
index = getattr(tool_call, "index", 0)
function = getattr(tool_call, "function", None)
if getattr(function, "arguments", None):
yield index, "arguments", function.arguments
if object_arguments := getattr(function, "arguments", None):
yield index, "arguments", object_arguments
custom = getattr(tool_call, "custom", None)
if getattr(custom, "input", None):
yield index, "custom_input", custom.input
if object_custom_input := getattr(custom, "input", None):
yield index, "custom_input", object_custom_input
@staticmethod
def _join_fragments_by_index_and_field(
@ -337,7 +388,7 @@ class ChunkProcessor:
)
def get_combined_tool_content(
self, tool_call_chunks: Sequence[Mapping[str, Any]]
self, tool_call_chunks: Sequence["_ToolCallChunk"]
) -> list[
ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall
]: # mutable-ok: assigned verbatim to Message.tool_calls, a list field
@ -364,7 +415,7 @@ class ChunkProcessor:
has_function = "function" in tool_call and tool_call["function"] is not None
has_custom = "custom" in tool_call and tool_call["custom"] is not None
else:
has_function = hasattr(tool_call, "function") and tool_call.function is not None
has_function = getattr(tool_call, "function", None) is not None
has_custom = getattr(tool_call, "custom", None) is not None
if not has_function and not has_custom:
@ -387,61 +438,67 @@ class ChunkProcessor:
# Extract id, type, and function data (handle both dict and object)
if isinstance(tool_call, dict):
if tool_call.get("id"):
tool_call_map[index]["id"] = tool_call["id"]
if tool_call.get("type"):
tool_call_map[index]["type"] = tool_call["type"]
if fragment_id := tool_call.get("id"):
tool_call_map[index]["id"] = fragment_id
if fragment_type := tool_call.get("type"):
tool_call_map[index]["type"] = fragment_type
function = tool_call.get("function", {})
if isinstance(function, dict):
if function.get("name"):
tool_call_map[index]["name"] = function["name"]
if fragment_name := function.get("name"):
tool_call_map[index]["name"] = fragment_name
else:
# function is an object
if hasattr(function, "name") and function.name:
tool_call_map[index]["name"] = function.name
if function_name := getattr(function, "name", None):
tool_call_map[index]["name"] = function_name
custom = tool_call.get("custom")
if isinstance(custom, dict):
if custom.get("name"):
tool_call_map[index]["custom_name"] = custom["name"]
if custom_name := custom.get("name"):
tool_call_map[index]["custom_name"] = custom_name
else:
# tool_call is an object
if hasattr(tool_call, "id") and tool_call.id:
tool_call_map[index]["id"] = tool_call.id
if hasattr(tool_call, "type") and tool_call.type:
tool_call_map[index]["type"] = tool_call.type
if hasattr(tool_call, "function"):
if hasattr(tool_call.function, "name") and tool_call.function.name:
tool_call_map[index]["name"] = tool_call.function.name
if object_function_name := getattr(getattr(tool_call, "function", None), "name", None):
tool_call_map[index]["name"] = object_function_name
custom = getattr(tool_call, "custom", None)
if custom is not None:
if getattr(custom, "name", None):
tool_call_map[index]["custom_name"] = custom.name
object_custom: ChatCompletionDeltaCustomToolCallPayload | None = getattr(
tool_call, "custom", None
)
if object_custom is not None:
if getattr(object_custom, "name", None):
tool_call_map[index]["custom_name"] = object_custom.name
# Preserve provider_specific_fields from streaming chunks
provider_fields = None
provider_fields: object = None
if isinstance(tool_call, dict):
provider_fields = tool_call.get("provider_specific_fields")
if not provider_fields and isinstance(tool_call.get("function"), dict):
provider_fields = tool_call["function"].get("provider_specific_fields")
if not provider_fields and isinstance(fragment_function := tool_call.get("function"), dict):
provider_fields = fragment_function.get("provider_specific_fields")
else:
if hasattr(tool_call, "provider_specific_fields") and tool_call.provider_specific_fields:
provider_fields = tool_call.provider_specific_fields
elif (
hasattr(tool_call, "function")
and hasattr(tool_call.function, "provider_specific_fields")
and tool_call.function.provider_specific_fields
):
provider_fields = tool_call.function.provider_specific_fields
object_provider_fields: object = getattr(tool_call, "provider_specific_fields", None)
if object_provider_fields:
provider_fields = object_provider_fields
else:
function_provider_fields: object = getattr(
getattr(tool_call, "function", None),
"provider_specific_fields",
None,
)
if function_provider_fields:
provider_fields = function_provider_fields
if provider_fields:
# Merge provider_specific_fields if multiple chunks have them
if tool_call_map[index]["provider_specific_fields"] is None:
tool_call_map[index]["provider_specific_fields"] = {}
merged_provider_fields = tool_call_map[index]["provider_specific_fields"]
if merged_provider_fields is None:
merged_provider_fields = {}
tool_call_map[index]["provider_specific_fields"] = merged_provider_fields
if isinstance(provider_fields, dict):
tool_call_map[index]["provider_specific_fields"].update(provider_fields)
merged_provider_fields.update(provider_fields)
joined_fragments: Final = self._join_fragments_by_index_and_field(
self._iter_tool_call_fragments(tool_call_chunks)
@ -762,19 +819,14 @@ class ChunkProcessor:
server_tool_use = usage_chunk.server_tool_use
else:
server_tool_use = ServerToolUse.model_validate(usage_chunk.server_tool_use)
if (
usage_chunk_dict["prompt_tokens_details"] is not None
and getattr(
if usage_chunk_dict["prompt_tokens_details"] is not None:
chunk_web_search_requests: int | None = getattr(
usage_chunk_dict["prompt_tokens_details"],
"web_search_requests",
None,
)
is not None
):
web_search_requests = getattr(
usage_chunk_dict["prompt_tokens_details"],
"web_search_requests",
)
if chunk_web_search_requests is not None:
web_search_requests = chunk_web_search_requests
prompt_tokens_details = usage_chunk_dict["prompt_tokens_details"] or prompt_tokens_details

View file

@ -6,7 +6,7 @@ import logging
import threading
import time
import traceback
from collections.abc import AsyncIterator, Callable, Iterator, Mapping, Sequence
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Final, NoReturn, Protocol, TypeVar, cast
@ -155,6 +155,33 @@ class _TextCompletionChoiceLike(Protocol):
finish_reason: str | None
class _VertexFunctionCallLike(Protocol):
name: str
args: Mapping[str, Iterable[object]]
class _VertexPartLike(Protocol):
function_call: _VertexFunctionCallLike
class _VertexContentLike(Protocol):
parts: Sequence[_VertexPartLike]
class _VertexFinishReasonLike(Protocol):
name: str
class _VertexCandidateLike(Protocol):
content: _VertexContentLike
finish_reason: _VertexFinishReasonLike
class _VertexChunkLike(Protocol):
text: str
candidates: Sequence[_VertexCandidateLike]
class CustomStreamWrapper:
def __init__(
self,
@ -291,13 +318,13 @@ class CustomStreamWrapper:
that has since taken over the same Task/thread's context.
"""
try:
logging_obj: Final = getattr(self, "logging_obj", None)
logging_obj: Final[object | None] = getattr(self, "logging_obj", None)
if logging_obj is None:
return
method_name: Final = (
"_restore_correlation_context_if_unclaimed" if guarded else "_restore_correlation_context"
)
restore: Final = getattr(logging_obj, method_name, None)
restore: Final[Callable[[], object] | None] = getattr(logging_obj, method_name, None)
if restore is not None:
restore()
except Exception as restore_error: # noqa: BLE001 # best-effort cleanup; must not raise into the caller
@ -1261,18 +1288,18 @@ class CustomStreamWrapper:
raise Exception("An unknown error occurred with the stream")
self.received_finish_reason = "stop"
elif self.custom_llm_provider == "vertex_ai" and not isinstance(chunk, ModelResponseStream):
chunk = cast(Any, chunk)
vertex_chunk: Final = cast(_VertexChunkLike, chunk)
import proto
if hasattr(chunk, "candidates") is True:
if hasattr(vertex_chunk, "candidates") is True:
try:
try:
completion_obj["content"] = chunk.text
completion_obj["content"] = vertex_chunk.text
except Exception as e:
original_exception: Final = e
if "Part has no text." in str(e):
## check for function calling
function_call: Final = chunk.candidates[0].content.parts[0].function_call
function_call: Final = vertex_chunk.candidates[0].content.parts[0].function_call
args_dict: Final = {}
@ -1311,15 +1338,15 @@ class CustomStreamWrapper:
else:
raise original_exception
if (
hasattr(chunk.candidates[0], "finish_reason")
and chunk.candidates[0].finish_reason.name != "FINISH_REASON_UNSPECIFIED"
hasattr(vertex_chunk.candidates[0], "finish_reason")
and vertex_chunk.candidates[0].finish_reason.name != "FINISH_REASON_UNSPECIFIED"
): # every non-final chunk in vertex ai has this
self.received_finish_reason = map_finish_reason(chunk.candidates[0].finish_reason.name)
self.received_finish_reason = map_finish_reason(vertex_chunk.candidates[0].finish_reason.name)
except Exception:
if chunk.candidates[0].finish_reason.name == "SAFETY":
raise Exception(f"The response was blocked by VertexAI. {chunk}")
if vertex_chunk.candidates[0].finish_reason.name == "SAFETY":
raise Exception(f"The response was blocked by VertexAI. {vertex_chunk}")
else:
completion_obj["content"] = str(chunk)
completion_obj["content"] = str(vertex_chunk)
elif self.custom_llm_provider == "petals":
if self.completion_stream is None or len(self.completion_stream) == 0:
if self.received_finish_reason is not None:
@ -1357,13 +1384,14 @@ class CustomStreamWrapper:
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
if response_obj["usage"] is not None:
_text_completion_usage: Final[Usage] = response_obj["usage"]
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].prompt_tokens,
completion_tokens=response_obj["usage"].completion_tokens,
total_tokens=response_obj["usage"].total_tokens,
prompt_tokens=_text_completion_usage.prompt_tokens,
completion_tokens=_text_completion_usage.completion_tokens,
total_tokens=_text_completion_usage.total_tokens,
),
)
elif self.custom_llm_provider == "text-completion-codestral":
@ -1395,15 +1423,17 @@ class CustomStreamWrapper:
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "cached_response":
chunk = cast(ModelResponseStream, chunk)
chunk_finish_reason: Final = chunk.choices[0].finish_reason
cached_chunk: Final = cast(ModelResponseStream, chunk)
chunk_finish_reason: Final = cached_chunk.choices[0].finish_reason
response_obj = {
"text": chunk.choices[0].delta.content,
"text": cached_chunk.choices[0].delta.content,
"is_finished": chunk_finish_reason is not None,
"finish_reason": chunk_finish_reason,
"original_chunk": chunk,
"original_chunk": cached_chunk,
"tool_calls": (
chunk.choices[0].delta.tool_calls if hasattr(chunk.choices[0].delta, "tool_calls") else None
cached_chunk.choices[0].delta.tool_calls
if hasattr(cached_chunk.choices[0].delta, "tool_calls")
else None
),
}
@ -1411,11 +1441,11 @@ class CustomStreamWrapper:
if response_obj["tool_calls"] is not None:
completion_obj["tool_calls"] = response_obj["tool_calls"]
print_verbose(f"completion obj content: {completion_obj['content']}")
if hasattr(chunk, "id"):
model_response.id = chunk.id
self.response_id = chunk.id
if hasattr(chunk, "system_fingerprint"):
self.system_fingerprint = chunk.system_fingerprint
if hasattr(cached_chunk, "id"):
model_response.id = cached_chunk.id
self.response_id = cached_chunk.id
if hasattr(cached_chunk, "system_fingerprint"):
self.system_fingerprint = cached_chunk.system_fingerprint
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
else: # openai / azure chat model
@ -1563,6 +1593,7 @@ class CustomStreamWrapper:
if self.stream_options is not None and self.stream_options["include_usage"] is True:
model_response.choices = []
return model_response
self._record_usage_only_chunk(model_response=model_response)
return
## CHECK FOR TOOL USE
@ -1789,6 +1820,16 @@ class CustomStreamWrapper:
model_response.choices[0].finish_reason = "tool_calls"
return model_response
def _record_usage_only_chunk(self, model_response: "ModelResponseStream") -> None:
"""
Keep provider usage-only chunks (e.g. OpenRouter's post-finish chunk, which carries a
provider-reported cost) available to cost tracking. They are never returned to the
caller; ``stream_options.include_usage`` only controls what the caller sees.
"""
if getattr(model_response, "usage", None) is None:
return
self.chunks.append(model_response.model_copy(update={"choices": []}))
@staticmethod
def _propagate_usage_cost_to_hidden_params(
response: "ModelResponse",
@ -2310,16 +2351,16 @@ class CustomStreamWrapper:
def _normalize_status_code(exc: Exception) -> int | None:
"""Best-effort status_code extraction."""
try:
code: Final = getattr(exc, "status_code", None)
code: Final[int | str | None] = getattr(exc, "status_code", None)
if code is not None:
return int(code)
except Exception:
pass
response: Final = getattr(exc, "response", None)
response: Final[object | None] = getattr(exc, "response", None)
if response is not None:
try:
status_code: Final = getattr(response, "status_code", None)
status_code: Final[int | str | None] = getattr(response, "status_code", None)
if status_code is not None:
return int(status_code)
except Exception:

View file

@ -13,7 +13,7 @@ Pattern Overview:
"""
import json
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from copy import deepcopy
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, cast
@ -61,6 +61,7 @@ if TYPE_CHECKING:
ModifyResponseException,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@ -123,7 +124,7 @@ class AnthropicMessagesHandler(BaseTranslation):
@staticmethod
def _build_streaming_usage_response(
responses_so_far: list[Any],
responses_so_far: list[object],
request_data: dict | None,
) -> ModelResponse | None:
chunks: Final = tuple(response for response in responses_so_far if isinstance(response, (str, bytes)))
@ -141,7 +142,7 @@ class AnthropicMessagesHandler(BaseTranslation):
self,
exc: "ModifyResponseException",
stream_started: bool = False,
responses_so_far: list[Any] | None = None,
responses_so_far: list[object] | None = None,
) -> list[bytes]:
"""
Build an Anthropic SSE sequence delivering the guardrail block message
@ -184,7 +185,7 @@ class AnthropicMessagesHandler(BaseTranslation):
)
return list(FakeAnthropicMessagesStreamIterator(response=block_response))
def _block_continuation_chunks(self, exc: "ModifyResponseException", responses_so_far: list[Any]) -> list[bytes]:
def _block_continuation_chunks(self, exc: "ModifyResponseException", responses_so_far: list[object]) -> list[bytes]:
"""Continue an already-started message: close the open content block,
append the block message as a new text block, then end the message --
without a second message_start."""
@ -234,7 +235,7 @@ class AnthropicMessagesHandler(BaseTranslation):
@staticmethod
def _content_block_state(
responses_so_far: list[Any],
responses_so_far: list[object],
) -> tuple[int | None, int | None]:
"""From the SSE chunks already sent to the client, return (open
content-block index or None, highest content-block index seen or None).
@ -260,7 +261,7 @@ class AnthropicMessagesHandler(BaseTranslation):
return open_index, max_index
@staticmethod
def _iter_sse_events(item: Any) -> list[dict]:
def _iter_sse_events(item: object) -> list[dict[str, object]]:
"""Yield the event-data dicts in one stream chunk.
Handles both formats this stream can carry (see
@ -271,14 +272,16 @@ class AnthropicMessagesHandler(BaseTranslation):
return [item]
if not isinstance(item, (bytes, bytearray)):
return []
events: Final[list[dict]] = []
events: Final[list[dict[str, object]]] = []
for block in item.decode("utf-8", errors="replace").split("\n\n"):
for line in block.split("\n"):
line = line.strip()
if not line.startswith("data:"):
continue
try:
parsed = json.loads(line[len("data:") :].strip())
parsed: str | int | float | bool | None | Sequence[object] | Mapping[str, object] = json.loads(
line[len("data:") :].strip()
)
except json.JSONDecodeError:
continue
if isinstance(parsed, dict):
@ -315,7 +318,7 @@ class AnthropicMessagesHandler(BaseTranslation):
self,
data: dict,
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Any | None = None,
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
) -> Any:
"""
Process input messages by applying guardrails to text content.
@ -467,8 +470,8 @@ class AnthropicMessagesHandler(BaseTranslation):
@staticmethod
def _openai_system_message_to_anthropic(
message: dict[str, Any],
) -> dict[str, Any] | None: # mutable-ok: API message payload
message: dict[str, object],
) -> dict[str, object] | None: # mutable-ok: API message payload
"""Convert an OpenAI system message to the client's Anthropic-shaped entry."""
content: Final = message.get("content")
if isinstance(content, str):
@ -477,14 +480,14 @@ class AnthropicMessagesHandler(BaseTranslation):
) # mutable-ok: API message payload
if not isinstance(content, list):
return None
blocks: Final[list[dict[str, Any]]] = [] # mutable-ok: API message payload
blocks: Final[list[dict[str, object]]] = [] # mutable-ok: API message payload
for block in content:
if not isinstance(block, dict) or block.get("type") != "text":
continue
text = block.get("text")
if not isinstance(text, str) or not text:
continue
anthropic_block: dict[str, Any] = { # mutable-ok: API message payload
anthropic_block: dict[str, object] = { # mutable-ok: API message payload
"type": "text",
"text": text,
} # mutable-ok: API message payload
@ -496,6 +499,39 @@ class AnthropicMessagesHandler(BaseTranslation):
{"role": "system", "content": blocks} if blocks else None # mutable-ok: API message payload
) # mutable-ok: API message payload
@staticmethod
def _fold_leading_systems_into_top_level(
data: dict[str, object], # mutable-ok: API message payload
leading_systems: Sequence[object],
include_existing_system: bool,
) -> None:
"""Deliver leading system rows through Anthropic's top-level system param, which rejects them in messages."""
existing: Final = data.get("system") if include_existing_system else None
existing_blocks: Final[list[object]] = ( # mutable-ok: API message payload
[{"type": "text", "text": existing}]
if isinstance(existing, str) and existing
else list(existing)
if isinstance(existing, list)
else []
)
converted_rows: Final = tuple(
AnthropicMessagesHandler._openai_system_message_to_anthropic(message)
for message in leading_systems
if isinstance(message, dict)
)
folded: Final[list[object]] = existing_blocks + [ # mutable-ok: API message payload
block
for row in converted_rows
if row is not None
for block in (
[{"type": "text", "text": row["content"]}] if isinstance(row["content"], str) else row["content"]
)
]
if folded:
data["system"] = folded # rebind-ok: write-back mutates the request payload in place
else:
data.pop("system", None)
@staticmethod
def _is_hoisted_top_level_system(message: object, hoisted_system_message: object) -> bool:
"""Match the hoisted prompt by identity, or by value after serialization."""
@ -572,9 +608,24 @@ class AnthropicMessagesHandler(BaseTranslation):
)
ordered: Final = AnthropicMessagesHandler._defer_systems_inside_tool_exchanges(structured_messages)
leading_count: Final = next(
(index for index, message in enumerate(ordered) if not _is_system(message)),
len(ordered),
)
leading_systems: Final = ordered[:leading_count]
hoisted_in_leading: Final = any(
AnthropicMessagesHandler._is_hoisted_top_level_system(message, hoisted_system_message)
for message in leading_systems
)
if leading_systems and not (leading_count == 1 and hoisted_in_leading):
AnthropicMessagesHandler._fold_leading_systems_into_top_level(
data,
leading_systems,
include_existing_system=hoisted_system_message is None,
)
run: Final[list] = [] # mutable-ok: API message payload
hoisted_dropped = False # rebind-ok: flips once the hoisted prompt is dropped
for message in ordered:
hoisted_dropped = hoisted_in_leading # rebind-ok: flips once the hoisted prompt is dropped
for message in ordered[leading_count:]:
if not _is_system(message):
run.append(message)
continue
@ -602,7 +653,7 @@ class AnthropicMessagesHandler(BaseTranslation):
@staticmethod
def _extract_midturn_system_text(
message: dict[str, Any], # mutable-ok: API message payload
message: Mapping[str, object],
msg_idx: int,
) -> ExtractedInput:
"""Match the adapter's filtering so positional guardrail write-back stays aligned."""
@ -636,7 +687,7 @@ class AnthropicMessagesHandler(BaseTranslation):
@classmethod
def _extract_input_text_and_images(
cls,
message: dict[str, Any],
message: Mapping[str, object],
msg_idx: int,
skip_system_message: bool = False,
skip_tool_message: bool = False,
@ -707,7 +758,7 @@ class AnthropicMessagesHandler(BaseTranslation):
@classmethod
def _extract_tool_result(
cls,
content_item: Mapping[str, Any],
content_item: Mapping[str, object],
msg_idx: int,
content_idx: int,
) -> ExtractedInput:
@ -736,7 +787,7 @@ class AnthropicMessagesHandler(BaseTranslation):
)
@staticmethod
def _image_sources(block: Mapping[str, Any]) -> tuple[str, ...]:
def _image_sources(block: Mapping[str, object]) -> tuple[str, ...]:
source: Final = block.get("source")
if not isinstance(source, Mapping):
return ()
@ -746,7 +797,7 @@ class AnthropicMessagesHandler(BaseTranslation):
async def _apply_guardrail_responses_to_input(
self,
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
responses: list[str],
scanned: tuple[ScannedText, ...],
) -> None:
@ -788,10 +839,10 @@ class AnthropicMessagesHandler(BaseTranslation):
self,
response: "AnthropicMessagesResponse",
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Any | None = None,
user_api_key_dict: Any | None = None,
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
) -> Any:
) -> "AnthropicMessagesResponse":
"""
Process output response by applying guardrails to text content and tool calls.
@ -869,8 +920,8 @@ class AnthropicMessagesHandler(BaseTranslation):
self,
responses_so_far: list[Any],
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Any | None = None,
user_api_key_dict: Any | None = None,
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
) -> list[Any]:
"""
@ -950,8 +1001,8 @@ class AnthropicMessagesHandler(BaseTranslation):
def _prepare_request_data(
self,
request_data: dict | None,
response: Any,
user_api_key_dict: Any | None,
response: object,
user_api_key_dict: "UserAPIKeyAuth | None",
key: str,
) -> dict:
"""Ensure request_data has the response/responses_so_far key and metadata."""
@ -968,7 +1019,7 @@ class AnthropicMessagesHandler(BaseTranslation):
return request_data
@staticmethod
def _get_response_content(response: Any) -> list[Any]:
def _get_response_content(response: object) -> list[Any]:
"""Extract content list from a dict or object response."""
if isinstance(response, dict):
return response.get("content", []) or []
@ -986,10 +1037,10 @@ class AnthropicMessagesHandler(BaseTranslation):
) -> None:
"""Extract text, images, and tool calls from content blocks."""
for content_idx, content_block in enumerate(response_content):
block_dict: dict[str, Any] = {}
block_dict: dict[str, object] = {}
if isinstance(content_block, dict):
block_type = content_block.get("type")
block_dict = cast(dict[str, Any], content_block)
block_dict = cast(dict[str, object], content_block)
elif hasattr(content_block, "type"):
block_type = getattr(content_block, "type", None)
if hasattr(content_block, "model_dump"):
@ -1017,7 +1068,7 @@ class AnthropicMessagesHandler(BaseTranslation):
texts_to_check: list[str],
images_to_check: list[str],
tool_calls_to_check: list["ChatCompletionToolCallChunk"],
response: Any,
response: object,
) -> "GenericGuardrailAPIInputs":
"""Build GenericGuardrailAPIInputs with optional images, tool calls, model."""
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
@ -1212,7 +1263,7 @@ class AnthropicMessagesHandler(BaseTranslation):
def _extract_output_text_and_images(
self,
content_block: dict[str, Any],
content_block: dict[str, object],
content_idx: int,
texts_to_check: list[str],
images_to_check: list[str],
@ -1282,7 +1333,7 @@ class AnthropicMessagesHandler(BaseTranslation):
# Handle both dict and Pydantic object content blocks
if isinstance(content_block, dict):
if content_block.get("type") == "text":
cast(dict[str, Any], content_block)["text"] = guardrail_response
cast(dict[str, object], content_block)["text"] = guardrail_response
elif hasattr(content_block, "type") and getattr(content_block, "type", None) == "text":
# Update Pydantic object's text attribute
if hasattr(content_block, "text"):

View file

@ -84,6 +84,7 @@ from litellm.types.llms.anthropic import (
AnthropicResponseContentBlockText,
AnthropicResponseContentBlockThinking,
AnthropicResponseContentBlockToolUse,
AnthropicThinkingParam,
AppliedEdit,
ContentBlockDelta,
ContentJsonBlockDelta,
@ -305,7 +306,7 @@ class LiteLLMAnthropicMessagesAdapter:
target["cache_control"] = cache_control
else:
# Fallback for non-dict objects (shouldn't happen in practice)
cast(dict[str, Any], target)["cache_control"] = cache_control
cast(dict[str, object], target)["cache_control"] = cache_control
def translatable_anthropic_params(self) -> list[str]:
"""
@ -323,7 +324,7 @@ class LiteLLMAnthropicMessagesAdapter:
"stop_sequences",
]
def _is_web_search_tool(self, tool: dict[str, Any]) -> bool:
def _is_web_search_tool(self, tool: Mapping[str, object]) -> bool:
"""
Check if a tool is an Anthropic web search tool.
@ -498,7 +499,7 @@ class LiteLLMAnthropicMessagesAdapter:
assistant_message_str = str(content)
elif isinstance(content, dict):
if content.get("type") == "text":
text_block: dict[str, Any] = {
text_block: dict[str, object] = {
"type": "text",
"text": content.get("text", ""),
}
@ -513,10 +514,12 @@ class LiteLLMAnthropicMessagesAdapter:
"name": tool_name,
"arguments": json.dumps(content.get("input", {})),
}
signature = self._extract_signature_from_tool_use_content(cast(dict[str, Any], content))
signature = self._extract_signature_from_tool_use_content(
cast(dict[str, object], content)
)
if signature:
provider_specific_fields: dict[str, Any] = (
provider_specific_fields: dict[str, object] = (
function_chunk.get("provider_specific_fields") or {}
)
provider_specific_fields["thought_signature"] = signature
@ -575,7 +578,7 @@ class LiteLLMAnthropicMessagesAdapter:
@staticmethod
def translate_anthropic_thinking_to_reasoning_effort(
thinking: dict[str, Any],
thinking: AnthropicThinkingParam,
) -> str | None:
"""
Translate Anthropic's thinking parameter to OpenAI's reasoning_effort.
@ -632,9 +635,9 @@ class LiteLLMAnthropicMessagesAdapter:
@staticmethod
def translate_thinking_for_model(
thinking: dict[str, Any],
thinking: AnthropicThinkingParam,
model: str,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Translate Anthropic thinking parameter based on the target model.
@ -670,7 +673,7 @@ class LiteLLMAnthropicMessagesAdapter:
@staticmethod
def _apply_reasoning_summary_wrapping(
reasoning_effort: str,
thinking: dict[str, Any],
thinking: Mapping[str, object],
) -> Any:
"""
Apply the reasoning_effort/summary wrapping rules shared by every
@ -731,6 +734,7 @@ class LiteLLMAnthropicMessagesAdapter:
"input_schema",
"description",
"cache_control",
"strict",
"type",
]
@ -760,6 +764,8 @@ class LiteLLMAnthropicMessagesAdapter:
function_chunk["parameters"] = tool["input_schema"]
if "description" in tool:
function_chunk["description"] = tool["description"]
if "strict" in tool:
function_chunk["strict"] = bool(tool["strict"])
for k, v in tool.items():
if k not in mapped_tool_params: # pass additional computer kwargs
@ -770,7 +776,7 @@ class LiteLLMAnthropicMessagesAdapter:
return new_tools, tool_name_mapping
def translate_anthropic_output_format_to_openai(self, output_format: Any) -> dict[str, Any] | None:
def translate_anthropic_output_format_to_openai(self, output_format: Any) -> dict[str, object] | None:
"""
Translate Anthropic's output_format to OpenAI's response_format.
@ -889,7 +895,7 @@ class LiteLLMAnthropicMessagesAdapter:
model_name: Final = anthropic_message_request.get("model", "")
for block in system_content:
if isinstance(block, dict) and block.get("type") == "text":
text_block: dict[str, Any] = {
text_block: dict[str, object] = {
"type": "text",
"text": block.get("text", ""),
}
@ -959,7 +965,7 @@ class LiteLLMAnthropicMessagesAdapter:
web_search_tools: Final[list[AllAnthropicToolsValues]] = []
regular_tools: Final[list[AllAnthropicToolsValues]] = []
for tool in tools:
cast_tool = cast(dict[str, Any], tool)
cast_tool = cast(dict[str, object], tool)
if self._is_web_search_tool(cast_tool):
web_search_tools.append(cast(AllAnthropicToolsValues, tool))
else:
@ -1007,7 +1013,7 @@ class LiteLLMAnthropicMessagesAdapter:
new_kwargs["output_config"] = effort_config # rebind-ok: out-param store like thinking above
return
reasoning_effort = self.translate_anthropic_thinking_to_reasoning_effort(cast(dict[str, Any], thinking))
reasoning_effort = self.translate_anthropic_thinking_to_reasoning_effort(cast(AnthropicThinkingParam, thinking))
if not reasoning_effort:
return
@ -1020,7 +1026,7 @@ class LiteLLMAnthropicMessagesAdapter:
reasoning_effort = output_config["effort"]
new_kwargs["reasoning_effort"] = self._apply_reasoning_summary_wrapping(
reasoning_effort, cast(dict[str, Any], thinking)
reasoning_effort, cast(dict[str, object], thinking)
)
def _translate_output_format_to_openai(
@ -1040,7 +1046,7 @@ class LiteLLMAnthropicMessagesAdapter:
``output_format`` takes precedence when both are provided.
"""
output_format: Any = anthropic_message_request.get("output_format")
output_format: object = anthropic_message_request.get("output_format")
if not output_format:
output_config: Final = anthropic_message_request.get("output_config")
if isinstance(output_config, dict):
@ -1407,7 +1413,7 @@ class LiteLLMAnthropicMessagesAdapter:
if THOUGHT_SIGNATURE_SEPARATOR in raw_id:
parts = raw_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)
thought_sig = parts[1] if len(parts) > 1 else None
tool_block: dict[str, Any] = {
tool_block: dict[str, object] = {
"type": "tool_use",
"id": normalize_anthropic_tool_use_id(raw_id),
"name": tool_name,

View file

@ -3,7 +3,7 @@
import json
import traceback
from collections import deque
from collections.abc import AsyncIterator
from collections.abc import AsyncIterator, Mapping
from typing import Any, Final
from litellm import verbose_logger
@ -68,6 +68,19 @@ class AnthropicResponsesStreamWrapper:
self._current_block_index += 1
return self._current_block_index
def _open_block(self, item_id: str | None, content_block: Mapping[str, Any]) -> int:
block_idx = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._chunk_queue.append(
{
"type": "content_block_start",
"index": block_idx,
"content_block": content_block,
}
)
return block_idx
def _process_event(self, event: Any) -> None:
"""Convert one Responses API event into zero or more Anthropic chunks queued for emission."""
event_type = getattr(event, "type", None)
@ -93,47 +106,22 @@ class AnthropicResponsesStreamWrapper:
item_id = getattr(item, "id", None) or (item.get("id") if isinstance(item, dict) else None)
if item_type == "message":
block_idx = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._chunk_queue.append(
{
"type": "content_block_start",
"index": block_idx,
"content_block": {"type": "text", "text": ""},
}
)
self._open_block(item_id, {"type": "text", "text": ""})
elif item_type == "function_call":
call_id: Final = (
getattr(item, "call_id", None) or (item.get("call_id") if isinstance(item, dict) else None) or ""
)
name = getattr(item, "name", None) or (item.get("name") if isinstance(item, dict) else None) or ""
block_idx = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._pending_tool_ids[item_id] = call_id
self._chunk_queue.append(
self._open_block(
item_id,
{
"type": "content_block_start",
"index": block_idx,
"content_block": {
"type": "tool_use",
"id": call_id,
"name": name,
"input": {},
},
}
)
elif item_type == "reasoning":
block_idx = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._chunk_queue.append(
{
"type": "content_block_start",
"index": block_idx,
"content_block": {"type": "thinking", "thinking": ""},
}
"type": "tool_use",
"id": call_id,
"name": name,
"input": {},
},
)
return
@ -146,16 +134,7 @@ class AnthropicResponsesStreamWrapper:
# Some providers (e.g. LMStudio) skip response.output_item.added,
# so no text block is open yet; synthesize content_block_start
# instead of emitting a delta with index -1
block_idx = self._next_block_index()
if item_id:
self._item_id_to_block_index[item_id] = block_idx
self._chunk_queue.append(
{
"type": "content_block_start",
"index": block_idx,
"content_block": {"type": "text", "text": ""},
}
)
block_idx = self._open_block(item_id, {"type": "text", "text": ""})
self._chunk_queue.append(
{
"type": "content_block_delta",
@ -169,11 +148,11 @@ class AnthropicResponsesStreamWrapper:
if event_type == "response.reasoning_summary_text.delta":
item_id = getattr(event, "item_id", None) or (event.get("item_id") if isinstance(event, dict) else None)
delta = getattr(event, "delta", "") or (event.get("delta", "") if isinstance(event, dict) else "")
block_idx = (
self._item_id_to_block_index.get(item_id, self._current_block_index)
if item_id
else self._current_block_index
)
block_idx = self._item_id_to_block_index.get(item_id, -1) if item_id else self._current_block_index
if block_idx < 0:
if not delta:
return
block_idx = self._open_block(item_id, {"type": "thinking", "thinking": ""})
self._chunk_queue.append(
{
"type": "content_block_delta",
@ -207,11 +186,9 @@ class AnthropicResponsesStreamWrapper:
item_id = (
getattr(item, "id", None) or (item.get("id") if isinstance(item, dict) else None) if item else None
)
block_idx = (
self._item_id_to_block_index.get(item_id, self._current_block_index)
if item_id
else self._current_block_index
)
block_idx = self._item_id_to_block_index.get(item_id, -1) if item_id else self._current_block_index
if block_idx < 0:
return
self._chunk_queue.append(
{
"type": "content_block_stop",

View file

@ -266,7 +266,13 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
if (isinstance(tool_type, str) and tool_type.startswith("web_search")) or tool_name == "web_search":
result.append({"type": "web_search_preview"})
continue
func_tool: dict[str, Any] = {"type": "function", "name": tool_name}
# Responses turns strict mode on when `strict` is omitted, silently rewriting
# `required` to every property. Anthropic tools are non-strict unless asked.
func_tool: dict[str, Any] = {
"type": "function",
"name": tool_name,
"strict": bool(tool_dict.get("strict")),
}
if "description" in tool_dict:
func_tool["description"] = tool_dict["description"]
if "input_schema" in tool_dict:

View file

@ -112,6 +112,16 @@ class AzureOpenAIConfig(BaseConfig):
"store",
]
@classmethod
def requires_max_completion_tokens(cls, model: str) -> bool:
"""Whether Azure rejects the legacy ``max_tokens`` key for this deployment.
Deliberately wider than ``AzureOpenAIGPT5Config.is_model_gpt_5_model``: the whole gpt-5
name family needs the rename, including the ``gpt-5-chat*`` models that are excluded from
the reasoning path by https://github.com/BerriAI/litellm/issues/13781.
"""
return "gpt-5" in model or "gpt5_series" in model
def _is_response_format_supported_model(self, model: str) -> bool:
"""
Determines if the model supports response_format.
@ -160,6 +170,7 @@ class AzureOpenAIConfig(BaseConfig):
api_version: str = "",
) -> dict:
supported_openai_params: Final = self.get_supported_openai_params(model)
renames_max_tokens: Final = self.requires_max_completion_tokens(model)
api_version_times: Final = api_version.split("-")
if len(api_version_times) >= 3:
@ -172,7 +183,9 @@ class AzureOpenAIConfig(BaseConfig):
api_version_day = None
for param, value in non_default_params.items():
if param == "tool_choice":
if param == "max_tokens" and renames_max_tokens:
optional_params.setdefault("max_completion_tokens", value)
elif param == "tool_choice":
"""
This parameter requires API version 2023-12-01-preview or later

View file

@ -22,10 +22,11 @@ import asyncio
import json
import time
import uuid
from collections.abc import AsyncIterator, Callable
from typing import TYPE_CHECKING, Any, Final
from collections.abc import AsyncIterator, Awaitable, Mapping
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeAlias, TypedDict
import httpx
from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
@ -33,7 +34,11 @@ from litellm.llms.azure_ai.agents.transformation import (
AzureAIAgentsConfig,
AzureAIAgentsError,
)
from litellm.types.utils import ModelResponse
from litellm.types.llms.openai import (
ChatCompletionAnnotation,
ChatCompletionAnnotationURLCitation,
)
from litellm.types.utils import ModelResponse, ModelResponseStream
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -46,6 +51,69 @@ else:
AsyncHTTPHandler = Any
class _AzureRawAnnotation(TypedDict, total=False):
type: ReadOnly[str]
text: ReadOnly[str]
start_index: ReadOnly[int]
end_index: ReadOnly[int]
url_citation: ReadOnly[ChatCompletionAnnotationURLCitation]
_TransformedAnnotation: TypeAlias = ChatCompletionAnnotation | _AzureRawAnnotation
class _AzureText(TypedDict, total=False):
value: ReadOnly[str]
annotations: ReadOnly[list[_AzureRawAnnotation]]
class _AzureContentItem(TypedDict, total=False):
type: ReadOnly[str]
text: ReadOnly[_AzureText]
class _AzureMessage(TypedDict, total=False):
role: ReadOnly[str]
content: ReadOnly[list[_AzureContentItem]]
class _AzureMessagesData(TypedDict, total=False):
data: ReadOnly[list[_AzureMessage]]
class _CreatedObject(TypedDict):
id: ReadOnly[str]
class _RunError(TypedDict, total=False):
message: ReadOnly[str]
class _RunStatus(TypedDict, total=False):
status: ReadOnly[str]
last_error: ReadOnly[_RunError]
class _SSEDelta(TypedDict, total=False):
content: ReadOnly[list[_AzureContentItem]]
class _SSEEventData(TypedDict, total=False):
id: ReadOnly[str]
content: ReadOnly[list[_AzureContentItem]]
delta: ReadOnly[_SSEDelta]
class _SyncAgentRequest(Protocol):
def __call__(self, method: str, url: str, json_data: Mapping[str, object] | None = None) -> httpx.Response: ...
class _AsyncAgentRequest(Protocol):
def __call__(
self, method: str, url: str, json_data: Mapping[str, object] | None = None
) -> Awaitable[httpx.Response]: ...
class AzureAIAgentsHandler:
"""
Handler for Azure AI Agent Service.
@ -89,7 +157,9 @@ class AzureAIAgentsHandler:
# -------------------------------------------------------------------------
# Response Helpers
# -------------------------------------------------------------------------
def _extract_content_from_messages(self, messages_data: dict) -> tuple[str, list[dict[str, Any]] | None]:
def _extract_content_from_messages(
self, messages_data: _AzureMessagesData
) -> tuple[str, list[_TransformedAnnotation] | None]:
"""Extract assistant content and annotations from the messages response.
Returns (content, annotations) where annotations is a list of
@ -108,8 +178,8 @@ class AzureAIAgentsHandler:
def _transform_annotations(
self,
raw_annotations: list[dict[str, Any]] | None,
) -> list[dict[str, Any]] | None:
raw_annotations: list[_AzureRawAnnotation] | None,
) -> list[_TransformedAnnotation] | None:
"""Transform Azure AI Foundry annotations to OpenAI-compatible format.
Azure AI returns annotations like:
@ -123,11 +193,11 @@ class AzureAIAgentsHandler:
if not raw_annotations:
return None
result: Final[list[dict[str, Any]]] = []
result: Final[list[_TransformedAnnotation]] = []
for ann in raw_annotations:
ann_type = ann.get("type")
if ann_type == "url_citation":
url_citation = dict(ann.get("url_citation", {}))
url_citation: ChatCompletionAnnotationURLCitation = {**ann.get("url_citation", {})}
# Azure puts start/end_index at annotation level; OpenAI
# expects them inside url_citation
if "start_index" in ann and "start_index" not in url_citation:
@ -147,8 +217,8 @@ class AzureAIAgentsHandler:
content: str,
model_response: ModelResponse,
thread_id: str,
messages: list[dict[str, Any]],
annotations: list[dict[str, Any]] | None = None,
messages: list[dict[str, object]],
annotations: list[_TransformedAnnotation] | None = None,
) -> ModelResponse:
"""Build the ModelResponse from agent output."""
from litellm.types.utils import Choices, Message, Usage
@ -201,7 +271,7 @@ class AzureAIAgentsHandler:
api_key: str,
optional_params: dict,
headers: dict | None,
) -> tuple:
) -> tuple[dict[str, str], str, str, str | None, str]:
"""Prepare common parameters for completion.
Azure Foundry Agents API uses Bearer token authentication:
@ -241,7 +311,7 @@ class AzureAIAgentsHandler:
def completion(
self,
model: str,
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
api_base: str,
api_key: str,
model_response: ModelResponse,
@ -266,7 +336,7 @@ class AzureAIAgentsHandler:
api_base,
) = self._prepare_completion_params(model, api_base, api_key, optional_params, headers)
def make_request(method: str, url: str, json_data: dict | None = None) -> httpx.Response:
def make_request(method: str, url: str, json_data: Mapping[str, object] | None = None) -> httpx.Response:
if method == "GET":
return client.get(url=url, headers=headers)
return client.post(
@ -290,14 +360,14 @@ class AzureAIAgentsHandler:
def _execute_agent_flow_sync(
self,
make_request: Callable,
make_request: _SyncAgentRequest,
api_base: str,
api_version: str,
agent_id: str,
thread_id: str | None,
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
optional_params: dict,
) -> tuple[str, str, list[dict[str, Any]] | None]:
) -> tuple[str, str, list[_TransformedAnnotation] | None]:
"""Execute the agent flow synchronously. Returns (thread_id, content, annotations)."""
# Step 1: Create thread if not provided
@ -305,7 +375,8 @@ class AzureAIAgentsHandler:
verbose_logger.debug("Creating thread at: %s", self._build_thread_url(api_base, api_version))
response = make_request("POST", self._build_thread_url(api_base, api_version), {})
self._check_response(response, [200, 201], "Failed to create thread")
thread_id = response.json()["id"]
thread_data: Final[_CreatedObject] = response.json()
thread_id = thread_data["id"]
verbose_logger.debug("Created thread: %s", thread_id)
# At this point thread_id is guaranteed to be a string
@ -325,7 +396,8 @@ class AzureAIAgentsHandler:
response = make_request("POST", self._build_runs_url(api_base, thread_id, api_version), run_payload)
self._check_response(response, [200, 201], "Failed to create run")
run_id: Final = response.json()["id"]
run_data: Final[_CreatedObject] = response.json()
run_id: Final = run_data["id"]
verbose_logger.debug("Created run: %s", run_id)
# Step 4: Poll for completion
@ -334,13 +406,15 @@ class AzureAIAgentsHandler:
response = make_request("GET", status_url)
self._check_response(response, [200], "Failed to get run status")
status = response.json().get("status")
status_data: _RunStatus = response.json()
status = status_data.get("status")
verbose_logger.debug("Run status: %s", status)
if status == "completed":
break
elif status in ["failed", "cancelled", "expired"]:
error_msg = response.json().get("last_error", {}).get("message", "Unknown error")
error_data: _RunStatus = response.json()
error_msg = error_data.get("last_error", {}).get("message", "Unknown error")
raise AzureAIAgentsError(status_code=500, message=f"Run {status}: {error_msg}")
time.sleep(self.config.POLL_INTERVAL_SECONDS)
@ -351,7 +425,8 @@ class AzureAIAgentsHandler:
response = make_request("GET", self._build_list_messages_url(api_base, thread_id, api_version))
self._check_response(response, [200], "Failed to get messages")
content, annotations = self._extract_content_from_messages(response.json())
messages_data: Final[_AzureMessagesData] = response.json()
content, annotations = self._extract_content_from_messages(messages_data)
return thread_id, content, annotations
# -------------------------------------------------------------------------
@ -360,7 +435,7 @@ class AzureAIAgentsHandler:
async def acompletion(
self,
model: str,
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
api_base: str,
api_key: str,
model_response: ModelResponse,
@ -389,7 +464,7 @@ class AzureAIAgentsHandler:
api_base,
) = self._prepare_completion_params(model, api_base, api_key, optional_params, headers)
async def make_request(method: str, url: str, json_data: dict | None = None) -> httpx.Response:
async def make_request(method: str, url: str, json_data: Mapping[str, object] | None = None) -> httpx.Response:
if method == "GET":
return await client.get(url=url, headers=headers)
return await client.post(
@ -413,14 +488,14 @@ class AzureAIAgentsHandler:
async def _execute_agent_flow_async(
self,
make_request: Callable,
make_request: _AsyncAgentRequest,
api_base: str,
api_version: str,
agent_id: str,
thread_id: str | None,
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
optional_params: dict,
) -> tuple[str, str, list[dict[str, Any]] | None]:
) -> tuple[str, str, list[_TransformedAnnotation] | None]:
"""Execute the agent flow asynchronously. Returns (thread_id, content, annotations)."""
# Step 1: Create thread if not provided
@ -428,7 +503,8 @@ class AzureAIAgentsHandler:
verbose_logger.debug("Creating thread at: %s", self._build_thread_url(api_base, api_version))
response = await make_request("POST", self._build_thread_url(api_base, api_version), {})
self._check_response(response, [200, 201], "Failed to create thread")
thread_id = response.json()["id"]
thread_data: Final[_CreatedObject] = response.json()
thread_id = thread_data["id"]
verbose_logger.debug("Created thread: %s", thread_id)
# At this point thread_id is guaranteed to be a string
@ -448,7 +524,8 @@ class AzureAIAgentsHandler:
response = await make_request("POST", self._build_runs_url(api_base, thread_id, api_version), run_payload)
self._check_response(response, [200, 201], "Failed to create run")
run_id: Final = response.json()["id"]
run_data: Final[_CreatedObject] = response.json()
run_id: Final = run_data["id"]
verbose_logger.debug("Created run: %s", run_id)
# Step 4: Poll for completion
@ -457,13 +534,15 @@ class AzureAIAgentsHandler:
response = await make_request("GET", status_url)
self._check_response(response, [200], "Failed to get run status")
status = response.json().get("status")
status_data: _RunStatus = response.json()
status = status_data.get("status")
verbose_logger.debug("Run status: %s", status)
if status == "completed":
break
elif status in ["failed", "cancelled", "expired"]:
error_msg = response.json().get("last_error", {}).get("message", "Unknown error")
error_data: _RunStatus = response.json()
error_msg = error_data.get("last_error", {}).get("message", "Unknown error")
raise AzureAIAgentsError(status_code=500, message=f"Run {status}: {error_msg}")
await asyncio.sleep(self.config.POLL_INTERVAL_SECONDS)
@ -474,7 +553,8 @@ class AzureAIAgentsHandler:
response = await make_request("GET", self._build_list_messages_url(api_base, thread_id, api_version))
self._check_response(response, [200], "Failed to get messages")
content, annotations = self._extract_content_from_messages(response.json())
messages_data: Final[_AzureMessagesData] = response.json()
content, annotations = self._extract_content_from_messages(messages_data)
return thread_id, content, annotations
# -------------------------------------------------------------------------
@ -483,7 +563,7 @@ class AzureAIAgentsHandler:
async def acompletion_stream(
self,
model: str,
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
api_base: str,
api_key: str,
logging_obj: LiteLLMLoggingObj,
@ -491,7 +571,7 @@ class AzureAIAgentsHandler:
litellm_params: dict,
timeout: float,
headers: dict | None = None,
) -> AsyncIterator:
) -> AsyncIterator[ModelResponseStream]:
"""Execute async streaming completion using Azure Agent Service with native SSE."""
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
@ -505,12 +585,12 @@ class AzureAIAgentsHandler:
) = self._prepare_completion_params(model, api_base, api_key, optional_params, headers)
# Build payload for create-thread-and-run with streaming
thread_messages: Final = []
thread_messages: Final[list[dict[str, object]]] = []
for msg in messages:
if msg.get("role") in ["user", "system"]:
thread_messages.append({"role": "user", "content": msg.get("content", "")})
payload: Final[dict[str, Any]] = {
payload: Final[dict[str, object]] = {
"assistant_id": agent_id,
"stream": True,
}
@ -552,14 +632,14 @@ class AzureAIAgentsHandler:
self,
response: httpx.Response,
model: str,
) -> AsyncIterator:
) -> AsyncIterator[ModelResponseStream]:
"""Process SSE stream and yield OpenAI-compatible streaming chunks."""
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
response_id: Final = f"chatcmpl-{uuid.uuid4().hex[:8]}"
created: Final = int(time.time())
thread_id = None
collected_annotations: list[dict[str, Any]] | None = None
collected_annotations: list[_TransformedAnnotation] | None = None
current_event = None
@ -597,7 +677,7 @@ class AzureAIAgentsHandler:
return
try:
data = json.loads(data_str)
data: _SSEEventData = json.loads(data_str)
except json.JSONDecodeError:
continue

View file

@ -11,6 +11,7 @@ The operation location must be polled until the analysis completes.
import asyncio
import re
import time
from collections.abc import Mapping
from typing import Any, Final
from urllib.parse import quote
@ -23,15 +24,19 @@ from litellm.constants import (
AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI,
AZURE_OPERATION_POLLING_TIMEOUT,
)
from litellm.exceptions import UnsupportedParamsError
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin, encode_url_path_segment
from litellm.llms.base_llm.ocr.transformation import (
OCR_REQUEST_FORMAT_PARAM,
BaseOCRConfig,
DocumentType,
OCRPage,
OCRPageDimensions,
OCRRequestData,
OCRRequestFormat,
OCRResponse,
OCRUsageInfo,
parse_ocr_request_format,
)
from litellm.secret_managers.main import get_secret_str
@ -97,8 +102,12 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
comma-separated string. Other Mistral-specific params (e.g.
`include_image_base64`) are not supported by Azure DI and are
ignored during transformation.
`req_format` selects the response shape: "litellm" (default) returns
the normalized OCR schema, "native" returns Azure DI's own analyze
operation payload as-is.
"""
return ["pages", "features"]
return ["pages", "features", OCR_REQUEST_FORMAT_PARAM]
def map_ocr_params(
self,
@ -117,14 +126,27 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
"""
pages: Final = non_default_params.get("pages")
features: Final = non_default_params.get("features")
request_format: Final = non_default_params.get(OCR_REQUEST_FORMAT_PARAM)
normalized_pages: Final = self._normalize_pages_param(pages) if pages is not None else ""
normalized_features: Final = self._normalize_features_param(features) if features is not None else ""
return {
**optional_params,
**({"pages": normalized_pages} if normalized_pages else {}),
**({"features": normalized_features} if normalized_features else {}),
**(
{OCR_REQUEST_FORMAT_PARAM: self._parse_request_format(request_format, model)}
if request_format is not None
else {}
),
}
@staticmethod
def _parse_request_format(request_format: object, model: str) -> OCRRequestFormat:
try:
return parse_ocr_request_format(request_format)
except ValueError as e:
raise UnsupportedParamsError(message=f"{e}", model=model, llm_provider="azure_ai") from e
@staticmethod
def _normalize_pages_param(pages: Any) -> str:
"""
@ -594,14 +616,33 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
poll_headers = {"Ocp-Apim-Subscription-Key": raw_response.request.headers.get("Ocp-Apim-Subscription-Key", "")}
return operation_url, poll_headers
def _transform_completed_response(self, model: str, raw_response: httpx.Response) -> OCRResponse:
@staticmethod
def _get_request_format(optional_params: object) -> OCRRequestFormat:
if not isinstance(optional_params, dict):
return "litellm"
request_format: Final = optional_params.get(OCR_REQUEST_FORMAT_PARAM)
if request_format is None:
return "litellm"
return parse_ocr_request_format(request_format)
def _transform_completed_response(
self,
model: str,
raw_response: httpx.Response,
request_format: OCRRequestFormat,
) -> OCRResponse:
"""
Transform a completed Azure Document Intelligence analyze operation
into the Mistral OCR response shape, preserving Azure-native
`analyzeResult` fields (`content`, `tables`, `keyValuePairs`) as
top-level response fields.
When `request_format` is "native", the untouched Azure operation
payload is attached to the response's hidden params so the proxy can
return it verbatim while cost tracking still reads `usage_info`.
"""
operation: Final = AzureDocumentIntelligenceOperation.model_validate(raw_response.json())
raw_operation: Final[Mapping[str, object]] = raw_response.json()
operation: Final = AzureDocumentIntelligenceOperation.model_validate(raw_operation)
verbose_logger.debug("Azure Document Intelligence response status: %s", operation.status)
@ -614,7 +655,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
mistral_pages: Final = [self._transform_azure_page(azure_page) for azure_page in analyze_result.pages]
usage_info: Final = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
return OCRResponse(
response: Final = OCRResponse(
pages=mistral_pages,
model=model,
usage_info=usage_info,
@ -624,6 +665,11 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
keyValuePairs=analyze_result.keyValuePairs,
)
if request_format == "native":
response.set_provider_native_response(raw_operation)
return response
def transform_ocr_response(
self,
model: str,
@ -681,8 +727,12 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
Returns:
OCRResponse in Mistral format
"""
request_format: Final = self._get_request_format(kwargs.get("optional_params"))
if raw_response.status_code != 202:
return self._transform_completed_response(model=model, raw_response=raw_response)
return self._transform_completed_response(
model=model, raw_response=raw_response, request_format=request_format
)
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation...")
operation_url, poll_headers = self._get_polling_target(raw_response)
@ -691,7 +741,9 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
headers=poll_headers,
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
)
return self._transform_completed_response(model=model, raw_response=completed_response)
return self._transform_completed_response(
model=model, raw_response=completed_response, request_format=request_format
)
async def async_transform_ocr_response(
self,
@ -714,8 +766,12 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
Returns:
OCRResponse in Mistral format
"""
request_format: Final = self._get_request_format(kwargs.get("optional_params"))
if raw_response.status_code != 202:
return self._transform_completed_response(model=model, raw_response=raw_response)
return self._transform_completed_response(
model=model, raw_response=raw_response, request_format=request_format
)
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation (async)...")
operation_url, poll_headers = self._get_polling_target(raw_response)
@ -724,4 +780,6 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
headers=poll_headers,
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
)
return self._transform_completed_response(model=model, raw_response=completed_response)
return self._transform_completed_response(
model=model, raw_response=completed_response, request_format=request_format
)

View file

@ -2,7 +2,8 @@
Base OCR transformation configuration.
"""
from typing import TYPE_CHECKING, Any
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
from pydantic import PrivateAttr
@ -21,6 +22,26 @@ else:
# File-type inputs are preprocessed to this format in litellm/ocr/main.py.
DocumentType = dict[str, str]
OCRRequestFormat = Literal["litellm", "native"]
OCR_REQUEST_FORMATS: Final[tuple[OCRRequestFormat, ...]] = ("litellm", "native")
OCR_REQUEST_FORMAT_PARAM: Final = "req_format"
OCR_REQUEST_FORMAT_HEADER: Final = "x-req-format"
PROVIDER_NATIVE_RESPONSE_KEY: Final = "provider_native_response"
def parse_ocr_request_format(value: object) -> OCRRequestFormat:
if value == "litellm":
return "litellm"
if value == "native":
return "native"
raise ValueError(
f"Invalid `{OCR_REQUEST_FORMAT_PARAM}`: {value!r}. Expected one of {', '.join(OCR_REQUEST_FORMATS)}."
)
class OCRPageDimensions(LiteLLMPydanticObjectBase):
"""Page dimensions from OCR response."""
@ -80,6 +101,15 @@ class OCRResponse(LiteLLMPydanticObjectBase):
# Define private attributes using PrivateAttr
_hidden_params: dict = PrivateAttr(default_factory=dict)
def set_provider_native_response(self, native_response: Mapping[str, object]) -> None:
"""Keep the provider's own response payload alongside the normalized one."""
self._hidden_params[PROVIDER_NATIVE_RESPONSE_KEY] = native_response
def get_provider_native_response(self) -> Mapping[str, object] | None:
"""The provider's own response payload, when `req_format=native` was requested."""
native_response: Final = self._hidden_params.get(PROVIDER_NATIVE_RESPONSE_KEY)
return native_response if isinstance(native_response, dict) else None
class OCRRequestData(LiteLLMPydanticObjectBase):
"""OCR request data structure."""

View file

@ -1,11 +1,14 @@
from datetime import datetime
from typing import Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, cast
from openai.types.batch import BatchRequestCounts
from openai.types.batch import Metadata as OpenAIBatchMetadata
from litellm.types.utils import LiteLLMBatch
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
# AWS Bedrock model-invocation-job statuses → OpenAI Batch statuses.
# Mirrors the mapping used by `BedrockBatchesConfig.transform_create_batch_response`
# so create / retrieve return consistent statuses.
@ -22,6 +25,8 @@ _BEDROCK_MIJ_STATUS_TO_OPENAI: Final = {
"Expired": "expired",
}
_CANCEL_IDEMPOTENT_STATUSES: Final = frozenset({"cancelling", "cancelled", "completed", "failed", "expired"})
def _extract_region_from_bedrock_arn(arn: str) -> str | None:
"""ARN shape: ``arn:aws:bedrock:<region>:<account>:<type>/<id>``"""
@ -82,6 +87,81 @@ class BedrockBatchesHandler:
E.g. Twelve Labs Embedding Async Invoke
"""
@staticmethod
def cancel_batch(
batch_id: str,
aws_region_name: str | None = None,
logging_obj: "LiteLLMLoggingObj | None" = None,
aws_access_key_id: str | None = None,
aws_secret_access_key: str | None = None,
aws_session_token: str | None = None,
aws_session_name: str | None = None,
aws_profile_name: str | None = None,
aws_role_name: str | None = None,
aws_web_identity_token: str | None = None,
aws_sts_endpoint: str | None = None,
aws_external_id: str | None = None,
**kwargs: object, # kwargs-ok: litellm.cancel_batch forwards arbitrary user kwargs verbatim
) -> "LiteLLMBatch":
try:
import boto3
from botocore.exceptions import ClientError
except ImportError as exc:
raise ImportError("Missing boto3/botocore to call bedrock. Run 'pip install boto3'.") from exc
region: Final = aws_region_name or _extract_region_from_bedrock_arn(batch_id) or "us-east-1"
from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig
creds: Final = BedrockBatchesConfig().get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=region,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
client: Final = boto3.client(
"bedrock",
region_name=region,
aws_access_key_id=creds.access_key,
aws_secret_access_key=creds.secret_key,
aws_session_token=creds.token,
)
def job_status() -> "LiteLLMBatch":
return BedrockBatchesHandler._handle_model_invocation_job_status(
batch_id=batch_id,
aws_region_name=region,
logging_obj=logging_obj,
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
try:
client.stop_model_invocation_job(jobIdentifier=batch_id)
except ClientError as e:
if e.response.get("Error", {}).get("Code") not in ("ValidationException", "ConflictException"):
raise
current_batch: Final = job_status()
if current_batch.status not in _CANCEL_IDEMPOTENT_STATUSES:
raise
return current_batch
return job_status()
@staticmethod
def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj=None, **kwargs) -> "LiteLLMBatch":
"""

View file

@ -6,6 +6,7 @@ import copy
import json
import time
import types
from collections.abc import Mapping
from typing import Final, Literal, cast, overload
import httpx
@ -39,6 +40,12 @@ from litellm.llms.anthropic.chat.transformation import (
AnthropicConfig,
)
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.bedrock.request_metadata import (
bedrock_request_metadata_headers,
bedrock_request_metadata_is_owned,
merge_bedrock_invoke_headers,
resolve_bedrock_request_metadata,
)
from litellm.types.llms.bedrock import *
from litellm.types.llms.openai import (
AllMessageValues,
@ -1652,6 +1659,13 @@ class AmazonConverseConfig(BaseConfig):
user_continue_message=litellm_params.pop("user_continue_message", None),
)
request_metadata: Final = resolve_bedrock_request_metadata(
litellm_params=litellm_params, caller_metadata=_data.get("requestMetadata")
)
if bedrock_request_metadata_is_owned():
_data.pop("requestMetadata", None)
if request_metadata is not None:
_data["requestMetadata"] = request_metadata
data: Final[RequestObject] = {"messages": bedrock_messages, **_data}
return data
@ -1705,6 +1719,13 @@ class AmazonConverseConfig(BaseConfig):
user_continue_message=litellm_params.pop("user_continue_message", None),
)
request_metadata: Final = resolve_bedrock_request_metadata(
litellm_params=litellm_params, caller_metadata=_data.get("requestMetadata")
)
if bedrock_request_metadata_is_owned():
_data.pop("requestMetadata", None)
if request_metadata is not None:
_data["requestMetadata"] = request_metadata
data: Final[RequestObject] = {"messages": bedrock_messages, **_data}
return data
@ -1770,7 +1791,43 @@ class AmazonConverseConfig(BaseConfig):
thinking_blocks_list.append(_redacted_block)
return thinking_blocks_list
def _transform_usage(
@staticmethod
def is_converse_usage_shape(usage_object: Mapping[str, object]) -> bool:
"""Converse-family models report camelCase token counts, not Anthropic's snake_case."""
return "inputTokens" in usage_object or "outputTokens" in usage_object
@staticmethod
def _usage_count(usage_object: Mapping[str, object], *keys: str) -> int:
for key in keys:
value = usage_object.get(key)
if isinstance(value, (int, float)) and not isinstance(value, bool):
return int(value)
return 0
def usage_from_batch_output(self, usage_object: Mapping[str, object]) -> Usage:
"""Read a Converse-shaped usage block out of a batch output line.
Batch output omits fields the live API always sends, so the block is
completed before going through the same transform, keeping a batch and an
equivalent non-batch call in agreement on tokens.
"""
input_tokens: Final = self._usage_count(usage_object, "inputTokens")
output_tokens: Final = self._usage_count(usage_object, "outputTokens")
cache_read: Final = self._usage_count(usage_object, "cacheReadInputTokens", "cacheReadInputTokenCount")
cache_write: Final = self._usage_count(usage_object, "cacheWriteInputTokens", "cacheWriteInputTokenCount")
return self.transform_usage(
ConverseTokenUsageBlock(
inputTokens=input_tokens,
outputTokens=output_tokens,
totalTokens=self._usage_count(usage_object, "totalTokens") or input_tokens + output_tokens,
cacheReadInputTokenCount=cache_read,
cacheReadInputTokens=cache_read,
cacheWriteInputTokenCount=cache_write,
cacheWriteInputTokens=cache_write,
)
)
def transform_usage(
self,
usage: ConverseTokenUsageBlock,
reasoning_content: str | None = None,
@ -2191,7 +2248,7 @@ class AmazonConverseConfig(BaseConfig):
chat_completion_message["tool_calls"] = filtered_tools
## CALCULATING USAGE - bedrock returns usage in the headers
usage: Final = self._transform_usage(
usage: Final = self.transform_usage(
completion_response["usage"],
reasoning_content=chat_completion_message.get("reasoning_content"),
)
@ -2258,7 +2315,8 @@ class AmazonConverseConfig(BaseConfig):
) -> dict:
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
return headers
owned_names, metadata_headers = bedrock_request_metadata_headers(litellm_params)
return merge_bedrock_invoke_headers(headers, (), metadata_headers, owned_names)
def should_fake_stream(
self,

View file

@ -559,7 +559,7 @@ class AWSEventStreamDecoder:
elif "stopReason" in chunk_data:
finish_reason = map_finish_reason(chunk_data.get("stopReason", "stop"))
elif "usage" in chunk_data:
usage = converse_config._transform_usage(chunk_data.get("usage", {}))
usage = converse_config.transform_usage(chunk_data.get("usage", {}))
model_response_provider_specific_fields: Final = {}
if "trace" in chunk_data:

View file

@ -13,6 +13,10 @@ import httpx
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.request_metadata import (
bedrock_request_metadata_headers,
merge_bedrock_invoke_headers,
)
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.passthrough.utils import CommonUtils
from litellm.types.llms.openai import AllMessageValues
@ -169,9 +173,12 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM):
"""
Validate the environment and return headers.
For Bedrock, we don't need Bearer token auth since we use AWS SigV4.
For Bedrock, we don't need Bearer token auth since we use AWS SigV4. This path signs the
same ``/model/{id}/invoke`` endpoint as ``AmazonInvokeConfig``, so it owns the request
metadata header on the same terms rather than letting a caller supply it.
"""
return headers
owned_names, metadata_headers = bedrock_request_metadata_headers(litellm_params)
return merge_bedrock_invoke_headers(headers, (), metadata_headers, owned_names)
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BedrockError:
"""Return the appropriate error class for Bedrock."""

View file

@ -20,6 +20,10 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.bedrock.chat.invoke_handler import make_call, make_sync_call
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.request_metadata import (
bedrock_request_metadata_headers,
merge_bedrock_invoke_headers,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -417,15 +421,13 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
api_base: str | None = None,
) -> dict:
raw_guardrail_config: Final = optional_params.pop("guardrailConfig", None)
if raw_guardrail_config is None:
return headers
existing_header_names: Final = frozenset(name.lower() for name in headers)
guardrail_headers: Final = {
name: value
for name, value in _bedrock_invoke_guardrail_headers(raw_guardrail_config).items()
if name.lower() not in existing_header_names
}
return {**headers, **guardrail_headers}
guardrail_headers: Final = (
()
if raw_guardrail_config is None
else tuple(_bedrock_invoke_guardrail_headers(raw_guardrail_config).items())
)
owned_names, metadata_headers = bedrock_request_metadata_headers(litellm_params)
return merge_bedrock_invoke_headers(headers, guardrail_headers, metadata_headers, owned_names)
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException:
return BedrockError(status_code=status_code, message=error_message)

View file

@ -3,6 +3,7 @@ import json
import os
import time
from collections.abc import Iterable, Mapping, MutableMapping, Sequence
from contextlib import suppress
from functools import cache
from itertools import chain
from types import MappingProxyType
@ -64,6 +65,10 @@ from ..common_utils import BedrockError, merge_bedrock_aws_request_params, resol
# Same pattern as the `upload_url` handoff in `transform_create_file_request`.
S3_SIGNED_GET_HEADERS_PARAM: Final = "_s3_signed_get_headers"
# litellm_params key carrying the size of the body uploaded to S3, handed from
# `transform_create_file_request` to `transform_create_file_response`.
UPLOAD_CONTENT_LENGTH_PARAM: Final = "_s3_upload_content_length"
def _frozen_mapping(items: Iterable[tuple[str, object]]) -> Mapping[str, object]:
return MappingProxyType(dict(items))
@ -145,11 +150,12 @@ class _BedrockS3RequestParams(BaseModel):
class _TrustedS3ModelCredentials(BaseModel):
"""The S3 bucket the server trusts file ids against, from the deployment snapshot."""
"""The S3 buckets the server trusts file ids against, from the deployment snapshot."""
model_config = ConfigDict(extra="ignore")
s3_bucket_name: str | None = None
s3_output_bucket_name: str | None = None
def extract_s3_uri_from_file_id(file_id: str) -> str:
@ -175,6 +181,18 @@ def extract_s3_uri_from_file_id(file_id: str) -> str:
raise ValueError("file_id must be a managed LiteLLM S3 file id")
_S3_BUCKET_REQUIRED_ERROR: Final = "S3 bucket_name is required. Set 's3_bucket_name' in proxy config or AWS_S3_BUCKET_NAME for Bedrock file content retrieval."
def _trusted_s3_model_credentials(litellm_params: Mapping[str, object]) -> _TrustedS3ModelCredentials:
trusted_model_credentials: Final = litellm_params.get("_litellm_internal_model_credentials")
if not isinstance(trusted_model_credentials, MappingProxyType):
return _TrustedS3ModelCredentials()
snapshot: Final[dict[str, object]] = {}
snapshot.update(trusted_model_credentials) # any-ok: untyped snapshot
return _TrustedS3ModelCredentials.model_validate(snapshot)
def get_configured_s3_bucket_name(litellm_params: Mapping[str, object]) -> str:
"""
Resolve the server-configured S3 bucket for Bedrock file operations.
@ -183,20 +201,62 @@ def get_configured_s3_bucket_name(litellm_params: Mapping[str, object]) -> str:
environment; never a request-supplied param, since the bucket is what
`validate_managed_cloud_file_id` checks file ids against.
"""
trusted_model_credentials: Final = litellm_params.get("_litellm_internal_model_credentials")
bucket_name: str | None = None
if isinstance(trusted_model_credentials, MappingProxyType):
snapshot: Final[dict[str, object]] = {}
snapshot.update(trusted_model_credentials) # any-ok: untyped snapshot
bucket_name = _TrustedS3ModelCredentials.model_validate(snapshot).s3_bucket_name
bucket_name = bucket_name or os.getenv("AWS_S3_BUCKET_NAME")
bucket_name: Final = _trusted_s3_model_credentials(litellm_params).s3_bucket_name or os.getenv("AWS_S3_BUCKET_NAME")
if not bucket_name:
raise ValueError(
"S3 bucket_name is required. Set 's3_bucket_name' in proxy config or AWS_S3_BUCKET_NAME for Bedrock file content retrieval."
)
raise ValueError(_S3_BUCKET_REQUIRED_ERROR)
return bucket_name
def get_configured_s3_bucket_names(litellm_params: Mapping[str, object]) -> tuple[str, ...]:
"""
Resolve the server-configured S3 buckets a Bedrock file id may live in.
Bedrock batch outputs land in ``s3_output_bucket_name`` when it differs from
the input bucket, so retrieval validates against both. Same trust rules as
``get_configured_s3_bucket_name``: only the immutable credential snapshot or
the environment, never a request param.
"""
trusted: Final = _trusted_s3_model_credentials(litellm_params)
input_bucket: Final = trusted.s3_bucket_name or os.getenv("AWS_S3_BUCKET_NAME")
output_bucket: Final = trusted.s3_output_bucket_name or os.getenv("AWS_S3_OUTPUT_BUCKET_NAME")
buckets: Final = tuple(dict.fromkeys(bucket for bucket in (input_bucket, output_bucket) if bucket))
if not buckets:
raise ValueError(_S3_BUCKET_REQUIRED_ERROR)
return buckets
def _validate_file_id_against_configured_buckets(
s3_uri: str,
configured_bucket_names: tuple[str, ...],
allow_legacy_cloud_file_ids: bool,
) -> tuple[str, str]:
def validate_against(configured_bucket_name: str) -> tuple[str, str]:
return validate_managed_cloud_file_id(
file_id=s3_uri,
scheme="s3://",
configured_bucket_name=configured_bucket_name,
allowed_object_prefixes=BEDROCK_MANAGED_S3_PREFIXES,
allow_legacy_cloud_file_ids=allow_legacy_cloud_file_ids,
)
for candidate_bucket_name in configured_bucket_names[:-1]:
with suppress(ValueError):
return validate_against(candidate_bucket_name)
return validate_against(configured_bucket_names[-1])
def _uploaded_object_size(litellm_params: Mapping[str, object], raw_response: Response) -> 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.
"""
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")
return int(response_content_length) if response_content_length.isdigit() else 0
class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
"""
Config for Bedrock Files - handles S3 uploads for Bedrock batch processing
@ -924,6 +984,8 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
)
litellm_params["upload_url"] = api_base
upload_content_length: Final = len(file_content.encode("utf-8"))
litellm_params[UPLOAD_CONTENT_LENGTH_PARAM] = upload_content_length # rebind-ok: same handoff as upload_url
# Return a dict that tells the HTTP handler exactly what to do
return {
@ -1081,12 +1143,6 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
"""
Transform S3 File upload response into OpenAI-style FileObject
"""
# For S3 uploads, we typically get an ETag and other metadata
response_headers: Final = raw_response.headers
# Extract S3 object information from the response
# S3 PUT object returns ETag and other metadata in headers
content_length: Final[str] = response_headers.get("Content-Length", "0")
# Use the actual upload URL that was used for the S3 upload
upload_url: Final = litellm_params.get("upload_url")
file_id: str = ""
@ -1101,7 +1157,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
filename=filename,
created_at=int(time.time()), # Current timestamp
status="uploaded",
bytes=int(content_length) if content_length.isdigit() else 0,
bytes=_uploaded_object_size(litellm_params=litellm_params, raw_response=raw_response),
object="file",
)
@ -1174,11 +1230,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
raise ValueError("file_id is required for Bedrock file content retrieval")
s3_uri: Final = extract_s3_uri_from_file_id(file_id)
bucket_name, object_key = validate_managed_cloud_file_id(
file_id=s3_uri,
scheme="s3://",
configured_bucket_name=get_configured_s3_bucket_name(litellm_params),
allowed_object_prefixes=BEDROCK_MANAGED_S3_PREFIXES,
bucket_name, object_key = _validate_file_id_against_configured_buckets(
s3_uri=s3_uri,
configured_bucket_names=get_configured_s3_bucket_names(litellm_params),
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(litellm_params),
)

View file

@ -1,4 +1,5 @@
from collections.abc import AsyncIterator
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, cast
import httpx
@ -37,6 +38,10 @@ from litellm.llms.bedrock.common_utils import (
normalize_tool_input_schema_types_for_bedrock_invoke,
pop_bedrock_invoke_output_config_format,
)
from litellm.llms.bedrock.request_metadata import (
bedrock_request_metadata_headers,
merge_bedrock_invoke_headers,
)
from litellm.types.llms.anthropic import (
ANTHROPIC_BETA_HEADER_VALUES,
ANTHROPIC_TOOL_SEARCH_BETA_HEADER,
@ -89,7 +94,8 @@ class AmazonAnthropicClaudeMessagesConfig(
api_key: str | None = None,
api_base: str | None = None,
) -> tuple[dict, str | None]:
return headers, api_base
owned_names, metadata_headers = bedrock_request_metadata_headers(litellm_params)
return merge_bedrock_invoke_headers(headers, (), metadata_headers, owned_names), api_base
def sign_request(
self,
@ -956,13 +962,32 @@ class AmazonAnthropicClaudeMessagesStreamDecoder(AWSEventStreamDecoder):
Bedrock returns usage metrics using camelCase keys. Convert these to
the Anthropic `/v1/messages` specification so callers receive a
consistent response shape when streaming.
Token counts already present in the chunk's own Anthropic usage block
win over the invocationMetrics-derived ones, and cache token fields
(``cache_read_input_tokens`` / ``cache_creation_input_tokens`` on
``message_stop.usage``, or ``cacheReadInputTokenCount`` /
``cacheWriteInputTokenCount`` inside the invocation metrics) are
preserved: ``invocationMetrics.inputTokenCount`` excludes cache reads
and writes, so replacing the whole usage block with input/output counts
alone drops the cache breakdown, ``_promote_message_stop_usage`` has
nothing left to promote, and cache tokens end up billed at $0.
"""
amazon_bedrock_invocation_metrics: Final = chunk_data.pop("amazon-bedrock-invocationMetrics", {})
if amazon_bedrock_invocation_metrics:
anthropic_usage: Final = {}
if "inputTokenCount" in amazon_bedrock_invocation_metrics:
anthropic_usage["input_tokens"] = amazon_bedrock_invocation_metrics["inputTokenCount"]
if "outputTokenCount" in amazon_bedrock_invocation_metrics:
anthropic_usage["output_tokens"] = amazon_bedrock_invocation_metrics["outputTokenCount"]
chunk_data["usage"] = anthropic_usage
existing_usage: Final = chunk_data.get("usage")
preserved_usage: Final = existing_usage if isinstance(existing_usage, dict) else MappingProxyType({})
metrics_usage: Final = MappingProxyType(
{
anthropic_key: amazon_bedrock_invocation_metrics[metrics_key]
for anthropic_key, metrics_key in (
("input_tokens", "inputTokenCount"),
("output_tokens", "outputTokenCount"),
("cache_read_input_tokens", "cacheReadInputTokenCount"),
("cache_creation_input_tokens", "cacheWriteInputTokenCount"),
)
if metrics_key in amazon_bedrock_invocation_metrics
}
)
chunk_data["usage"] = {**metrics_usage, **preserved_usage}
return chunk_data

View file

@ -0,0 +1,199 @@
"""
Resolve AWS Bedrock ``requestMetadata`` from LiteLLM proxy identity and caller metadata.
Bedrock attaches request metadata to CloudTrail records and to the dimension AWS Cost
Explorer groups on, so everything here is opt-in: nothing is forwarded unless the operator
sets ``litellm.bedrock_request_metadata_fields`` (``litellm_settings`` on the proxy).
Two properties are load-bearing for that billing record and are asserted by the tests:
proxy identity is resolved first so it can never be evicted by caller-supplied pairs, and the
whole ``user_api_key_`` prefix is reserved so a caller cannot write a proxy-authoritative
looking key. Values that break Bedrock's constraints are dropped rather than sanitised or
rejected, because an operator flipping this setting on must not turn a working request into a
400 and a silently rewritten attribution key is worse than an absent one.
"""
from __future__ import annotations
import json
import re
from collections.abc import Mapping
from typing import Final
import litellm
BEDROCK_REQUEST_METADATA_HEADER: Final = "X-Amzn-Bedrock-Request-Metadata"
BEDROCK_REQUEST_METADATA_MAX_PAIRS: Final = 16
BEDROCK_REQUEST_METADATA_IDENTITY_PREFIX: Final = "user_api_key_"
BEDROCK_REQUEST_METADATA_CLIENT_FIELD: Final = "spend_logs_metadata"
_METADATA_PARAM_NAMES: Final[tuple[str, ...]] = ("metadata", "litellm_metadata")
_KEY_PATTERN: Final = re.compile(r"^[a-zA-Z0-9\s:_@$#=/+,.-]{1,256}$")
_VALUE_PATTERN: Final = re.compile(r"^[a-zA-Z0-9\s:_@$#=/+,.-]{0,256}$")
_OWNED_HEADER_NAMES: Final[frozenset[str]] = frozenset((BEDROCK_REQUEST_METADATA_HEADER.lower(),))
def _is_forwardable(key: str, value: str) -> bool:
return _KEY_PATTERN.match(key) is not None and _VALUE_PATTERN.match(value) is not None
def _text_pairs(source: object) -> tuple[tuple[str, str], ...]:
if not isinstance(source, Mapping):
return ()
return tuple((key, value) for key, value in source.items() if isinstance(key, str) and isinstance(value, str))
def _allowed_fields() -> tuple[str, ...]:
"""
The operator allow-list, deduplicated so a field repeated in config cannot consume a second
reserved slot and shrink the client budget for nothing. First occurrence wins, which keeps
the operator's declared precedence intact.
"""
configured: Final[object] = litellm.bedrock_request_metadata_fields
if not isinstance(configured, (list, tuple)):
return ()
fields: Final = tuple(str(field) for field in configured)
return tuple(field for index, field in enumerate(fields) if field not in fields[:index])
def _metadata_sources(litellm_params: Mapping[str, object] | None) -> tuple[Mapping[str, object], ...]:
"""``metadata`` on /v1/chat/completions, ``litellm_metadata`` on the LITELLM_METADATA_ROUTES."""
if litellm_params is None:
return ()
return tuple(
source
for name in _METADATA_PARAM_NAMES
for source in (litellm_params.get(name),)
if isinstance(source, Mapping)
)
def _identity_pairs(
sources: tuple[Mapping[str, object], ...],
allowed_fields: tuple[str, ...],
) -> tuple[tuple[str, str], ...]:
return tuple(
(field, value)
for field in allowed_fields
if field.startswith(BEDROCK_REQUEST_METADATA_IDENTITY_PREFIX)
for value in (_first_text(sources, field),)
if value is not None and _is_forwardable(field, value)
)[:BEDROCK_REQUEST_METADATA_MAX_PAIRS]
def _first_text(sources: tuple[Mapping[str, object], ...], field: str) -> str | None:
return next((value for source in sources if isinstance(value := source.get(field), str)), None)
def _client_pairs(
sources: tuple[Mapping[str, object], ...],
allowed_fields: tuple[str, ...],
caller_metadata: object,
budget: int,
) -> tuple[tuple[str, str], ...]:
spend_logs_pairs: Final = (
tuple(pair for source in sources for pair in _text_pairs(source.get(BEDROCK_REQUEST_METADATA_CLIENT_FIELD)))
if BEDROCK_REQUEST_METADATA_CLIENT_FIELD in allowed_fields
else ()
)
candidates: Final = tuple(
(key, value)
for key, value in (*_text_pairs(caller_metadata), *spend_logs_pairs)
if not key.startswith(BEDROCK_REQUEST_METADATA_IDENTITY_PREFIX) and _is_forwardable(key, value)
)
return tuple(
pair
for index, pair in enumerate(candidates)
if pair[0] not in tuple(earlier for earlier, _ in candidates[:index])
)[:budget]
def resolve_bedrock_request_metadata(
litellm_params: Mapping[str, object] | None,
caller_metadata: object = None,
) -> dict[str, str] | None:
"""
Resolve the ``requestMetadata`` pairs to send to Bedrock, or ``None`` when the feature is
off or nothing survives Bedrock's constraints. The result is a plain dict because it is
written straight onto the Converse body, which Bedrock types as ``dict[str, str]``.
``caller_metadata`` is any ``requestMetadata`` the caller passed explicitly. It has already
been validated (and rejected with a 400) by the Converse transformation, so it is only
filtered here for the reserved identity prefix and the remaining slot budget.
"""
allowed_fields: Final = _allowed_fields()
if not allowed_fields:
return None
sources: Final = _metadata_sources(litellm_params)
identity: Final = _identity_pairs(sources, allowed_fields)
client: Final = _client_pairs(
sources=sources,
allowed_fields=allowed_fields,
caller_metadata=caller_metadata,
budget=BEDROCK_REQUEST_METADATA_MAX_PAIRS - len(identity),
)
resolved: Final = {key: value for key, value in (*identity, *client)}
return resolved or None
def bedrock_request_metadata_is_owned() -> bool:
"""
Whether the proxy OWNS the request-metadata field and header name for this request.
Ownership follows the operator's opt-in alone, never whether anything resolved, because a
caller can suppress the resolver by omitting the allow-listed fields or by sending values
that all fail Bedrock's rules. Owned-but-empty has to mean "absent on the wire" rather than
"fall back to whatever the caller supplied", or the reserved-prefix guarantee is bypassable
by anyone who can make the resolver produce nothing.
"""
return bool(_allowed_fields())
def bedrock_request_metadata_headers(
litellm_params: Mapping[str, object] | None,
) -> tuple[frozenset[str], tuple[tuple[str, str], ...]]:
"""
The signed ``X-Amzn-Bedrock-Request-Metadata`` header for the Invoke paths, which have no
body field for request metadata.
Returns the header names the proxy OWNS and, separately, the pairs to send. Ownership is
reported whenever forwarding is enabled, including when nothing resolves, because a caller
can suppress the resolver (omit the allow-listed fields, or send values that all fail
Bedrock's rules) and an owned-but-empty result must still evict the caller's header rather
than fall back to it.
"""
if not bedrock_request_metadata_is_owned():
return frozenset(), ()
resolved: Final = resolve_bedrock_request_metadata(litellm_params)
if resolved is None:
return _OWNED_HEADER_NAMES, ()
return _OWNED_HEADER_NAMES, ((BEDROCK_REQUEST_METADATA_HEADER, json.dumps(resolved, separators=(",", ":"))),)
def merge_bedrock_invoke_headers(
headers: dict[str, str],
caller_owned: tuple[tuple[str, str], ...],
proxy_owned: tuple[tuple[str, str], ...],
proxy_owned_names: frozenset[str],
) -> dict[str, str]:
"""
Merge the ``X-Amzn-*`` headers the Invoke paths derive from params.
``caller_owned`` (the guardrail headers) defers to a header the caller already set, which is
the long-standing behaviour for those. ``proxy_owned_names`` are dropped from the caller's
headers unconditionally and re-supplied only from ``proxy_owned``, because those names carry
proxy-authenticated identity into an AWS billing record that the caller must not be able to
write. Names are compared case-insensitively so a caller cannot leave a second spelling in
the dict and let the transport pick the winner.
"""
if not caller_owned and not proxy_owned and not proxy_owned_names:
return headers
existing_names: Final = frozenset(name.lower() for name in headers)
return {
name: value
for name, value in (
*((n, v) for n, v in headers.items() if n.lower() not in proxy_owned_names),
*((n, v) for n, v in caller_owned if n.lower() not in existing_names),
*proxy_owned,
)
}

View file

@ -49,6 +49,7 @@ from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
@ -1555,12 +1556,14 @@ class BaseLLMHTTPHandler:
model: str,
response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
optional_params: Mapping[str, object],
) -> OCRResponse:
"""Shared logic for transforming OCR responses."""
return provider_config.transform_ocr_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
optional_params=optional_params,
)
def ocr(
@ -1636,6 +1639,7 @@ class BaseLLMHTTPHandler:
model=model,
response=response,
logging_obj=logging_obj,
optional_params=optional_params,
)
async def async_ocr(
@ -1698,6 +1702,7 @@ class BaseLLMHTTPHandler:
model=model,
raw_response=response,
logging_obj=logging_obj,
optional_params=optional_params,
)
def search(
@ -5930,10 +5935,10 @@ class BaseLLMHTTPHandler:
self,
api_base: str,
api_key: str,
request_data: dict[str, Any],
request_data: dict[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
provider_config: Any | None = None,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
extra_headers: dict[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
@ -5963,10 +5968,10 @@ class BaseLLMHTTPHandler:
self,
api_base: str,
api_key: str,
request_data: dict[str, Any],
request_data: dict[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
provider_config: Any | None = None,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
extra_headers: dict[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
@ -5992,7 +5997,7 @@ class BaseLLMHTTPHandler:
endpoint: Literal["client_secrets", "transcription_sessions"],
api_base: str,
api_key: str,
request_data: dict[str, Any],
request_data: dict[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
provider_config: Any | None = None,
@ -11077,7 +11082,7 @@ class BaseLLMHTTPHandler:
client: HTTPHandler | AsyncHTTPHandler | None = None,
stream: bool = False,
litellm_metadata: dict[str, object] | None = None,
system_instruction: Any | None = None,
system_instruction: object | None = None,
) -> Any:
"""
Handles Google GenAI generate content requests.
@ -11208,7 +11213,7 @@ class BaseLLMHTTPHandler:
client: AsyncHTTPHandler | None = None,
stream: bool = False,
litellm_metadata: dict[str, object] | None = None,
system_instruction: Any | None = None,
system_instruction: object | None = None,
) -> Any:
"""
Async version of the generate content handler.

View file

@ -19424,20 +19424,20 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.6-flash": {
"cache_read_input_token_cost": 1.5e-07,
"cache_read_input_token_cost_flex": 7.5e-08,
"input_cost_per_token": 1.5e-06,
"input_cost_per_token_batches": 7.5e-07,
"input_cost_per_token_flex": 7.5e-07,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_batches": 3.75e-07,
"input_cost_per_token_flex": 3.75e-07,
"litellm_provider": "vertex_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_reasoning_token": 7.5e-06,
"output_cost_per_token": 7.5e-06,
"output_cost_per_token_batches": 3.75e-06,
"output_cost_per_token_flex": 3.75e-06,
"output_cost_per_reasoning_token": 3.75e-06,
"output_cost_per_token": 3.75e-06,
"output_cost_per_token_batches": 1.875e-06,
"output_cost_per_token_flex": 1.875e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -19467,9 +19467,9 @@
"supports_vision": true,
"supports_web_search": true,
"supports_native_streaming": true,
"input_cost_per_token_priority": 2.7e-06,
"output_cost_per_token_priority": 1.35e-05,
"cache_read_input_token_cost_priority": 2.7e-07,
"input_cost_per_token_priority": 1.35e-06,
"output_cost_per_token_priority": 6.75e-06,
"cache_read_input_token_cost_priority": 1.35e-07,
"search_context_cost_per_query": {
"search_context_size_low": 0.014,
"search_context_size_medium": 0.014,
@ -21150,20 +21150,20 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.6-flash": {
"cache_read_input_token_cost": 1.5e-07,
"cache_read_input_token_cost_flex": 7.5e-08,
"input_cost_per_token": 1.5e-06,
"input_cost_per_token_batches": 7.5e-07,
"input_cost_per_token_flex": 7.5e-07,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_batches": 3.75e-07,
"input_cost_per_token_flex": 3.75e-07,
"litellm_provider": "gemini",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_reasoning_token": 7.5e-06,
"output_cost_per_token": 7.5e-06,
"output_cost_per_token_batches": 3.75e-06,
"output_cost_per_token_flex": 3.75e-06,
"output_cost_per_reasoning_token": 3.75e-06,
"output_cost_per_token": 3.75e-06,
"output_cost_per_token_batches": 1.875e-06,
"output_cost_per_token_flex": 1.875e-06,
"rpm": 2000,
"source": "https://ai.google.dev/pricing/gemini-3",
"supported_endpoints": [
@ -21196,9 +21196,9 @@
"supports_web_search": true,
"supports_native_streaming": true,
"tpm": 800000,
"input_cost_per_token_priority": 2.7e-06,
"output_cost_per_token_priority": 1.35e-05,
"cache_read_input_token_cost_priority": 2.7e-07,
"input_cost_per_token_priority": 1.35e-06,
"output_cost_per_token_priority": 6.75e-06,
"cache_read_input_token_cost_priority": 1.35e-07,
"search_context_cost_per_query": {
"search_context_size_low": 0.014,
"search_context_size_medium": 0.014,
@ -21544,20 +21544,20 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.6-flash": {
"cache_read_input_token_cost": 1.5e-07,
"cache_read_input_token_cost_flex": 7.5e-08,
"input_cost_per_token": 1.5e-06,
"input_cost_per_token_batches": 7.5e-07,
"input_cost_per_token_flex": 7.5e-07,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_batches": 3.75e-07,
"input_cost_per_token_flex": 3.75e-07,
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_reasoning_token": 7.5e-06,
"output_cost_per_token": 7.5e-06,
"output_cost_per_token_batches": 3.75e-06,
"output_cost_per_token_flex": 3.75e-06,
"output_cost_per_reasoning_token": 3.75e-06,
"output_cost_per_token": 3.75e-06,
"output_cost_per_token_batches": 1.875e-06,
"output_cost_per_token_flex": 1.875e-06,
"source": "https://ai.google.dev/pricing/gemini-3",
"supported_endpoints": [
"/v1/chat/completions",
@ -21588,9 +21588,9 @@
"supports_vision": true,
"supports_web_search": true,
"supports_native_streaming": true,
"input_cost_per_token_priority": 2.7e-06,
"output_cost_per_token_priority": 1.35e-05,
"cache_read_input_token_cost_priority": 2.7e-07,
"input_cost_per_token_priority": 1.35e-06,
"output_cost_per_token_priority": 6.75e-06,
"cache_read_input_token_cost_priority": 1.35e-07,
"search_context_cost_per_query": {
"search_context_size_low": 0.014,
"search_context_size_medium": 0.014,

View file

@ -21,7 +21,12 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.llms.azure_ai.ocr.common_utils import (
is_azure_document_intelligence_model,
)
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
from litellm.llms.base_llm.ocr.transformation import (
OCR_REQUEST_FORMAT_PARAM,
BaseOCRConfig,
OCRResponse,
parse_ocr_request_format,
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.types.router import GenericLiteLLMParams
@ -124,6 +129,24 @@ def _prepare_ocr_request(
litellm_params: Final = GenericLiteLLMParams.model_validate(kwargs)
supported_params: Final = ocr_provider_config.get_supported_ocr_params(model=model)
requested_format: Final = kwargs.get(OCR_REQUEST_FORMAT_PARAM)
if requested_format is not None:
try:
parsed_format: Final = parse_ocr_request_format(requested_format)
except ValueError as e:
raise litellm.exceptions.UnsupportedParamsError(
message=f"{e}", model=model, llm_provider=custom_llm_provider
) from e
if OCR_REQUEST_FORMAT_PARAM not in supported_params and parsed_format == "native":
raise litellm.exceptions.UnsupportedParamsError(
message=(
f"`{OCR_REQUEST_FORMAT_PARAM}='native'` is not supported for provider: {custom_llm_provider}, "
f"model: {model}"
),
model=model,
llm_provider=custom_llm_provider,
)
non_default_params: Final = {}
for param in supported_params:
if param in kwargs:
@ -166,6 +189,8 @@ def _prepare_ocr_request(
def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
return False
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS

View file

@ -4,7 +4,7 @@ import hashlib
import json
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, TypeVar, cast
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -45,10 +45,47 @@ from litellm.types.mcp import MCPCredentials
if TYPE_CHECKING:
from prisma import models as prisma_db_models
from prisma import types as prisma_db_types
from prisma.actions import LiteLLM_MCPUserCredentialsActions, LiteLLM_MCPUserEnvVarsActions
from litellm.types.mcp_server.mcp_server_manager import MCPServer
_RowT = TypeVar("_RowT")
class _TableActions(Protocol[_RowT]):
async def find_unique(
self, where: Mapping[str, object], include: Mapping[str, object] | None = None
) -> _RowT | None: ...
async def find_many(
self,
take: int | None = None,
where: Mapping[str, object] | None = None,
order: Mapping[str, object] | None = None,
) -> list[_RowT]: ...
async def create(self, data: Mapping[str, object]) -> _RowT: ...
async def upsert(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT: ...
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT | None: ...
async def delete(self, where: Mapping[str, object]) -> _RowT | None: ...
async def delete_many(self, where: Mapping[str, object] | None = None) -> int: ...
class _UserEnvVarsTransactionClient(Protocol):
litellm_mcpuserenvvars: "_TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]"
async def execute_raw(self, query: str, *args: object) -> int: ...
class _UserEnvVarsTransaction(Protocol):
async def __aenter__(self) -> _UserEnvVarsTransactionClient: ...
async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ...
_AUTH_FLOW_SCOPED_FIELDS: Final["frozenset[str]"] = frozenset(
{
"issuer",
@ -434,23 +471,54 @@ def _credentials_blob_to_mutable_dict(blob: str | Mapping[str, object]) -> dict[
return parsed_blob
def _mcp_server_table_actions(
prisma_client: PrismaClient,
) -> "_TableActions[prisma_db_models.LiteLLM_MCPServerTable]":
table: Final[_TableActions[prisma_db_models.LiteLLM_MCPServerTable]] = MCPServerRepository(prisma_client).table
return table
def _verification_token_table_actions(
prisma_client: PrismaClient,
) -> "_TableActions[prisma_db_models.LiteLLM_VerificationToken]":
table: Final[_TableActions[prisma_db_models.LiteLLM_VerificationToken]] = VerificationTokenRepository(
prisma_client
).table
return table
def _team_table_actions(
prisma_client: PrismaClient,
) -> "_TableActions[prisma_db_models.LiteLLM_TeamTable]":
table: Final[_TableActions[prisma_db_models.LiteLLM_TeamTable]] = TeamRepository(prisma_client).table
return table
def _oauth_client_table_actions(
prisma_client: PrismaClient,
) -> "_TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]":
table: Final[_TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = MCPServerOAuthClientRepository(
prisma_client
).table
return table
def _db_transaction_manager(prisma_client: PrismaClient) -> _UserEnvVarsTransaction:
manager: Final[_UserEnvVarsTransaction] = prisma_client.db.tx()
return manager
async def _db_find_mcp_server_rows(
prisma_client: PrismaClient,
where: "prisma_db_types.LiteLLM_MCPServerTableWhereInput | None" = None,
) -> "list[prisma_db_models.LiteLLM_MCPServerTable]":
rows: list[prisma_db_models.LiteLLM_MCPServerTable] = await MCPServerRepository(prisma_client).table.find_many(
where=where
)
return rows
return await _mcp_server_table_actions(prisma_client).find_many(where=where)
async def _db_find_mcp_server_row(
prisma_client: PrismaClient, server_id: str
) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
row: prisma_db_models.LiteLLM_MCPServerTable | None = await MCPServerRepository(prisma_client).table.find_unique(
where={"server_id": server_id}
)
return row
return await _mcp_server_table_actions(prisma_client).find_unique(where={"server_id": server_id})
async def _db_update_mcp_server_row(
@ -467,19 +535,17 @@ async def _db_update_mcp_server_row(
def _user_credential_actions(
prisma_client: PrismaClient,
) -> "LiteLLM_MCPUserCredentialsActions[prisma_db_models.LiteLLM_MCPUserCredentials]":
table: Final[LiteLLM_MCPUserCredentialsActions[prisma_db_models.LiteLLM_MCPUserCredentials]] = (
MCPUserCredentialsRepository(prisma_client).table
)
) -> "_TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]":
table: Final[_TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]] = MCPUserCredentialsRepository(
prisma_client
).table
return table
def _user_env_var_actions(
prisma_client: PrismaClient,
) -> "LiteLLM_MCPUserEnvVarsActions[prisma_db_models.LiteLLM_MCPUserEnvVars]":
table: Final[LiteLLM_MCPUserEnvVarsActions[prisma_db_models.LiteLLM_MCPUserEnvVars]] = (
prisma_client.db.litellm_mcpuserenvvars
)
) -> "_TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]":
table: Final[_TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]] = prisma_client.db.litellm_mcpuserenvvars
return table
@ -501,7 +567,7 @@ async def _db_find_user_credential_rows(
async def _db_upsert_user_credential_row(
prisma_client: PrismaClient, user_id: str, server_id: str, credential_b64: str
) -> None:
await MCPUserCredentialsRepository(prisma_client).table.upsert(
await _user_credential_actions(prisma_client).upsert(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}},
data={
"create": {
@ -592,9 +658,9 @@ async def get_mcp_servers(prisma_client: PrismaClient, server_ids: Iterable[str]
"""
Returns the matching mcp servers from the db with the server_ids
"""
_mcp_servers: Final[list[prisma_db_models.LiteLLM_MCPServerTable]] = await MCPServerRepository(
_mcp_servers: Final[list[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions(
prisma_client
).table.find_many(
).find_many(
where={
"server_id": {"in": server_ids},
}
@ -612,9 +678,9 @@ async def get_mcp_servers_by_verificationtoken(prisma_client: PrismaClient, toke
"""
Returns the mcp servers from the db for the verification token
"""
verification_token_record: prisma_db_models.LiteLLM_VerificationToken | None = await VerificationTokenRepository(
prisma_client
).table.find_unique(
verification_token_record: (
prisma_db_models.LiteLLM_VerificationToken | None
) = await _verification_token_table_actions(prisma_client).find_unique(
where={
"token": token,
},
@ -633,7 +699,7 @@ async def get_mcp_servers_by_team(prisma_client: PrismaClient, team_id: str) ->
"""
Returns the mcp servers from the db for the team id
"""
team_record: prisma_db_models.LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique(
team_record: prisma_db_models.LiteLLM_TeamTable | None = await _team_table_actions(prisma_client).find_unique(
where={
"team_id": team_id,
},
@ -760,9 +826,9 @@ async def delete_mcp_server(
if deleted_server is not None:
credential_user_ids: list[str] = []
try:
credential_rows: Sequence[
prisma_db_models.LiteLLM_MCPUserCredentials
] = await prisma_client.db.litellm_mcpusercredentials.find_many(where={"server_id": server_id})
credential_rows: Sequence[prisma_db_models.LiteLLM_MCPUserCredentials] = await _user_credential_actions(
prisma_client
).find_many(where={"server_id": server_id})
credential_user_ids = [row.user_id for row in credential_rows]
except Exception as e: # noqa: BLE001 - enumeration is best-effort; cached tokens expire by TTL
verbose_proxy_logger.warning(
@ -771,9 +837,9 @@ async def delete_mcp_server(
e,
)
for model, label in (
(prisma_client.db.litellm_mcpusercredentials, "credential"),
(prisma_client.db.litellm_mcpuserenvvars, "env var"),
(prisma_client.db.litellm_mcpserveroauthclient, "OAuth client"),
(_user_credential_actions(prisma_client), "credential"),
(_user_env_var_actions(prisma_client), "env var"),
(_oauth_client_table_actions(prisma_client), "OAuth client"),
):
try:
await model.delete_many(where={"server_id": server_id})
@ -1042,9 +1108,9 @@ async def get_mcp_server_oauth_client_credentials(prisma_client: PrismaClient, s
LiteLLM_MCPServerTable row, so their dynamically registered client lives here keyed
by server_id. The returned value is the raw credentials blob for
``_get_persisted_dcr_credentials`` to parse."""
row: Final[prisma_db_models.LiteLLM_MCPServerOAuthClient | None] = await MCPServerOAuthClientRepository(
row: Final[prisma_db_models.LiteLLM_MCPServerOAuthClient | None] = await _oauth_client_table_actions(
prisma_client
).table.find_unique(where={"server_id": server_id})
).find_unique(where={"server_id": server_id})
if row is None:
return None
return row.credentials
@ -1062,7 +1128,7 @@ async def upsert_mcp_server_oauth_client_credentials(
encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=_get_salt_key())
blob: Final = safe_dumps(encrypted)
await MCPServerOAuthClientRepository(prisma_client).table.upsert(
await _oauth_client_table_actions(prisma_client).upsert(
where={"server_id": server_id},
data={
"create": {"server_id": server_id, "credentials": blob},
@ -1109,21 +1175,21 @@ async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient,
continue
update_data["updated_by"] = touched_by
await MCPServerRepository(prisma_client).table.update(
await _mcp_server_table_actions(prisma_client).update(
where={"server_id": mcp_server.server_id},
data=update_data,
)
updated += 1
oauth_clients: Final[list[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = await MCPServerOAuthClientRepository(
oauth_clients: Final[list[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = await _oauth_client_table_actions(
prisma_client
).table.find_many()
).find_many()
oauth_updated = 0
for oauth_client in oauth_clients:
rotated_credentials = _reencrypt_mcp_credentials_blob(oauth_client.credentials, new_master_key)
if rotated_credentials is None:
continue
await MCPServerOAuthClientRepository(prisma_client).table.update(
await _oauth_client_table_actions(prisma_client).update(
where={"server_id": oauth_client.server_id},
data={"credentials": rotated_credentials},
)
@ -1813,7 +1879,9 @@ async def get_mcp_submissions(
along with a summary count breakdown by approval_status.
Mirrors get_guardrail_submissions() from guardrail_endpoints.py.
"""
rows: list[prisma_db_models.LiteLLM_MCPServerTable] = await MCPServerRepository(prisma_client).table.find_many(
rows: Final[list[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions(
prisma_client
).find_many(
where={"submitted_at": {"not": None}},
order={"submitted_at": "desc"},
take=500, # safety cap; paginate if needed in a future iteration
@ -1915,7 +1983,7 @@ async def merge_user_env_vars(
"big",
signed=True,
)
async with prisma_client.db.tx() as tx:
async with _db_transaction_manager(prisma_client) as tx:
await tx.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key)
row: Final[prisma_db_models.LiteLLM_MCPUserEnvVars | None] = await tx.litellm_mcpuserenvvars.find_unique(
where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}

View file

@ -2397,6 +2397,7 @@ def _build_oauth_authorization_server_response(
request_base_url: Final = get_request_base_url(request)
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
explicitly_named: Final = mcp_server_name is not None
# When no server name provided, try to resolve the single OAuth2 server
if mcp_server_name is None:
@ -2415,8 +2416,10 @@ def _build_oauth_authorization_server_response(
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth authorization server")
issuer: Final = f"{request_base_url}/{mcp_server_name}" if explicitly_named else request_base_url
return {
"issuer": request_base_url, # point to your proxy
"issuer": issuer,
"authorization_endpoint": authorization_endpoint,
"token_endpoint": token_endpoint,
"response_types_supported": ["code"],

View file

@ -13,7 +13,7 @@ import json
import os
import re
import time
from collections.abc import AsyncIterator, Callable, Sequence
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast
from urllib.parse import ParseResult, urlparse
@ -46,6 +46,9 @@ from litellm.constants import (
)
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth
from litellm.integrations.custom_guardrail import (
_sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic
)
from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
@ -162,6 +165,7 @@ if TYPE_CHECKING:
from mcp.types import CreateMessageRequestParams
from litellm.caching.caching import InMemoryCache
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.mcp_server.mcp_toolset import MCPToolset
try:
@ -1233,6 +1237,35 @@ def _create_elicitation_callback():
return _elicitation_callback
def _record_mcp_guardrail_evaluations(
synthetic_llm_data: dict[str, Any], # mutable-ok: `_sync_guardrail_info_to_logging_obj` takes a concrete dict
litellm_logging_obj: "LiteLLMLoggingObj | None",
) -> None:
"""Bridge guardrail decision records off an MCP synthetic request onto the request's logger.
MCP guardrails run against a throwaway LLM-shaped dict from
``ProxyLogging._convert_mcp_to_llm_format``, so ``@log_guardrail_information``
files ``standard_logging_guardrail_information`` in that dict's metadata bucket,
which ``get_standard_logging_object_payload`` never reads. Native (non-unified)
guardrails receive no ``logging_obj`` kwarg, so the decorator cannot bridge on
their behalf; this calls the same helper it would have.
Only the decision records move. The synthetic request's messages and tool
arguments stay behind: they can carry end-user data, and the monitor needs none
of it.
"""
if litellm_logging_obj is None:
return
try:
_sync_guardrail_info_to_logging_obj(synthetic_llm_data, litellm_logging_obj)
except Exception as e: # noqa: BLE001 # callers run this from a `finally` on the block path
# The breadth is the point. Narrowing to the knowable AttributeError/TypeError
# would let an unexpected type escape that ``finally`` and replace the guardrail's
# block with a bookkeeping error.
verbose_logger.warning("Failed to record MCP guardrail evaluation for logging: %s", e)
class MCPServerManager:
_STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$")
@ -4573,6 +4606,7 @@ class MCPServerManager:
proxy_logging_obj: ProxyLogging | None,
server: MCPServer,
raw_headers: dict[str, str] | None = None,
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
) -> dict[str, Any]:
"""
Run pre-call checks and guardrail hooks for an MCP tool call.
@ -4582,6 +4616,10 @@ class MCPServerManager:
present. An absent logger must never be able to turn an authorization
decision into a no-op.
``litellm_logging_obj`` is the request's logger, and it is what lands a
``pre_mcp_call`` evaluation (or a block) on the spend-log row the Guardrails
Monitor counts. It stays optional so callers that do no logging are unchanged.
Returns a dict that may contain:
- "arguments": hook-modified tool arguments (only if changed)
- "extra_headers": headers injected by pre_mcp_call guardrail hooks
@ -4640,8 +4678,13 @@ class MCPServerManager:
# Create MCP request object for processing
mcp_request_obj: Final = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs)
# Convert to LLM format for existing guardrail compatibility
# Convert to LLM format for existing guardrail compatibility.
# Unified guardrails read the seeded logger off the request dict and pass it
# into ``apply_guardrail``, so ``@log_guardrail_information`` bridges their
# evaluations itself; the ``finally`` below covers native guardrails, which
# never receive it. Same seeding the pass-through routes do.
synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs)
synthetic_llm_data["litellm_logging_obj"] = litellm_logging_obj
try:
# Use standard pre_call_hook
@ -4666,6 +4709,12 @@ class MCPServerManager:
# Re-raise guardrail exceptions to properly fail the MCP call
verbose_logger.error("Guardrail blocked MCP tool call pre call: %s", e)
raise e
finally:
# ``finally`` rather than after the ``try``: a block raises straight out of
# here, and the failure spend-log row that "Total Blocked" counts is built
# from this logger further up the stack, so the record has to be attached
# before the exception leaves this frame.
_record_mcp_guardrail_evaluations(synthetic_llm_data, litellm_logging_obj)
return hook_result
@ -4677,8 +4726,14 @@ class MCPServerManager:
user_api_key_auth: UserAPIKeyAuth | None,
proxy_logging_obj: ProxyLogging,
start_time: datetime.datetime,
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
):
"""Create and return a during hook task for MCP tool calls."""
"""Create and return a during hook task for MCP tool calls.
``litellm_logging_obj`` is the request's logger; see ``pre_call_tool_check``.
The task is awaited before the tool call's success logging runs, so a
``during_mcp_call`` evaluation recorded on it is serialized with that call.
"""
from litellm.types.llms.base import HiddenParams
from litellm.types.mcp import MCPDuringCallRequestObject
@ -4697,15 +4752,23 @@ class MCPServerManager:
"user_api_key_auth": user_api_key_auth,
}
# Seeded for the same reason as in ``pre_call_tool_check``.
synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs)
synthetic_llm_data["litellm_logging_obj"] = litellm_logging_obj
return asyncio.create_task(
proxy_logging_obj.during_call_hook(
user_api_key_dict=user_api_key_auth,
data=synthetic_llm_data,
call_type=CallTypes.call_mcp_tool.value,
)
)
# Wrapped so the bridge runs inside the task: the caller only holds the task and
# gathers it later, so there is no other point that still sees a block here.
async def _run_during_call_hook() -> Mapping[str, Any] | None:
try:
return await proxy_logging_obj.during_call_hook(
user_api_key_dict=user_api_key_auth,
data=synthetic_llm_data,
call_type=CallTypes.call_mcp_tool.value,
)
finally:
_record_mcp_guardrail_evaluations(synthetic_llm_data, litellm_logging_obj)
return asyncio.create_task(_run_during_call_hook())
def _get_call_semaphore(self, mcp_server: MCPServer) -> asyncio.Semaphore | None:
limit: Final = mcp_server.max_concurrent_requests
@ -5234,6 +5297,7 @@ class MCPServerManager:
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
host_progress_callback: Callable | None = None,
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
) -> CallToolResult:
"""
Call a tool with the given name and arguments
@ -5246,6 +5310,9 @@ class MCPServerManager:
mcp_auth_header: MCP auth header (deprecated)
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
proxy_logging_obj: Optional ProxyLogging object for hook integration
litellm_logging_obj: Optional request logger the guardrail hooks record
their evaluations onto, so MCP guardrail activity reaches the
Guardrails Monitor. See ``pre_call_tool_check``
Returns:
@ -5276,6 +5343,7 @@ class MCPServerManager:
proxy_logging_obj=proxy_logging_obj,
server=mcp_server,
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
)
if "arguments" in hook_result:
arguments = hook_result["arguments"]
@ -5290,6 +5358,7 @@ class MCPServerManager:
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
start_time=start_time,
litellm_logging_obj=litellm_logging_obj,
)
tasks.append(during_hook_task)

View file

@ -2824,6 +2824,7 @@ if MCP_AVAILABLE:
proxy_logging_obj=proxy_logging_obj,
server=mcp_server,
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
)
# `pre_call_tool_check` may return guardrail-modified
# arguments; honor them on the local path too.
@ -2962,6 +2963,7 @@ if MCP_AVAILABLE:
proxy_logging_obj=proxy_logging_obj,
server=prefix_server,
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
)
if "arguments" in hook_result:
arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args
@ -3149,6 +3151,20 @@ if MCP_AVAILABLE:
traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG)
from litellm.proxy.proxy_server import proxy_logging_obj
# Ordering is load-bearing. ``_ProxyDBLogger.async_post_call_failure_hook``,
# reached below, writes the failure spend-log row from this logger's
# ``standard_logging_object``, which only exists once the failure handlers
# have run. Flush them first or the row lands with
# ``guardrail_information=None`` and a guardrail block is never counted.
#
# Not double-logged: both handlers gate on ``should_run_logging`` and then
# mark it, so the ``@client`` wrapper's own post-raise logging no-ops on this
# logger, same as ``_fire_mcp_tool_call_logging`` does for ``isError=True``.
if litellm_logging_obj is not None:
end_time: Final = datetime.now() # noqa: DTZ005 # naive to match `start_time`, which it is subtracted from
litellm_logging_obj.failure_handler(e, traceback_str, start_time, end_time)
await litellm_logging_obj.async_failure_handler(e, traceback_str, start_time, end_time)
if proxy_logging_obj and user_api_key_auth:
await proxy_logging_obj.post_call_failure_hook(
request_data=kwargs,
@ -3326,6 +3342,7 @@ if MCP_AVAILABLE:
raw_headers=raw_headers,
proxy_logging_obj=proxy_logging_obj,
host_progress_callback=host_progress_callback,
litellm_logging_obj=litellm_logging_obj,
)
verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result)
return call_tool_result

View file

@ -11349,6 +11349,40 @@
"title": "UpdateGuardrailRequest",
"type": "object"
},
"UsageChartPoint": {
"properties": {
"blocked": {
"title": "Blocked",
"type": "integer"
},
"date": {
"title": "Date",
"type": "string"
},
"passed": {
"title": "Passed",
"type": "integer"
},
"score": {
"anyOf": [
{
"type": "number"
},
{
"type": "null"
}
],
"title": "Score"
}
},
"required": [
"date",
"passed",
"blocked"
],
"title": "UsageChartPoint",
"type": "object"
},
"UsageDetailResponse": {
"properties": {
"avgLatency": {
@ -11410,8 +11444,7 @@
},
"time_series": {
"items": {
"additionalProperties": true,
"type": "object"
"$ref": "#/components/schemas/UsageChartPoint"
},
"title": "Time Series",
"type": "array"
@ -11423,6 +11456,40 @@
"type": {
"title": "Type",
"type": "string"
},
"usage_units": {
"additionalProperties": {
"type": "integer"
},
"title": "Usage Units",
"type": "object"
},
"usage_units_by_key": {
"additionalProperties": {
"additionalProperties": {
"type": "integer"
},
"type": "object"
},
"title": "Usage Units By Key",
"type": "object"
},
"usage_units_by_team": {
"additionalProperties": {
"additionalProperties": {
"type": "integer"
},
"type": "object"
},
"title": "Usage Units By Team",
"type": "object"
},
"usage_units_daily": {
"items": {
"$ref": "#/components/schemas/UsageUnitsDailyPoint"
},
"title": "Usage Units Daily",
"type": "array"
}
},
"required": [
@ -11437,7 +11504,11 @@
"status",
"trend",
"description",
"time_series"
"time_series",
"usage_units",
"usage_units_daily",
"usage_units_by_team",
"usage_units_by_key"
],
"title": "UsageDetailResponse",
"type": "object"
@ -11572,8 +11643,7 @@
"properties": {
"chart": {
"items": {
"additionalProperties": true,
"type": "object"
"$ref": "#/components/schemas/UsageChartPoint"
},
"title": "Chart",
"type": "array"
@ -11596,6 +11666,13 @@
"totalRequests": {
"title": "Totalrequests",
"type": "integer"
},
"totalUsageUnits": {
"additionalProperties": {
"type": "integer"
},
"title": "Totalusageunits",
"type": "object"
}
},
"required": [
@ -11603,7 +11680,8 @@
"chart",
"totalRequests",
"totalBlocked",
"passRate"
"passRate",
"totalUsageUnits"
],
"title": "UsageOverviewResponse",
"type": "object"
@ -11663,6 +11741,13 @@
"type": {
"title": "Type",
"type": "string"
},
"usageUnits": {
"additionalProperties": {
"type": "integer"
},
"title": "Usageunits",
"type": "object"
}
},
"required": [
@ -11675,11 +11760,33 @@
"avgScore",
"avgLatency",
"status",
"trend"
"trend",
"usageUnits"
],
"title": "UsageOverviewRow",
"type": "object"
},
"UsageUnitsDailyPoint": {
"properties": {
"date": {
"title": "Date",
"type": "string"
},
"units": {
"additionalProperties": {
"type": "integer"
},
"title": "Units",
"type": "object"
}
},
"required": [
"date",
"units"
],
"title": "UsageUnitsDailyPoint",
"type": "object"
},
"ValidationError": {
"properties": {
"loc": {
@ -21477,6 +21584,13 @@
"totalRequests": {
"title": "Totalrequests",
"type": "integer"
},
"totalUsageUnits": {
"additionalProperties": {
"type": "integer"
},
"title": "Totalusageunits",
"type": "object"
}
},
"required": [
@ -21484,7 +21598,8 @@
"chart",
"totalRequests",
"totalBlocked",
"passRate"
"passRate",
"totalUsageUnits"
],
"title": "UsageOverviewResponse",
"type": "object"
@ -21544,6 +21659,13 @@
"type": {
"title": "Type",
"type": "string"
},
"usageUnits": {
"additionalProperties": {
"type": "integer"
},
"title": "Usageunits",
"type": "object"
}
},
"required": [
@ -21556,7 +21678,8 @@
"avgScore",
"avgLatency",
"status",
"trend"
"trend",
"usageUnits"
],
"title": "UsageOverviewRow",
"type": "object"

View file

@ -451,6 +451,7 @@ class LiteLLMRoutes(enum.Enum):
mapped_pass_through_routes = [
"/bedrock",
"/comprehendmedical",
"/vertex-ai",
"/vertex_ai",
"/cohere",
@ -680,6 +681,7 @@ class LiteLLMRoutes(enum.Enum):
# permitted teams exactly like /spend/logs/ui — it belongs to the same
# access tier, not to customer management.
"/management/v1/spend_logs/end_users",
"/management/v1/spend_logs/users",
"/cost/estimate",
]
@ -872,12 +874,13 @@ class LiteLLMRoutes(enum.Enum):
# PROXY_ADMIN_VIEW_ONLY — the route gate must match).
"/customer/list",
"/customer/info",
# UI Logs page detail drawer (single + session) and the end-user filter
# facet. The list endpoint `/spend/logs/ui` is covered via
# UI Logs page detail drawer (single + session) and the filter facets.
# The list endpoint `/spend/logs/ui` is covered via
# spend_tracking_routes below.
"/spend/logs/ui/{logId}",
"/spend/logs/session/ui",
"/management/v1/spend_logs/end_users",
"/management/v1/spend_logs/users",
# Settings / observability read endpoints exposed in admin-only
# sidebar groups (Logging & Alerts, Admin Settings, Budgets,
# Invitations).

View file

@ -193,6 +193,9 @@ async def anthropic_response(
)
verbose_proxy_logger.exception("litellm.proxy.proxy_server.anthropic_response(): Exception occured - %s", e)
if isinstance(e, ProxyException):
raise
# Extract model_id from request metadata (same as success path)
litellm_metadata: Final = data.get("litellm_metadata", {}) or {}
model_info: Final = litellm_metadata.get("model_info", {}) or {}

View file

@ -13,7 +13,7 @@ import asyncio
import math
import re
import time
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
@ -30,6 +30,9 @@ from litellm.constants import (
DEFAULT_IN_MEMORY_TTL,
DEFAULT_MAX_RECURSE_DEPTH,
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE,
END_USER_RESTRICTED_REGISTRY_MAX_SIZE,
REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
TAG_REGISTRY_MAX_SIZE,
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
@ -74,9 +77,15 @@ from litellm.proxy.common_utils.http_parsing_utils import (
)
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.common_utils.user_api_key_cache import (
END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL,
TAG_REGISTRY_OVERFLOW_SENTINEL,
UserApiKeyCache,
end_user_cache_key,
end_user_restricted_registry_cache_key,
get_management_object_ttl,
object_permission_cache_key,
tag_cache_key,
tag_registry_cache_key,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.guardrails.tool_name_extraction import (
@ -163,7 +172,7 @@ class _PrismaAuthTable(Protocol[RowT_co]):
async def find_many(
self,
*,
where: Mapping[str, object],
where: Mapping[str, object] | None = None,
include: Mapping[str, object] | None = None,
take: int | None = None,
) -> Sequence[RowT_co]: ...
@ -220,6 +229,16 @@ def _tag_table(repo: _PrismaTableHolder[_PrismaTagRow]) -> _PrismaAuthTable[_Pri
return repo.table
class _PrismaEndUserRow(Protocol):
user_id: str
def dict(self) -> Mapping[str, object]: ...
def _end_user_table(repo: _PrismaTableHolder[_PrismaEndUserRow]) -> _PrismaAuthTable[_PrismaEndUserRow]:
return repo.table
class _RawCacheRead(Protocol):
async def async_get_cache(self, *, key: str) -> object: ...
@ -1284,6 +1303,191 @@ async def _check_end_user_budget(
)
#: Columns whose non-null value makes an end-user row restrict something auth enforces. ``blocked``
#: is separate: it restricts when true rather than when merely set.
_RESTRICTED_COLUMNS: Final = ("budget_id", "allowed_model_region", "default_model", "object_permission_id")
def _column_is_set(column: str) -> Mapping[str, object]:
"""``column IS NOT NULL`` as a plain dict, which is the only shape prisma's builder accepts."""
return {column: {"not": None}} # mutable-ok: prisma's query builder isinstance-checks for dict
def _restricted_end_user_where() -> Mapping[str, object]:
"""Prisma filter selecting every end-user row that carries a restriction auth enforces."""
return {"OR": [{"blocked": True}, *map(_column_is_set, _RESTRICTED_COLUMNS)]} # mutable-ok: prisma needs dict/list
class _RegistryNotCached:
"""No cached registry answer, as distinct from the cached answer ``None`` (registry unusable)."""
_REGISTRY_NOT_CACHED: Final = _RegistryNotCached()
#: One lock per registry; module-level because the stampede to collapse is worker-wide.
_TAG_REGISTRY_LOAD_LOCK: Final = asyncio.Lock()
_END_USER_REGISTRY_LOAD_LOCK: Final = asyncio.Lock()
async def _cached_registry(
cache_key: str,
overflow_sentinel: str,
user_api_key_cache: UserApiKeyCache,
) -> frozenset[str] | None | _RegistryNotCached:
"""The cached registry answer, or ``_REGISTRY_NOT_CACHED`` when the caller has to query."""
cached: Final = await _raw_cache(user_api_key_cache).async_get_cache(key=cache_key)
if cached == overflow_sentinel:
return None
# Memory hands back the tuple that was written; Redis round-trips it through JSON as a list.
if isinstance(cached, (list, tuple)):
return frozenset(entry for entry in cached if isinstance(entry, str))
return _REGISTRY_NOT_CACHED
async def _cache_registry_answer(
cache_key: str,
value: tuple[str, ...] | str,
ttl: float,
user_api_key_cache: UserApiKeyCache,
) -> None:
"""Best-effort: a cache backend failure must not turn a registry load into a failed request."""
try:
await user_api_key_cache.async_set_cache(key=cache_key, value=value, ttl=ttl)
except Exception as e: # noqa: BLE001 # best-effort cache write: auth must survive a cache backend error
verbose_proxy_logger.warning("Failed to cache registry %s: %s", cache_key, e)
async def _fetch_and_cache_registry(
cache_key: str,
overflow_sentinel: str,
max_size: int,
fetch_ids: Callable[[], Awaitable[tuple[str, ...]]],
user_api_key_cache: UserApiKeyCache,
) -> frozenset[str] | None:
"""The registry as the database has it, cached whole, or ``None`` when it is unusable."""
try:
registry_ids: Final = await fetch_ids()
except Exception as e: # noqa: BLE001 # fail-safe: any registry load error must degrade to per-id lookups, never break auth
verbose_proxy_logger.warning(
"Registry %s could not be loaded from the database, so per-id lookups will run and the "
"registry query is suppressed for %ss: %s",
cache_key,
REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
e,
)
await _cache_registry_answer(
cache_key=cache_key,
value=overflow_sentinel,
ttl=REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
user_api_key_cache=user_api_key_cache,
)
return None
if len(registry_ids) > max_size:
await _cache_registry_answer(
cache_key=cache_key,
value=overflow_sentinel,
ttl=get_management_object_ttl(user_api_key_cache),
user_api_key_cache=user_api_key_cache,
)
return None
await _cache_registry_answer(
cache_key=cache_key,
value=registry_ids,
ttl=get_management_object_ttl(user_api_key_cache),
user_api_key_cache=user_api_key_cache,
)
return frozenset(registry_ids)
async def _load_bounded_registry(
cache_key: str,
overflow_sentinel: str,
max_size: int,
load_lock: asyncio.Lock,
fetch_ids: Callable[[], Awaitable[tuple[str, ...]]],
user_api_key_cache: UserApiKeyCache,
) -> frozenset[str] | None:
"""
A bounded id set under one cache key, so an id outside it costs no DB read.
``None`` = unusable (overflow or recent DB error): fall back to per-id lookups. An empty
frozenset is a real, cacheable answer. Loads are single-flighted to stop TTL-expiry stampedes.
"""
cached: Final = await _cached_registry(cache_key, overflow_sentinel, user_api_key_cache)
if not isinstance(cached, _RegistryNotCached):
return cached
async with load_lock:
# The request that held the lock has since cached an answer for everyone waiting on it.
cached_after_wait: Final = await _cached_registry(cache_key, overflow_sentinel, user_api_key_cache)
if not isinstance(cached_after_wait, _RegistryNotCached):
return cached_after_wait
return await _fetch_and_cache_registry(
cache_key=cache_key,
overflow_sentinel=overflow_sentinel,
max_size=max_size,
fetch_ids=fetch_ids,
user_api_key_cache=user_api_key_cache,
)
async def _load_end_user_restricted_registry(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> frozenset[str] | None:
"""The set of end-user ids whose ``LiteLLM_EndUserTable`` row carries a restriction."""
async def fetch_ids() -> tuple[str, ...]:
restricted_rows: Final = await _end_user_table(EndUserRepository(prisma_client)).find_many(
where=_restricted_end_user_where(),
take=END_USER_RESTRICTED_REGISTRY_MAX_SIZE + 1,
)
return tuple(row.user_id for row in restricted_rows)
return await _load_bounded_registry(
cache_key=end_user_restricted_registry_cache_key(),
overflow_sentinel=END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL,
max_size=END_USER_RESTRICTED_REGISTRY_MAX_SIZE,
load_lock=_END_USER_REGISTRY_LOAD_LOCK,
fetch_ids=fetch_ids,
user_api_key_cache=user_api_key_cache,
)
async def _end_user_is_known_unrestricted(
end_user_id: str,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
token_end_user_max_budget: float | None,
) -> bool:
"""
True when the cached registry proves the id restricts nothing, so its row need not be read.
Every field ``get_end_user_object`` callers consume (budget, spend under that budget, region,
default model, object permission, blocked) is part of the registry predicate, so an id outside
it is indistinguishable from one with no row at all. The skip is off whenever mere existence of
the row is meaningful: ``max_end_user_budget_id`` grafts a default budget onto any row that
exists, ``validate_end_user_id_in_db`` rejects ids that resolve to no row, and a token-supplied
``end_user_max_budget`` (a ``user_custom_auth`` callable can set one against an otherwise
unrestricted row) is enforced against the row's recorded spend.
"""
if (
litellm.max_end_user_budget_id is not None
or litellm.validate_end_user_id_in_db
or token_end_user_max_budget is not None
):
return False
registry: Final = await _load_end_user_restricted_registry(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
return registry is not None and end_user_id not in registry
@log_db_metrics
async def get_end_user_object(
end_user_id: str | None,
@ -1292,6 +1496,7 @@ async def get_end_user_object(
route: str | None = "",
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
token_end_user_max_budget: float | None = None,
) -> LiteLLM_EndUserTable | None:
"""
Returns end user object from database or cache.
@ -1306,6 +1511,9 @@ async def get_end_user_object(
route: The request route
parent_otel_span: Optional OpenTelemetry span for tracing
proxy_logging_obj: Optional proxy logging object
token_end_user_max_budget: ``valid_token.end_user_max_budget``, when the caller holds a
token. Budget enforcement reads the row's spend, so a row that restricts nothing on
its own must still be loaded when the token carries a budget for it.
Returns:
LiteLLM_EndUserTable if found, None otherwise
@ -1316,7 +1524,7 @@ async def get_end_user_object(
if end_user_id is None:
return None
_key: Final = f"end_user_id:{end_user_id}"
_key: Final = end_user_cache_key(end_user_id)
# Check cache first
cached_user_obj: Final = await user_api_key_cache.async_get_cache(
@ -1335,6 +1543,14 @@ async def get_end_user_object(
return return_obj
if await _end_user_is_known_unrestricted(
end_user_id=end_user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
token_end_user_max_budget=token_end_user_max_budget,
):
return None
# Fetch from database
try:
response: Final = await _dictable_table(EndUserRepository(prisma_client)).find_unique(
@ -1358,9 +1574,10 @@ async def get_end_user_object(
# Save to cache
await user_api_key_cache.async_set_cache(
key=f"end_user_id:{end_user_id}",
key=_key,
value=_response,
model_type=LiteLLM_EndUserTable,
ttl=get_management_object_ttl(user_api_key_cache),
)
return _response
@ -1480,6 +1697,67 @@ async def _end_user_id_exists_in_db(
return False
async def _load_tag_registry(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> frozenset[str] | None:
"""The set of tag names that have a row in ``LiteLLM_TagTable``."""
async def fetch_ids() -> tuple[str, ...]:
registry_rows: Final = await _tag_table(TagRepository(prisma_client)).find_many(
take=TAG_REGISTRY_MAX_SIZE + 1,
)
return tuple(row.tag_name for row in registry_rows)
return await _load_bounded_registry(
cache_key=tag_registry_cache_key(),
overflow_sentinel=TAG_REGISTRY_OVERFLOW_SENTINEL,
max_size=TAG_REGISTRY_MAX_SIZE,
load_lock=_TAG_REGISTRY_LOAD_LOCK,
fetch_ids=fetch_ids,
user_api_key_cache=user_api_key_cache,
)
async def _fetch_uncached_tags(
uncached_tags: Sequence[str],
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
) -> tuple[tuple[str, LiteLLM_TagTable], ...]:
"""Rows for the tags a cache probe missed; names absent from the registry never reach the DB."""
if not uncached_tags:
return ()
registry: Final = await _load_tag_registry(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
tags_to_fetch: Final = (
tuple(uncached_tags) if registry is None else tuple(tag for tag in uncached_tags if tag in registry)
)
if not tags_to_fetch:
return ()
try:
db_tags: Final = await _tag_table(TagRepository(prisma_client)).find_many(
where={"tag_name": {"in": list(tags_to_fetch)}},
include={"litellm_budget_table": True},
)
fetched: Final = tuple((db_tag.tag_name, LiteLLM_TagTable.model_validate(db_tag.dict())) for db_tag in db_tags)
for fetched_name, fetched_obj in fetched:
await user_api_key_cache.async_set_cache(
key=tag_cache_key(fetched_name),
value=fetched_obj,
model_type=LiteLLM_TagTable,
ttl=get_management_object_ttl(user_api_key_cache),
)
except Exception as e: # noqa: BLE001 # fail-safe: a tag fetch error must yield "no budget objects", never break auth
verbose_proxy_logger.debug("Error batch fetching tags from database: %s", e)
return ()
else:
return fetched
@log_db_metrics
async def get_tag_objects_batch(
tag_names: list[str],
@ -1492,8 +1770,9 @@ async def get_tag_objects_batch(
Batch fetch multiple tag objects from cache and db.
Optimizes for latency by:
1. Fetching all cached tags in parallel
2. Batch fetching uncached tags in one DB query
1. Serving already-cached tags without touching the DB
2. Skipping tags that no ``LiteLLM_TagTable`` row exists for, via the cached name registry
3. Batch fetching the remaining uncached tags in one DB query
Args:
tag_names: List of tag names to fetch
@ -1505,50 +1784,22 @@ async def get_tag_objects_batch(
Returns:
Dictionary mapping tag_name to LiteLLM_TagTable object
"""
if prisma_client is None:
if prisma_client is None or not tag_names:
return {}
if not tag_names:
return {}
tag_objects: Final = dict[str, LiteLLM_TagTable]()
uncached_tags: Final = list[str]()
# Try to get all tags from cache first
for tag_name in tag_names:
cache_key = f"tag:{tag_name}"
cached_tag = await user_api_key_cache.async_get_cache(
key=cache_key,
model_type=LiteLLM_TagTable,
probed: Final = [
(
tag_name,
await user_api_key_cache.async_get_cache(key=tag_cache_key(tag_name), model_type=LiteLLM_TagTable),
)
if cached_tag is not None:
tag_objects[tag_name] = cached_tag
else:
uncached_tags.append(tag_name)
# Batch fetch uncached tags from DB in one query
if uncached_tags:
try:
db_tags: Final = await _tag_table(TagRepository(prisma_client)).find_many(
where={"tag_name": {"in": uncached_tags}},
include={"litellm_budget_table": True},
)
# Cache and add to tag_objects
for db_tag in db_tags:
tag_name = db_tag.tag_name
cache_key = f"tag:{tag_name}"
_tag_obj = LiteLLM_TagTable.model_validate(db_tag.dict())
await user_api_key_cache.async_set_cache(
key=cache_key,
value=_tag_obj,
model_type=LiteLLM_TagTable,
)
tag_objects[tag_name] = _tag_obj
except Exception as e:
verbose_proxy_logger.debug("Error batch fetching tags from database: %s", e)
return tag_objects
for tag_name in tag_names
]
fetched: Final = await _fetch_uncached_tags(
uncached_tags=tuple(tag_name for tag_name, tag_obj in probed if tag_obj is None),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
return {tag_name: tag_obj for tag_name, tag_obj in (*probed, *fetched) if tag_obj is not None}
@log_db_metrics
@ -4573,25 +4824,15 @@ async def delete_cached_project_object(
user_api_key_cache: UserApiKeyCache,
) -> None:
"""
Every endpoint that mutates litellm_projecttable must call this: get_project_object
serves auth cache-first with no freshness check, so without invalidation a stale
project (e.g. a pre-update empty model allowlist) keeps being enforced until the
TTL expires (LIT-3803). Best-effort on both steps: the DB write has already
committed, so a cache backend error must not fail the endpoint; the stale entry
then expires via TTL.
Every endpoint that mutates litellm_projecttable must call this, or a stale project (e.g. a
pre-update empty model allowlist) keeps being enforced until the TTL expires (LIT-3803).
"""
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
cache_key: Final = _project_cache_key(project_id)
try:
await user_api_key_cache.async_delete_cache(key=cache_key)
except Exception as e: # noqa: BLE001 # best-effort eviction: any cache backend error must not fail the mutation
verbose_proxy_logger.warning(
"Failed to evict cached project entry %s; a stale project may be served until its TTL expires: %s",
cache_key,
e,
)
await publish_auth_cache_invalidation(cache_key=cache_key)
await evict_and_broadcast(
cache_keys=(_project_cache_key(project_id),),
user_api_key_cache=user_api_key_cache,
)
async def _organization_max_budget_check(

View file

@ -2307,6 +2307,7 @@ async def _run_centralized_common_checks(
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
token_end_user_max_budget=user_api_key_auth_obj.end_user_max_budget,
),
)
)
@ -2841,6 +2842,7 @@ async def _lookup_end_user_and_apply_budget(
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
token_end_user_max_budget=valid_token.end_user_max_budget,
)
if end_user_object is not None:
end_user_params = {

View file

@ -0,0 +1,18 @@
from litellm.proxy._types import ProxyException
def validate_batch_list_limit(limit: int | None) -> None:
if limit is None or 0 <= limit <= 100:
return
bound, expected, openai_code = (
("below minimum", ">= 0", "integer_below_min_value")
if limit < 0
else ("above maximum", "<= 100", "integer_above_max_value")
)
raise ProxyException(
message=f"Invalid 'limit': integer {bound} value. Expected a value {expected}, but got {limit} instead.",
type="invalid_request_error",
param="limit",
code=400,
openai_code=openai_code,
)

View file

@ -5,6 +5,8 @@
######################################################################
import asyncio
import os
from collections.abc import Mapping
from typing import Any, Final, cast
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response
@ -14,6 +16,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.batches.main import CancelBatchRequest, RetrieveBatchRequest
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.batches_endpoints.common_utils import validate_batch_list_limit
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.callback_utils import sanitize_openai_provider_metadata
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
@ -40,6 +43,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
update_batch_in_database,
validate_managed_id_requirement,
)
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
from litellm.repositories.table_repositories import ManagedFileRepository
from litellm.types.llms.openai import LiteLLMBatchCreateRequest
@ -47,6 +51,23 @@ from litellm.types.llms.openai import LiteLLMBatchCreateRequest
router: Final = APIRouter()
def _raise_not_found_when_openai_fallback_unservable(
requested_provider: "str | None",
data: Mapping[str, object],
not_found_message: str,
) -> None:
if requested_provider is not None:
return
if data.get("api_key") or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY"):
return
raise ProxyException(
message=not_found_message,
type="invalid_request_error",
param=None,
code=404,
)
async def _resolve_managed_input_file_storage_url(input_file_id: str) -> "str | None":
"""Resolve a managed (unified) input_file_id to its backend storage_url.
@ -140,6 +161,8 @@ async def create_batch(
)
data["metadata"] = sanitize_openai_provider_metadata(data.get("metadata"))
raise_if_required_body_param_missing(route_type="acreate_batch", data=data)
## check if model is a loadbalanced model
router_model: str | None = None
is_router_model = False
@ -147,12 +170,12 @@ async def create_batch(
router_model = data.get("model", None)
is_router_model = is_known_model(model=router_model, llm_router=llm_router)
custom_llm_provider: Final = (
requested_provider: Final = (
provider
or data.pop("custom_llm_provider", None)
or get_custom_llm_provider_from_request_headers(request=request)
or "openai"
)
custom_llm_provider: Final = requested_provider or "openai"
_create_batch_data: Final = LiteLLMBatchCreateRequest(**data)
# Apply team-level batch output expiry enforcement
@ -314,6 +337,11 @@ async def create_batch(
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
_raise_not_found_when_openai_fallback_unservable(
requested_provider=requested_provider,
data=cast(dict, _create_batch_data), # cast-ok: TypedDict is a dict at runtime
not_found_message=f"No such File object: {input_file_id}",
)
response = await litellm.acreate_batch(
custom_llm_provider=custom_llm_provider,
**_create_batch_data,
@ -563,18 +591,23 @@ async def retrieve_batch(
# SCENARIO 3: Fallback to custom_llm_provider (uses env variables)
else:
custom_llm_provider: Final = (
requested_provider: Final = (
provider
or get_custom_llm_provider_from_request_headers(request=request)
or get_custom_llm_provider_from_request_query(request=request)
or "openai"
)
custom_llm_provider: Final = requested_provider or "openai"
apply_team_provider_credentials(
data=data,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
_raise_not_found_when_openai_fallback_unservable(
requested_provider=requested_provider,
data=data,
not_found_message=f"No batch found with id '{batch_id}'.",
)
response = await litellm.aretrieve_batch(
custom_llm_provider=custom_llm_provider,
**data,
@ -679,6 +712,7 @@ async def list_batches(
```
"""
validate_batch_list_limit(limit)
from litellm.proxy.proxy_server import (
general_settings,
llm_router,
@ -967,13 +1001,13 @@ async def cancel_batch(
# SCENARIO 3: Fallback to custom_llm_provider (uses env variables)
else:
body_custom_llm_provider = data.pop("custom_llm_provider", None)
custom_llm_provider: Final = (
requested_provider: Final = (
provider
or body_custom_llm_provider
or get_custom_llm_provider_from_request_headers(request=request)
or get_custom_llm_provider_from_request_query(request=request)
or "openai"
)
custom_llm_provider: Final = requested_provider or "openai"
# Extract batch_id from data to avoid "multiple values for keyword argument" error
# data was cast from CancelBatchRequest which already contains batch_id
data.pop("batch_id", None)
@ -983,6 +1017,11 @@ async def cancel_batch(
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)
_raise_not_found_when_openai_fallback_unservable(
requested_provider=requested_provider,
data=data,
not_found_message=f"No batch found with id '{batch_id}'.",
)
_cancel_batch_data: Final = CancelBatchRequest(batch_id=batch_id, **data)
response = await litellm.acancel_batch(
custom_llm_provider=custom_llm_provider,

View file

@ -2669,10 +2669,10 @@ class ProxyBaseLLMRequestProcessing:
streaming pipeline (including unified_guardrail end-of-stream blocks)
has completed.
Guardrails with apply_guardrail are skipped — they already ran via
unified_guardrail's streaming iterator. Only guardrails that override
async_post_call_success_hook directly (without apply_guardrail) run
here.
Guardrails routed through unified_guardrail are skipped, since they already ran
via its streaming iterator. Guardrails that override
async_post_call_success_hook directly run here, including those that implement
apply_guardrail but keep their native lifecycle hooks.
This is audit-only — content has already been delivered to the client.
@ -2695,8 +2695,8 @@ class ProxyBaseLLMRequestProcessing:
continue
try:
guardrail_result = None
if "apply_guardrail" in type(cb).__dict__:
# Skip — apply_guardrail guardrails already ran via
if "apply_guardrail" in type(cb).__dict__ and not cb.use_native_lifecycle_hooks:
# Skip — unified-routed guardrails already ran via
# unified_guardrail's end-of-stream block in the
# streaming iterator pipeline. Running them again
# here would duplicate the guardrail API call

View file

@ -1,5 +1,6 @@
import asyncio
import json
from collections.abc import Sequence
from dataclasses import asdict, dataclass
from typing import TYPE_CHECKING, Final
@ -72,6 +73,27 @@ async def publish_auth_cache_invalidation(cache_key: str) -> None:
verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e)
async def evict_and_broadcast(cache_keys: Sequence[str], user_api_key_cache: "UserApiKeyCache") -> None:
"""
Drop cached management objects here and on every other worker.
Every endpoint that mutates a cached object must call this: auth serves those objects
cache-first with no freshness check, so a mutation that leaves the entry in place keeps the
stale object enforced until its TTL expires (LIT-3803). Best-effort on both steps: the DB write
has already committed, so a cache backend error must not fail the endpoint.
"""
for cache_key in cache_keys:
try:
await user_api_key_cache.async_delete_cache(key=cache_key)
except Exception as e: # noqa: BLE001 # best-effort eviction: any cache backend error must not fail the mutation
verbose_proxy_logger.warning(
"Failed to evict cached entry %s; a stale object may be served until its TTL expires: %s",
cache_key,
e,
)
await publish_auth_cache_invalidation(cache_key=cache_key)
class AuthCacheInvalidationSubscriber:
__slots__ = ("_redis_cache", "_task", "_user_api_key_cache")

View file

@ -0,0 +1,226 @@
from __future__ import annotations
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import date, datetime, timezone
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
import litellm
from litellm._logging import verbose_logger
from litellm.types.proxy.model_deprecation import (
DEFAULT_DEPRECATION_WARN_DAYS,
DeprecationStatus,
ModelDeprecationInfo,
ModelDeprecationResponse,
)
if TYPE_CHECKING:
from litellm.router import Router
_NO_MODEL_METADATA: Final[Mapping[str, object]] = MappingProxyType({})
@dataclass(frozen=True, slots=True)
class _ResolvedDeprecation:
deprecation_date: date
litellm_model: str | None
litellm_provider: str | None
def _parse_deprecation_date(raw_value: object) -> date | None:
if isinstance(raw_value, datetime):
return raw_value.date()
if isinstance(raw_value, date):
return raw_value
if not isinstance(raw_value, str):
return None
try:
return date.fromisoformat(raw_value.strip())
except ValueError:
return None
def _cost_map_lookup(model_key: object) -> _ResolvedDeprecation | None:
if not isinstance(model_key, str) or not model_key:
return None
entry: Final = litellm.model_cost.get(model_key)
if not isinstance(entry, Mapping):
return None
parsed: Final = _parse_deprecation_date(entry.get("deprecation_date"))
if parsed is None:
return None
provider: Final = entry.get("litellm_provider")
return _ResolvedDeprecation(
deprecation_date=parsed,
litellm_model=model_key,
litellm_provider=provider if isinstance(provider, str) else None,
)
def _mapping_field(deployment: Mapping[str, object], key: str) -> Mapping[str, object]:
value: Final = deployment.get(key)
return value if isinstance(value, Mapping) else _NO_MODEL_METADATA
def _resolve_deployment_deprecation(
deployment: Mapping[str, object],
) -> _ResolvedDeprecation | None:
"""Resolve a deployment's deprecation date, preferring its explicit override"""
model_info: Final = _mapping_field(deployment, "model_info")
raw_model: Final = _mapping_field(deployment, "litellm_params").get("model")
override: Final = _parse_deprecation_date(model_info.get("deprecation_date"))
if override is not None:
provider: Final = model_info.get("litellm_provider")
return _ResolvedDeprecation(
deprecation_date=override,
litellm_model=raw_model if isinstance(raw_model, str) else None,
litellm_provider=provider if isinstance(provider, str) else None,
)
unprefixed: Final = raw_model.split("/", 1)[1] if isinstance(raw_model, str) and "/" in raw_model else None
return next(
(
resolved
for resolved in (
_cost_map_lookup(model_info.get("base_model")),
_cost_map_lookup(raw_model),
_cost_map_lookup(unprefixed),
)
if resolved is not None
),
None,
)
def _classify(days_until: int, warn_within_days: int) -> DeprecationStatus:
if days_until < 0:
return "deprecated"
if days_until <= warn_within_days:
return "imminent"
return "upcoming"
def _build_info(deployment: Mapping[str, object], today: date, warn_within_days: int) -> ModelDeprecationInfo | None:
model_name: Final = deployment.get("model_name")
if not isinstance(model_name, str) or not model_name:
return None
resolved: Final = _resolve_deployment_deprecation(deployment)
if resolved is None:
return None
days_until: Final = (resolved.deprecation_date - today).days
return ModelDeprecationInfo(
model_name=model_name,
litellm_model=resolved.litellm_model,
deprecation_date=resolved.deprecation_date,
days_until_deprecation=days_until,
status=_classify(days_until, warn_within_days),
litellm_provider=resolved.litellm_provider,
)
def _dedupe(
models: Sequence[ModelDeprecationInfo],
) -> tuple[ModelDeprecationInfo, ...]:
"""Report a model group carrying the same date on several deployments once"""
ordered: Final = sorted(models, key=lambda model: (model.model_name, model.deprecation_date))
return tuple(
next(group) for _, group in groupby(ordered, key=lambda model: (model.model_name, model.deprecation_date))
)
def _bucket(models: Sequence[ModelDeprecationInfo], status: DeprecationStatus) -> tuple[ModelDeprecationInfo, ...]:
return tuple(
sorted(
(model for model in models if model.status == status),
key=lambda model: model.deprecation_date,
)
)
def collect_model_deprecations(
llm_router: Router | None,
warn_within_days: int = DEFAULT_DEPRECATION_WARN_DAYS,
today: date | None = None,
) -> ModelDeprecationResponse:
"""Bucket every deployment carrying a deprecation date by how urgent it is"""
snapshot_time: Final = datetime.now(timezone.utc)
effective_today: Final = today or snapshot_time.date()
deployments: Final = (llm_router.get_model_list() or ()) if llm_router is not None else ()
deduped: Final = _dedupe(
tuple(
info
for info in (_build_info(deployment, effective_today, warn_within_days) for deployment in deployments)
if info is not None
)
)
verbose_logger.debug(
"model_deprecation: %d/%d deployments carry a deprecation date",
len(deduped),
len(deployments),
)
return ModelDeprecationResponse(
deprecated=_bucket(deduped, "deprecated"),
imminent=_bucket(deduped, "imminent"),
upcoming=_bucket(deduped, "upcoming"),
warn_within_days=warn_within_days,
checked_at=snapshot_time,
)
def _escape_slack_mrkdwn(value: str) -> str:
"""Neutralize Slack control characters so a model name cannot forge a mention or link"""
return value.replace("&", "&amp;").replace("<", "&lt;").replace(">", "&gt;")
def _format_entry(info: ModelDeprecationInfo) -> str:
suffix: Final = (
f"already deprecated {abs(info.days_until_deprecation)}d ago"
if info.days_until_deprecation < 0
else f"in {info.days_until_deprecation}d"
)
return (
f"• `{_escape_slack_mrkdwn(info.model_name)}` "
f"(provider: {_escape_slack_mrkdwn(info.litellm_provider) if info.litellm_provider else 'unknown'}, "
f"deprecates {info.deprecation_date.isoformat()}, {suffix})"
)
def format_deprecation_alert_message(
snapshot: ModelDeprecationResponse,
) -> str | None:
"""Render the alert for the deprecated and imminent buckets, None when both are empty
Upcoming models are left out of the alert to keep it actionable.
"""
if not snapshot.deprecated and not snapshot.imminent:
return None
deprecated_section: Final = (
("\n*Already deprecated:*", *(_format_entry(i) for i in snapshot.deprecated)) if snapshot.deprecated else ()
)
imminent_section: Final = (
(
f"\n*Deprecating within {snapshot.warn_within_days} days:*",
*(_format_entry(i) for i in snapshot.imminent),
)
if snapshot.imminent
else ()
)
return "\n".join(
(
"*⚠️ Model Deprecation Warning*",
*deprecated_section,
*imminent_section,
"\nPlan migrations to a supported model. See "
"https://docs.litellm.ai/docs/proxy/model_management for guidance.",
)
)

View file

@ -28,6 +28,7 @@ from litellm.proxy.common_utils.timezone_utils import (
compute_budget_reset_at,
get_budget_reset_settings,
)
from litellm.proxy.common_utils.user_api_key_cache import tag_cache_key
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.prisma_protocols import ReadOnlyTable, SpendLinkedTable
@ -112,7 +113,7 @@ def _tag_counter_key(row: _TagRow) -> str:
def _tag_cache_keys(row: _TagRow) -> tuple[str, ...]:
return (f"tag:{row.tag_name}",)
return (tag_cache_key(row.tag_name),)
def _budget_link_where(

View file

@ -170,6 +170,36 @@ def object_permission_cache_key(object_permission_id: str) -> str:
return f"object_permission_id:{object_permission_id}"
#: Cached under ``tag_registry_cache_key`` when the table exceeds ``TAG_REGISTRY_MAX_SIZE``:
#: registry unusable, fall back to the per-tag lookup.
TAG_REGISTRY_OVERFLOW_SENTINEL: Final = "__tag_registry_overflow__"
def tag_cache_key(tag_name: str) -> str:
"""Cache key one tag row is stored under; shared so its five reader/writer modules cannot drift."""
return f"tag:{tag_name}"
def tag_registry_cache_key() -> str:
"""Cache key for the set of tag names that exist in ``LiteLLM_TagTable``."""
return "tag_registry"
#: Cached under ``end_user_restricted_registry_cache_key`` when the restricted set exceeds
#: ``END_USER_RESTRICTED_REGISTRY_MAX_SIZE``: registry unusable, fall back to the per-id fetch.
END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL: Final = "__end_user_restricted_registry_overflow__"
def end_user_cache_key(end_user_id: str) -> str:
"""Cache key one end-user row is stored under; shared so auth and spend tracking cannot drift."""
return f"end_user_id:{end_user_id}"
def end_user_restricted_registry_cache_key() -> str:
"""Cache key for the set of end-user ids whose row carries a restriction auth enforces."""
return "end_user_restricted_registry"
def get_management_object_ttl(cache: DualCache) -> float:
"""
In-memory TTL for management-object cache writes (keys, teams, users, budgets, ...).

View file

@ -13,7 +13,7 @@ import random
import time
import traceback
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, cast, overload
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload
import litellm
from litellm._logging import verbose_proxy_logger
@ -64,6 +64,7 @@ from litellm.proxy.spend_tracking.savings import (
extract_cache_read_tokens,
)
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.repositories.prisma_protocols import BatchTable
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient, ProxyLogging
@ -72,6 +73,37 @@ else:
ProxyLogging = Any
class _SpendBatch(Protocol):
litellm_usertable: BatchTable
litellm_verificationtoken: BatchTable
litellm_teamtable: BatchTable
litellm_teammembership: BatchTable
litellm_organizationtable: BatchTable
litellm_tagtable: BatchTable
litellm_agentstable: BatchTable
class _SpendBatchManager(Protocol):
async def __aenter__(self) -> _SpendBatch: ...
async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ...
class _SpendTransaction(Protocol):
def batch_(self) -> _SpendBatchManager: ...
class _SpendTransactionManager(Protocol):
async def __aenter__(self) -> _SpendTransaction: ...
async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ...
def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager:
tx: Final[_SpendTransactionManager] = prisma_client.db.tx(timeout=timedelta(seconds=60))
return tx
def _get_llm_router():
"""The proxy's router, or None outside a running proxy.
@ -1195,7 +1227,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction:
async with _spend_update_tx(prisma_client) as transaction:
async with transaction.batch_() as batcher:
# Sort by ID for consistent lock ordering across pods to prevent deadlocks.
# batch_() issues statements sequentially within the tx, so iteration
@ -1237,7 +1269,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction:
async with _spend_update_tx(prisma_client) as transaction:
async with transaction.batch_() as batcher:
# Sort by token for consistent lock ordering across pods to prevent deadlocks.
for token, response_cost in sorted(key_list_transactions.items()):
@ -1270,7 +1302,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction:
async with _spend_update_tx(prisma_client) as transaction:
async with transaction.batch_() as batcher:
# Sort by team_id for consistent lock ordering across pods to prevent deadlocks.
for team_id, response_cost in sorted(team_list_transactions.items()):
@ -1311,7 +1343,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction:
async with _spend_update_tx(prisma_client) as transaction:
async with transaction.batch_() as batcher:
# Sort by composite key for consistent lock ordering across pods to prevent deadlocks.
# Key format "team_id::<v>::user_id::<v>" makes the string sort equivalent to sorting by (team_id, user_id).
@ -1362,7 +1394,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction:
async with _spend_update_tx(prisma_client) as transaction:
async with transaction.batch_() as batcher:
# Sort by org_id for consistent lock ordering across pods to prevent deadlocks.
for org_id, response_cost in sorted(org_list_transactions.items()):
@ -1420,7 +1452,7 @@ class DBSpendUpdateWriter:
async def _update_entity_spend_in_db(
entity_name: str,
transactions: dict[str, float] | None,
table_accessor: Any,
table_accessor: Literal["litellm_tagtable", "litellm_agentstable"],
where_field: str,
n_retry_times: int,
prisma_client: PrismaClient,
@ -1445,7 +1477,7 @@ class DBSpendUpdateWriter:
for i in range(n_retry_times + 1):
start_time = time.time()
try:
async with prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction:
async with _spend_update_tx(prisma_client) as transaction:
async with transaction.batch_() as batcher:
# Sort by entity_id for consistent lock ordering across pods to prevent deadlocks.
for entity_id, response_cost in sorted(transactions.items()):

View file

@ -6,8 +6,9 @@ Admins use the management endpoints to read and update input_policy / output_pol
"""
import uuid
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ToolDiscoveryQueueItem
@ -20,8 +21,41 @@ from litellm.types.tool_management import (
)
if TYPE_CHECKING:
from prisma import models as prisma_db_models
from litellm.proxy.utils import PrismaClient
_RowT_co: Final = TypeVar("_RowT_co", covariant=True)
class _TableActions(Protocol[_RowT_co]):
async def find_unique(self, where: Mapping[str, object]) -> _RowT_co | None: ...
async def find_many(
self,
where: Mapping[str, object] | None = None,
order: Mapping[str, object] | None = None,
include: Mapping[str, object] | None = None,
) -> Sequence[_RowT_co]: ...
async def upsert(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT_co: ...
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT_co | None: ...
def _tool_table_actions(prisma_client: "PrismaClient") -> "_TableActions[prisma_db_models.LiteLLM_ToolTable]":
table: Final[_TableActions[prisma_db_models.LiteLLM_ToolTable]] = ToolRepository(prisma_client).table
return table
def _object_permission_table_actions(
prisma_client: "PrismaClient",
) -> "_TableActions[prisma_db_models.LiteLLM_ObjectPermissionTable]":
table: Final[_TableActions[prisma_db_models.LiteLLM_ObjectPermissionTable]] = ObjectPermissionRepository(
prisma_client
).table
return table
def _row_to_model(row: dict | Any) -> LiteLLM_ToolTableRow:
"""Convert a Prisma model instance or dict to LiteLLM_ToolTableRow."""
@ -87,7 +121,7 @@ async def batch_upsert_tools(
if not data:
return
now: Final = datetime.now(timezone.utc)
table: Final = ToolRepository(prisma_client).table
table: Final = _tool_table_actions(prisma_client)
for item in data:
tool_name = item.get("tool_name", "")
origin = item.get("origin") or "user_defined"
@ -132,8 +166,8 @@ async def list_tools(
) -> list[LiteLLM_ToolTableRow]:
"""Return all tools, optionally filtered by input_policy."""
try:
where: Final = {"input_policy": input_policy} if input_policy is not None else {}
rows: Final = await ToolRepository(prisma_client).table.find_many(
where: Final[Mapping[str, str]] = {"input_policy": input_policy} if input_policy is not None else {}
rows: Final = await _tool_table_actions(prisma_client).find_many(
where=where,
order={"created_at": "desc"},
)
@ -149,7 +183,7 @@ async def get_tool(
) -> LiteLLM_ToolTableRow | None:
"""Return a single tool row by tool_name."""
try:
row: Final = await ToolRepository(prisma_client).table.find_unique(
row: Final = await _tool_table_actions(prisma_client).find_unique(
where={"tool_name": tool_name},
)
if row is None:
@ -172,7 +206,7 @@ async def update_tool_policy(
_updated_by: Final = updated_by or "system"
now: Final = datetime.now(timezone.utc)
create_data: Final[dict] = {
create_data: Final[dict[str, object]] = {
"tool_id": str(uuid.uuid4()),
"tool_name": tool_name,
"input_policy": input_policy or "untrusted",
@ -182,7 +216,7 @@ async def update_tool_policy(
"created_at": now,
"updated_at": now,
}
update_data: Final[dict] = {
update_data: Final[dict[str, object]] = {
"updated_by": _updated_by,
"updated_at": now,
}
@ -191,7 +225,7 @@ async def update_tool_policy(
if output_policy is not None:
update_data["output_policy"] = output_policy
await ToolRepository(prisma_client).table.upsert(
await _tool_table_actions(prisma_client).upsert(
where={"tool_name": tool_name},
data={
"create": create_data,
@ -214,7 +248,7 @@ async def get_tools_by_names(
if not tool_names:
return {}
try:
rows: Final = await ToolRepository(prisma_client).table.find_many(
rows: Final = await _tool_table_actions(prisma_client).find_many(
where={"tool_name": {"in": tool_names}},
)
return {
@ -239,7 +273,7 @@ async def list_overrides_for_tool(
"""
out: Final[list[ToolPolicyOverrideRow]] = []
try:
perms: Final = await ObjectPermissionRepository(prisma_client).table.find_many(
perms: Final = await _object_permission_table_actions(prisma_client).find_many(
where={"blocked_tools": {"has": tool_name}},
include={
"verification_tokens": True,
@ -302,7 +336,7 @@ class ToolPolicyRegistry:
try:
tools: Final = await call_with_db_reconnect_retry(
prisma_client,
lambda: ToolRepository(prisma_client).table.find_many(),
lambda: _tool_table_actions(prisma_client).find_many(),
reason="sync_tool_policy_from_db_tools_lookup_failure",
)
self._tool_input_policies = {
@ -314,7 +348,7 @@ class ToolPolicyRegistry:
perms: Final = await call_with_db_reconnect_retry(
prisma_client,
lambda: ObjectPermissionRepository(prisma_client).table.find_many(),
lambda: _object_permission_table_actions(prisma_client).find_many(),
reason="sync_tool_policy_from_db_perms_lookup_failure",
)
self._blocked_tools_by_op_id = {}
@ -352,7 +386,7 @@ class ToolPolicyRegistry:
"""
if not tool_names:
return {}
blocked: Final[set] = set()
blocked: Final[set[str]] = set()
for op_id in (object_permission_id, team_object_permission_id):
if op_id and op_id.strip():
blocked.update(self._blocked_tools_by_op_id.get(op_id.strip(), []))
@ -385,7 +419,7 @@ async def add_tool_to_object_permission_blocked(
if not object_permission_id or not tool_name:
return False
try:
row: Final = await ObjectPermissionRepository(prisma_client).table.find_unique(
row: Final = await _object_permission_table_actions(prisma_client).find_unique(
where={"object_permission_id": object_permission_id},
)
if row is None:
@ -394,7 +428,7 @@ async def add_tool_to_object_permission_blocked(
if tool_name in current:
return True
current.append(tool_name)
await ObjectPermissionRepository(prisma_client).table.update(
await _object_permission_table_actions(prisma_client).update(
where={"object_permission_id": object_permission_id},
data={"blocked_tools": current},
)
@ -413,7 +447,7 @@ async def remove_tool_from_object_permission_blocked(
if not object_permission_id or not tool_name:
return False
try:
row: Final = await ObjectPermissionRepository(prisma_client).table.find_unique(
row: Final = await _object_permission_table_actions(prisma_client).find_unique(
where={"object_permission_id": object_permission_id},
)
if row is None:
@ -422,7 +456,7 @@ async def remove_tool_from_object_permission_blocked(
if tool_name not in current:
return False
current = [t for t in current if t != tool_name]
await ObjectPermissionRepository(prisma_client).table.update(
await _object_permission_table_actions(prisma_client).update(
where={"object_permission_id": object_permission_id},
data={"blocked_tools": current},
)

View file

@ -3,7 +3,7 @@
Azure Prompt Shield Native Guardrail Integrationfor LiteLLM
"""
from typing import TYPE_CHECKING, Any, Final, cast
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast
from fastapi import HTTPException
@ -13,11 +13,12 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypesLiteral
from litellm.types.utils import CallTypesLiteral, GenericGuardrailAPIInputs
from .base import AzureGuardrailBase
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_prompt_shield import (
@ -40,6 +41,8 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
default_on: Whether to enable by default
"""
use_native_lifecycle_hooks: ClassVar[bool] = True
def __init__(
self,
guardrail_name: str,
@ -103,6 +106,19 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
assert last_response is not None
return last_response
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: "LiteLLMLoggingObj | None" = None,
) -> GenericGuardrailAPIInputs:
for text in inputs.get("texts") or ():
if text:
await self.async_make_request(user_prompt=text)
return inputs
@log_guardrail_information
async def async_pre_call_hook(
self,

View file

@ -3,7 +3,7 @@
Azure Text Moderation Native Guardrail Integrationfor LiteLLM
"""
from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Union, cast
from fastapi import HTTPException
@ -14,11 +14,12 @@ from litellm.integrations.custom_guardrail import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypesLiteral
from litellm.types.utils import CallTypesLiteral, GenericGuardrailAPIInputs
from .base import AzureGuardrailBase
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import (
AzureTextModerationGuardrailResponse,
@ -41,6 +42,8 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
default_on: Whether to enable by default
"""
use_native_lifecycle_hooks: ClassVar[bool] = True
default_severity_threshold: int = 2
@classmethod
@ -147,6 +150,19 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
assert last_response is not None
return last_response
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: "LiteLLMLoggingObj | None" = None,
) -> GenericGuardrailAPIInputs:
for text in inputs.get("texts") or ():
if text:
await self.async_make_request(text=text)
return inputs
def check_severity_threshold(self, response: "AzureTextModerationGuardrailResponse") -> Literal[True]:
"""
- Check if threshold set by category

View file

@ -2053,6 +2053,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
bedrock_action: Final = response.get("action")
if isinstance(bedrock_action, str):
tracing_detail["guardrail_action"] = bedrock_action
usage: Final = response.get("usage")
if isinstance(usage, dict):
usage_units: Final = { # mutable-ok: json.dumps'd into spend log metadata downstream
key: value for key, value in usage.items() if isinstance(value, int)
}
if usage_units:
tracing_detail["guardrail_usage"] = usage_units
return tracing_detail
def _extract_violation_category_names(self, response: BedrockGuardrailResponse) -> list[str]:

View file

@ -53,7 +53,7 @@ class LassoResponse(TypedDict):
violations_detected: bool
deputies: dict[str, bool]
findings: dict[str, list[dict[str, Any]]]
findings: dict[str, list[dict[str, object]]]
messages: list[dict[str, str]] | None
@ -120,7 +120,7 @@ class LassoGuardrail(CustomGuardrail):
super().__init__(**kwargs)
@staticmethod
def _get_field(obj: Any, field: str, default: Any = None) -> Any:
def _get_field(obj: Any, field: str, default: object = None) -> Any:
"""Get a field from either a dict or a Pydantic object."""
if isinstance(obj, dict):
return obj.get(field, default)
@ -129,7 +129,7 @@ class LassoGuardrail(CustomGuardrail):
@staticmethod
def _extract_tool_call_fields(
call: Any,
) -> tuple[str | None, str | None, dict[str, Any] | None]:
) -> tuple[str | None, str | None, dict[str, object] | None]:
"""Extract (call_id, name, parsed_input) from a tool call.
Handles both dict-style and Pydantic object-style tool_calls.
@ -142,7 +142,7 @@ class LassoGuardrail(CustomGuardrail):
return call_id, None, None
name: Final = get(func, "name")
args_str: Final = get(func, "arguments")
input_data: dict[str, Any] | None = None
input_data: dict[str, object] | None = None
if args_str:
try:
parsed = json.loads(args_str)
@ -248,7 +248,7 @@ class LassoGuardrail(CustomGuardrail):
# Extract messages from the response for validation
if isinstance(response, litellm.ModelResponse):
response_messages: Final[list[dict[str, Any]]] = []
response_messages: Final[list[dict[str, object]]] = []
for choice in response.choices:
if not hasattr(choice, "message"):
continue
@ -392,7 +392,7 @@ class LassoGuardrail(CustomGuardrail):
LassoGuardrailAPIError: If the Lasso API call fails
HTTPException: If blocking violations are detected
"""
raw_messages: Final[list[dict[str, Any]]] = data.get("messages") or []
raw_messages: Final[list[dict[str, object]]] = data.get("messages") or []
messages: list[dict[str, Any]] = self._expand_messages_for_classification(raw_messages) if raw_messages else []
messages_count: Final = len(messages)
if data.get("input") is not None:
@ -417,7 +417,7 @@ class LassoGuardrail(CustomGuardrail):
data: dict,
cache: DualCache,
message_type: Literal["PROMPT", "COMPLETION"],
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
) -> dict:
"""Handle classification without masking."""
try:
@ -435,7 +435,7 @@ class LassoGuardrail(CustomGuardrail):
data: dict,
cache: DualCache,
message_type: Literal["PROMPT", "COMPLETION"],
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
messages_count: int,
) -> dict:
"""Handle masking with classifix endpoint.
@ -477,7 +477,7 @@ class LassoGuardrail(CustomGuardrail):
self,
original_messages: list[dict[str, Any]],
masked_messages: list[dict[str, Any]],
) -> list[dict[str, Any]]:
) -> list[dict[str, object]]:
"""Map Lasso-format masked messages back onto the original OpenAI-format messages.
Lasso receives expanded messages (tool_use / tool_result blocks) and returns them
@ -487,7 +487,7 @@ class LassoGuardrail(CustomGuardrail):
while preserving the original structure.
"""
# Index masked content by type so we can look up by id without caring about order.
masked_tool_use: Final[dict[str, dict[str, Any]]] = {}
masked_tool_use: Final[dict[str, dict[str, object]]] = {}
masked_tool_result: Final[dict[str, str]] = {}
masked_text: Final[list[str]] = []
@ -524,7 +524,7 @@ class LassoGuardrail(CustomGuardrail):
},
)
result: Final[list[dict[str, Any]]] = []
result: Final[list[dict[str, object]]] = []
text_cursor = 0
for orig_msg in original_messages:
@ -563,9 +563,9 @@ class LassoGuardrail(CustomGuardrail):
def _update_tool_calls_from_masked(
self,
tool_calls: list[Any],
masked_tool_use: dict[str, dict[str, Any]],
) -> list[Any]:
tool_calls: list[object],
masked_tool_use: dict[str, dict[str, object]],
) -> list[object]:
"""Replace tool_call arguments with masked values returned by Lasso."""
updated: Final = []
for call in tool_calls:
@ -745,11 +745,11 @@ class LassoGuardrail(CustomGuardrail):
def _prepare_payload(
self,
messages: list[dict[str, Any]],
messages: list[dict[str, object]],
data: dict,
cache: DualCache,
message_type: Literal["PROMPT", "COMPLETION"] = "PROMPT",
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Prepare the payload for the Lasso API request.
@ -759,7 +759,7 @@ class LassoGuardrail(CustomGuardrail):
data: Request data (used for conversation_id generation and tools extraction)
cache: Cache instance for storing conversation_id (optional for post-call)
"""
payload: Final[dict[str, Any]] = {
payload: Final[dict[str, object]] = {
"messages": messages,
"messageType": message_type,
# Drives the "Used By" badge on Lasso Application API Keys: every call from this
@ -776,7 +776,7 @@ class LassoGuardrail(CustomGuardrail):
payload["sessionId"] = conversation_id
# Map OpenAI ChatCompletionToolParam array → ToolDefinition array
tools_data: Final[list[dict[str, Any]]] = data.get("tools") or []
tools_data: Final[list[dict[str, object]]] = data.get("tools") or []
if tools_data:
get: Final = self._get_field
tool_definitions: Final = []
@ -787,7 +787,7 @@ class LassoGuardrail(CustomGuardrail):
name = get(func, "name")
if not name:
continue
td: dict[str, Any] = {"name": name}
td: dict[str, object] = {"name": name}
description = get(func, "description")
if description:
td["description"] = description
@ -803,7 +803,7 @@ class LassoGuardrail(CustomGuardrail):
async def _call_lasso_api(
self,
headers: dict[str, str],
payload: dict[str, Any],
payload: dict[str, object],
api_url: str | None = None,
) -> LassoResponse:
"""Call the Lasso API and return the response."""
@ -921,7 +921,7 @@ class LassoGuardrail(CustomGuardrail):
) -> None:
"""Apply masking to the actual model response when mask=True and masked content is available."""
# Index masked tool_use blocks by id for O(1) lookup.
masked_tool_use: Final[dict[str, dict[str, Any]]] = {}
masked_tool_use: Final[dict[str, dict[str, object]]] = {}
masked_text: Final[list[str]] = []
for masked_msg in masked_messages:
content = masked_msg.get("content")

View file

@ -36,6 +36,13 @@ _AIDR_SCAN_ENDPOINT: Final = "/litellm/guardrail"
_INTERVENED_INPUT_FIELDS: Final = ("texts", "images", "tools", "tool_calls")
_DEFAULT_API_BASE_HOSTNAME: Final = urlparse(_DEFAULT_API_BASE).hostname
_KEYS_DUPLICATING_SCAN_INPUTS: Final = ("messages", "input")
_LOGGING_KEYS_DUPLICATING_SCAN_INPUTS: Final = _KEYS_DUPLICATING_SCAN_INPUTS + (
"additional_args",
"standard_logging_object",
"original_response",
)
class _Action(str, enum.Enum):
BLOCKED = "BLOCKED"
@ -131,9 +138,20 @@ class NomaV2Guardrail(CustomGuardrail):
logging_obj: Optional["LiteLLMLoggingObj"],
application_id: str | None,
) -> dict:
payload_request_data: Final = self._sanitize_payload_for_transport(request_data)
payload_request_data: Final = self._sanitize_payload_for_transport(
{key: value for key, value in request_data.items() if key not in _KEYS_DUPLICATING_SCAN_INPUTS}
)
if logging_obj is not None:
payload_request_data["litellm_logging_obj"] = getattr(logging_obj, "model_call_details", None)
model_call_details: Final = getattr(logging_obj, "model_call_details", None)
payload_request_data["litellm_logging_obj"] = (
{
key: value
for key, value in model_call_details.items()
if key not in _LOGGING_KEYS_DUPLICATING_SCAN_INPUTS
}
if isinstance(model_call_details, dict)
else model_call_details
)
payload: Final[dict[str, Any]] = {
"inputs": inputs,

View file

@ -14,9 +14,10 @@ import threading
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, cast
import aiohttp
from typing_extensions import NotRequired, ReadOnly
import litellm
from litellm import get_secret
@ -53,9 +54,18 @@ from litellm.utils import (
)
class _PresidioAnonymizeItem(TypedDict, total=False):
entity_type: ReadOnly[str | None]
class _PresidioAnonymizeResponse(TypedDict):
text: ReadOnly[str]
items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]]
class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
user_api_key_cache = None
ad_hoc_recognizers = None
ad_hoc_recognizers: list[str] | None = None
@classmethod
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
@ -72,7 +82,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
def __init__(
self,
mock_testing: bool = False,
mock_redacted_text: dict | None = None,
mock_redacted_text: _PresidioAnonymizeResponse | None = None,
presidio_analyzer_api_base: str | None = None,
presidio_anonymizer_api_base: str | None = None,
output_parse_pii: bool | None = False,
@ -91,7 +101,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
super().__init__(**kwargs)
self.guardrail_provider = "presidio"
self.pii_tokens: dict = {} # mapping of PII token to original text - only used with Presidio `replace` operation
self.pii_tokens: dict[
str, str
] = {} # mapping of PII token to original text - only used with Presidio `replace` operation
self.mock_redacted_text = mock_redacted_text
self.output_parse_pii = output_parse_pii or False
self.apply_to_output = apply_to_output
@ -265,7 +277,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
text: str,
presidio_config: PresidioPerRequestConfig | None,
request_data: dict,
) -> list[PresidioAnalyzeResponseItem] | dict:
) -> list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse:
"""
Send text to the Presidio analyzer endpoint and get analysis results
"""
@ -385,7 +397,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
# contain API keys or other secrets) in error responses.
raise Exception(f"Presidio PII analysis failed: {type(e).__name__}") from e
async def _post_presidio_anonymize(self, text: str, analyze_results: Any) -> Any:
async def _post_presidio_anonymize(
self,
text: str,
analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse,
) -> _PresidioAnonymizeResponse | None:
"""POST to Presidio anonymize; returns parsed JSON body."""
# Use shared session to prevent memory leak (issue #14540)
async with self._get_session_iterator() as session:
@ -417,7 +433,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
def _finalize_presidio_anonymize_simple(
self,
redacted_text: dict[str, Any],
redacted_text: _PresidioAnonymizeResponse,
masked_entity_count: dict[str, int],
) -> str:
# No need to build numbered tokens — just use Presidio's
@ -483,7 +499,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
async def anonymize_text(
self,
text: str,
analyze_results: Any,
analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse,
output_parse_pii: bool,
masked_entity_count: dict[str, int],
request_data: dict | None = None,
@ -517,8 +533,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
raise Exception(f"Presidio PII anonymization failed: {type(e).__name__}") from e
def filter_analyze_results_by_score(
self, analyze_results: list[PresidioAnalyzeResponseItem] | dict
) -> list[PresidioAnalyzeResponseItem] | dict:
self, analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse
) -> list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse:
"""
Drop detections that fall below configured per-entity score thresholds
or match an entity type in the deny list.
@ -556,7 +572,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
return filtered_results
def raise_exception_if_blocked_entities_detected(self, analyze_results: list[PresidioAnalyzeResponseItem] | dict):
def raise_exception_if_blocked_entities_detected(
self, analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse
):
"""
Raise an exception if blocked entities are detected
"""
@ -590,7 +608,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
Calls Presidio Analyze + Anonymize endpoints for PII Analysis + Masking
"""
start_time: Final = datetime.now()
analyze_results: list[PresidioAnalyzeResponseItem] | dict | None = None
analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse | None = None
status: GuardrailStatus = "success"
masked_entity_count: Final[dict[str, int]] = {}
exception_str: str = ""
@ -895,7 +913,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
return text
@staticmethod
def _is_anthropic_message_response(response: Any) -> bool:
def _is_anthropic_message_response(
response: ModelResponse | EmbeddingResponse | ImageResponse | dict[str, object],
) -> bool:
"""Check if the response is an Anthropic native message dict."""
return (
isinstance(response, dict)
@ -1283,8 +1303,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
@staticmethod
def _preserve_usage_from_last_chunk(
assembled_model_response: Any,
chunks: list[Any],
assembled_model_response: ModelResponse,
chunks: list[ModelResponseStream],
) -> None:
"""Copy usage metadata from the last chunk when stream_chunk_builder misses it."""
if not getattr(assembled_model_response, "usage", None) and chunks:

View file

@ -4,18 +4,21 @@ GET /guardrails/usage/overview, /guardrails/usage/detail/:id, /guardrails/usage/
"""
import json
from collections.abc import Mapping, Sequence
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import datetime, timedelta, timezone
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, overload
from fastapi import APIRouter, Depends, Query
from pydantic import BaseModel
from typing_extensions import NotRequired, TypedDict
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.repositories.table_repositories import (
DailyGuardrailMetricsRepository,
DailyGuardrailUsageUnitsRepository,
DailyPolicyMetricsRepository,
GuardrailsRepository,
PolicyRepository,
@ -26,7 +29,13 @@ from litellm.repositories.table_repositories import (
if TYPE_CHECKING:
from prisma import models as prisma_models
from prisma import types as prisma_types
from prisma.actions import LiteLLM_GuardrailsTableActions, LiteLLM_PolicyTableActions
from prisma.actions import (
LiteLLM_DailyGuardrailMetricsActions,
LiteLLM_DailyGuardrailUsageUnitsActions,
LiteLLM_DailyPolicyMetricsActions,
LiteLLM_GuardrailsTableActions,
LiteLLM_PolicyTableActions,
)
from litellm.proxy.utils import PrismaClient
from litellm.types.guardrails import Guardrail
@ -36,6 +45,8 @@ if TYPE_CHECKING:
router: Final = APIRouter()
_EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({})
def _guardrails_table(
prisma_client: "PrismaClient",
@ -55,9 +66,86 @@ def _policies_table(
return policies_table
def _daily_guardrail_metrics_table(
prisma_client: "PrismaClient",
) -> "LiteLLM_DailyGuardrailMetricsActions[prisma_models.LiteLLM_DailyGuardrailMetrics]":
metrics_table: Final[LiteLLM_DailyGuardrailMetricsActions[prisma_models.LiteLLM_DailyGuardrailMetrics]] = (
DailyGuardrailMetricsRepository(prisma_client).table
)
return metrics_table
def _daily_policy_metrics_table(
prisma_client: "PrismaClient",
) -> "LiteLLM_DailyPolicyMetricsActions[prisma_models.LiteLLM_DailyPolicyMetrics]":
metrics_table: Final[LiteLLM_DailyPolicyMetricsActions[prisma_models.LiteLLM_DailyPolicyMetrics]] = (
DailyPolicyMetricsRepository(prisma_client).table
)
return metrics_table
async def _find_daily_guardrail_metrics(
prisma_client: "PrismaClient",
where: "prisma_types.LiteLLM_DailyGuardrailMetricsWhereInput",
) -> "Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]":
return await _daily_guardrail_metrics_table(prisma_client).find_many(where=where)
async def _find_daily_policy_metrics(
prisma_client: "PrismaClient",
where: "prisma_types.LiteLLM_DailyPolicyMetricsWhereInput",
) -> "Sequence[prisma_models.LiteLLM_DailyPolicyMetrics]":
return await _daily_policy_metrics_table(prisma_client).find_many(where=where)
def _daily_guardrail_usage_units_table(
prisma_client: "PrismaClient",
) -> "LiteLLM_DailyGuardrailUsageUnitsActions[prisma_models.LiteLLM_DailyGuardrailUsageUnits]":
units_table: Final[LiteLLM_DailyGuardrailUsageUnitsActions[prisma_models.LiteLLM_DailyGuardrailUsageUnits]] = (
DailyGuardrailUsageUnitsRepository(prisma_client).table
)
return units_table
async def _find_daily_guardrail_usage_units(
prisma_client: "PrismaClient",
where: "prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereInput",
) -> "Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits]":
return await _daily_guardrail_usage_units_table(prisma_client).find_many(where=where)
def _counter_name(row: "prisma_models.LiteLLM_DailyGuardrailUsageUnits") -> str:
return row.usage_unit
def _sum_counter_units(rows: "Iterable[prisma_models.LiteLLM_DailyGuardrailUsageUnits]") -> Mapping[str, int]:
ordered: Final = sorted(rows, key=_counter_name)
return MappingProxyType(
{name: sum(int(r.units) for r in group) for name, group in groupby(ordered, key=_counter_name)}
)
def _units_by(
rows: "Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits]",
key_of: "Callable[[prisma_models.LiteLLM_DailyGuardrailUsageUnits], str]",
) -> Mapping[str, Mapping[str, int]]:
ordered: Final = sorted(rows, key=key_of)
return MappingProxyType({key: _sum_counter_units(group) for key, group in groupby(ordered, key=key_of)})
# --- Response models ---
class _GuardrailRunInfo(TypedDict, total=False):
guardrail_id: ReadOnly[str | None]
guardrail_name: ReadOnly[str | None]
guardrail_status: ReadOnly[str | None]
duration: ReadOnly[float | None]
confidence_score: ReadOnly[float | None]
risk_score: ReadOnly[float | None]
guardrail_response: ReadOnly[str | Mapping[str, object] | Sequence[Mapping[str, object]] | None]
class UsageChartPoint(TypedDict):
date: str
passed: int
@ -93,6 +181,7 @@ class UsageOverviewRow(BaseModel):
avgLatency: float | None
status: str # healthy | warning | critical
trend: str # up | down | stable
usageUnits: Mapping[str, int]
class UsageOverviewResponse(BaseModel):
@ -101,6 +190,12 @@ class UsageOverviewResponse(BaseModel):
totalRequests: int
totalBlocked: int
passRate: float
totalUsageUnits: Mapping[str, int]
class UsageUnitsDailyPoint(BaseModel):
date: str
units: Mapping[str, int]
class UsageDetailResponse(BaseModel):
@ -116,6 +211,10 @@ class UsageDetailResponse(BaseModel):
trend: str
description: str | None
time_series: list[UsageChartPoint]
usage_units: Mapping[str, int]
usage_units_daily: Sequence[UsageUnitsDailyPoint]
usage_units_by_team: Mapping[str, Mapping[str, int]]
usage_units_by_key: Mapping[str, Mapping[str, int]]
class UsageLogEntry(BaseModel):
@ -231,6 +330,7 @@ def _guardrail_overview_rows(
guardrails: "Sequence[_DbOrConfigGuardrail]",
agg: Mapping[str, _MetricTotals],
prev_agg: Mapping[str, float],
units_agg: Mapping[str, Mapping[str, int]],
) -> list[UsageOverviewRow]:
rows: Final[list[UsageOverviewRow]] = []
covered_keys: Final[set[str]] = set()
@ -256,6 +356,7 @@ def _guardrail_overview_rows(
prev_fail = float(prev_agg.get(k, 0.0) or 0.0)
break
trend = _trend_from_comparison(fail_rate, prev_fail)
row_units: Mapping[str, int] = next((units_agg[k] for k in lookup_keys if k in units_agg), _EMPTY_UNITS)
rows.append(
UsageOverviewRow(
id=gid,
@ -268,6 +369,7 @@ def _guardrail_overview_rows(
avgLatency=None,
status=_status_from_fail_rate(fail_rate),
trend=trend,
usageUnits=row_units,
)
)
# Add rows for guardrails with metrics but not in guardrails table (e.g. MCP, config)
@ -290,6 +392,7 @@ def _guardrail_overview_rows(
avgLatency=None,
status=_status_from_fail_rate(fail_rate),
trend=trend,
usageUnits=units_agg.get(agg_key, _EMPTY_UNITS),
)
)
return rows
@ -319,6 +422,7 @@ def _policy_overview_rows(
avgLatency=None,
status=_status_from_fail_rate(fail_rate),
trend=trend,
usageUnits=_EMPTY_UNITS,
)
)
return rows
@ -339,7 +443,9 @@ async def guardrails_usage_overview(
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return UsageOverviewResponse(rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0)
return UsageOverviewResponse(
rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0, totalUsageUnits=_EMPTY_UNITS
)
now: Final = datetime.now(timezone.utc)
end: Final = end_date or now.strftime("%Y-%m-%d")
@ -356,29 +462,38 @@ async def guardrails_usage_overview(
guardrails: Final[Sequence[_DbOrConfigGuardrail]] = [*db_guardrails, *config_guardrails]
# Daily metrics in range
metrics: Final[Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]] = await DailyGuardrailMetricsRepository(
prisma_client
).table.find_many(where={"date": {"gte": start, "lte": end}})
metrics: Final[Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]] = await _find_daily_guardrail_metrics(
prisma_client, where={"date": {"gte": start, "lte": end}}
)
# Previous period for trend
start_prev: Final = (datetime.strptime(start, "%Y-%m-%d") - timedelta(days=7)).strftime("%Y-%m-%d")
metrics_prev: Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics] = await DailyGuardrailMetricsRepository(
prisma_client
).table.find_many(where={"date": {"gte": start_prev, "lt": start}})
metrics_prev: Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics] = await _find_daily_guardrail_metrics(
prisma_client, where={"date": {"gte": start_prev, "lt": start}}
)
units_where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereInput] = {
"date": {"gte": start, "lte": end}
}
units_rows: Final[
Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits]
] = await _find_daily_guardrail_usage_units(prisma_client, where=units_where)
agg: Final = _aggregate_daily_metrics(metrics, "guardrail_id")
prev_agg: Final = _prev_fail_rates(metrics_prev, "guardrail_id")
units_agg: Final = _units_by(units_rows, lambda r: r.guardrail_id)
chart: Final = _chart_from_metrics(metrics)
total_requests: Final = sum(a["requests"] for a in agg.values())
total_blocked: Final = sum(a["blocked"] for a in agg.values())
pass_rate: Final = (100.0 * (total_requests - total_blocked) / total_requests) if total_requests else 100.0
rows: Final = _guardrail_overview_rows(guardrails, agg, prev_agg)
rows: Final = _guardrail_overview_rows(guardrails, agg, prev_agg, units_agg)
return UsageOverviewResponse(
rows=rows,
chart=chart,
totalRequests=total_requests,
totalBlocked=total_blocked,
passRate=round(pass_rate, 1),
totalUsageUnits=_sum_counter_units(units_rows),
)
except Exception as e:
from litellm.proxy.utils import handle_exception_on_proxy
@ -424,22 +539,27 @@ async def guardrails_usage_detail(
logical_id: Final = _get_guardrail_field(guardrail, "guardrail_name")
metric_ids: Final = [i for i in (logical_id, guardrail_id) if i]
metrics: Final[Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]] = await DailyGuardrailMetricsRepository(
prisma_client
).table.find_many(
metrics: Final[Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]] = await _find_daily_guardrail_metrics(
prisma_client,
where={
"guardrail_id": {"in": metric_ids},
"date": {"gte": start, "lte": end},
}
},
)
metrics_prev: Final[Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]] = await DailyGuardrailMetricsRepository(
prisma_client
).table.find_many(
metrics_prev: Final[Sequence[prisma_models.LiteLLM_DailyGuardrailMetrics]] = await _find_daily_guardrail_metrics(
prisma_client,
where={
"guardrail_id": {"in": metric_ids},
"date": {"lt": start},
}
},
)
units_where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereInput] = {
"guardrail_id": {"in": metric_ids},
"date": {"gte": start, "lte": end},
}
units_rows: Final[
Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits]
] = await _find_daily_guardrail_usage_units(prisma_client, where=units_where)
requests: Final = sum(int(m.requests_evaluated or 0) for m in metrics)
blocked: Final = sum(int(m.blocked_count or 0) for m in metrics)
@ -465,6 +585,8 @@ async def guardrails_usage_detail(
litellm_params: Final = _to_dict(_get_guardrail_field(guardrail, "litellm_params"))
guardrail_info: Final = _to_dict(_get_guardrail_field(guardrail, "guardrail_info"))
_guardrail_name: Final = _get_guardrail_field(guardrail, "guardrail_name")
daily_unit_sums: Final = sorted(_units_by(units_rows, lambda r: r.date).items())
units_daily: Final = tuple(UsageUnitsDailyPoint(date=d, units=units) for d, units in daily_unit_sums)
return UsageDetailResponse(
guardrail_id=guardrail_id,
@ -479,6 +601,10 @@ async def guardrails_usage_detail(
trend=trend,
description=guardrail_info.get("description"),
time_series=time_series,
usage_units=_sum_counter_units(units_rows),
usage_units_daily=units_daily,
usage_units_by_team=_units_by(units_rows, lambda r: r.team_id),
usage_units_by_key=_units_by(units_rows, lambda r: r.api_key),
)
@ -510,7 +636,9 @@ def _build_usage_logs_where(
def _usage_log_entry_from_row(
r: "prisma_models.LiteLLM_SpendLogGuardrailIndex", sl: Any, action_filter: str | None
r: "prisma_models.LiteLLM_SpendLogGuardrailIndex",
sl: "prisma_models.LiteLLM_SpendLogs",
action_filter: str | None,
) -> UsageLogEntry | None:
meta = sl.metadata
if isinstance(meta, str):
@ -518,8 +646,8 @@ def _usage_log_entry_from_row(
meta = json.loads(meta)
except Exception:
meta = {}
guardrail_info_list: Final = (meta or {}).get("guardrail_information") or []
entry_for_guardrail = None
guardrail_info_list: Final[Sequence[_GuardrailRunInfo]] = (meta or {}).get("guardrail_information") or []
entry_for_guardrail: _GuardrailRunInfo | None = None
for gi in guardrail_info_list:
if (gi.get("guardrail_id") or gi.get("guardrail_name")) == r.guardrail_id:
entry_for_guardrail = gi
@ -567,13 +695,12 @@ def _snippet(text: Any, max_len: int = 200) -> str | None:
if isinstance(text, str):
s = text
elif isinstance(text, list):
parts: Final = []
for item in text:
if isinstance(item, dict) and "content" in item:
c = item["content"]
parts.append(c if isinstance(c, str) else str(c))
else:
parts.append(str(item))
parts: Final[Sequence[str]] = [
(c if isinstance(c := item["content"], str) else str(c))
if isinstance(item, dict) and "content" in item
else str(item)
for item in text
]
s = " ".join(parts)
else:
s = str(text)
@ -697,7 +824,9 @@ async def policies_usage_overview(
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return UsageOverviewResponse(rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0)
return UsageOverviewResponse(
rows=[], chart=[], totalRequests=0, totalBlocked=0, passRate=100.0, totalUsageUnits=_EMPTY_UNITS
)
now: Final = datetime.now(timezone.utc)
end: Final = end_date or now.strftime("%Y-%m-%d")
@ -705,18 +834,17 @@ async def policies_usage_overview(
try:
policies: Final = await _policies_table(prisma_client).find_many()
metrics: Final[Sequence[prisma_models.LiteLLM_DailyPolicyMetrics]] = await DailyPolicyMetricsRepository(
prisma_client
).table.find_many(where={"date": {"gte": start, "lte": end}})
metrics_prev: Final[Sequence[prisma_models.LiteLLM_DailyPolicyMetrics]] = await DailyPolicyMetricsRepository(
prisma_client
).table.find_many(
metrics: Final[Sequence[prisma_models.LiteLLM_DailyPolicyMetrics]] = await _find_daily_policy_metrics(
prisma_client, where={"date": {"gte": start, "lte": end}}
)
metrics_prev: Final[Sequence[prisma_models.LiteLLM_DailyPolicyMetrics]] = await _find_daily_policy_metrics(
prisma_client,
where={
"date": {
"gte": (datetime.strptime(start, "%Y-%m-%d") - timedelta(days=7)).strftime("%Y-%m-%d"),
"lt": start,
}
}
},
)
agg: Final = _aggregate_daily_metrics(metrics, "policy_id")
prev_agg: Final = _prev_fail_rates(metrics_prev, "policy_id")
@ -731,6 +859,7 @@ async def policies_usage_overview(
totalRequests=total_requests,
totalBlocked=total_blocked,
passRate=round(pass_rate, 1),
totalUsageUnits=_EMPTY_UNITS,
)
except Exception as e:
from litellm.proxy.utils import handle_exception_on_proxy

View file

@ -3,18 +3,82 @@ Track guardrail and policy usage for the dashboard: upsert daily metrics and
insert into SpendLogGuardrailIndex when spend logs are written.
"""
import asyncio
import json
from collections import defaultdict
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from datetime import datetime, timezone
from typing import Any, Final
from functools import partial
from itertools import groupby
from operator import itemgetter
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NamedTuple, TypeVar
from litellm._logging import verbose_proxy_logger
from litellm.proxy.utils import PrismaClient
from litellm.repositories.table_repositories import (
DailyGuardrailMetricsRepository,
DailyGuardrailUsageUnitsRepository,
SpendLogGuardrailIndexRepository,
)
if TYPE_CHECKING:
from prisma import types as prisma_types
_UPSERT_RETRY_TIMES: Final = 3
_RowKey = TypeVar("_RowKey")
_RowValue = TypeVar("_RowValue")
class _UsageUnitKey(NamedTuple):
guardrail_id: str
date: str
team_id: str
api_key: str
usage_unit: str
class _MetricsKey(NamedTuple):
guardrail_id: str
date: str
async def _attempt_upsert(
upsert_row: Callable[[_RowKey, _RowValue], Awaitable[None]], key: _RowKey, value: _RowValue
) -> Exception | None:
try:
await upsert_row(key, value)
except Exception as error:
return error
return None
async def _upsert_rows_with_retry(
rows: Mapping[_RowKey, _RowValue],
upsert_row: Callable[[_RowKey, _RowValue], Awaitable[None]],
label: str,
sleep: Callable[[float], Awaitable[None]],
retries_left: int = _UPSERT_RETRY_TIMES,
) -> None:
outcomes: Final = {key: await _attempt_upsert(upsert_row, key, value) for key, value in rows.items()}
failed: Final = MappingProxyType({key: rows[key] for key, error in outcomes.items() if error is not None})
if not failed:
return
if retries_left == 0:
for key in failed:
verbose_proxy_logger.warning(
"Guardrail usage tracking: %s upsert failed for %s after %d retries (non-fatal): %s",
label,
key,
_UPSERT_RETRY_TIMES,
outcomes[key],
)
return
await sleep(2 ** (_UPSERT_RETRY_TIMES - retries_left))
await _upsert_rows_with_retry(failed, upsert_row, label, sleep, retries_left - 1)
def _guardrail_status_to_action(status: str | None) -> str:
"""Map StandardLogging guardrail_status to blocked/passed/flagged."""
@ -28,7 +92,7 @@ def _guardrail_status_to_action(status: str | None) -> str:
return "passed"
def _parse_guardrail_info_from_payload(payload: dict[str, Any]) -> list[dict[str, Any]]:
def _parse_guardrail_info_from_payload(payload: Mapping[str, Any]) -> Sequence[Mapping[str, Any]]:
"""Extract guardrail_information from spend log payload metadata."""
meta = payload.get("metadata")
if not meta:
@ -53,9 +117,95 @@ def _date_str(dt: datetime) -> str:
return dt.astimezone(timezone.utc).strftime("%Y-%m-%d")
def _parse_payload_start_time(payload: Mapping[str, Any]) -> datetime | None:
start_time: Final = payload.get("startTime")
if isinstance(start_time, datetime):
return start_time
if not isinstance(start_time, str):
return None
try:
return datetime.fromisoformat(start_time.replace("Z", "+00:00"))
except (ValueError, TypeError):
return None
def _iter_usage_unit_increments(logs_to_process: Sequence[Mapping[str, Any]]) -> Iterator[tuple[_UsageUnitKey, int]]:
for payload in logs_to_process:
start_time = _parse_payload_start_time(payload)
if not payload.get("request_id") or start_time is None:
continue
date_key = _date_str(start_time)
team_id = str(payload.get("team_id") or "")
api_key = str(payload.get("api_key") or "")
for entry in _parse_guardrail_info_from_payload(payload):
guardrail_id = str(entry.get("guardrail_id") or entry.get("guardrail_name") or "")
usage = entry.get("guardrail_usage")
if not guardrail_id or not isinstance(usage, dict):
continue
for unit_name, units in usage.items():
if isinstance(units, int) and not isinstance(units, bool) and units > 0:
yield _UsageUnitKey(guardrail_id, date_key, team_id, api_key, str(unit_name)), units
def _sum_usage_unit_increments(logs_to_process: Sequence[Mapping[str, Any]]) -> Mapping[_UsageUnitKey, int]:
ordered: Final = sorted(_iter_usage_unit_increments(logs_to_process), key=itemgetter(0))
return MappingProxyType(
{key: sum(units for _, units in group) for key, group in groupby(ordered, key=itemgetter(0))}
)
async def _upsert_usage_unit_row(prisma_client: PrismaClient, key: _UsageUnitKey, units: int) -> None:
row: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsCreateInput] = {
"guardrail_id": key.guardrail_id,
"date": key.date,
"team_id": key.team_id,
"api_key": key.api_key,
"usage_unit": key.usage_unit,
"units": units,
}
where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereUniqueInput] = {
"guardrail_id_date_team_id_api_key_usage_unit": {
"guardrail_id": key.guardrail_id,
"date": key.date,
"team_id": key.team_id,
"api_key": key.api_key,
"usage_unit": key.usage_unit,
}
}
data: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsUpsertInput] = {
"create": row,
"update": {"units": {"increment": units}},
}
await DailyGuardrailUsageUnitsRepository(prisma_client).table.upsert(where=where, data=data)
async def _upsert_metrics_row(prisma_client: PrismaClient, key: _MetricsKey, agg: Mapping[str, int]) -> None:
n: Final = int(agg["requests_evaluated"])
await DailyGuardrailMetricsRepository(prisma_client).table.upsert(
where={"guardrail_id_date": {"guardrail_id": key.guardrail_id, "date": key.date}},
data={
"create": {
"guardrail_id": key.guardrail_id,
"date": key.date,
"requests_evaluated": n,
"passed_count": int(agg["passed_count"]),
"blocked_count": int(agg["blocked_count"]),
"flagged_count": int(agg["flagged_count"]),
},
"update": {
"requests_evaluated": {"increment": n},
"passed_count": {"increment": int(agg["passed_count"])},
"blocked_count": {"increment": int(agg["blocked_count"])},
"flagged_count": {"increment": int(agg["flagged_count"])},
},
},
)
async def process_spend_logs_guardrail_usage(
prisma_client: PrismaClient,
logs_to_process: list[dict[str, Any]],
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
) -> None:
"""
After spend logs are written: update DailyGuardrailMetrics and insert
@ -64,7 +214,7 @@ async def process_spend_logs_guardrail_usage(
if not logs_to_process:
return
# Aggregate daily metrics by (guardrail_id, date). Latency/score metrics dropped.
daily_guardrail: Final[dict[tuple, dict[str, Any]]] = defaultdict(
daily_guardrail: Final[dict[_MetricsKey, dict[str, Any]]] = defaultdict(
lambda: {
"requests_evaluated": 0,
"passed_count": 0,
@ -76,21 +226,16 @@ async def process_spend_logs_guardrail_usage(
for payload in logs_to_process:
request_id = payload.get("request_id")
start_time = payload.get("startTime")
if not request_id or not start_time:
start_time = _parse_payload_start_time(payload)
if not request_id or start_time is None:
continue
if isinstance(start_time, str):
try:
start_time = datetime.fromisoformat(start_time.replace("Z", "+00:00"))
except (ValueError, TypeError):
continue
date_key = _date_str(start_time)
for entry in _parse_guardrail_info_from_payload(payload):
guardrail_id = entry.get("guardrail_id") or entry.get("guardrail_name") or ""
if not guardrail_id:
continue
key = (guardrail_id, date_key)
key = _MetricsKey(guardrail_id, date_key)
daily_guardrail[key]["requests_evaluated"] += 1
action = _guardrail_status_to_action(entry.get("guardrail_status"))
if action == "passed":
@ -109,64 +254,29 @@ async def process_spend_logs_guardrail_usage(
}
)
if not daily_guardrail and not index_rows:
usage_unit_totals: Final = _sum_usage_unit_increments(logs_to_process)
if not daily_guardrail and not index_rows and not usage_unit_totals:
return
try:
# Insert index rows (skip duplicates by request_id + guardrail_id)
if index_rows:
index_data: Final = []
for r in index_rows:
st = r["start_time"]
if isinstance(st, str):
try:
st = datetime.fromisoformat(st.replace("Z", "+00:00"))
except (ValueError, TypeError):
continue
index_data.append(
{
"request_id": r["request_id"],
"guardrail_id": r["guardrail_id"],
"policy_id": r.get("policy_id"),
"start_time": st,
}
)
try:
await SpendLogGuardrailIndexRepository(prisma_client).table.create_many(
data=index_data,
data=index_rows,
skip_duplicates=True,
)
except Exception as e:
verbose_proxy_logger.debug("Guardrail usage tracking: index create_many skipped: %s", e)
# Upsert daily guardrail metrics (counts only; latency/score dropped)
for (guardrail_id, date_key), agg in daily_guardrail.items():
n = int(agg["requests_evaluated"])
if n == 0:
continue
await DailyGuardrailMetricsRepository(prisma_client).table.upsert(
where={
"guardrail_id_date": {
"guardrail_id": guardrail_id,
"date": date_key,
}
},
data={
"create": {
"guardrail_id": guardrail_id,
"date": date_key,
"requests_evaluated": n,
"passed_count": int(agg["passed_count"]),
"blocked_count": int(agg["blocked_count"]),
"flagged_count": int(agg["flagged_count"]),
},
"update": {
"requests_evaluated": {"increment": n},
"passed_count": {"increment": int(agg["passed_count"])},
"blocked_count": {"increment": int(agg["blocked_count"])},
"flagged_count": {"increment": int(agg["flagged_count"])},
},
},
)
metrics_rows: Final = MappingProxyType(
{key: agg for key, agg in daily_guardrail.items() if int(agg["requests_evaluated"]) > 0}
)
await _upsert_rows_with_retry(metrics_rows, partial(_upsert_metrics_row, prisma_client), "daily metrics", sleep)
await _upsert_rows_with_retry(
usage_unit_totals, partial(_upsert_usage_unit_row, prisma_client), "usage unit", sleep
)
except Exception as e:
verbose_proxy_logger.warning("Guardrail usage tracking failed (non-fatal): %s", e)

View file

@ -125,6 +125,15 @@ class _ProxyDBLogger(CustomLogger):
existing_metadata: Final[dict] = request_data.get("metadata", None) or {}
existing_metadata.update(_metadata)
litellm_metadata_bucket: Final = request_data.get("litellm_metadata")
if (
isinstance(litellm_metadata_bucket, dict)
and "standard_logging_guardrail_information" not in existing_metadata
):
guardrail_info: Final = litellm_metadata_bucket.get("standard_logging_guardrail_information")
if guardrail_info is not None:
existing_metadata["standard_logging_guardrail_information"] = guardrail_info
if "litellm_params" not in request_data:
request_data["litellm_params"] = {}

View file

@ -38,6 +38,8 @@ from litellm.proxy._types import (
CommonProxyErrors,
LitellmDataForBackendLLMCall,
LitellmUserRoles,
ProxyErrorTypes,
ProxyException,
SpecialHeaders,
TeamCallbackMetadata,
UserAPIKeyAuth,
@ -348,6 +350,36 @@ def reject_url_valued_destination(field: str, value: str) -> None:
)
_METADATA_JSON_TYPE_NAMES: Final[Mapping[type, str]] = MappingProxyType(
{bool: "a boolean", int: "an integer", float: "a number", str: "a string", list: "an array"}
)
def _invalid_metadata_type_error(field: str, value: object) -> ProxyException:
received_type: Final = _METADATA_JSON_TYPE_NAMES.get(type(value), f"a {type(value).__name__}")
return ProxyException(
message=f"Invalid type for '{field}': expected an object, but got {received_type} instead.",
type=ProxyErrorTypes.bad_request_error,
param=field,
code=400,
)
def _normalized_metadata_object(field: str, value: object) -> Mapping[str, Any]:
"""Return ``value`` as a metadata object or raise a 400 like OpenAI does.
A JSON string that parses to an object is accepted because multipart/form-data
and ``extra_body`` callers can only send metadata as a string. The caller pops
the raw value from the request body before validating so the failure-logging
hooks that inspect the body afterwards don't crash on it and mask the 400 as a 500.
"""
if isinstance(value, dict):
return value
if isinstance(value, str) and isinstance((parsed := safe_json_loads(value)), dict):
return parsed
raise _invalid_metadata_type_error(field=field, value=value)
def _strip_untrusted_request_header_controls(
headers: Any,
*,
@ -1572,6 +1604,13 @@ async def add_litellm_data_to_request(
continue
data.pop(_internal_key, None)
_reject_url_valued_destinations(data)
_raw_metadata_by_field: Final = {
_metadata_field: data.pop(_metadata_field)
for _metadata_field in ("metadata", "litellm_metadata")
if data.get(_metadata_field) is not None
}
for _metadata_field, _raw_metadata in _raw_metadata_by_field.items():
data[_metadata_field] = _normalized_metadata_object(_metadata_field, _raw_metadata)
# Strip spoofable auth metadata from user-supplied metadata dict
_user_metadata = data.get("metadata")
if isinstance(_user_metadata, dict):
@ -1711,29 +1750,10 @@ async def add_litellm_data_to_request(
verbose_proxy_logger.debug("receiving data: %s", data)
# Parse metadata if it's a string (e.g., from multipart/form-data)
if "metadata" in data and data["metadata"] is not None:
if isinstance(data["metadata"], str):
data["metadata"] = safe_json_loads(data["metadata"])
if not isinstance(data["metadata"], dict):
verbose_proxy_logger.warning(
"Failed to parse 'metadata' as JSON dict. Received value: %s", data["metadata"]
)
# requester_metadata is snapshotted AFTER the strip below so
# downstream consumers (e.g. PANW guardrail reading user_ip /
# profile_id) don't see attacker-injected admin slots preserved in
# the deepcopy.
# Parse litellm_metadata if it's a string (e.g., from multipart/form-data or extra_body)
if "litellm_metadata" in data and data["litellm_metadata"] is not None:
if isinstance(data["litellm_metadata"], str):
parsed_litellm_metadata: Final = safe_json_loads(data["litellm_metadata"])
if not isinstance(parsed_litellm_metadata, dict):
verbose_proxy_logger.warning(
"Failed to parse 'litellm_metadata' as JSON dict. Received value: %s", data["litellm_metadata"]
)
else:
data["litellm_metadata"] = parsed_litellm_metadata
# requester_metadata is snapshotted AFTER the strip below so
# downstream consumers (e.g. PANW guardrail reading user_ip /
# profile_id) don't see attacker-injected admin slots preserved in
# the deepcopy.
# Strip internal pipeline state and admin-injection slots from user input.
# Runs AFTER the string-to-dict parse above so JSON-string metadata (sent

View file

@ -574,6 +574,33 @@ def _slices(rows: Sequence[_AttemptAggRow]) -> tuple[ShadowEvalSlice, ...]:
)
_NO_KEY_LABELS: Final[tuple[str | None, str | None]] = (None, None)
async def _with_key_labels(
prisma_client: "PrismaClient", responses: Sequence[ShadowEvalJobResponse]
) -> tuple[ShadowEvalJobResponse, ...]:
"""Resolve each job's key hash to the key's alias and masked name in one batched read,
so the UI can say whose traffic a job shadows. Deleted keys resolve to None."""
if not responses:
return ()
key_rows: Final = await prisma_client.db.litellm_verificationtoken.find_many(
where={"token": {"in": sorted({response.api_key_id for response in responses})}} # mutable-ok: Prisma filter
)
labels: Final[Mapping[str, tuple[str | None, str | None]]] = {
row.token: (row.key_alias, row.key_name) for row in key_rows or ()
}
return tuple(
response.model_copy(
update={ # mutable-ok: pydantic update payload
"key_alias": labels.get(response.api_key_id, _NO_KEY_LABELS)[0],
"key_name": labels.get(response.api_key_id, _NO_KEY_LABELS)[1],
}
)
for response in responses
)
async def _shadow_eval_results(prisma_client: "PrismaClient", job_id: str) -> ShadowEvalResult | None:
"""Both stratifications of one job's verdicts. Tier answers "where does the router do
well"; the model stratification groups by whichever model served the real arm, so it
@ -686,7 +713,9 @@ async def start_shadow_eval(
f"Key already has an active {data.direction} shadow eval job (started concurrently). Stop it first."
),
) from e
return ShadowEvalJobResponse.model_validate(job, from_attributes=True)
return ShadowEvalJobResponse.model_validate(job, from_attributes=True).model_copy(
update={"key_alias": key_row.key_alias, "key_name": key_row.key_name} # mutable-ok: pydantic update payload
)
@router.get(
@ -711,7 +740,10 @@ async def list_shadow_eval_jobs(
order={"created_at": "desc"}, # mutable-ok: Prisma order
take=limit,
)
return tuple(ShadowEvalJobResponse.model_validate(record, from_attributes=True) for record in records or ())
return await _with_key_labels(
prisma_client,
tuple(ShadowEvalJobResponse.model_validate(record, from_attributes=True) for record in records or ()),
)
@router.get(
@ -742,7 +774,10 @@ async def get_shadow_eval_job(
where={"job_id": job_id, "outcome": "error"}, # mutable-ok: Prisma filter
order={"created_at": "desc"}, # mutable-ok: Prisma order
)
return ShadowEvalJobResponse.model_validate(record, from_attributes=True).model_copy(
labeled: Final = await _with_key_labels(
prisma_client, (ShadowEvalJobResponse.model_validate(record, from_attributes=True),)
)
return labeled[0].model_copy(
update={ # mutable-ok: pydantic update payload
"judged_count": totals[0].judged_count if totals else 0,
"error_count": totals[0].error_count if totals else 0,
@ -781,4 +816,7 @@ async def stop_shadow_eval_job(
where={"id": job_id}, # mutable-ok: Prisma filter
data={"stopped_at": datetime.now(timezone.utc)}, # mutable-ok: Prisma payload
)
return ShadowEvalJobResponse.model_validate(updated, from_attributes=True)
labeled: Final = await _with_key_labels(
prisma_client, (ShadowEvalJobResponse.model_validate(updated, from_attributes=True),)
)
return labeled[0]

View file

@ -1,12 +1,13 @@
import asyncio
import json
import os
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from typing import Any, Final
from typing import TYPE_CHECKING, Any, Final, Protocol
from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import TypeAdapter
from pydantic import BaseModel, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm._uuid import uuid
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -38,13 +39,32 @@ from litellm.types.proxy.management_endpoints.config_overrides import (
HashicorpVaultConfig,
)
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
router: Final = APIRouter()
class _ConfigOverrideRow(Protocol):
config_value: str | Mapping[str, object] | None
class _ConfigOverridesTableClient(Protocol):
async def find_unique(self, where: Mapping[str, str]) -> _ConfigOverrideRow | None: ...
async def upsert(self, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> object: ...
async def delete(self, where: Mapping[str, str]) -> object: ...
def _config_overrides_table(prisma_client: "PrismaClient") -> _ConfigOverridesTableClient:
return ConfigOverridesRepository(prisma_client).table
_AUDIT_REDACTED: Final = "***REDACTED***"
def _redact_config(config: Mapping[str, Any] | None) -> dict[str, Any]:
def _redact_config(config: Mapping[str, object] | None) -> dict[str, str]:
"""Strip values from a config snapshot before audit-log emission.
Hashicorp Vault config carries ``vault_token``, ``approle_secret_id``,
@ -68,8 +88,8 @@ def _log_audit_task_exception(task: "asyncio.Task[None]") -> None:
async def _emit_hashicorp_vault_audit_log(
*,
action: AUDIT_ACTIONS,
before_config: Mapping[str, Any] | None,
after_config: Mapping[str, Any] | None,
before_config: Mapping[str, object] | None,
after_config: Mapping[str, object] | None,
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: str | None,
) -> None:
@ -136,9 +156,9 @@ _sensitive_masker: Final = SensitiveDataMasker()
# --- Shared helpers ---
def _mask_sensitive_fields(data: dict[str, Any], sensitive_fields: set[str]) -> dict[str, Any]:
def _mask_sensitive_fields(data: Mapping[str, object], sensitive_fields: set[str]) -> dict[str, object]:
"""Mask sensitive fields for API responses. Non-sensitive fields are left as-is."""
masked: Final = {}
masked: Final[dict[str, object]] = {}
for key, value in data.items():
if value is not None and key in sensitive_fields and isinstance(value, str):
masked[key] = _sensitive_masker._mask_value(value)
@ -147,7 +167,7 @@ def _mask_sensitive_fields(data: dict[str, Any], sensitive_fields: set[str]) ->
return masked
def _get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, Any]:
def _get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, str | None]:
"""Read current env var values as fallback when no DB record exists."""
values: Final = {}
for field_name, env_var_name in env_var_mapping.items():
@ -156,7 +176,13 @@ def _get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, Any]:
return values
def _extract_field_type(field_info: dict[str, Any]) -> str:
class _JsonSchemaField(TypedDict, total=False):
type: ReadOnly[str]
anyOf: ReadOnly[Sequence["_JsonSchemaField"]]
description: ReadOnly[str]
def _extract_field_type(field_info: _JsonSchemaField) -> str:
"""Extract the non-null type from a Pydantic v2 JSON schema field."""
if "type" in field_info:
return field_info["type"]
@ -166,11 +192,12 @@ def _extract_field_type(field_info: dict[str, Any]) -> str:
return "string"
def _build_field_schema(model_class: type) -> dict[str, Any]:
def _build_field_schema(model_class: type[BaseModel]) -> dict[str, object]:
"""Build field_schema dict from a Pydantic model for UI rendering."""
schema: Final = TypeAdapter(model_class).json_schema(by_alias=True)
raw_properties: Final[Mapping[str, _JsonSchemaField]] = schema.get("properties", {})
properties: Final = {}
for field_name, field_info in schema.get("properties", {}).items():
for field_name, field_info in raw_properties.items():
properties[field_name] = {
"description": field_info.get("description", ""),
"type": _extract_field_type(field_info),
@ -181,14 +208,14 @@ def _build_field_schema(model_class: type) -> dict[str, Any]:
}
def _parse_config_value(raw: Any) -> dict[str, Any]:
def _parse_config_value(raw: str | Mapping[str, object]) -> dict[str, object]:
"""Parse a config_value from DB (may be JSON string or dict)."""
if isinstance(raw, str):
return safe_json_loads(raw, default={})
return dict(raw)
def _set_env_vars(config_data: dict[str, Any]) -> None:
def _set_env_vars(config_data: Mapping[str, object]) -> None:
"""Set HCP_VAULT_* env vars from config data. Unsets vars for missing/None/empty fields."""
for field_name, env_var_name in HASHICORP_ENV_VAR_MAPPING.items():
value = config_data.get(field_name)
@ -242,15 +269,15 @@ async def update_hashicorp_vault_config(
detail=CommonProxyErrors.db_not_connected_error.value,
)
config_data = config.model_dump(exclude_none=True)
config_data: dict[str, object] = config.model_dump(exclude_none=True)
# Merge ALL fields the user didn't send: try DB first, fall back to env vars.
# Omitted field = keep existing; empty string = clear/remove the field.
existing_record: Final = await ConfigOverridesRepository(prisma_client).table.find_unique(
existing_record: Final = await _config_overrides_table(prisma_client).find_unique(
where={"config_type": "hashicorp_vault"}
)
existing_decrypted: dict[str, Any] | None = None
env_values: dict[str, Any] = {}
existing_decrypted: dict[str, object] | None = None
env_values: dict[str, str | None] = {}
if existing_record is not None and existing_record.config_value is not None:
existing_data: Final = _parse_config_value(existing_record.config_value)
existing_decrypted = proxy_config._decrypt_db_variables(existing_data)
@ -307,7 +334,7 @@ async def update_hashicorp_vault_config(
# Only persist to DB after successful init
encrypted_data: Final = proxy_config._encrypt_env_variables(config_data)
config_value: Final = safe_dumps(encrypted_data)
await ConfigOverridesRepository(prisma_client).table.upsert(
await _config_overrides_table(prisma_client).upsert(
where={"config_type": "hashicorp_vault"},
data={
"create": {
@ -377,7 +404,7 @@ async def get_hashicorp_vault_config(
field_schema: Final = _build_field_schema(HashicorpVaultConfig)
# Try to load from DB
db_record: Final = await ConfigOverridesRepository(prisma_client).table.find_unique(
db_record: Final = await _config_overrides_table(prisma_client).find_unique(
where={"config_type": "hashicorp_vault"}
)
@ -385,7 +412,7 @@ async def get_hashicorp_vault_config(
config_data: Final = _parse_config_value(db_record.config_value)
# Decrypt then mask sensitive fields so plaintext secrets are never sent to the UI
decrypted_data: Final = proxy_config._decrypt_db_variables(config_data)
decrypted_data: Final[Mapping[str, object]] = proxy_config._decrypt_db_variables(config_data)
masked_data: Final = _mask_sensitive_fields(decrypted_data, HASHICORP_SENSITIVE_FIELDS)
return ConfigOverrideSettingsResponse(
@ -434,10 +461,10 @@ async def delete_hashicorp_vault_config(
# Capture the prior config before delete so the audit-log row can
# show *what* was removed (keys only — values get redacted).
existing_record: Final = await ConfigOverridesRepository(prisma_client).table.find_unique(
existing_record: Final = await _config_overrides_table(prisma_client).find_unique(
where={"config_type": "hashicorp_vault"}
)
before_config: dict[str, Any] | None = None
before_config: dict[str, object] | None = None
if existing_record is not None and existing_record.config_value is not None:
try:
before_config = proxy_config._decrypt_db_variables(_parse_config_value(existing_record.config_value))
@ -447,7 +474,7 @@ async def delete_hashicorp_vault_config(
# Delete DB record if it exists — ignore if not found
deleted = False
try:
await ConfigOverridesRepository(prisma_client).table.delete(where={"config_type": "hashicorp_vault"})
await _config_overrides_table(prisma_client).delete(where={"config_type": "hashicorp_vault"})
deleted = True
except RecordNotFoundError:
verbose_proxy_logger.debug("No existing Hashicorp Vault config record to delete")
@ -502,7 +529,7 @@ async def test_hashicorp_vault_connection(
# Step 1: Authenticate (exercises AppRole login, TLS cert login, or direct token)
try:
headers: Final = await asyncio.to_thread(client._get_request_headers)
headers: Final[dict[str, str]] = await asyncio.to_thread(client._get_request_headers)
except Exception as e:
raise HTTPException(
status_code=502,

View file

@ -29,6 +29,10 @@ from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.user_api_key_cache import (
end_user_cache_key,
end_user_restricted_registry_cache_key,
)
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
from litellm.proxy.management_endpoints.common_utils import validate_budget_duration
from litellm.proxy.management_helpers.object_permission_utils import (
@ -99,6 +103,25 @@ def _typed_table(repo: EndUserRepository | BudgetRepository) -> object:
router: Final = APIRouter()
async def _evict_end_user_cache_keys(cache_keys: Sequence[str]) -> None:
"""
Every endpoint that mutates an end-user row must call this, or a newly blocked or budgeted
customer keeps being served unrestricted until the TTL expires: auth reads end users
cache-first, and the cached restricted-id registry decides whether the row is read at all.
"""
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
evict_and_broadcast,
)
from litellm.proxy.proxy_server import user_api_key_cache
await evict_and_broadcast(cache_keys=cache_keys, user_api_key_cache=user_api_key_cache)
def _end_user_cache_keys(user_ids: Sequence[str]) -> tuple[str, ...]:
"""The per-id entries plus the registry, which any restriction change can move ids in or out of."""
return (*(end_user_cache_key(user_id) for user_id in user_ids), end_user_restricted_registry_cache_key())
def _to_customer_response(record: BaseModel) -> CustomerResponse:
"""Validate a raw end-user DB row into the typed customer response.
@ -152,6 +175,7 @@ async def block_user(data: BlockUsers):
},
)
records.append(record)
await _evict_end_user_cache_keys(_end_user_cache_keys(data.user_ids))
else:
raise HTTPException(
status_code=500,
@ -448,6 +472,8 @@ async def new_end_user(
include={"litellm_budget_table": True, "object_permission": True},
)
await _evict_end_user_cache_keys(_end_user_cache_keys((data.user_id,)))
return _to_customer_response(end_user_record)
except Exception as e:
verbose_proxy_logger.exception(
@ -691,6 +717,8 @@ async def update_end_user(
raise ValueError(f"Failed updating customer data. User ID does not exist passed user_id={data.user_id}")
verbose_proxy_logger.debug("received response from updating prisma client. response=%s", response)
await _evict_end_user_cache_keys(_end_user_cache_keys((data.user_id,)))
return _to_customer_response(response)
else:
raise ValueError(f"user_id is required, passed user_id = {data.user_id}")
@ -764,6 +792,9 @@ async def delete_end_user(
where={"user_id": {"in": data.user_ids}}
)
verbose_proxy_logger.debug("received response from updating prisma client. response=%s", response)
await _evict_end_user_cache_keys(_end_user_cache_keys(data.user_ids))
return DeleteCustomersResponse(
deleted_customers=response,
message="Successfully deleted customers with ids: " + str(data.user_ids),

View file

@ -17,7 +17,7 @@ import json
import traceback
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from typing import Any, Final, Literal, cast
from typing import Any, Final, Literal, Protocol, cast
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -88,6 +88,7 @@ if TYPE_CHECKING:
from prisma.actions import (
LiteLLM_InvitationLinkActions,
LiteLLM_OrganizationMembershipActions,
LiteLLM_OrganizationTableActions,
LiteLLM_TeamMembershipActions,
LiteLLM_TeamTableActions,
LiteLLM_UserTableActions,
@ -142,6 +143,15 @@ def _invitation_link_table(
return invitation_table
def _organization_table(
prisma_client: "PrismaClient | None",
) -> "LiteLLM_OrganizationTableActions[prisma_models.LiteLLM_OrganizationTable]":
organization_table: Final[LiteLLM_OrganizationTableActions[prisma_models.LiteLLM_OrganizationTable]] = (
OrganizationRepository(prisma_client).table
)
return organization_table
def _team_membership_table(
prisma_client: "PrismaClient | None",
) -> "LiteLLM_TeamMembershipActions[prisma_models.LiteLLM_TeamMembership]":
@ -234,7 +244,7 @@ async def _check_duplicate_user_field(
if case_insensitive:
where_clause[field_name]["mode"] = "insensitive"
existing_user: Final = await UserRepository(prisma_client).table.find_first(where=where_clause)
existing_user: Final[object] = await UserRepository(prisma_client).table.find_first(where=where_clause)
if existing_user is not None:
existing_value: Final = getattr(existing_user, field_name, value)
@ -737,11 +747,11 @@ async def _get_user_info_teams(
user_id: str | None,
user_info: Any | None,
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[list[Any], list[Any] | None]:
) -> tuple[list[TeamListResponseObject], list[TeamListResponseObject] | None]:
"""Fetch and merge teams from membership + user.teams field."""
from litellm.proxy.management_endpoints.team_endpoints import list_team
team_list: list[Any] = []
team_list: list[TeamListResponseObject] = []
team_id_list: list[str] = []
teams_1: Final = await list_team(
@ -756,7 +766,7 @@ async def _get_user_info_teams(
team_list = teams_1
team_id_list = [team.team_id for team in teams_1]
teams_2: list[Any] | None = None
teams_2: list[TeamListResponseObject] | None = None
target_team_ids: Final = getattr(user_info, "teams", None)
if target_team_ids and isinstance(target_team_ids, list):
@ -766,7 +776,7 @@ async def _get_user_info_teams(
query_type="find_all",
)
elif user_api_key_dict.user_id is not None and user_id is None:
caller_user_info: Final = await prisma_client.get_data(user_id=user_api_key_dict.user_id)
caller_user_info: Final[object] = await prisma_client.get_data(user_id=user_api_key_dict.user_id)
caller_team_ids: Final = getattr(caller_user_info, "teams", None)
if caller_team_ids:
teams_2 = await prisma_client.get_data(
@ -805,8 +815,8 @@ def _build_user_info_response(
user_id: str | None,
user_info: Any | None,
keys: list[LiteLLM_VerificationToken] | None,
team_list: list[Any],
teams_1: list[Any] | None,
team_list: list[TeamListResponseObject],
teams_1: list[TeamListResponseObject] | None,
) -> UserInfoResponse:
"""Create UserInfoResponse while filtering sensitive fields."""
if user_info is None and keys is not None:
@ -1085,7 +1095,7 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth):
verbose_proxy_logger.debug("results_keys: %s", results)
_keys_in_db: Final[list] = results[0]["keys"] or []
_keys_in_db: Final[Sequence[dict[str, object]]] = results[0]["keys"] or []
# cast all keys to LiteLLM_VerificationToken
keys_in_db: Final = []
for key in _keys_in_db:
@ -1094,7 +1104,7 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth):
keys_in_db.append(LiteLLM_VerificationToken.model_validate(key))
# cast all teams to LiteLLM_TeamTable
_teams_in_db: list = results[0]["teams"] or []
_teams_in_db: list[LiteLLM_TeamTable] = results[0]["teams"] or []
_teams_in_db = [LiteLLM_TeamTable.model_validate(team) for team in _teams_in_db]
_teams_in_db.sort(key=lambda x: getattr(x, "team_alias", "") or "")
returned_keys: Final = _process_keys_for_user_info(keys=keys_in_db, all_teams=_teams_in_db)
@ -1885,7 +1895,7 @@ async def get_user_key_counts(
# Get count for each user_id individually
for user_id in user_ids:
count = await VerificationTokenRepository(prisma_client).table.count(
count = await _verification_token_table(prisma_client).count(
where={
"user_id": user_id,
"OR": [
@ -2166,6 +2176,13 @@ async def get_users(
}
class _DeleteTeamRow(Protocol):
team_id: str
members_with_roles: object
def model_dump(self) -> Mapping[str, object]: ...
@router.post(
"/user/delete",
tags=["Internal User management"],
@ -2308,7 +2325,9 @@ async def delete_user(
)
## CLEANUP MEMBERS_WITH_ROLES
fetch_all_teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_row.teams}})
fetch_all_teams: Sequence[_DeleteTeamRow] = await TeamRepository(prisma_client).table.find_many(
where={"team_id": {"in": user_row.teams}}
)
teams_to_update = []
for team in fetch_all_teams:
removed_team_members, new_team_members = _cleanup_members_with_roles(
@ -2363,7 +2382,7 @@ async def add_internal_user_to_organization(
user_id: str,
organization_id: str,
user_role: LitellmUserRoles,
):
) -> "prisma_models.LiteLLM_OrganizationMembership":
"""
Helper function to add an internal user to an organization
@ -2382,14 +2401,16 @@ async def add_internal_user_to_organization(
try:
# Check if organization_id exists
organization_row: Final = await OrganizationRepository(prisma_client).table.find_unique(
organization_row: Final = await _organization_table(prisma_client).find_unique(
where={"organization_id": organization_id}
)
if organization_row is None:
raise Exception(f"Organization not found, passed organization_id={organization_id}")
# Create a new organization membership entry
new_membership: Final = await OrganizationMembershipRepository(prisma_client).table.create(
new_membership: Final[prisma_models.LiteLLM_OrganizationMembership] = await OrganizationMembershipRepository(
prisma_client
).table.create(
data={
"user_id": user_id,
"organization_id": organization_id,

View file

@ -998,7 +998,7 @@ async def _common_key_generation_helper(
)
new_budget: Final = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
_budget: Final = await BudgetRepository(prisma_client).table.create(
_budget: Final[LiteLLM_BudgetTable] = await BudgetRepository(prisma_client).table.create(
data={
**new_budget,
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
@ -4755,7 +4755,9 @@ async def _execute_virtual_key_regeneration(
grace_period=data.grace_period if data else None,
)
updated_token: Final[Mapping[str, object] | None] = await VerificationTokenRepository(prisma_client).table.update(
updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table(
VerificationTokenRepository(prisma_client)
).update(
where={"token": hashed_api_key},
data=with_settings_updated_at(jsonified_update_data),
)
@ -5307,7 +5309,9 @@ async def validate_key_list_check(
if key_hash:
try:
key_info: Final = await VerificationTokenRepository(prisma_client).table.find_unique(
key_info: Final[LiteLLM_VerificationToken] = await VerificationTokenRepository(
prisma_client
).table.find_unique(
where={"token": key_hash},
)
except Exception:

View file

@ -1,7 +1,7 @@
"""`/management/v1/spend_logs` facets."""
from datetime import datetime, timezone
from typing import Annotated, Any, Final
from typing import Annotated, Any, Final, Literal
from fastapi import APIRouter, Depends, Query, Request
@ -35,7 +35,7 @@ def _as_utc(value: datetime) -> datetime:
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc)
async def _end_user_scope_clause(
async def _spend_log_scope_clause(
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
next_param_index: int,
@ -43,8 +43,8 @@ async def _end_user_scope_clause(
"""SQL predicate restricting the facet to spend logs this caller may read.
Returns ``(None, ())`` for a proxy admin. Mirrors the scoping ``/spend/logs/ui``
applies, so the dropdown can never offer an end user whose rows the caller
could not open.
applies, so a dropdown can never offer a value from a row the caller could
not open.
"""
from litellm.proxy.spend_tracking.spend_management_endpoints import (
_get_permitted_team_ids_for_spend_logs,
@ -77,6 +77,98 @@ async def _end_user_scope_clause(
return f"({' OR '.join(clauses)})", params
async def _list_spend_log_facet(
request: Request,
user_api_key_dict: UserAPIKeyAuth,
start_time: datetime,
end_time: datetime,
q: str | None,
page: int,
page_size: int,
column: Literal["end_user", "user"],
) -> FacetListResponse:
try:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise ManagementProblem(
ProblemDetail(
type=f"{PROBLEM_TYPE_BASE}database-not-connected",
title="Database not connected",
status=503,
detail=CommonProxyErrors.db_not_connected_error.value,
)
)
column_sql: Final = "end_user" if column == "end_user" else '"user"'
window_params: Final[tuple[Any, ...]] = (_as_utc(start_time), _as_utc(end_time))
search_params: Final[tuple[Any, ...]] = (f"%{escape_like(q)}%",) if q else ()
search_clause: Final = (f"{column_sql} ILIKE ${len(window_params) + 1} ESCAPE '\\'",) if q else ()
scope_clause, scope_params = await _spend_log_scope_clause(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
next_param_index=len(window_params) + len(search_params) + 1,
)
where_parts: Final = (
(
"\"startTime\" >= ($1::timestamptz AT TIME ZONE 'UTC')",
"\"startTime\" <= ($2::timestamptz AT TIME ZONE 'UTC')",
f"{column_sql} IS NOT NULL",
f"{column_sql} != ''",
)
+ search_clause
+ ((scope_clause,) if scope_clause is not None else ())
)
# The inner LIMIT walks the startTime index newest first and bounds the
# rows DISTINCT can inspect. request_id makes the cut-off deterministic,
# and page_size + 1 reveals has_more without a COUNT(*).
params: Final = (
window_params
+ search_params
+ scope_params
+ (SPEND_LOGS_FACET_SCAN_CAP, page_size + 1, (page - 1) * page_size)
)
scan_idx: Final = len(params) - 2
facet_sql: Final = (
f"SELECT DISTINCT {column_sql} FROM ("
f" SELECT {column_sql}"
f' FROM "LiteLLM_SpendLogs"'
f" WHERE {' AND '.join(where_parts)}"
f' ORDER BY "startTime" DESC, request_id DESC'
f" LIMIT ${scan_idx}"
f") recent"
f" ORDER BY {column_sql} ASC"
f" LIMIT ${scan_idx + 1} OFFSET ${scan_idx + 2}"
)
rows: Final = await prisma_client.db.query_raw(facet_sql, *params)
values: Final[list[str]] = [row[column] for row in rows if row.get(column)]
has_more: Final = len(values) > page_size
return FacetListResponse(
data=values[:page_size],
meta=PageMeta(page=page, page_size=page_size, has_more=has_more),
links=build_page_links(request=request, page=page, has_more=has_more),
)
except ManagementProblem:
raise
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.management_endpoints.management_v1.spend_logs._list_spend_log_facet(): Exception occured - %s",
e,
)
raise ManagementProblem(
ProblemDetail(
type=f"{PROBLEM_TYPE_BASE}internal-server-error",
title="Internal server error",
status=500,
detail=f"Failed to list spend log {column.replace('_', ' ')}s.",
)
)
@router.get(
"/spend_logs/end_users",
tags=["Budget & Spend Tracking"],
@ -116,85 +208,47 @@ async def list_spend_log_end_users(
--header 'Authorization: Bearer sk-1234'
```
"""
try:
from litellm.proxy.proxy_server import prisma_client
return await _list_spend_log_facet(
request=request,
user_api_key_dict=user_api_key_dict,
start_time=start_time,
end_time=end_time,
q=q,
page=page,
page_size=page_size,
column="end_user",
)
if prisma_client is None:
raise ManagementProblem(
ProblemDetail(
type=f"{PROBLEM_TYPE_BASE}database-not-connected",
title="Database not connected",
status=503,
detail=CommonProxyErrors.db_not_connected_error.value,
)
)
window_params: Final[tuple[Any, ...]] = (_as_utc(start_time), _as_utc(end_time))
search_params: Final[tuple[Any, ...]] = (f"%{escape_like(q)}%",) if q else ()
search_clause: Final = (f"end_user ILIKE ${len(window_params) + 1} ESCAPE '\\'",) if q else ()
scope_clause, scope_params = await _end_user_scope_clause(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
next_param_index=len(window_params) + len(search_params) + 1,
)
where_parts: Final = (
(
"\"startTime\" >= ($1::timestamptz AT TIME ZONE 'UTC')",
"\"startTime\" <= ($2::timestamptz AT TIME ZONE 'UTC')",
"end_user IS NOT NULL",
"end_user != ''",
)
+ search_clause
+ ((scope_clause,) if scope_clause is not None else ())
)
# The inner LIMIT is the safety bound: it walks the startTime index newest
# first and stops, so DISTINCT never runs over an unbounded row set.
# request_id breaks startTime ties so the cut-off row is deterministic and
# successive OFFSET pages agree on the set they are paging through.
# page_size + 1: one row beyond the page reveals has_more without a COUNT(*).
params: Final = (
window_params
+ search_params
+ scope_params
+ (SPEND_LOGS_FACET_SCAN_CAP, page_size + 1, (page - 1) * page_size)
)
scan_idx: Final = len(params) - 2
facet_sql: Final = (
f"SELECT DISTINCT end_user FROM ("
f" SELECT end_user"
f' FROM "LiteLLM_SpendLogs"'
f" WHERE {' AND '.join(where_parts)}"
f' ORDER BY "startTime" DESC, request_id DESC'
f" LIMIT ${scan_idx}"
f") recent"
f" ORDER BY end_user ASC"
f" LIMIT ${scan_idx + 1} OFFSET ${scan_idx + 2}"
)
rows: Final = await prisma_client.db.query_raw(facet_sql, *params)
end_users: Final[list[str]] = [row["end_user"] for row in rows if row.get("end_user")]
has_more: Final = len(end_users) > page_size
return FacetListResponse(
data=end_users[:page_size],
meta=PageMeta(page=page, page_size=page_size, has_more=has_more),
links=build_page_links(request=request, page=page, has_more=has_more),
)
except ManagementProblem:
raise
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.management_endpoints.management_v1.spend_logs.list_spend_log_end_users(): Exception occured - %s",
e,
)
raise ManagementProblem(
ProblemDetail(
type=f"{PROBLEM_TYPE_BASE}internal-server-error",
title="Internal server error",
status=500,
detail="Failed to list spend log end users.",
)
)
@router.get(
"/spend_logs/users",
tags=["Budget & Spend Tracking"],
dependencies=[Depends(user_api_key_auth), Depends(reject_unknown_query_params)],
response_model=FacetListResponse,
)
async def list_spend_log_users(
request: Request,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
start_time: Annotated[
datetime,
Query(alias="filter[startTime][gte]", description="Window start (UTC when no offset is given)"),
],
end_time: Annotated[
datetime,
Query(alias="filter[startTime][lte]", description="Window end (UTC when no offset is given)"),
],
q: Annotated[str | None, Query(description="Case-insensitive partial match on the internal user id")] = None,
page: Annotated[int, Query(ge=1, description="Page number")] = 1,
page_size: Annotated[int, Query(ge=1, le=100, description="Page size")] = 50,
) -> FacetListResponse:
"""The distinct internal users appearing in spend logs the caller can read."""
return await _list_spend_log_facet(
request=request,
user_api_key_dict=user_api_key_dict,
start_time=start_time,
end_time=end_time,
q=q,
page=page,
page_size=page_size,
column="user",
)

View file

@ -19,10 +19,10 @@ import functools
import importlib
import json
import os
from collections.abc import Iterable
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Final, Literal, Protocol
from fastapi import (
APIRouter,
@ -36,6 +36,7 @@ from fastapi import (
status,
)
from fastapi.responses import JSONResponse
from typing_extensions import ReadOnly, TypedDict
try:
from prisma.errors import RecordNotFoundError, UniqueViolationError
@ -77,7 +78,11 @@ TEMPORARY_MCP_SERVER_TTL_SECONDS: Final = 300
TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX: Final = "litellm:mcp:temporary_server"
def does_mcp_server_exist(mcp_server_records: Iterable[Any], mcp_server_id: str) -> bool:
class _HasServerId(Protocol):
server_id: str
def does_mcp_server_exist(mcp_server_records: Iterable[_HasServerId], mcp_server_id: str) -> bool:
"""
Check if the mcp server with the given id exists in the iterable of mcp servers.
@ -93,6 +98,8 @@ def does_mcp_server_exist(mcp_server_records: Iterable[Any], mcp_server_id: str)
DEFAULT_MCP_REGISTRY_VERSION: Final = "1.0.0"
if TYPE_CHECKING:
from prisma import models as prisma_models
from litellm.proxy.utils import PrismaClient
try:
@ -111,7 +118,7 @@ if MCP_AVAILABLE:
class _ToolNameValidationResult(BaseModel):
is_valid: bool = True
warnings: list = []
warnings: list[str] = []
def validate_tool_name(name: str) -> _ToolNameValidationResult:
return _ToolNameValidationResult()
@ -263,7 +270,7 @@ if MCP_AVAILABLE:
_VALID_MCP_REQUIRED_FIELDS: Final[frozenset] = frozenset(NewMCPServerRequest.model_fields)
def _validate_mcp_required_fields(payload: Any) -> None:
def _validate_mcp_required_fields(payload: NewMCPServerRequest) -> None:
"""Validate submission payload against admin-configured mcp_required_fields."""
from litellm.proxy.proxy_server import (
general_settings as proxy_general_settings,
@ -329,7 +336,18 @@ if MCP_AVAILABLE:
return server.server_name
return server.server_id
def _build_mcp_registry_entry_for_server(server: MCPServer, base_url: str) -> dict[str, Any]:
class _McpRegistryRemote(TypedDict):
type: ReadOnly[str]
url: ReadOnly[str]
class _McpRegistryEntry(TypedDict):
name: ReadOnly[str]
title: ReadOnly[str]
description: ReadOnly[str]
version: ReadOnly[str]
remotes: ReadOnly[Sequence[_McpRegistryRemote]]
def _build_mcp_registry_entry_for_server(server: MCPServer, base_url: str) -> _McpRegistryEntry:
server_name: Final = _build_mcp_registry_server_name(server)
title: Final = server_name
description: Final = server_name
@ -353,7 +371,7 @@ if MCP_AVAILABLE:
],
}
def _build_builtin_registry_entry(base_url: str) -> dict[str, Any]:
def _build_builtin_registry_entry(base_url: str) -> _McpRegistryEntry:
remote_url: Final = _build_registry_remote_url(base_url, "/mcp")
return {
"name": LITELLM_MCP_SERVER_NAME,
@ -400,7 +418,7 @@ if MCP_AVAILABLE:
if cache_backend is None or not hasattr(cache_backend, "async_set_cache"):
return
payload: Final[dict[str, Any]] = server.model_dump(mode="json")
payload: Final[dict[str, object]] = server.model_dump(mode="json")
payload_json: Final = json.dumps(payload)
try:
encrypted_payload: Final = encrypt_value_helper(payload_json)
@ -464,7 +482,7 @@ if MCP_AVAILABLE:
return None
if not isinstance(loaded, dict):
return None
payload_dict: Final[dict[str, Any]] = loaded
payload_dict: Final[dict[str, object]] = loaded
try:
return MCPServer.model_validate(payload_dict)
@ -725,7 +743,7 @@ if MCP_AVAILABLE:
one, so a form that round-trips it must not read as "credentials supplied"."""
if not credentials:
return False
as_dict: Final[dict[str, Any]] = dict(credentials)
as_dict: Final[dict[str, object]] = dict(credentials)
return any(value for key, value in as_dict.items() if key not in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS)
def _inherit_credentials_from_existing_server(
@ -738,7 +756,7 @@ if MCP_AVAILABLE:
if existing_server is None:
return payload
inherited_credentials: dict[str, Any] = {
inherited_credentials: dict[str, object] = {
credential_key: value
for server_attr, credential_key in _INHERITED_CREDENTIAL_FIELDS
if (value := getattr(existing_server, server_attr, None))
@ -755,7 +773,7 @@ if MCP_AVAILABLE:
except AttributeError:
pass
payload_dict: dict[str, Any]
payload_dict: dict[str, object]
try:
payload_dict = payload.model_dump()
except AttributeError:
@ -888,7 +906,9 @@ if MCP_AVAILABLE:
# Get from DB
if prisma_client is not None:
try:
mcp_servers: Final = await MCPServerRepository(prisma_client).table.find_many()
mcp_servers: Final[Sequence[prisma_models.LiteLLM_MCPServerTable]] = await MCPServerRepository(
prisma_client
).table.find_many()
for server in mcp_servers:
if hasattr(server, "mcp_access_groups") and server.mcp_access_groups:
access_groups.update(server.mcp_access_groups)
@ -930,7 +950,7 @@ if MCP_AVAILABLE:
verbose_proxy_logger.debug("MCP registry request from IP=%s", client_ip)
base_url: Final = get_request_base_url(request)
registry_servers: Final[list[dict[str, Any]]] = []
registry_servers: Final[list[dict[str, _McpRegistryEntry]]] = []
registry_servers.append({"server": _build_builtin_registry_entry(base_url)})
# Centralized IP-based filtering: external callers only see public servers
@ -1126,7 +1146,9 @@ if MCP_AVAILABLE:
if user_id and _byok_prisma_client is not None:
byok_server_ids: Final = [s.server_id for s in redacted_mcp_servers if getattr(s, "is_byok", False)]
if byok_server_ids:
cred_rows: Final = await MCPUserCredentialsRepository(_byok_prisma_client).table.find_many(
cred_rows: Final[
Sequence[prisma_models.LiteLLM_MCPUserCredentials]
] = await MCPUserCredentialsRepository(_byok_prisma_client).table.find_many(
where={"user_id": user_id, "server_id": {"in": byok_server_ids}}
)
cred_set: Final = {r.server_id for r in cred_rows}
@ -1680,7 +1702,7 @@ if MCP_AVAILABLE:
options={"verify_exp": False, "verify_aud": False},
)
if decoded.get("login_method") in ("sso", "username_password"):
cookie_key: Final = decoded.get("key", "")
cookie_key: Final[str] = decoded.get("key", "")
if cookie_key:
api_key = f"Bearer {cookie_key}"
except _jwt.InvalidTokenError:
@ -1707,7 +1729,7 @@ if MCP_AVAILABLE:
get_request_route,
)
server_id: Final = request.path_params.get("server_id", "")
server_id: Final[str] = request.path_params.get("server_id", "")
if server_id:
_s = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if not _s:
@ -2324,7 +2346,7 @@ if MCP_AVAILABLE:
required: Final[list[MCPUserEnvVarSpec]] = []
missing_count = 0
for spec in user_specs:
name = spec["name"]
name: str = spec["name"]
if name not in blocking:
continue
value = stored_values.get(name)
@ -2672,16 +2694,16 @@ if MCP_AVAILABLE:
"mcp_registry.json",
)
_mcp_registry_cache: dict[str, Any] | None = None
_mcp_registry_cache: Mapping[str, Sequence[Mapping[str, str]]] | None = None
def _load_mcp_registry() -> dict[str, Any]:
def _load_mcp_registry() -> Mapping[str, Sequence[Mapping[str, str]]]:
"""Load the curated MCP registry from disk. Cached after first read."""
global _mcp_registry_cache
if _mcp_registry_cache is not None:
return _mcp_registry_cache
try:
with open(_MCP_REGISTRY_PATH, "r") as f:
data: dict[str, Any] = json.load(f)
data: Mapping[str, Sequence[Mapping[str, str]]] = json.load(f)
except Exception as e:
verbose_proxy_logger.warning("Failed to load MCP registry from %s: %s", _MCP_REGISTRY_PATH, e)
data = {"servers": []}
@ -2747,9 +2769,9 @@ if MCP_AVAILABLE:
)
@functools.lru_cache(maxsize=1)
def _load_openapi_registry() -> dict[str, Any]:
def _load_openapi_registry() -> dict[str, object]:
with open(_OPENAPI_REGISTRY_PATH, "r") as f:
data: Final[dict[str, Any]] = json.load(f)
data: Final[dict[str, object]] = json.load(f)
return data
@router.get(

View file

@ -7,7 +7,7 @@ Endpoints here:
import json
from collections.abc import Mapping, Sequence
from typing import Any, Final
from typing import TYPE_CHECKING, Any, Final, Protocol
from fastapi import APIRouter, Depends, HTTPException
@ -33,10 +33,31 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import
UpdateModelGroupRequest,
)
if TYPE_CHECKING:
from litellm import Router
router: Final = APIRouter()
def validate_models_exist(model_names: list[str], llm_router) -> tuple[bool, list[str]]:
class _DeploymentRow(Protocol):
model_id: str
model_name: str
model_info: object
class _ModelTableClient(Protocol):
async def find_many(self, where: Mapping[str, object] | None = None) -> Sequence[_DeploymentRow]: ...
async def find_unique(self, where: Mapping[str, object]) -> _DeploymentRow | None: ...
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ...
def _model_table(prisma_client: PrismaClient) -> _ModelTableClient:
return ModelRepository(prisma_client).table
def validate_models_exist(model_names: list[str], llm_router: "Router | None") -> tuple[bool, list[str]]:
"""
Validate that all requested model names exist in the router.
Checks only exact model name matches.
@ -117,7 +138,7 @@ async def _tag_deployment_with_access_group(
)
if not was_modified:
return None
await ModelRepository(prisma_client).table.update(
await _model_table(prisma_client).update(
where={"model_id": model_id},
data={"model_info": json.dumps(updated_model_info)},
)
@ -150,7 +171,7 @@ async def _strip_access_group_from_deployment(
)
if not was_modified:
return None
await ModelRepository(prisma_client).table.update(
await _model_table(prisma_client).update(
where={"model_id": model_id},
data={"model_info": json.dumps(updated_model_info)},
)
@ -174,7 +195,7 @@ async def update_deployments_with_access_group(
The (model_id, updated model_info) pair of every deployment actually written,
so callers can verify each one survived the post-write reload
"""
deployments: Final = await ModelRepository(prisma_client).table.find_many(where={"model_name": {"in": model_names}})
deployments: Final = await _model_table(prisma_client).find_many(where={"model_name": {"in": model_names}})
verbose_proxy_logger.debug("Found %s deployments for model_names: %s", len(deployments), model_names)
found_names: Final = {deployment.model_name for deployment in deployments}
@ -225,8 +246,8 @@ async def update_specific_deployments_with_access_group(
return tuple(pair for pair in tagged if pair is not None)
async def _find_deployment_or_400(model_id: str, prisma_client: PrismaClient) -> Mapping[str, object] | None:
deployment: Final = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id})
async def _find_deployment_or_400(model_id: str, prisma_client: PrismaClient) -> object:
deployment: Final = await _model_table(prisma_client).find_unique(where={"model_id": model_id})
if deployment is None:
raise HTTPException(
status_code=400,
@ -646,7 +667,7 @@ async def update_access_group(
try:
# Step 1: Remove access group from ALL DB deployments (skip config models)
all_deployments: Final = await ModelRepository(prisma_client).table.find_many()
all_deployments: Final = await _model_table(prisma_client).find_many()
stripped: Final = [
await _strip_access_group_from_deployment(
@ -764,7 +785,7 @@ async def delete_access_group(
try:
# Remove access group from all DB deployments (skip config models)
all_deployments: Final = await ModelRepository(prisma_client).table.find_many()
all_deployments: Final = await _model_table(prisma_client).find_many()
removed: Final = [
await _strip_access_group_from_deployment(

View file

@ -21,6 +21,10 @@ from fastapi import APIRouter, Depends, HTTPException, Query
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth, user_api_key_has_admin_view
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.user_api_key_cache import (
tag_cache_key,
tag_registry_cache_key,
)
from litellm.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
get_daily_activity,
@ -133,6 +137,20 @@ def _table(
return prisma_table
async def _evict_tag_cache_keys(cache_keys: Sequence[str]) -> None:
"""
Every endpoint that mutates a tag row must call this, or a deleted tag keeps its budget
enforced and a newly created one stays invisible to the cached name registry until the TTL
expires: auth reads tags cache-first, with no freshness check.
"""
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
evict_and_broadcast,
)
from litellm.proxy.proxy_server import user_api_key_cache
await evict_and_broadcast(cache_keys=cache_keys, user_api_key_cache=user_api_key_cache)
async def _get_internal_user_api_keys(
prisma_client: "PrismaClient",
user_api_key_dict: UserAPIKeyAuth,
@ -294,6 +312,8 @@ async def new_tag(
}
)
await _evict_tag_cache_keys((tag_cache_key(tag.name), tag_registry_cache_key()))
# Update models with new tag
if tag.models:
tasks: Final = []
@ -440,6 +460,8 @@ async def update_tag(
data=update_data,
)
await _evict_tag_cache_keys((tag_cache_key(tag.name),))
# Build response
tag_config: Final = TagConfig(
name=updated_tag_record.tag_name,
@ -689,6 +711,8 @@ async def delete_tag(
# Delete tag from database
await _table(TagRepository(prisma_client)).delete(where={"tag_name": data.name})
await _evict_tag_cache_keys((tag_cache_key(data.name), tag_registry_cache_key()))
return {"message": f"Tag {data.name} deleted successfully"}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))

View file

@ -479,7 +479,7 @@ def _is_safe_cli_sso_metadata_dest_key(dest_key: str) -> bool:
return not any(fragment in lowered for fragment in _CLI_SSO_SECRET_KEY_FRAGMENTS)
def _is_safe_cli_sso_scalar_claim_value(value: Any) -> bool:
def _is_safe_cli_sso_scalar_claim_value(value: object) -> bool:
if not isinstance(value, _CLI_SSO_SCALAR_TYPES):
return False
if isinstance(value, str):
@ -490,17 +490,17 @@ def _is_safe_cli_sso_scalar_claim_value(value: Any) -> bool:
return True
def _sso_result_to_dict(result: CustomOpenID | OpenID | dict) -> dict[str, Any]:
def _sso_result_to_dict(result: CustomOpenID | OpenID | dict[str, object]) -> dict[str, object]:
if isinstance(result, dict):
return result
if hasattr(result, "model_dump"):
dumped: Final = result.model_dump()
if isinstance(dumped, dict):
return cast(dict[str, Any], dumped)
return dumped
return {}
def _get_nested_claim_value(data: dict[str, Any], claim_path: str) -> Any:
def _get_nested_claim_value(data: Mapping[str, object], claim_path: str) -> object:
"""Resolve a dot-notation claim path against an SSO result dict.
Unlike ``get_nested_value``, this does not strip a leading ``metadata.``
@ -514,7 +514,7 @@ def _get_nested_claim_value(data: dict[str, Any], claim_path: str) -> Any:
placeholder: Final = "\x00"
parts = claim_path.replace("\\.", placeholder).split(".")
parts = [p.replace(placeholder, ".") for p in parts]
current: Any = data
current: object = data
for part in parts:
if isinstance(current, dict) and part in current:
current = current[part]
@ -523,7 +523,7 @@ def _get_nested_claim_value(data: dict[str, Any], claim_path: str) -> Any:
return current
def _extract_sso_claim_value(result: CustomOpenID | OpenID | dict, claim_path: str) -> Any:
def _extract_sso_claim_value(result: CustomOpenID | OpenID | dict[str, object], claim_path: str) -> object:
extra_fields: Final = getattr(result, "extra_fields", None)
if isinstance(extra_fields, dict):
if claim_path in extra_fields:
@ -539,7 +539,7 @@ def _extract_sso_claim_value(result: CustomOpenID | OpenID | dict, claim_path: s
return _get_nested_claim_value(result_dict, claim_path)
def _set_nested_metadata_value(metadata: dict[str, Any], key_path: str, value: Any) -> None:
def _set_nested_metadata_value(metadata: dict[str, object], key_path: str, value: object) -> None:
placeholder: Final = "\x00"
parts = key_path.replace("\\.", placeholder).split(".")
parts = [p.replace(placeholder, ".") for p in parts]
@ -554,24 +554,25 @@ def _set_nested_metadata_value(metadata: dict[str, Any], key_path: str, value: A
def _flatten_cli_sso_metadata_for_poll(
metadata: dict[str, Any],
metadata: Mapping[str, object],
) -> dict[str, str | int | float | bool]:
"""Expose scalar attribution metadata as a flat dict for CLI poll responses."""
flattened: Final[dict[str, str | int | float | bool]] = {}
stack: Final[list[tuple[str, Any]]] = [("", metadata)]
stack: Final[list[tuple[str, object]]] = [("", metadata)]
while stack:
prefix, value = stack.pop()
if isinstance(value, dict):
for key, nested in value.items():
nested_items: Mapping[str, object] = value
for key, nested in nested_items.items():
nested_prefix = f"{prefix}.{key}" if prefix else key
stack.append((nested_prefix, nested))
elif _is_safe_cli_sso_scalar_claim_value(value):
elif isinstance(value, (str, int, float, bool)) and _is_safe_cli_sso_scalar_claim_value(value):
flattened[prefix] = value
return flattened
def build_cli_sso_attribution_metadata(
result: CustomOpenID | OpenID | dict,
result: CustomOpenID | OpenID | dict[str, object],
) -> dict[str, object]:
"""
Build allowlisted, non-secret scalar attribution metadata from an SSO result.
@ -599,8 +600,8 @@ def build_cli_sso_attribution_metadata(
def _merge_cli_sso_attribution_metadata(
existing_metadata: dict[str, Any], attribution_metadata: dict[str, Any]
) -> dict[str, Any]:
existing_metadata: dict[str, object], attribution_metadata: dict[str, object]
) -> dict[str, object]:
"""Merge attribution metadata into existing user metadata in-place.
Preserves original value types (in particular, string claim values that
@ -608,7 +609,7 @@ def _merge_cli_sso_attribution_metadata(
are merged iteratively so attribution claims do not clobber unrelated keys
under the same parent.
"""
pending: Final[list[tuple[dict[str, Any], dict[str, Any]]]] = [(existing_metadata, attribution_metadata)]
pending: Final[list[tuple[dict[str, object], dict[str, object]]]] = [(existing_metadata, attribution_metadata)]
while pending:
target, source = pending.pop()
for key, value in source.items():
@ -656,7 +657,7 @@ async def _persist_cli_sso_user_metadata(
def _cli_poll_attribution_metadata_from_session(
session_data: dict[str, Any],
session_data: Mapping[str, object],
) -> dict[str, str | int | float | bool]:
stored: Final = session_data.get("attribution_metadata")
if isinstance(stored, dict):
@ -960,11 +961,12 @@ def process_sso_jwt_access_token(
# Try role_mappings first (group-based role determination)
if role_mappings is not None and role_mappings.roles:
group_claim: Final = role_mappings.group_claim
user_groups_raw: Final[Any] = get_nested_value(access_token_payload, group_claim)
user_groups_raw: Final[object] = get_nested_value(access_token_payload, group_claim)
user_groups: list[str] = []
if isinstance(user_groups_raw, list):
user_groups = [str(g) for g in user_groups_raw]
raw_groups: Final[Sequence[object]] = user_groups_raw
user_groups = [str(g) for g in raw_groups]
elif isinstance(user_groups_raw, str):
user_groups = [g.strip() for g in user_groups_raw.split(",") if g.strip()]
elif user_groups_raw is not None:
@ -1214,12 +1216,13 @@ def generic_response_convertor(
]:
# Use role_mappings to determine role from groups
group_claim: Final = role_mappings.group_claim
user_groups_raw: Final[Any] = get_nested_value(response, group_claim)
user_groups_raw: Final[object] = get_nested_value(response, group_claim)
# Handle different formats: could be a list, string (comma-separated), or single value
user_groups: list[str] = []
if isinstance(user_groups_raw, list):
user_groups = [str(g) for g in user_groups_raw]
raw_groups: Final[Sequence[object]] = user_groups_raw
user_groups = [str(g) for g in raw_groups]
elif isinstance(user_groups_raw, str):
# Handle comma-separated string
user_groups = [g.strip() for g in user_groups_raw.split(",") if g.strip()]
@ -3093,7 +3096,7 @@ class SSOAuthenticationHandler:
def _get_generic_sso_redirect_params(
state: str | None = None,
generic_authorization_endpoint: str | None = None,
) -> tuple[dict, str | None]:
) -> tuple[dict[str, str], str | None]:
"""
Get redirect parameters for Generic SSO with proper state priority handling.
Optionally generates PKCE parameters if GENERIC_CLIENT_USE_PKCE is enabled.

View file

@ -92,6 +92,7 @@ _LLM_ROUTE_EXACT: Final[tuple[str, ...]] = (
"/v1/messages",
"/interactions", # Google Interactions create; /{id} reads and /cancel do not match
"/v1beta/interactions",
"/comprehendmedical", # AWS-SDK-shaped passthrough: the operation rides in the X-Amz-Target header
)
# Provider passthrough prefixes (e.g. /bedrock/..., /vertex-ai/...) carry real

View file

@ -1,13 +1,20 @@
#### OCR Endpoints #####
import json
from collections.abc import Mapping
from typing import Any, Final, cast
import orjson
from fastapi import APIRouter, Depends, Request, Response, UploadFile
from fastapi import APIRouter, Depends, HTTPException, Request, Response, UploadFile
from fastapi.responses import ORJSONResponse
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.ocr.transformation import (
OCR_REQUEST_FORMAT_HEADER,
OCR_REQUEST_FORMAT_PARAM,
OCRResponse,
parse_ocr_request_format,
)
from litellm.ocr.main import convert_file_document_to_url_document, get_mime_type
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
@ -41,6 +48,48 @@ def _build_document_from_upload(
)
def _with_request_format(data: Mapping[str, Any], request: Request) -> Mapping[str, Any]:
"""
Resolve the requested response format from the body or the `x-req-format` header.
An explicit `req_format` in the body wins over the header.
"""
body_value: Final = data.get(OCR_REQUEST_FORMAT_PARAM)
header_value: Final = request.headers.get(OCR_REQUEST_FORMAT_HEADER)
raw_value: Final = body_value if body_value is not None else header_value
if raw_value is None:
return data
try:
request_format: Final = parse_ocr_request_format(
raw_value.strip().lower() if isinstance(raw_value, str) else raw_value
)
except ValueError as e:
raise HTTPException(status_code=400, detail={"error": f"{e}"})
return {**data, OCR_REQUEST_FORMAT_PARAM: request_format}
def _native_response(response: object, fastapi_response: Response) -> Response | None:
"""
Return the provider's native payload when the caller asked for
`req_format=native` and the provider config captured it, carrying over the
LiteLLM response headers (cost, call id, etc.) built for the normalized response.
"""
if not isinstance(response, OCRResponse):
return None
native_payload: Final = response.get_provider_native_response()
if native_payload is None:
return None
return Response(
content=orjson.dumps(native_payload),
media_type="application/json",
headers={
key: value
for key, value in fastapi_response.headers.items()
if key.lower() not in ("content-length", "content-type")
},
)
async def _parse_multipart_form(request: Request) -> dict[str, Any]:
"""
Extract OCR data from a multipart form request.
@ -105,7 +154,12 @@ async def _parse_multipart_form(request: Request) -> dict[str, Any]:
return data
async def _parse_ocr_request(request: Request) -> dict[str, Any]:
async def _parse_ocr_request(request: Request) -> Mapping[str, Any]:
"""Parse an OCR request and apply the `x-req-format` header, if any."""
return _with_request_format(await _parse_ocr_request_body(request), request)
async def _parse_ocr_request_body(request: Request) -> dict[str, Any]:
"""
Parse an OCR request, supporting both JSON and multipart form data.
@ -238,6 +292,11 @@ async def ocr(
-F "model=mistral-ocr" \
-F "file=@document.pdf"
```
Response format is normalized to the LiteLLM OCR schema by default. Providers
that support it (Azure Document Intelligence) can return their own payload
instead, with cost tracking unchanged, via `x-req-format: native` (or
`"req_format": "native"` in the body).
"""
from litellm.proxy.proxy_server import (
general_settings,
@ -256,12 +315,12 @@ async def ocr(
data: dict = {}
try:
# Parse request body (JSON or multipart form)
data = await _parse_ocr_request(request)
data = dict(await _parse_ocr_request(request))
# Process request using ProxyBaseLLMRequestProcessing
processor = ProxyBaseLLMRequestProcessing(data=data)
return await processor.base_process_llm_request(
response: Final = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -279,6 +338,8 @@ async def ocr(
user_api_base=user_api_base,
version=version,
)
return _native_response(response, fastapi_response) or response
except Exception as e:
processor = ProxyBaseLLMRequestProcessing(data=data)
raise await processor._handle_llm_api_exception(

View file

@ -9,7 +9,8 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
import json
import os
import re
from typing import Any, Final, cast
from types import MappingProxyType
from typing import Annotated, Any, Final, cast
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
@ -27,7 +28,7 @@ from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth, user_api_key_auth_websocket
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
@ -1079,6 +1080,130 @@ async def bedrock_proxy_route(
return received_value
COMPREHEND_MEDICAL_TARGET_PREFIX: Final = "ComprehendMedical_20181030"
def _resolve_comprehend_medical_region() -> str | None:
region_candidates: Final = (
get_secret_str(secret_name="AWS_REGION_NAME"),
get_secret_str(secret_name="AWS_REGION"),
get_secret_str(secret_name="AWS_DEFAULT_REGION"),
)
return next((region for region in region_candidates if region), None)
@router.post(
"/comprehendmedical/{operation}",
tags=["AWS Comprehend Medical Pass-through", "pass-through"], # mutable-ok: fastapi route tags must be a list
)
async def comprehend_medical_proxy_route(
operation: str,
request: Request,
fastapi_response: Response,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
):
"""
Pass-through for Amazon Comprehend Medical, e.g. `POST /comprehendmedical/DetectEntitiesV2`.
The request body is forwarded as-is to the AWS JSON 1.1 API and signed with SigV4
using the proxy's AWS credentials.
[Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical)
"""
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call comprehendmedical. Run 'pip install boto3'.")
from .llm_provider_handlers.comprehend_medical_passthrough_logging_handler import (
COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS,
)
if operation not in COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS:
raise HTTPException(
status_code=400,
detail=(
f"Unsupported Comprehend Medical operation: {operation}. "
f"Supported operations: {', '.join(sorted(COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS))}"
),
)
aws_region_name: Final = _resolve_comprehend_medical_region()
if aws_region_name is None:
raise HTTPException(
status_code=400,
detail="AWS region not found. Set AWS_REGION_NAME in the proxy environment.",
)
try:
data: Final = await request.json()
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
if not isinstance(data, dict):
raise HTTPException(status_code=400, detail="Request body must be a JSON object")
if "stream" in data:
raise HTTPException(status_code=400, detail="'stream' is not a Comprehend Medical request member")
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
credentials: Final[Credentials] = BaseAWSLLM().get_credentials(aws_region_name=aws_region_name)
sigv4: Final = SigV4Auth(credentials, "comprehendmedical", aws_region_name)
headers: Final = MappingProxyType(
{
"Content-Type": "application/x-amz-json-1.1",
"X-Amz-Target": f"{COMPREHEND_MEDICAL_TARGET_PREFIX}.{operation}",
}
)
target_url: Final = f"https://comprehendmedical.{aws_region_name}.amazonaws.com/"
_request: Final = AWSRequest(method="POST", url=target_url, data=json.dumps(data), headers=headers)
sigv4.add_auth(_request)
prepped: Final = _request.prepare()
endpoint_func: Final = create_pass_through_route(
endpoint=operation,
target=str(prepped.url),
custom_headers=prepped.headers,
custom_llm_provider="comprehendmedical",
)
setattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, data)
setattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, prepped.body)
return await endpoint_func(request, fastapi_response, user_api_key_dict)
@router.post(
"/comprehendmedical",
tags=["AWS Comprehend Medical Pass-through", "pass-through"], # mutable-ok: fastapi route tags must be a list
)
async def comprehend_medical_sdk_proxy_route(
request: Request,
fastapi_response: Response,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
):
"""
AWS-SDK-shaped pass-through for Amazon Comprehend Medical: point the SDK's
`endpoint_url` at `/comprehendmedical` and the operation is read from the
`X-Amz-Target` header, per the AWS JSON 1.1 protocol.
[Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical)
"""
target_header: Final = request.headers.get("x-amz-target", "")
target_prefix, _, operation = target_header.partition(".")
if target_prefix != COMPREHEND_MEDICAL_TARGET_PREFIX or not operation:
raise HTTPException(
status_code=400,
detail=f"Expected an X-Amz-Target header of the form {COMPREHEND_MEDICAL_TARGET_PREFIX}.<Operation>",
)
return await comprehend_medical_proxy_route(
operation=operation,
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
)
def _resolve_vertex_model_from_router(
model_id: str,
llm_router: litellm.Router | None,
@ -1972,6 +2097,104 @@ async def openai_proxy_route(
)
def _join_url_paths(base_url: httpx.URL, path: str, custom_llm_provider: litellm.LlmProviders) -> str:
"""
Properly joins a base URL with a path, preserving any existing path in the base URL.
"""
# Combine paths via the shared helper so any '..' in the path cannot
# climb above the configured base path.
joined_path_str = str(
base_url.copy_with(path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, path))
)
# Apply OpenAI-specific path handling for both branches
if custom_llm_provider == litellm.LlmProviders.OPENAI and "/v1/" not in joined_path_str:
# Insert v1 after api.openai.com for OpenAI requests
joined_path_str = joined_path_str.replace("api.openai.com/", "api.openai.com/v1/")
return joined_path_str
_OPENAI_WS_ALL_MODEL_ACCESS: Final = frozenset(
{
SpecialModelNames.all_proxy_models.value,
SpecialModelNames.all_team_models.value,
"*",
}
)
def _key_has_model_restrictions(user_api_key_dict: UserAPIKeyAuth) -> bool:
scoped_models: Final = (*user_api_key_dict.models, *user_api_key_dict.team_models)
return any(str(model) not in _OPENAI_WS_ALL_MODEL_ACCESS for model in scoped_models)
@router.websocket("/openai_passthrough/{endpoint:path}")
@router.websocket("/openai/{endpoint:path}")
async def openai_websocket_proxy_route(
websocket: WebSocket,
endpoint: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)],
) -> None:
"""WebSocket passthrough for OpenAI prefixes (realtime / responses.connect)."""
if _key_has_model_restrictions(user_api_key_dict):
await websocket.close(
code=1008,
reason="Keys with model restrictions cannot use OpenAI websocket passthrough",
)
return
base_target_url: Final = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/"
openai_api_key: Final = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.OPENAI.value,
region_name=None,
)
if openai_api_key is None:
await websocket.close(
code=1011,
reason="Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.",
)
return
raw_path: Final = httpx.URL(endpoint).path
encoded_endpoint: Final = raw_path if raw_path.startswith("/") else f"/{raw_path}"
base_url: Final = httpx.URL(base_target_url)
updated_url: Final = _join_url_paths(
base_url=base_url,
path=encoded_endpoint,
custom_llm_provider=litellm.LlmProviders.OPENAI,
)
wss_base: Final = (
"wss://" + updated_url[len("https://") :]
if updated_url.startswith("https://")
else "ws://" + updated_url[len("http://") :]
if updated_url.startswith("http://")
else updated_url
)
query_string: Final = websocket.url.query
wss_target: Final = f"{wss_base}{'&' if '?' in wss_base else '?'}{query_string}" if query_string else wss_base
custom_headers: Final = { # mutable-ok: websocket_passthrough_request requires a plain dict of upstream headers
"Authorization": f"Bearer {openai_api_key}"
}
requested_subprotocols: Final = tuple(
protocol.strip()
for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",")
if protocol.strip()
)
await websocket.accept(subprotocol=requested_subprotocols[0] if requested_subprotocols else None)
await websocket_passthrough_request(
websocket=websocket,
target=wss_target,
custom_headers=custom_headers,
user_api_key_dict=user_api_key_dict,
forward_headers=False,
endpoint=websocket.url.path,
accept_websocket=False,
)
class BaseOpenAIPassThroughHandler:
@staticmethod
async def _base_openai_pass_through_handler(
@ -1991,7 +2214,7 @@ class BaseOpenAIPassThroughHandler:
# Construct the full target URL by properly joining the base URL and endpoint path
base_url: Final = httpx.URL(base_target_url)
updated_url: Final = BaseOpenAIPassThroughHandler._join_url_paths(
updated_url: Final = _join_url_paths(
base_url=base_url,
path=encoded_endpoint,
custom_llm_provider=custom_llm_provider,
@ -2050,24 +2273,6 @@ class BaseOpenAIPassThroughHandler:
request=request,
)
@staticmethod
def _join_url_paths(base_url: httpx.URL, path: str, custom_llm_provider: litellm.LlmProviders) -> str:
"""
Properly joins a base URL with a path, preserving any existing path in the base URL.
"""
# Combine paths via the shared helper so any '..' in the path cannot
# climb above the configured base path.
joined_path_str = str(
base_url.copy_with(path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, path))
)
# Apply OpenAI-specific path handling for both branches
if custom_llm_provider == litellm.LlmProviders.OPENAI and "/v1/" not in joined_path_str:
# Insert v1 after api.openai.com for OpenAI requests
joined_path_str = joined_path_str.replace("api.openai.com/", "api.openai.com/v1/")
return joined_path_str
@router.api_route(
"/cursor/{endpoint:path}",

View file

@ -0,0 +1,102 @@
import math
from collections.abc import Mapping
from datetime import datetime
from types import MappingProxyType
from typing import Final
import httpx
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import (
get_standard_logging_object_payload,
)
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.types.utils import StandardPassThroughResponseObject
COMPREHEND_MEDICAL_CHARS_PER_UNIT: Final = 100
COMPREHEND_MEDICAL_COST_PER_UNIT_USD: Final[Mapping[str, float]] = MappingProxyType(
{
"DetectEntitiesV2": 0.01,
"DetectPHI": 0.0014,
"InferICD10CM": 0.0005,
"InferRxNorm": 0.00025,
"InferSNOMEDCT": 0.0075,
}
)
COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS: Final = frozenset(COMPREHEND_MEDICAL_COST_PER_UNIT_USD)
class ComprehendMedicalPassthroughLoggingHandler:
@staticmethod
def _operation_from_response(httpx_response: httpx.Response) -> str:
target: Final = httpx_response.request.headers.get("x-amz-target", "")
return target.split(".")[-1]
@staticmethod
def get_cost_for_operation(operation: str, text: str) -> float:
cost_per_unit: Final = COMPREHEND_MEDICAL_COST_PER_UNIT_USD.get(operation)
if cost_per_unit is None:
return 0.0
units: Final = max(1, math.ceil(len(text) / COMPREHEND_MEDICAL_CHARS_PER_UNIT))
return units * cost_per_unit
@staticmethod
def comprehend_medical_passthrough_handler(
httpx_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
url_route: str,
result: str,
start_time: datetime,
end_time: datetime,
cache_hit: bool,
request_body: Mapping[str, object],
**kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler
) -> PassThroughEndpointLoggingTypedDict:
"""
Prices a Comprehend Medical sync operation from the request text length
(billed per started 100-character unit, 1-unit minimum) and records
model, provider, and cost on the logging payload.
"""
try:
operation: Final = ComprehendMedicalPassthroughLoggingHandler._operation_from_response(httpx_response)
text: Final = request_body.get("Text")
response_cost: Final = ComprehendMedicalPassthroughLoggingHandler.get_cost_for_operation(
operation=operation,
text=text if isinstance(text, str) else "",
)
model_name: Final = f"comprehendmedical/{operation}"
updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict
**kwargs,
"model": model_name,
"custom_llm_provider": "comprehendmedical",
"response_cost": response_cost,
}
logging_obj.model_call_details.update(
model=model_name,
custom_llm_provider="comprehendmedical",
response_cost=response_cost,
)
standard_logging_object: Final = get_standard_logging_object_payload(
kwargs=updated_kwargs,
init_response_obj=StandardPassThroughResponseObject(response=result),
start_time=start_time,
end_time=end_time,
logging_obj=logging_obj,
status="success",
)
handler_payload: Final[PassThroughEndpointLoggingTypedDict] = {
"result": StandardPassThroughResponseObject(response=result),
"kwargs": {**updated_kwargs, "standard_logging_object": standard_logging_object},
}
except Exception as e:
verbose_proxy_logger.exception("Error in Comprehend Medical passthrough logging handler: %s", e)
fallback_payload: Final[PassThroughEndpointLoggingTypedDict] = {
"result": StandardPassThroughResponseObject(response=result),
"kwargs": kwargs,
}
return fallback_payload
return handler_payload

View file

@ -45,6 +45,7 @@ from litellm.llms.base_llm.managed_resources.isolation import (
can_access_resource,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.batches_endpoints.common_utils import validate_batch_list_limit
from litellm.repositories.table_repositories import (
ManagedFileRepository,
ManagedObjectRepository,
@ -1029,12 +1030,16 @@ async def list_passthrough_ids_from_db(
if resource_kind is None:
return None
raw_limit, fetch_limit = _parse_list_limit(query_params)
if resource_kind == "batches":
validate_batch_list_limit(raw_limit)
if raw_limit == 0:
return _empty_list_response()
owner_filter: Final = build_owner_filter(user_api_key_dict)
if owner_filter is None:
verbose_proxy_logger.warning("managed_id_rewriter: list denied — caller has no user_id or team_id")
return _empty_list_response()
raw_limit, fetch_limit = _parse_list_limit(query_params)
where, fetch_order = await _build_list_where_with_cursor(
prisma_client, resource_kind, provider, owner_filter, query_params
)

View file

@ -1538,6 +1538,8 @@ async def pass_through_request(
#########################################################
if isinstance(e, ProxyException):
raise
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(getattr(e, "detail", str(e)))),
@ -2091,8 +2093,8 @@ async def websocket_passthrough_request(
raw_response = await upstream_ws.recv(decode=False)
# Ensure raw_response is bytes before decoding
if isinstance(raw_response, str):
raw_response = raw_response.encode("ascii")
setup_response: Final[Mapping[str, object]] = json.loads(raw_response.decode("ascii"))
raw_response = raw_response.encode("utf-8")
setup_response: Final[Mapping[str, object]] = json.loads(raw_response.decode("utf-8"))
verbose_proxy_logger.debug("Setup response: %s", setup_response)
# Extract model and provider from setup response for Vertex AI Live

View file

@ -236,6 +236,26 @@ class PassThroughEndpointLogging:
)
standard_logging_response_object = cursor_passthrough_logging_handler_result["result"]
kwargs = cursor_passthrough_logging_handler_result["kwargs"]
elif self.is_comprehend_medical_route(custom_llm_provider):
from .llm_provider_handlers.comprehend_medical_passthrough_logging_handler import (
ComprehendMedicalPassthroughLoggingHandler,
)
comprehend_medical_handler_result: Final = (
ComprehendMedicalPassthroughLoggingHandler.comprehend_medical_passthrough_handler(
httpx_response=httpx_response,
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
request_body=request_body,
**kwargs,
)
)
standard_logging_response_object = comprehend_medical_handler_result["result"] # rebind-ok: elif-chain
kwargs = comprehend_medical_handler_result["kwargs"] # rebind-ok: elif-chain contract
elif self.is_vertex_ai_live_route(url_route):
from .llm_provider_handlers.vertex_ai_live_passthrough_logging_handler import (
VertexAILivePassthroughLoggingHandler,
@ -364,6 +384,9 @@ class PassThroughEndpointLogging:
return True
return False
def is_comprehend_medical_route(self, custom_llm_provider: str | None) -> bool:
return custom_llm_provider == "comprehendmedical"
def is_langfuse_route(self, url_route: str):
parsed_url: Final = urlparse(url_route)
for route in self.TRACKED_LANGFUSE_ROUTES:

View file

@ -174,7 +174,9 @@ class PipelineExecutor:
# Use unified_guardrail path if callback implements apply_guardrail
target: CustomLogger = callback
use_unified: Final = "apply_guardrail" in type(callback).__dict__
use_unified: Final = (
"apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks
)
if use_unified:
data["guardrail_to_apply"] = callback
target = UnifiedLLMGuardrails()

View file

@ -916,6 +916,7 @@ class ProxyInitializationHelpers:
"path that can cause schema thrashing during rolling deploys where two "
"LiteLLM versions contend for the same DB. Default is the v1 resolver."
),
envvar="USE_V2_MIGRATION_RESOLVER",
)
@click.option(
"--reload",

View file

@ -328,6 +328,7 @@ from litellm.proxy.common_utils.load_config_utils import (
get_config_file_contents_from_gcs,
get_file_contents_from_s3,
)
from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations
from litellm.proxy.common_utils.model_listing_utils import TeamModelNameTranslator
from litellm.proxy.common_utils.openai_endpoint_utils import (
remove_sensitive_info_from_deployment,
@ -358,7 +359,9 @@ from litellm.proxy.common_utils.timezone_utils import (
)
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
end_user_cache_key,
get_management_object_ttl,
tag_cache_key,
)
from litellm.proxy.config_resolvers import resolve_fields
from litellm.proxy.config_resolvers.alerting import (
@ -648,6 +651,10 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
LiteLLM_UpperboundKeyGenerateParams,
)
from litellm.types.proxy.model_deprecation import (
DEFAULT_DEPRECATION_WARN_DAYS,
ModelDeprecationResponse,
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.router import (
DeploymentTypedDict,
@ -864,7 +871,7 @@ async def _flush_spend_logs_queue_on_shutdown() -> None:
verbose_proxy_logger.exception("Error flushing spend logs queue on shutdown: %s", e)
async def proxy_shutdown_event():
async def proxy_shutdown_event() -> None:
global prisma_client, master_key, user_custom_auth, user_custom_key_generate, user_custom_key_update
verbose_proxy_logger.info("Shutting down LiteLLM Proxy Server")
if prisma_client:
@ -958,7 +965,7 @@ async def _initialize_shared_aiohttp_session():
@asynccontextmanager
async def proxy_startup_event(app: FastAPI):
async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
global \
prisma_client, \
master_key, \
@ -2780,7 +2787,7 @@ async def _increment_end_user_and_tag_spend_counters(
if end_user_id is not None:
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:end_user:{end_user_id}",
source_cache_key=f"end_user_id:{end_user_id}",
source_cache_key=end_user_cache_key(end_user_id),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
@ -2795,7 +2802,7 @@ async def _increment_end_user_and_tag_spend_counters(
seen_tags.add(tag_name)
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:tag:{tag_name}",
source_cache_key=f"tag:{tag_name}",
source_cache_key=tag_cache_key(tag_name),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
@ -3134,7 +3141,7 @@ async def update_cache(
if end_user_id is None or response_cost is None:
return
_id: Final = f"end_user_id:{end_user_id}"
_id: Final = end_user_cache_key(end_user_id)
try:
# Fetch the existing cost for the given user
cached_end_user: Final = await user_api_key_cache.async_get_cache(key=_id)
@ -3226,7 +3233,7 @@ async def update_cache(
if not tag_name or not isinstance(tag_name, str):
continue
cache_key = f"tag:{tag_name}"
cache_key = tag_cache_key(tag_name)
# Fetch the existing tag object from cache
cached_tag = await user_api_key_cache.async_get_cache(key=cache_key)
if cached_tag is None:
@ -3732,11 +3739,11 @@ _DB_OVERLAY_REMOTE_MODULE_LIST_FIELDS: Final[dict[str, tuple[str, ...]]] = {
}
def _is_remote_module_url(value: Any) -> bool:
def _is_remote_module_url(value: object) -> bool:
return isinstance(value, str) and (value.startswith("s3://") or value.startswith("gcs://"))
def _scrub_guardrail_inner(inner: dict[str, Any]) -> None:
def _scrub_guardrail_inner(inner: dict[str, JsonValue]) -> None:
"""Strip remote-URL entries from a guardrail's ``callbacks`` list
and ``guardrail`` (v2 module-path) field. Mutates in place."""
cbs: Final = inner.get("callbacks")
@ -3756,7 +3763,7 @@ def _scrub_guardrail_inner(inner: dict[str, Any]) -> None:
inner["guardrail"] = None
def _scrub_db_overlay_remote_module_loads(section: str, db_value: Any) -> Any:
def _scrub_db_overlay_remote_module_loads(section: str, db_value: JsonValue) -> JsonValue:
"""Strip ``s3://`` / ``gcs://`` entries from the DB-overlay value for
fields whose contents reach ``get_instance_fn``. The same scheme is
allowed from a YAML config (the documented operator flow) but a
@ -4064,8 +4071,8 @@ class ProxyConfig:
def __init__(self) -> None:
self.config: dict[str, Any] = {}
self._last_semantic_filter_config: dict[str, Any] | None = None
self._last_hashicorp_vault_config: dict[str, Any] | None = None
self._last_semantic_filter_config: dict[str, object] | None = None
self._last_hashicorp_vault_config: dict[str, object] | None = None
self.worker_registry: list[WorkerRegistryEntry] = []
self.config_sync_subscriber: ConfigSyncSubscriber | None = None
self.auth_cache_invalidation_subscriber: AuthCacheInvalidationSubscriber | None = None
@ -5955,7 +5962,7 @@ class ProxyConfig:
)
@staticmethod
def _parse_router_settings_value(value: Any) -> dict | None:
def _parse_router_settings_value(value: object) -> dict | None:
"""
Parse a router_settings value that may be a dict or a JSON/YAML string.
@ -6499,7 +6506,7 @@ class ProxyConfig:
as "all models deleted" and must not evict existing router deployments.
"""
try:
new_models: Final = await ModelRepository(prisma_client).table.find_many()
new_models: Final[list[_ModelTableRow]] = await ModelRepository(prisma_client).table.find_many()
return new_models
except Exception as e:
verbose_proxy_logger.exception(
@ -7563,9 +7570,9 @@ def _get_client_requested_model_for_streaming(request_data: dict) -> str:
return requested_model if isinstance(requested_model, str) else ""
def _is_positive_int_like(value: Any) -> bool:
def _is_positive_int_like(value: str | float | None) -> bool:
try:
return int(value) > 0
return value is not None and int(value) > 0
except (TypeError, ValueError):
return False
@ -7832,7 +7839,7 @@ _STREAM_KEEPALIVE: Final = object()
_KEEPALIVE_MIN_SECONDS: Final = 1.0
_KEEPALIVE_MAX_SECONDS: Final = 300.0
_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({})
_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
async def _iter_with_keepalive(
@ -7887,7 +7894,7 @@ async def _iter_with_keepalive(
class _DeploymentKeepaliveConfig(NamedTuple):
keepalive_seconds: Any
keepalive_seconds: object
allow_client_override: bool
@ -7945,7 +7952,7 @@ def _is_explicit_keepalive_disable(raw: object) -> bool:
return False
def _resolve_keepalive_seconds(request_data: Mapping[str, Any], response: object = None) -> float:
def _resolve_keepalive_seconds(request_data: Mapping[str, object], response: object = None) -> float:
deployment_config: Final = _keepalive_from_deployment_config(request_data, response)
deployment_raw: Final = deployment_config.keepalive_seconds if deployment_config is not None else None
allow_client_override: Final = deployment_config.allow_client_override if deployment_config is not None else False
@ -7992,7 +7999,7 @@ def _resolve_keepalive_seconds(request_data: Mapping[str, Any], response: object
_KEEPALIVE_CACHE_TTL_SECONDS: Final = 5.0
def _make_keepalive_resolver(request_data: Mapping[str, Any]) -> Callable[[object], float]:
def _make_keepalive_resolver(request_data: Mapping[str, object]) -> Callable[[object], float]:
"""Wrap `_resolve_keepalive_seconds` with a memo keyed on the serving
deployment's model_id. The steady-state case (no mid-stream fallback, the
overwhelming majority of streams) sees the same model_id on every chunk, so
@ -9784,7 +9791,7 @@ async def model_info(
)
def _blocked_response_usage(original_response: Any | None) -> "litellm.Usage":
def _blocked_response_usage(original_response: object | None) -> "litellm.Usage":
"""
Token usage for a synthetic guardrail-blocked response.
@ -10371,6 +10378,8 @@ async def moderations(
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
)
verbose_proxy_logger.exception("litellm.proxy.proxy_server.moderations(): Exception occured - %s", e)
if isinstance(e, ProxyException):
raise
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e)),
@ -12303,7 +12312,7 @@ def _enrich_model_info_with_litellm_data(
async def _get_caller_byok_team_scope(
user_api_key_dict: UserAPIKeyAuth | None,
prisma_client: Any | None,
prisma_client: PrismaClient | None,
) -> set[str] | None:
"""
Return the team IDs whose BYOK rows the caller is allowed to see via
@ -12338,7 +12347,7 @@ async def _get_caller_byok_team_scope(
return key_team_scope | set(user_row.teams or [])
def _byok_row_outside_caller_teams(model_info_dict: dict[str, Any], allowed_team_ids: set[str] | None) -> bool:
def _byok_row_outside_caller_teams(model_info_dict: dict[str, JsonValue], allowed_team_ids: set[str] | None) -> bool:
"""Whether a team BYOK row belongs to a team the caller is not a member of.
`team_id` is only set on team BYOK rows; non-team rows fall through
@ -12360,15 +12369,15 @@ _SORTED_SEARCH_DB_FETCH_CAP: Final = 500
async def _fetch_db_models_for_search(
prisma_client: Any,
proxy_config: Any,
prisma_client: PrismaClient,
proxy_config: ProxyConfig,
search_lower: str,
db_model_ids_in_router: set[str],
router_models_count: int,
page: int,
size: int,
sort_by: str | None,
is_byok_outside_caller_teams: Callable[[dict[str, Any]], bool],
is_byok_outside_caller_teams: Callable[[dict[str, JsonValue]], bool],
) -> tuple[list[dict[str, Any]], int]:
"""
Run the bounded DB query that backs `/v2/model/info?search=`. Returns
@ -12414,7 +12423,7 @@ async def _fetch_db_models_for_search(
if not is_byok_outside_caller_teams(m.model_info if isinstance(m.model_info, dict) else {})
]
decrypted: Final[list[dict[str, Any]]] = []
decrypted: Final[list[dict[str, object]]] = []
for db_model in matching_db_rows:
decrypted_models = proxy_config.decrypt_model_list_from_db([db_model])
if decrypted_models:
@ -12426,8 +12435,8 @@ async def _fetch_db_models_for_search(
async def _apply_search_filter_to_models(
all_models: list[dict[str, Any]],
search: str,
prisma_client: Any | None,
proxy_config: Any,
prisma_client: PrismaClient | None,
proxy_config: ProxyConfig,
user_api_key_dict: UserAPIKeyAuth | None = None,
page: int = 1,
size: int = 50,
@ -12466,7 +12475,7 @@ async def _apply_search_filter_to_models(
prisma_client=prisma_client,
)
def _is_byok_outside_caller_teams(model_info_dict: dict[str, Any]) -> bool:
def _is_byok_outside_caller_teams(model_info_dict: dict[str, JsonValue]) -> bool:
return _byok_row_outside_caller_teams(model_info_dict, allowed_team_ids)
def _model_matches_search(m: dict[str, Any]) -> bool:
@ -12532,7 +12541,7 @@ async def _apply_search_filter_to_models(
return filtered_router_models + db_models, search_total_count
def _normalize_datetime_for_sorting(dt: Any) -> datetime | None:
def _normalize_datetime_for_sorting(dt: object) -> datetime | None:
"""
Normalize a datetime value to a timezone-aware UTC datetime for sorting.
@ -12685,7 +12694,7 @@ def _paginate_models_response(
size: int,
total_count: int | None,
search: str | None,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Paginate models and return response dictionary.
@ -12724,7 +12733,7 @@ def _paginate_models_response(
}
def _team_models_resolve_to_names(team_models: list[str], access_groups: dict[str, Any]) -> list[str]:
def _team_models_resolve_to_names(team_models: list[str], access_groups: Mapping[str, Sequence[str]]) -> list[str]:
"""Expand team model entries (including access group names) to concrete model names."""
resolved: Final[list[str]] = []
for name in team_models:
@ -13600,7 +13609,7 @@ async def model_metrics_exceptions(
return {"data": response, "exception_types": list(exception_types)}
def _deployment_matches_allowed_model_names(model: dict[str, Any], allowed_model_names: set[str]) -> bool:
def _deployment_matches_allowed_model_names(model: dict[str, JsonValue], allowed_model_names: set[str]) -> bool:
"""Match a router deployment against allowed public model names.
Team-scoped rows store an internal routing key in ``model_name``; callers
@ -13923,6 +13932,48 @@ async def model_info_v1(
return {"data": all_models}
@router.get(
"/model/deprecations",
tags=("model management",),
dependencies=(Depends(user_api_key_auth),),
response_model=ModelDeprecationResponse,
)
@router.get(
"/v1/model/deprecations",
tags=("model management",),
dependencies=(Depends(user_api_key_auth),),
response_model=ModelDeprecationResponse,
)
async def model_deprecations(
warn_within_days: int = DEFAULT_DEPRECATION_WARN_DAYS,
) -> ModelDeprecationResponse:
"""List models with known deprecation/sunset dates, bucketed by urgency.
Reads `deprecation_date` metadata from `model_prices_and_context_window.json`
(and any per-deployment `model_info.deprecation_date` overrides) for the
models configured on this proxy.
Parameters:
warn_within_days: Window (in days) used to bucket "imminent" models,
30 by default.
Returns:
A payload with three lists of `ModelDeprecationInfo` entries:
- `deprecated`: deprecation date is in the past, so these requests may
fail at any time.
- `imminent`: deprecation date is within `warn_within_days` from today.
- `upcoming`: deprecation date is further out.
Example:
```shell
curl -X GET 'http://localhost:4000/model/deprecations' \\
-H 'Authorization: Bearer sk-1234'
```
"""
return collect_model_deprecations(llm_router=llm_router, warn_within_days=warn_within_days)
def _get_model_group_info(
llm_router: Router, all_models_str: list[str], model_group: str | None
) -> list[ModelGroupInfoProxy]:
@ -14860,7 +14911,7 @@ async def _rollback_onboarding_invite_claim(
verbose_proxy_logger.exception("Failed to roll back onboarding invitation after session key mint failed.")
async def _generate_onboarding_ui_session_token(user_obj: Any) -> str:
async def _generate_onboarding_ui_session_token(user_obj: _UserTableRow) -> str:
global master_key, general_settings
response: Final = await generate_key_helper_fn(
@ -15975,7 +16026,7 @@ def _general_settings_ui_litellm_default(
return False if spec["type"] == "Boolean" else None
def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) -> GeneralSettingsUILiteLLMValue:
def _validate_general_settings_ui_litellm_value(field_name: str, value: object) -> GeneralSettingsUILiteLLMValue:
spec: Final = _GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name]
field_type: Final = spec["type"]
if value is None or value == "":
@ -16015,7 +16066,7 @@ def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) ->
async def _persist_general_settings_ui_litellm_field(
field_name: str, value: Any, user_api_key_dict: UserAPIKeyAuth
field_name: str, value: object, user_api_key_dict: UserAPIKeyAuth
) -> dict:
validated: Final = _validate_general_settings_ui_litellm_value(field_name, value)
config: Final = await proxy_config.get_config()

View file

@ -28,6 +28,7 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import
)
from litellm.types.proxy.public_endpoints.public_endpoints import (
AgentCreateInfo,
ComplexityScorerDefaults,
ProviderCreateInfo,
PublicModelHubInfo,
SupportedEndpointsResponse,
@ -398,6 +399,28 @@ async def get_provider_fields() -> list[ProviderCreateInfo]:
return provider_create_fields
@router.get(
"/public/complexity_router/scorer_defaults",
tags=["public", "auto router"],
response_model=ComplexityScorerDefaults,
)
async def get_complexity_scorer_defaults() -> ComplexityScorerDefaults:
"""
Return the complexity router's shipped heuristic scorer defaults, for the dashboard to prefill with.
"""
from litellm.router_strategy.complexity_router.config import (
DEFAULT_DIMENSION_WEIGHTS,
DEFAULT_TIER_BOUNDARIES,
DEFAULT_TOKEN_THRESHOLDS,
)
return ComplexityScorerDefaults(
tier_boundaries=DEFAULT_TIER_BOUNDARIES,
token_thresholds=DEFAULT_TOKEN_THRESHOLDS,
dimension_weights=DEFAULT_DIMENSION_WEIGHTS,
)
@router.get(
"/public/litellm_model_cost_map",
tags=["public", "model management"],

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