mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
commit
cd5d414c13
430 changed files with 18593 additions and 3092 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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/",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -83,6 +83,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/azure_ai/",
|
||||
"/aws/",
|
||||
"/bedrock/",
|
||||
"/comprehendmedical",
|
||||
"/cohere/",
|
||||
"/gemini/",
|
||||
"/google/",
|
||||
|
|
|
|||
|
|
@ -119,4 +119,7 @@ spec:
|
|||
{{- end }}
|
||||
ttlSecondsAfterFinished: {{ .Values.migrationJob.ttlSecondsAfterFinished }}
|
||||
backoffLimit: {{ .Values.migrationJob.backoffLimit }}
|
||||
{{- with .Values.migrationJob.activeDeadlineSeconds }}
|
||||
activeDeadlineSeconds: {{ . }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
#
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
199
litellm/llms/bedrock/request_metadata.py
Normal file
199
litellm/llms/bedrock/request_metadata.py
Normal 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,
|
||||
)
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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}}
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
18
litellm/proxy/batches_endpoints/common_utils.py
Normal file
18
litellm/proxy/batches_endpoints/common_utils.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
226
litellm/proxy/common_utils/model_deprecation.py
Normal file
226
litellm/proxy/common_utils/model_deprecation.py
Normal 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("&", "&").replace("<", "<").replace(">", ">")
|
||||
|
||||
|
||||
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.",
|
||||
)
|
||||
)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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, ...).
|
||||
|
|
|
|||
|
|
@ -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()):
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"] = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue