Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_legacy_hook_streaming_pipeline_step

This commit is contained in:
mateo-berri 2026-09-09 13:30:39 -07:00
commit 8ecd9c16cd
294 changed files with 22909 additions and 5698 deletions

View file

@ -113,7 +113,7 @@ jobs:
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 8
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra mongodb
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]'
- name: Cache Prisma binaries

View file

@ -66,7 +66,7 @@ jobs:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
full_suite() { npm run test -- --run --pool forks --poolOptions.forks.maxForks=14; }
full_suite() { npm run test -- --run --pool forks --maxWorkers=14; }
if [ -z "$BASE_SHA" ]; then
echo "Push to $GITHUB_REF_NAME: running the full suite"
@ -95,4 +95,4 @@ jobs:
echo "Pull request: running tests related to ${#changed_files[@]} changed UI files"
npm run test -- related "${changed_files[@]}" --run --passWithNoTests \
--pool forks --poolOptions.forks.maxForks=14
--pool forks --maxWorkers=14

View file

@ -67,7 +67,6 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13
# Copy full source tree
@ -90,7 +89,6 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -105,7 +105,7 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38271
"limit": 38269
},
"reportUnknownParameterType": {
"limit": 19584

View file

@ -65,7 +65,6 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13
# Copy full source tree
@ -88,7 +87,6 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -71,7 +71,6 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13
# Copy full source tree
@ -100,7 +99,6 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13 \
--no-sources-package litellm-proxy-extras; \
else \
@ -111,7 +109,6 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13; \
fi

View file

@ -1801,7 +1801,16 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Remove conflicting keys from data to avoid duplicate keyword arguments
filtered_data = {k: v for k, v in data.items() if k not in ("model", "file_id")}
for model_id, model_file_id in specific_model_file_id_mapping.items():
delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **filtered_data) # type: ignore
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
delete_data = {
**{k: v for k, v in filtered_data.items() if k != "_litellm_internal_model_credentials"},
**(
{"_litellm_internal_model_credentials": MappingProxyType(dict(credentials))}
if credentials is not None
else {}
),
}
delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
stored_file_object = await self.delete_unified_file_id(file_id, litellm_parent_otel_span)
@ -1812,7 +1821,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
prom_logger.record_managed_file_deleted(result="success")
if stored_file_object:
return stored_file_object
return OpenAIFileObject.model_validate(stored_file_object).model_copy(update={"id": file_id})
elif delete_response:
delete_response.id = file_id
return delete_response

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.65"
version = "0.1.66"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.65"
version = "0.1.66"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -47,7 +47,6 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra extra_proxy \
--extra semantic-router \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13
# Stage 2 — copy source and install the project + workspace members.
@ -60,7 +59,6 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra extra_proxy \
--extra semantic-router \
--extra bedrock-realtime \
--extra mongodb \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -326,6 +326,9 @@ ssl_certificate: Optional[str] = None
user_url_validation: bool = True
user_url_allowed_hosts: List[str] = []
provider_url_destination_allowed_hosts: List[str] = []
#: "override" (default) or "additive": whether a key or team destination replaces
#: the operator's exporter for that backend or exports alongside it.
otel_tenant_destination_mode: str | None = None
ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance
disable_streaming_logging: bool = False
disable_token_counter: bool = False
@ -543,7 +546,7 @@ _key_management_system: Optional["KeyManagementSystem"] = None
#### PII MASKING ####
output_parse_pii: bool = False
#############################################
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map, mark_litellm_import_complete
model_cost = get_model_cost_map(url=model_cost_map_url)
cost_discount_config: Dict[str, float] = {} # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount
@ -2402,3 +2405,5 @@ def __getattr__(name: str) -> Any:
# ALL_LITELLM_RESPONSE_TYPES is lazy-loaded via __getattr__ to avoid loading utils at import time
mark_litellm_import_complete()

View file

@ -6,9 +6,33 @@ be settable from user input. Context variables are scoped to the current
asyncio task and cannot be injected via HTTP request bodies.
"""
from collections.abc import Generator
from contextlib import contextmanager
from contextvars import ContextVar
from datetime import datetime, timezone
from typing import Final
# When True, suppresses async logging and billing for internal sub-calls
# (e.g., emulated file-search steps that make nested LLM calls).
is_internal_call: Final[ContextVar[bool]] = ContextVar("is_internal_call", default=False)
# One request prices its totals, its per-token-type lines and the rates it reports on
# separate code paths. Each reads the clock for off-peak pricing, so without a pinned
# moment they can land on either side of a window boundary and disagree with each other.
_billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", default=None)
@contextmanager
def pinned_billing_time(moment: datetime) -> Generator[None]:
"""Price every rate lookup inside this block at ``moment`` rather than at each one's own clock read."""
token: Final = _billing_time.set(moment)
try:
yield
finally:
_billing_time.reset(token)
def current_billing_time() -> datetime:
"""The pinned billing moment, or now in UTC outside a pinned block."""
pinned: Final = _billing_time.get()
return pinned if pinned is not None else datetime.now(timezone.utc)

View file

@ -143,6 +143,7 @@ DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD: Final = float(
os.getenv("DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD", 0.3)
)
MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH: Final = int(os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150))
MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH: Final = 2048
DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS: Final = 2000
@ -197,6 +198,7 @@ LITELLM_UI_ALLOW_HEADERS: Final = [
"x-litellm-adaptive-router-model",
"x-litellm-applied-guardrails",
"x-litellm-guardrail-scan-id",
"x-litellm-guardrail-scan-metadata",
"x-litellm-cache-key",
]
@ -333,6 +335,7 @@ DEFAULT_SSL_CIPHERS: Final = os.getenv(
########### v2 Architecture constants for managing writing updates to the database ###########
REDIS_UPDATE_BUFFER_KEY: Final = "litellm_spend_update_buffer"
REDIS_GATEWAY_REQUESTS_BUFFER_KEY: Final = "litellm_gateway_requests_buffer"
REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_spend_update_buffer"
REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_team_spend_update_buffer"
REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_org_spend_update_buffer"
@ -1374,6 +1377,7 @@ bedrock_embedding_models: Final[set] = set(
"cohere.embed-multilingual-v3",
"cohere.embed-v4:0",
"twelvelabs.marengo-embed-2-7-v1:0",
"twelvelabs.marengo-embed-3-0-v1:0",
]
)
@ -1766,6 +1770,10 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [
SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"]
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
DEFAULT_ACCESS_GROUP_CACHE_TTL: Final = int(os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600))
SPEND_LOG_KEY_METADATA_CACHE_TTL: Final = 600
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL: Final = 30
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS: Final = 10000
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS: Final = 5000
# Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated
# callers from forcing a DB query per request for unknown names, while bounding
# staleness so a transient DB error (which surfaces as an empty list) cannot

View file

@ -25,6 +25,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
TranscriptionUsageObjectTransformation,
)
from litellm.litellm_core_utils.llm_cost_calc.utils import (
BilledTokenRates,
CostCalculatorUtils,
_generic_cost_per_character,
_get_regional_uplift_multiplier,
@ -45,6 +46,9 @@ from litellm.llms.azure.cost_calculation import (
from litellm.llms.azure_ai.cost_calculator import (
cost_per_token as azure_ai_cost_per_token,
)
from litellm.llms.azure_ai.cost_calculator import (
is_azure_model_router as azure_ai_is_model_router_name,
)
from litellm.llms.base_llm.search.transformation import SearchResponse
from litellm.llms.bedrock.cost_calculation import (
cost_per_token as bedrock_cost_per_token,
@ -1122,6 +1126,7 @@ def _store_cost_breakdown_in_logging_obj(
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
billed_token_rates: BilledTokenRates | None = None,
) -> None:
"""
Helper function to store cost breakdown in the logging object.
@ -1166,6 +1171,7 @@ def _store_cost_breakdown_in_logging_obj(
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
billed_token_rates=billed_token_rates,
)
except Exception as breakdown_error:
@ -1659,11 +1665,10 @@ def completion_cost(
data_residency=data_residency,
vertex_location=vertex_location,
response=completion_response,
request_model=request_model_for_cost,
)
# Get additional costs from provider (e.g., routing fees, infrastructure costs)
if custom_llm_provider == "azure_ai":
if custom_llm_provider == "azure_ai" and not azure_ai_is_model_router_name(model):
model_for_additional_costs = request_model_for_cost
if completion_response is not None:
hidden_params = getattr(completion_response, "_hidden_params", None) or {}
@ -1735,6 +1740,7 @@ def completion_cost(
_reasoning_cost: float | None = None
_cache_read_cost: float | None = None
_cache_creation_cost: float | None = None
_billed_token_rates: BilledTokenRates | None = None
if cost_per_token_usage_object is not None and model:
_breakdown_provider: str | None = (
custom_llm_provider if isinstance(custom_llm_provider, str) else None
@ -1746,10 +1752,12 @@ def completion_cost(
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
custom_cost_per_token=custom_cost_per_token,
)
_reasoning_cost = _token_type_breakdown.reasoning_cost
_cache_read_cost = _token_type_breakdown.cache_read_cost
_cache_creation_cost = _token_type_breakdown.cache_creation_cost
_billed_token_rates = _token_type_breakdown.rates
_store_cost_breakdown_in_logging_obj(
litellm_logging_obj=litellm_logging_obj,
prompt_tokens_cost_usd_dollar=prompt_tokens_cost_usd_dollar,
@ -1769,6 +1777,7 @@ def completion_cost(
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
billed_token_rates=_billed_token_rates,
)
return _final_cost

View file

@ -18,6 +18,7 @@ from mcp import ClientSession, McpError, ReadResourceResult, Resource, StdioServ
from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client
from mcp.shared.message import SessionMessage
from mcp.shared.session import RequestResponder
from typing_extensions import Unpack
_TransportStreams: TypeAlias = tuple[
@ -56,10 +57,13 @@ def missing_streamable_http_client_error() -> ImportError:
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import (
ClientResult,
GetPromptRequestParams,
GetPromptResult,
Prompt,
ResourceTemplate,
ServerNotification,
ServerRequest,
TextContent,
)
from mcp.types import Tool as MCPTool
@ -146,8 +150,8 @@ _SDK_READ_TIMEOUT_CODE: Final = int(httpx.codes.REQUEST_TIMEOUT)
otherwise carries JSON-RPC error codes."""
def _as_read_timeout(exc: BaseException) -> TimeoutError | None:
"""The session read timeout elapsing, re-expressed as a ``TimeoutError``, or ``None``.
def as_mcp_read_timeout(exc: BaseException) -> TimeoutError | None:
"""Normalize an MCP SDK read timeout for client and gateway diagnostics, or return ``None``.
The SDK reports its own elapsed read timeout as ``McpError`` carrying an HTTP status code in a
field that otherwise holds JSON-RPC error codes, and it relays an upstream's JSON-RPC error
@ -442,6 +446,18 @@ class MCPClient:
in_flight_error: BaseException | None = None
try:
read_stream, write_stream = transport[0], transport[1]
stream_error: Final[asyncio.Future[Exception]] = asyncio.get_running_loop().create_future()
async def receive_message(
message: RequestResponder[ServerRequest, ClientResult] | ServerNotification | Exception,
) -> None:
if not isinstance(message, (ValueError, httpx.RequestError, OSError)):
return
if not stream_error.done():
stream_error.set_result(message)
# The SDK closes pending requests when its message handler raises.
raise RuntimeError("MCP response stream failed")
# Build session kwargs with optional callbacks
session_kwargs: Final[dict[str, Any]] = {}
if self._sampling_callback is not None:
@ -456,6 +472,7 @@ class MCPClient:
read_stream,
write_stream,
read_timeout_seconds=timedelta(seconds=self.timeout),
message_handler=receive_message,
**session_kwargs,
)
session: Final = await session_ctx.__aenter__()
@ -467,6 +484,10 @@ class MCPClient:
if isinstance(ins, str) and ins.strip():
self._last_initialize_instructions = ins.strip()
return await operation(session)
except McpError:
if stream_error.done():
raise stream_error.result()
raise
finally:
try:
await session_ctx.__aexit__(None, None, None)
@ -501,11 +522,10 @@ class MCPClient:
transport_ctx, http_client = self._create_transport_context()
return await self._execute_session_operation(transport_ctx, operation)
except Exception as e:
read_timeout: Final = _as_read_timeout(e)
read_timeout: Final = as_mcp_read_timeout(e)
if read_timeout is not None:
verbose_logger.warning(
"MCP client timed out after %ss waiting for %s to answer; the server accepted the "
"request and ended its response stream without a JSON-RPC reply",
"MCP client timed out after %ss waiting for a valid MCP response from %s",
self.timeout,
self.server_url or "stdio",
)

View file

@ -31,7 +31,7 @@ FileCreateProvider = Literal[
FileRetrieveProvider = Literal[
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic"
]
FileDeleteProvider = Literal["openai", "azure", "gemini", "litellm_proxy", "manus", "anthropic"]
FileDeleteProvider = Literal["openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic"]
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic"]
import litellm
from litellm import get_secret_str

View file

@ -21,6 +21,8 @@ from types import MappingProxyType
from typing import Final, TypeVar
from urllib.parse import urlparse
import httpx
from litellm._logging import verbose_logger
from litellm.integrations.batch_utils import (
BatchSendCancelled,
@ -418,7 +420,7 @@ class AzureSentinelLogger(CustomBatchLogger):
"Content-Type": "application/json",
}
async def _send_batch(batch: Sequence[_QueuedPayload]):
async def _send_batch(batch: Sequence[_QueuedPayload]) -> httpx.Response:
body: Final = safe_dumps(batch)
return await self.async_httpx_client.post(
url=api_endpoint,

View file

@ -367,6 +367,13 @@
"ui_name": "Headers",
"description": "Headers for OTEL exporter (e.g., x-honeycomb-team=YOUR_API_KEY)",
"required": false
},
"otel_exporter_otlp_protocol": {
"type": "select",
"ui_name": "Export Protocol",
"description": "OTLP wire format for trace exports. Use http/json for collectors that cannot decode protobuf",
"options": ["http/protobuf", "http/json"],
"required": false
}
},
"description": "OpenTelemetry Logging Integration"

View file

@ -850,20 +850,24 @@ class CustomGuardrail(CustomLogger):
if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True:
return None
# CHECK IF GUARDRAIL REJECTS THE REQUEST
target: Final = self._deployment_hook_target()
hook_request_data: Final = {**request_data, "guardrail_to_apply": self} if target is not self else request_data
result: Final = await target.async_post_call_success_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id=request_data.get("user_api_key_user_id"),
team_id=request_data.get("user_api_key_team_id"),
end_user_id=request_data.get("user_api_key_end_user_id"),
api_key=request_data.get("user_api_key_hash"),
request_route=request_data.get("user_api_key_request_route"),
),
data=hook_request_data,
response=response,
)
try:
if target is not self:
request_data["guardrail_to_apply"] = self # rebind-ok: dispatch consumes this key
result: Final = await target.async_post_call_success_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id=request_data.get("user_api_key_user_id"),
team_id=request_data.get("user_api_key_team_id"),
end_user_id=request_data.get("user_api_key_end_user_id"),
api_key=request_data.get("user_api_key_hash"),
request_route=request_data.get("user_api_key_request_route"),
),
data=request_data,
response=response,
)
finally:
if target is not self:
request_data.pop("guardrail_to_apply", None)
if not self._is_valid_response_type(result):
return None

View file

@ -133,17 +133,17 @@ class MlflowLogger(CustomLogger):
if final_response:
end_time_ns: Final = int(end_time.timestamp() * 1e9)
self._extract_and_set_chat_attributes(span, kwargs, final_response)
self._end_span_or_trace(
span=span,
outputs=final_response,
status=SpanStatusCode.OK,
end_time_ns=end_time_ns,
)
# Remove the stream_id from the map
with self._lock:
self._stream_id_to_span.pop(litellm_call_id)
try:
self._extract_and_set_chat_attributes(span, kwargs, final_response)
self._end_span_or_trace(
span=span,
outputs=final_response,
status=SpanStatusCode.OK,
end_time_ns=end_time_ns,
)
finally:
with self._lock:
self._stream_id_to_span.pop(litellm_call_id, None)
def _add_chunk_events(self, span, response_obj):
from mlflow.entities import SpanEvent
@ -282,15 +282,15 @@ class MlflowLogger(CustomLogger):
"""End an MLflow span or a trace."""
if span.parent_id is None:
self._client.end_trace(
trace_id=span.request_id,
span.request_id,
outputs=outputs,
status=status,
end_time_ns=end_time_ns,
)
else:
self._client.end_span(
trace_id=span.request_id,
span_id=span.span_id,
span.request_id,
span.span_id,
outputs=outputs,
status=status,
end_time_ns=end_time_ns,

View file

@ -16,9 +16,11 @@ from opentelemetry.trace import (
Span,
Tracer,
get_current_span,
get_tracer_provider,
set_span_in_context,
use_span,
)
from opentelemetry.trace import TracerProvider as ApiTracerProvider
import litellm
from litellm._logging import verbose_logger
@ -63,6 +65,7 @@ from litellm.integrations.otel.plumbing.metrics import (
create_genai_metrics,
)
from litellm.integrations.otel.plumbing.providers import (
attach_tenant_fan_out,
build_tracer_provider,
get_event_logger,
get_meter,
@ -85,6 +88,7 @@ if TYPE_CHECKING:
)
LITELLM_TRACER_NAME: Final = "litellm"
_published_v2_provider: ApiTracerProvider | None = None
def _span_error_from_exception(
@ -180,7 +184,9 @@ class OpenTelemetryV2(CustomLogger):
self.config: OpenTelemetryV2Config = config or OpenTelemetryV2Config(**kwargs)
self.callback_name = callback_name
self._tracer_provider: TracerProvider = (
tracer_provider if tracer_provider is not None else build_tracer_provider(self.config)
tracer_provider
if tracer_provider is not None
else build_tracer_provider(self.config, tenant_overrides=True)
)
self.tracer: Tracer = get_tracer(self._tracer_provider, LITELLM_TRACER_NAME)
self._metrics_recorder = self._init_metrics(meter_provider)
@ -195,6 +201,11 @@ class OpenTelemetryV2(CustomLogger):
self._open_llm_calls: OrderedDict[str, _LLMCallSpan] = OrderedDict()
self._init_otel_logger_on_litellm_proxy()
@property
def tracer_provider(self) -> TracerProvider:
"""The provider this logger emits through, read-only to its callers."""
return self._tracer_provider
def _init_metrics(self, meter_provider: "MeterProvider | None") -> "GenAIMetricRecorder | None":
"""Create the six GenAI histograms when metrics are enabled, else ``None``.
@ -863,12 +874,33 @@ def publish_global_otel_v2_provider(
``opentelemetry.trace.set_tracer_provider``) are injected so the publish step is
unit-testable without reading or mutating real global OTel state. Returns the
logger whose provider was published.
The published provider is also the one that fans spans out to key/team
destinations, because it is the only provider the whole request tree passes
through; see :func:`attach_tenant_fan_out`. It is remembered for
:func:`fan_out_provider` because neither the OTel global (``set_tracer_provider``
keeps the first provider it was ever handed) nor
``proxy_server.open_telemetry_logger`` (a legacy v1 logger can hold that slot)
reliably leads back to it.
"""
global _published_v2_provider
logger: Final = select_global_otel_v2_logger(in_memory_loggers, registered=registered)
set_global_provider(logger._tracer_provider)
attach_tenant_fan_out(logger.tracer_provider, *_v2_configs(in_memory_loggers, logger))
set_global_provider(logger.tracer_provider)
_published_v2_provider = logger.tracer_provider # rebind-ok: startup records the one provider carrying the fan-out
return logger
def _v2_configs(in_memory_loggers: Sequence[object], logger: "OpenTelemetryV2") -> tuple[OpenTelemetryV2Config, ...]:
"""Every v2 logger's config, the published logger's first.
Each preset keeps its own provider and exporters, so the accounts the operator
writes to are spread over all of them, not held by the published logger alone.
"""
others: Final = tuple(cb.config for cb in in_memory_loggers if isinstance(cb, OpenTelemetryV2) and cb is not logger)
return (logger.config, *others)
def _registered_v2_logger() -> "OpenTelemetryV2 | None":
try:
from litellm.proxy import proxy_server
@ -904,6 +936,25 @@ def seed_request_identity(user_api_key_dict: object, model: str | None = None) -
logger.seed_request_identity(user_api_key_dict, model=model)
def fan_out_provider() -> ApiTracerProvider:
"""The provider :func:`publish_global_otel_v2_provider` gave the tenant fan-out.
Read off the publish itself, not the OTel global and not the registered logger:
the global keeps whichever provider claimed it first (auto-instrumentation, a
legacy logger), and the registered slot can hold a v1 logger while the publish
picked a v2 one from ``_in_memory_loggers``. Either detour lands on a provider
with no fan-out and drops every destination at auth.
"""
published: Final = _published_v2_provider
if published is not None:
return published
logger: Final = _registered_v2_logger()
if logger is not None:
attach_tenant_fan_out(logger.tracer_provider, logger.config)
return logger.tracer_provider
return get_tracer_provider()
@contextmanager
def phase_span(name: str) -> "Iterator[Span | None]":
logger: Final = _registered_v2_logger()

View file

@ -23,6 +23,7 @@ from litellm.integrations.otel.model.payloads import (
ServiceSpanData,
ToolDefinition,
)
from litellm.integrations.otel.model.semconv import Error
# Attribute keys in the semconv-ai / Traceloop vocabulary.
_LEGACY_SYSTEM: Final = "gen_ai.system"
@ -36,7 +37,7 @@ _LEGACY_PRESENCE_PENALTY: Final = "llm.presence_penalty"
_LEGACY_STOP_SEQUENCES: Final = "llm.chat.stop_sequences"
_LEGACY_SERVICE: Final = "service"
_LEGACY_CALL_TYPE: Final = "call_type"
_LEGACY_ERROR: Final = "error"
_LEGACY_ERROR: Final = Error.MESSAGE_LEGACY
class LegacyMapper:

View file

@ -69,7 +69,7 @@ class ExporterSpec(BaseModel):
kind: str = Field(
default="console",
description="console | in_memory | otlp_http | otlp_grpc | <factory kind>",
description="console | in_memory | otlp_http | http/json | otlp_grpc | <factory kind>",
)
endpoint: str | None = None
traces_endpoint: str | None = Field(
@ -269,7 +269,9 @@ class OpenTelemetryV2Config(BaseSettings):
if (self.endpoint or self.traces_endpoint) and self.exporter == "console":
self.exporter = "otlp_http"
# When no explicit destinations are given, fold the single-destination
# shorthand into one spec so the provider always has a destination.
# shorthand into one spec so the provider always has a destination. A spec
# with no fields set is how the presets tell "nothing configured" from an
# operator who asked for the console by name.
if not self.exporters:
self.exporters = [
ExporterSpec(
@ -278,6 +280,8 @@ class OpenTelemetryV2Config(BaseSettings):
traces_endpoint=self.traces_endpoint,
headers=self.headers,
)
if not self.model_fields_set.isdisjoint(("exporter", "endpoint", "headers"))
else ExporterSpec()
]
# Ensure ``genai`` is always present and first.
names = list(self.mapper_names)

View file

@ -0,0 +1,49 @@
"""The resolved OTLP destination a request's traces export to.
Backend-agnostic on purpose: every OTEL backend reduces to an endpoint plus auth
headers. The per-backend field mapping lives in ``presets.destinations``.
"""
from collections.abc import Mapping
from typing import Final
from urllib.parse import quote
from pydantic import BaseModel, ConfigDict, Field
class OtelDestination(BaseModel):
model_config = ConfigDict(frozen=True)
endpoint: str
headers: Mapping[str, str] = Field(default_factory=dict)
resource_attributes: Mapping[str, str] = Field(default_factory=dict)
callback_name: str | None = None
protocol: str | None = Field(
default=None,
description=(
"OTLP transport, defaulting to the backend's own. Not derivable from the "
"scheme: Arize's ``https://otlp.arize.com/v1`` is gRPC."
),
)
def header_string(self) -> str:
"""Render headers as the ``k=v,k2=v2`` form an ``ExporterSpec`` expects.
Values are percent-encoded because ``providers.parse_headers`` decodes them
with the SDK's W3C-Baggage parser: a value carrying a ``,`` or ``=`` (a
Langfuse project name, a base64 Authorization payload ending in ``==``)
would otherwise be split into bogus pairs on the way back out.
"""
return ",".join(f"{key}={quote(value, safe='')}" for key, value in self.headers.items())
def cache_key(self) -> tuple[str, tuple[tuple[str, str], ...], tuple[tuple[str, str], ...], str | None]:
"""Identity for processor reuse, so one destination means one exporter."""
return (
self.endpoint,
tuple(sorted(self.headers.items())),
tuple(sorted(self.resource_attributes.items())),
self.protocol,
)
NO_DESTINATIONS: Final[tuple[OtelDestination, ...]] = ()

View file

@ -204,6 +204,9 @@ class Error:
TYPE: Final = "error.type"
MESSAGE: Final = "error.message"
# The same text under the bare key the semconv-ai / Traceloop vocabulary uses
# (see ``LegacyMapper``), so anything reading or redacting error text covers both.
MESSAGE_LEGACY: Final = "error"
class LiteLLMError:

View file

@ -1,8 +1,9 @@
"""Trace-context + Baggage helpers."""
import os
from collections.abc import Mapping
from contextvars import ContextVar, Token
from typing import Final
from typing import TYPE_CHECKING, Final
from opentelemetry import baggage
from opentelemetry.context import Context, get_current
@ -21,6 +22,9 @@ from opentelemetry.trace.propagation.tracecontext import (
from litellm.integrations.otel.model.semconv import HTTP
if TYPE_CHECKING:
from litellm.integrations.otel.model.destination import OtelDestination
_PROPAGATOR: Final = TraceContextTextMapPropagator()
# The request's root span — the FastAPI-owned SERVER span — captured ONCE when the
@ -304,3 +308,65 @@ def extract_traceparent(headers: Mapping[str, str]) -> Context | None:
return None
carrier: Final = {str(key).lower(): value for key, value in headers.items()}
return _PROPAGATOR.extract(carrier)
# The OTLP destinations this request's key or team pointed its traces at, resolved
# once during auth. A ``ContextVar`` for the same reason the root span above is one:
# it rides the request task's context into the ``asyncio.create_task`` children that
# close the LLM span, and it is visible to every ``SpanProcessor.on_end`` that fires
# on the request task. Stateful MCP handlers set and reset it per message; the
# request-task value otherwise dies with that task.
_request_destinations: Final['ContextVar[tuple["OtelDestination", ...]]'] = ContextVar(
"litellm_otel_request_destinations", default=()
)
def set_request_destinations(destinations: 'tuple["OtelDestination", ...]') -> "Token[tuple[OtelDestination, ...]]":
"""Anchor the destinations this request exports to and return a reset token."""
return _request_destinations.set(destinations)
def reset_request_destinations(token: "Token[tuple[OtelDestination, ...]]") -> None:
_request_destinations.reset(token)
def request_destinations() -> 'tuple["OtelDestination", ...]':
"""The destinations resolved for this request, empty outside a proxy request."""
return _request_destinations.get()
#: ``litellm_settings: otel_tenant_destination_mode`` and its env equivalent.
ADDITIVE_DESTINATION_MODE: Final = "additive"
OTEL_TENANT_DESTINATION_MODE_ENV: Final = "LITELLM_OTEL_TENANT_DESTINATION_MODE"
def tenant_destinations_are_additive() -> bool:
"""Whether a tenant destination exports alongside the operator's own exporter.
Override is the default: the tenant's traffic reaches the tenant's account and
nowhere else. Operators running one org-wide backend across every team set this
to ``additive`` so the same trace lands in both places.
"""
import litellm
configured: Final = litellm.otel_tenant_destination_mode or os.environ.get(OTEL_TENANT_DESTINATION_MODE_ENV)
return isinstance(configured, str) and configured.strip().lower() == ADDITIVE_DESTINATION_MODE
def destination_backends() -> frozenset[str]:
"""Backends this request resolved a tenant destination for.
The fan-out already carries the whole trace to those destinations, so the
per-request tracer route must never send a second copy, in either mode.
"""
return frozenset(d.callback_name for d in _request_destinations.get() if d.callback_name)
def suppressed_backends() -> frozenset[str]:
"""Backends whose operator-level exporters this request must NOT reach.
Empty under ``additive``, where the operator keeps its copy of every span.
"""
if tenant_destinations_are_additive():
return frozenset()
return destination_backends()

View file

@ -0,0 +1,70 @@
"""OTLP/HTTP span exporter that sends the OTLP/JSON encoding instead of protobuf.
The SDK only ships a protobuf OTLP/HTTP exporter; this reuses its transport and
retry loop and swaps the payload for OTLP/JSON (enums as integers, ids as hex).
"""
import base64
import json
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final, TypeAlias
from google.protobuf.json_format import MessageToDict
from opentelemetry.exporter.otlp.proto.common.trace_encoder import encode_spans
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
from opentelemetry.sdk.trace import ReadableSpan
JSON_CONTENT_TYPE: Final = "application/json"
_HEX_ID_KEYS: Final = frozenset({"traceId", "spanId", "parentSpanId"})
_JsonValue: TypeAlias = "Mapping[str, _JsonValue] | Sequence[_JsonValue] | str | int | float | bool | None"
_JsonObject: TypeAlias = Mapping[str, "_JsonValue"]
def _objects(node: _JsonObject, key: str) -> tuple[_JsonObject, ...]:
items: Final = node.get(key)
if isinstance(items, str) or not isinstance(items, Sequence):
return ()
return tuple(item for item in items if isinstance(item, Mapping))
def _hex_ids(node: _JsonObject) -> _JsonObject:
return MappingProxyType(
{
key: base64.b64decode(item).hex() if key in _HEX_ID_KEYS and isinstance(item, str) else item
for key, item in node.items()
}
)
def _hex_span(span: _JsonObject) -> _JsonObject:
links: Final = _objects(span, "links")
if not links:
return _hex_ids(span)
return MappingProxyType({**_hex_ids(span), "links": tuple(_hex_ids(link) for link in links)})
def _hex_scope_spans(scope: _JsonObject) -> _JsonObject:
return MappingProxyType({**scope, "spans": tuple(_hex_span(span) for span in _objects(scope, "spans"))})
def _hex_resource_spans(resource: _JsonObject) -> _JsonObject:
scope_spans: Final = tuple(_hex_scope_spans(scope) for scope in _objects(resource, "scopeSpans"))
return MappingProxyType({**resource, "scopeSpans": scope_spans})
def encode_spans_json(spans: Sequence[ReadableSpan]) -> bytes:
payload: Final[_JsonObject] = MessageToDict(encode_spans(spans), use_integers_for_enums=True)
resource_spans: Final = tuple(_hex_resource_spans(resource) for resource in _objects(payload, "resourceSpans"))
hexed: Final[_JsonObject] = MappingProxyType({**payload, "resourceSpans": resource_spans})
return json.dumps(hexed, default=dict, separators=(",", ":")).encode()
class OTLPJsonSpanExporter(OTLPSpanExporter):
def __init__(self, endpoint: str | None, headers: dict[str, str]) -> None: # mutable-ok: SDK __init__ takes Dict
super().__init__(endpoint=endpoint, headers=headers)
self._session.headers["Content-Type"] = JSON_CONTENT_TYPE
def _serialize_spans(self, spans: Sequence[ReadableSpan]) -> bytes:
return encode_spans_json(spans)

View file

@ -1,9 +1,14 @@
"""Provider / exporter factory + the Baggage span processor."""
from collections.abc import Callable, Iterable
import queue
import threading
import time
from collections import OrderedDict
from collections.abc import Callable, Iterable, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
from opentelemetry import _logs, baggage, metrics
from opentelemetry import _logs, baggage, metrics, trace
from opentelemetry._events import EventLogger
from opentelemetry._logs import LoggerProvider, NoOpLoggerProvider
from opentelemetry.context import Context
@ -19,7 +24,8 @@ from opentelemetry.sdk._logs.export import (
)
from opentelemetry.sdk.metrics import MeterProvider as SDKMeterProvider
from opentelemetry.sdk.resources import Resource
from opentelemetry.sdk.trace import ReadableSpan, SpanProcessor, TracerProvider
from opentelemetry.sdk.trace import Event, ReadableSpan, SpanProcessor, TracerProvider
from opentelemetry.sdk.trace import Span as SDKSpan
from opentelemetry.sdk.trace.export import (
BatchSpanProcessor,
ConsoleSpanExporter,
@ -29,18 +35,35 @@ from opentelemetry.sdk.trace.export import (
from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
InMemorySpanExporter,
)
from opentelemetry.trace import Span, SpanKind, Tracer
from opentelemetry.trace import Span, SpanKind, Status, Tracer
from opentelemetry.util.re import parse_env_headers
from opentelemetry.util.types import Attributes, AttributeValue
from litellm._logging import verbose_logger
from litellm._version import version as litellm_version
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.integrations.otel.model.semconv import LiteLLM
from litellm.integrations.otel.model.semconv import (
DB,
MCP,
Error,
ExceptionEvent,
GenAI,
LiteLLM,
LiteLLMError,
Server,
)
from litellm.integrations.otel.model.spans import LiteLLMSpanKind
from litellm.integrations.otel.plumbing.context import (
request_destinations,
suppressed_backends,
)
if TYPE_CHECKING:
from opentelemetry.metrics import Meter
from opentelemetry.sdk.metrics.export import MetricReader
from litellm.integrations.otel.model.destination import OtelDestination
_SPAN_KIND_BY_ROLE_KIND: Final[dict[LiteLLMSpanKind, SpanKind]] = {
LiteLLMSpanKind.SERVER: SpanKind.SERVER,
LiteLLMSpanKind.CLIENT: SpanKind.CLIENT,
@ -136,7 +159,8 @@ def parse_headers(raw: str | None) -> dict[str, str]:
_IN_MEMORY_KINDS: Final = ("in_memory", "inmemory", "memory")
_OTLP_HTTP_KINDS: Final = ("otlp_http", "http", "http/protobuf", "http/json")
_OTLP_HTTP_JSON_KINDS: Final = ("http/json",)
_OTLP_HTTP_KINDS: Final = ("otlp_http", "http", "http/protobuf", *_OTLP_HTTP_JSON_KINDS)
_OTLP_GRPC_KINDS: Final = ("otlp_grpc", "grpc")
@ -164,6 +188,13 @@ def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter:
return factory(spec)
if kind in _IN_MEMORY_KINDS:
return InMemorySpanExporter()
if kind in _OTLP_HTTP_JSON_KINDS:
from litellm.integrations.otel.plumbing.otlp_json import OTLPJsonSpanExporter
return OTLPJsonSpanExporter(
endpoint=spec.traces_endpoint or _otlp_traces_endpoint(spec.endpoint),
headers=parse_headers(spec.headers),
)
if kind in _OTLP_HTTP_KINDS:
from opentelemetry.exporter.otlp.proto.http.trace_exporter import (
OTLPSpanExporter as HTTPExporter,
@ -194,6 +225,555 @@ def _processor_for(exporter: SpanExporter, use_simple: bool | None) -> SpanProce
return SimpleSpanProcessor(exporter) if use_simple else BatchSpanProcessor(exporter)
#: Distinct tenant destinations whose exporters stay alive. Each holds a connection
#: pool and a batch thread, so the cache is bounded and evicts least-recently-used.
_MAX_CACHED_DESTINATION_PROCESSORS: Final = 32
#: Workers closing shed destination processors, bounding the threads a tenant can
#: create by cycling its destination config.
_DRAIN_WORKERS: Final = 2
#: Shed processors waiting to be closed before the fan-out stops building new ones.
#: Each still owns a batch thread until its close returns, and a collector that never
#: answers makes every close take the exporter's full timeout, so past this many the
#: operator's exporter keeps the span instead (see ``deliverable``).
_MAX_PENDING_DRAINS: Final = 64
#: How long ``shutdown`` waits for spans already being forwarded, so teardown closes
#: no processor under one. Bounded: an exporter that never returns must not hold the
#: proxy open.
_SHUTDOWN_DRAIN_SECONDS: Final = 5.0
#: An exporter's account: its normalized endpoint and the credentials it presents.
_SinkKey = tuple[str, tuple[tuple[str, str], ...]]
#: Header names that spell one credential two ways. Arize's operator exporter sends
#: ``space_id`` where a tenant destination sends ``arize-space-id``.
_CREDENTIAL_ALIASES: Final = MappingProxyType({"arize_space_id": "space_id"})
class _DrainPool:
"""Closes shed destination processors off the span-export path.
``shutdown`` flushes over the network and is reached from ``on_end``, so closing
one inline would let a single unreachable tenant collector stall every other
tenant's spans behind it. A fixed set of workers rather than a thread per
processor means a tenant cycling its destination config cannot spawn threads as
fast as it can send requests; slow shutdowns queue behind each other.
The workers are daemons and belong to the fan-out that sheds the processors, so
neither an unreachable collector nor a lazily built process-wide singleton can
hold the proxy open on the way down.
"""
def __init__(
self,
workers: int = _DRAIN_WORKERS,
pending: "queue.Queue[SpanProcessor | None] | None" = None,
capacity: int = _MAX_PENDING_DRAINS,
) -> None:
self._workers: Final = workers
self._capacity: Final = capacity
self._lock: Final = threading.Lock()
self._closed = False
self._backlog = 0 # guarded by ``_lock``: submitted processors whose close has not returned
self._pending: Final[queue.Queue[SpanProcessor | None]] = pending if pending is not None else queue.Queue()
self._threads: Final = tuple(
threading.Thread(target=self._drain_until_closed, daemon=True, name="litellm-otel-destination-drain")
for _ in range(workers)
)
for worker in self._threads:
worker.start()
def submit(self, processor: SpanProcessor) -> None:
"""Queue ``processor`` for closing, or hand it off once the pool is retired.
The check and the put share one lock. Reading a closed flag on its own leaves
room for :meth:`close` to run in between, and the processor would land behind
the sentinels every worker has already exited on.
Past close there is no worker left to take it, and the caller is whichever
thread just ended a span, so closing it inline would park that thread on a
network flush the shutdown deadline has already stopped waiting for. The extra
thread is bounded by the same close: the fan-out stops handing processors out
at that point, so only the ones already exporting when it happened arrive here.
"""
with self._lock:
if not self._closed:
self._backlog += 1
self._pending.put(processor)
return
threading.Thread(
target=_shutdown_quietly,
args=(processor,),
daemon=True,
name="litellm-otel-destination-drain-straggler",
).start()
def saturated(self) -> bool:
"""Whether enough closes are outstanding that building another processor must wait.
The workers close in order and each close blocks for as long as its exporter
does, so a collector that stopped answering would otherwise turn every new
destination into one more batch thread parked behind them, for as long as the
tenants keep rotating. Holding the count here rather than reading the queue
keeps the two processors a worker is mid-close on in the total.
"""
with self._lock:
return self._backlog >= self._capacity
def close(self, timeout: float | None = None) -> None:
"""Retire the workers once they have closed everything already queued.
A proxy that rebuilds its telemetry builds another fan-out, so workers that
outlive the one that started them are two more threads per reload, forever.
``timeout`` bounds how long the caller waits for that draining to finish. The
workers are daemons, so whatever is still flushing when it expires is dropped
by the interpreter rather than holding it open.
"""
with self._lock:
if self._closed:
return
self._closed = True
for _ in range(self._workers):
self._pending.put(None)
if timeout is None:
return
deadline: Final = time.monotonic() + timeout
for worker in self._threads:
worker.join(timeout=max(0.0, deadline - time.monotonic()))
def _drain_until_closed(self) -> None:
while True:
processor: SpanProcessor | None = self._pending.get() # rebind-ok: loop variable
if processor is None:
return
_shutdown_quietly(processor)
with self._lock:
self._backlog -= 1
_NO_ATTRIBUTES: Final[Mapping[str, AttributeValue]] = MappingProxyType({})
_DB_SYSTEM_KEYS: Final = frozenset({DB.SYSTEM_NAME, DB.SYSTEM_LEGACY})
# Keys on a database span that describe the proxy's own datastore: its host, its
# port, and its schema.
_DATASTORE_ENDPOINT_KEYS: Final = frozenset({Server.ADDRESS, Server.PORT, DB.NAMESPACE})
# A span carrying one of these describes the tenant's own call (the model call, the
# MCP call, the guardrail), so its error text is theirs to see. Every other span is
# the proxy's own work, whose error text names the operator's infrastructure.
_TENANT_OWNED_KEYS: Final = frozenset({GenAI.OPERATION_NAME, MCP.METHOD_NAME, LiteLLM.GUARDRAIL_NAME})
_PROXY_ERROR_TEXT_KEYS: Final = frozenset({Error.MESSAGE, Error.MESSAGE_LEGACY})
# A guardrail that never answered carries the exception it raised as its response,
# which names the operator's guardrail endpoint. The second spelling is the legacy
# status the request-level logger still maps.
_GUARDRAIL_UNREACHABLE_STATUSES: Final = frozenset({"guardrail_failed_to_respond", "failure"})
# Attribute prefixes the FastAPI instrumentor uses for headers the operator opted to
# capture (``OTEL_INSTRUMENTATION_HTTP_CAPTURE_HEADERS_SERVER_*``). The request
# side carries the caller's bearer token verbatim.
_CAPTURED_HEADER_PREFIXES: Final = ("http.request.header.", "http.response.header.")
# The instrumentor stamps the request URL on the server span with its query string,
# under the old convention and the new one, and litellm accepts a virtual key as a
# ``?key=`` query parameter.
_URL_KEYS: Final = frozenset({"http.url", "http.target", "url.full"})
_URL_QUERY_KEY: Final = "url.query"
class _TenantSpanView(ReadableSpan):
"""A ``ReadableSpan`` view for one destination, leaving the operator's own span alone."""
def __init__(
self,
inner: ReadableSpan,
resource: Resource,
attributes: Attributes,
events: Sequence[Event],
status: Status,
) -> None:
super().__init__(
name=inner.name,
context=inner.context,
parent=inner.parent,
resource=resource,
attributes=attributes,
events=events,
links=inner.links,
kind=inner.kind,
status=status,
start_time=inner.start_time,
end_time=inner.end_time,
instrumentation_scope=inner.instrumentation_scope,
)
def _is_database_span(attributes: Mapping[str, AttributeValue]) -> bool:
return any(key in attributes for key in _DB_SYSTEM_KEYS)
def _is_tenant_owned_span(attributes: Mapping[str, AttributeValue]) -> bool:
return any(key in attributes for key in _TENANT_OWNED_KEYS)
def _guardrail_unreachable(attributes: Mapping[str, AttributeValue]) -> bool:
return attributes.get(LiteLLM.GUARDRAIL_STATUS) in _GUARDRAIL_UNREACHABLE_STATUSES
def _tenant_visible(key: str, database: bool, owned: bool, unreachable_guardrail: bool) -> bool:
if key.startswith(_CAPTURED_HEADER_PREFIXES) or key in (LiteLLMError.STACK_TRACE, _URL_QUERY_KEY):
return False
if database and key in _DATASTORE_ENDPOINT_KEYS:
return False
if unreachable_guardrail and key == LiteLLM.GUARDRAIL_RESPONSE:
return False
return owned or key not in _PROXY_ERROR_TEXT_KEYS
def _without_query(key: str, value: AttributeValue) -> AttributeValue:
if key not in _URL_KEYS or not isinstance(value, str):
return value
return value.partition("?")[0]
def _same_attributes(kept: Mapping[str, AttributeValue], attributes: Mapping[str, AttributeValue]) -> bool:
return len(kept) == len(attributes) and all(kept[key] is value for key, value in attributes.items())
def _without_stack_trace(event: Event) -> Event:
attributes: Final = event.attributes or _NO_ATTRIBUTES
if ExceptionEvent.STACKTRACE not in attributes:
return event
return Event(
name=event.name,
attributes=MappingProxyType(
{key: value for key, value in attributes.items() if key != ExceptionEvent.STACKTRACE}
),
timestamp=event.timestamp,
)
def _for_destination(span: ReadableSpan, destination: "OtelDestination") -> ReadableSpan:
"""The view of ``span`` a tenant destination receives.
A span the tenant's own call produced keeps its error text. Every other span is
the proxy's own work (the request root, auth, the database), and its error text,
its events and its status description come off, since a Prisma failure there
spells out the operator's Postgres endpoint. A database span loses that endpoint
too, and a guardrail that failed to respond loses its response text, which is the
exception it raised and names the operator's guardrail endpoint. Stack traces walk
the operator's install and come off every span, as do the headers the operator
captures on the server span, whose request side holds the caller's bearer token,
and the query string of the request URL, which can hold the same key. The span
itself stays, so the tenant still gets the whole trace tree.
"""
extra: Final = destination.resource_attributes
attributes: Final = span.attributes or _NO_ATTRIBUTES
database: Final = _is_database_span(attributes)
owned: Final = _is_tenant_owned_span(attributes)
unreachable: Final = _guardrail_unreachable(attributes)
kept: Final = MappingProxyType(
{
key: _without_query(key, value)
for key, value in attributes.items()
if _tenant_visible(key, database, owned, unreachable)
}
)
recorded: Final = span.events
events: Final = tuple(_without_stack_trace(event) for event in recorded) if owned else ()
unchanged: Final = owned and _same_attributes(kept, attributes) and all(a is b for a, b in zip(events, recorded))
if not extra and unchanged:
return span
resource: Final = span.resource.merge(Resource(extra)) if extra else span.resource
status: Final = span.status if owned else Status(span.status.status_code)
return _TenantSpanView(span, resource, kept, events, status)
class TenantFanOutSpanProcessor(SpanProcessor):
"""Export every finished span to each destination this request resolved.
Destinations ride a request-scoped ``ContextVar`` set during auth, so concurrent
requests stay isolated. The forwarded view keeps the original trace and parent
ids, so the tenant gets the same tree the operator would have received.
Exactly one provider carries this processor, the one published as the OTel global
(see :func:`attach_tenant_fan_out`). That provider is the only one every span
passes through: the FastAPI server span, the auth span and the post-call database
spans are emitted on the global, while a second v2 logger's provider sees only
that logger's own gen-AI span. Attaching the fan-out per logger would hand a
tenant a one-span trace whenever its backend is not the global one, and two
copies of the model call whenever it is.
"""
def __init__(
self,
processor_factory: 'Callable[["OtelDestination"], SpanProcessor | None] | None' = None,
shutdown_drain_seconds: float = _SHUTDOWN_DRAIN_SECONDS,
operator_sinks: frozenset[_SinkKey] = frozenset(),
pending_drains: int = _MAX_PENDING_DRAINS,
drain_pool: _DrainPool | None = None,
) -> None:
self._operator_sinks: Final = operator_sinks
self._drain_seconds: Final = shutdown_drain_seconds
self._lock: Final = threading.Condition()
self._closed = False # guarded by ``_lock``: an unlocked read races the teardown it gates
self._build: Final = processor_factory if processor_factory is not None else _destination_processor
self._processors: OrderedDict[object, SpanProcessor] = OrderedDict() # mutable-ok: bounded LRU
self._retired: OrderedDict[int, SpanProcessor] = OrderedDict() # mutable-ok: drains as exports finish
self._exporting: dict[int, int] = {} # mutable-ok: per-processor in-flight export count
self._drain: Final = drain_pool if drain_pool is not None else _DrainPool(capacity=pending_drains)
def on_start(self, span: SDKSpan, parent_context: Context | None = None) -> None:
return None
def on_end(self, span: ReadableSpan) -> None:
suppressed: Final = suppressed_backends()
for destination in request_destinations():
if self._operator_already_writes(destination, suppressed):
continue
processor = self._acquire(destination) # rebind-ok: loop variable; pyright forbids Final in a loop
if processor is None:
continue
try:
processor.on_end(_for_destination(span, destination))
except Exception as exc: # noqa: BLE001 # one destination's failure must not cost the others their span
verbose_logger.debug("OTel V2 fan-out: forwarding to %s failed: %s", destination.endpoint, exc)
finally:
self._release(processor)
def _operator_already_writes(self, destination: "OtelDestination", suppressed: frozenset[str]) -> bool:
"""Whether the operator's own exporter is sending this span to the same account.
Only reachable under ``additive``, where nothing is suppressed: a team that
names the operator's own project would otherwise have every span written
there twice, once by the operator's exporter and once by the fan-out.
"""
return (
destination.callback_name not in suppressed
and _sink_key(destination.endpoint, destination.headers) in self._operator_sinks
)
def shutdown(self) -> None:
"""Close every destination processor, once the spans in flight have landed.
``on_end`` runs on whichever thread ends a span and can reach this fan-out
while the SDK is tearing the provider down, so closing blind would drop a
trace mid-forward and would hand the next caller a fresh exporter nothing
will ever close. Refusing new work and then waiting out the in-flight ones
keeps both from happening. A straggler past the bound is retired instead of
closed: the thread still exporting it closes it through the drain as soon as
its export returns, so no span is dropped mid-forward.
Every close then goes to the drain rather than running here. Closing a
destination processor flushes it over the network and the SDK joins its own
worker with no timeout of its own, so one tenant collector that answers but
never finishes a response would otherwise hold process teardown open for as
long as it likes. The drain's workers are daemons, and the whole teardown
shares one deadline.
"""
deadline: Final = time.monotonic() + self._drain_seconds
with self._lock:
self._closed = True
self._lock.wait_for(lambda: not self._exporting, timeout=self._drain_seconds)
live: Final = tuple((id(p), p) for p in (*self._processors.values(), *self._retired.values()))
closing: Final = tuple(p for ident, p in live if ident not in self._exporting)
self._processors.clear()
self._retired = OrderedDict( # mutable-ok: the same bounded map, keeping only what is still exporting
(ident, p) for ident, p in live if ident in self._exporting
)
for processor in closing:
self._drain.submit(processor)
self._drain.close(timeout=max(0.0, deadline - time.monotonic()))
def force_flush(self, timeout_millis: int = 30000) -> bool:
results: Final = tuple(self._flush_one(processor, timeout_millis) for processor in self._snapshot())
return all(results)
def _snapshot(self) -> tuple[SpanProcessor, ...]:
with self._lock:
return (*self._processors.values(), *self._retired.values())
@staticmethod
def _flush_one(processor: SpanProcessor, timeout_millis: int) -> bool:
try:
return processor.force_flush(timeout_millis)
except Exception: # noqa: BLE001 # one exporter's flush failure must not fail the whole flush
return False
def deliverable(self, destinations: Iterable["OtelDestination"]) -> tuple["OtelDestination", ...]:
"""The subset of ``destinations`` this fan-out can actually export to.
A destination whose exporter will not build (a protocol whose package is not
installed, a malformed endpoint) has to be dropped before the request anchors
it, not when its first span ends. By then the operator's own exporter has been
told to hold that backend's spans back for this request, so dropping there
loses the span outright instead of leaving it where it would have gone with no
override at all.
"""
return tuple(destination for destination in destinations if self._buildable(destination))
def _buildable(self, destination: "OtelDestination") -> bool:
"""Whether a processor for ``destination`` exists or can be built right now."""
with self._lock:
if self._closed:
return False
built: Final = self._cached_or_built_locked(destination, anchored=False)
drained: Final = self._drainable_locked()
for shed in drained:
self._drain.submit(shed)
return built is not None
def _acquire(self, destination: "OtelDestination") -> SpanProcessor | None:
"""The processor for ``destination``, marked busy until ``_release``.
The build happens under the same lock that reads the cache, so a cold cache
met by a burst of concurrent requests yields one exporter rather than one per
thread with all but the winner shed. Building an exporter opens no connection,
so the cost of holding the lock is a constructor, once per destination.
"""
with self._lock:
if self._closed:
return None
processor: Final = self._cached_or_built_locked(destination, anchored=True)
if processor is None:
return None
self._exporting[id(processor)] = self._exporting.get(id(processor), 0) + 1
drained: Final = self._drainable_locked()
for shed in drained:
self._drain.submit(shed)
return processor
def _cached_or_built_locked(self, destination: "OtelDestination", *, anchored: bool) -> SpanProcessor | None:
"""The cached processor for ``destination``, or a new one if the drain can take it.
Every build past the cache cap sheds one processor into the drain, so while the
shed ones are stuck closing against a collector that stopped answering, a
destination that is not yet anchored is refused rather than parked behind them:
``deliverable`` then leaves its spans with the operator's exporter until the
drain catches up. One the request already anchored is rebuilt regardless. The
operator's exporter has stood down for it, so refusing here would drop the span,
and other tenants' auths can evict it in the meantime, with that eviction being
what tips the drain over. Eviction holds while the drain is saturated, so such a
rebuild costs the cache one entry rather than shedding another processor, and
the total stays at one per destination in flight.
"""
key: Final = destination.cache_key()
if (cached := self._processors.get(key)) is not None:
self._processors.move_to_end(key)
self._retire_overflow_locked()
return cached
if not anchored and self._drain.saturated():
verbose_logger.debug("OTel V2 fan-out: drain saturated, not building for %s", destination.endpoint)
return None
return self._build_locked(destination, key)
def _build_locked(self, destination: "OtelDestination", key: object) -> SpanProcessor | None:
built: Final = self._build(destination)
if built is None:
return None
self._processors[key] = built
self._retire_overflow_locked()
return built
def _release(self, processor: SpanProcessor) -> None:
with self._lock:
remaining: Final = self._exporting.get(id(processor), 1) - 1
if remaining > 0:
self._exporting[id(processor)] = remaining
else:
self._exporting.pop(id(processor), None)
if not self._exporting:
self._lock.notify_all()
drained: Final = self._drainable_locked()
for retired in drained:
self._drain.submit(retired)
def _retire_overflow_locked(self) -> None:
"""Move the LRU processor out of the cache once it is past the cap, drain permitting.
Eviction is what feeds the drain, and a destination a request already anchored
is rebuilt on its next span, which would shed another one. While the shed ones
are stuck closing against a collector that stopped answering, evicting would
churn the cache at one more processor, and one more batch thread, per span.
Holding above the cap instead keeps the total at one processor per destination
in flight, since ``deliverable`` anchors no new destination while the drain is
saturated. Once it has room again, every hit and build trims one entry.
"""
if len(self._processors) <= _MAX_CACHED_DESTINATION_PROCESSORS or self._drain.saturated():
return
_, evicted = self._processors.popitem(last=False)
self._retired[id(evicted)] = evicted
def _drainable_locked(self) -> tuple[SpanProcessor, ...]:
"""Retired processors no thread is exporting through, removed from the list.
``on_end`` holds a processor across an export, so closing an evicted one there
drops the span it is holding. A retiree is out of the cache and can never be
handed out again, so once its export count reaches zero it stays there.
"""
idle: Final = tuple(key for key in self._retired if self._exporting.get(key, 0) == 0)
return tuple(self._retired.pop(key) for key in idle)
def _destination_processor(destination: "OtelDestination") -> SpanProcessor | None:
"""A batching OTLP processor aimed at ``destination``, or ``None`` if unbuildable.
A protocol that resolves to a headerless exporter is unbuildable too: the
console fallback would swallow the tenant's credentials and print its spans to
the proxy's stdout while the operator's exporter stands down for them.
"""
kind: Final = destination.protocol or "otlp_http"
if exporter_transport(kind) == "headerless":
verbose_logger.debug("OTel V2 fan-out: no OTLP transport for protocol %r at %s", kind, destination.endpoint)
return None
try:
spec: Final = ExporterSpec(
kind=kind,
endpoint=destination.endpoint,
headers=destination.header_string(),
owner=None,
)
return _processor_for(_exporter_from_spec(spec), use_simple=False)
except Exception as exc: # noqa: BLE001 # a malformed destination must not break the request or the other destinations
verbose_logger.debug("OTel V2 fan-out: no processor for %s: %s", destination.endpoint, exc)
return None
def _shutdown_quietly(processor: SpanProcessor) -> None:
try:
processor.shutdown()
except Exception as exc: # noqa: BLE001 # defensive: shedding a spare processor must not raise
verbose_logger.debug("OTel V2 fan-out: discarding processor failed: %s", exc)
class _OverriddenBackendFilter(SpanProcessor):
"""Hold a span back from ``owner``'s operator-level exporter when the request
pointed ``owner`` at a tenant's own account.
Wrapping is the only place this works: ``SynchronousMultiSpanProcessor.on_end``
ignores return values, so a sibling processor can never veto the export.
Under ``additive`` mode nothing is suppressed, so the wrapper passes every span
straight through and the operator keeps its copy.
"""
def __init__(self, inner: SpanProcessor, owner: str) -> None:
self._inner: Final = inner
self._owner: Final = owner
def on_start(self, span: SDKSpan, parent_context: Context | None = None) -> None:
self._inner.on_start(span, parent_context)
def on_end(self, span: ReadableSpan) -> None:
if self._owner in suppressed_backends():
return
self._inner.on_end(span)
def shutdown(self) -> None:
self._inner.shutdown()
def force_flush(self, timeout_millis: int = 30000) -> bool:
return self._inner.force_flush(timeout_millis)
def build_span_exporter(config: OpenTelemetryV2Config) -> SpanExporter:
"""Build a single exporter from the top-level config fields.
@ -444,6 +1024,7 @@ def build_tracer_provider(
exporter: SpanExporter | None = None,
baggage_processor: SpanProcessor | None = None,
use_simple_processor: bool | None = None,
tenant_overrides: bool = False,
) -> TracerProvider:
"""Build the shared :class:`TracerProvider`.
@ -452,6 +1033,13 @@ def build_tracer_provider(
``config.exporters`` entry this is what fans spans out to multiple
backends. ``exporter`` and ``use_simple_processor`` are explicit overrides:
pass a single exporter to attach exactly that one (used by tests).
``tenant_overrides`` wraps each owned exporter so a request that pointed that
backend at a key's or team's own account skips it. Every v2 logger's provider
wants it, since any of them may own the overridden backend; delivering to the
tenant is a separate job, done once by :func:`attach_tenant_fan_out`. The
per-tenant providers this same function builds must leave it off, or they would
filter out the very spans they exist to carry.
"""
provider: Final = TracerProvider(resource=build_resource(config))
if baggage_processor is None:
@ -468,15 +1056,107 @@ def build_tracer_provider(
if spec.requires_headers and not spec.headers:
continue
exp = _exporter_from_spec(spec)
processor = _processor_for(
exp,
(spec.use_simple_processor if spec.use_simple_processor is not None else use_simple_processor),
)
owner = spec.owner.value if spec.owner is not None else None
provider.add_span_processor(
_processor_for(
exp,
(spec.use_simple_processor if spec.use_simple_processor is not None else use_simple_processor),
)
_OverriddenBackendFilter(processor, owner) if tenant_overrides and owner is not None else processor
)
return provider
_FAN_OUT_ATTACH_LOCK: Final = threading.Lock()
def attach_tenant_fan_out(provider: TracerProvider, *configs: OpenTelemetryV2Config) -> None:
"""Give ``provider`` the fan-out that delivers spans to key/team destinations.
Called on the one provider published as the OTel global, and idempotent so a
second publish (a test, a re-initialized proxy) cannot double-export. Concurrent
first calls (requests racing to anchor before any publish) serialize on one lock
so exactly one fan-out lands. ``configs`` name the operator's own exporters, one
config per v2 logger since each keeps its own provider and still writes its
account, so an additive destination pointing at any of them is delivered once
rather than twice.
"""
with _FAN_OUT_ATTACH_LOCK:
if any(isinstance(processor, TenantFanOutSpanProcessor) for processor in _attached_processors(provider)):
return
provider.add_span_processor(TenantFanOutSpanProcessor(operator_sinks=operator_sink_keys(*configs)))
def deliverable_destinations(
destinations: Iterable["OtelDestination"],
provider: trace.TracerProvider | None = None,
) -> tuple["OtelDestination", ...]:
"""The destinations a request can anchor, given what is published to carry them.
Anchoring a destination is what tells the operator's own exporter to stand down
for that backend, so one nothing can deliver has to be dropped here: with no
fan-out attached, or with an exporter that will not build, the request keeps
exactly the routing it would have had without any override.
"""
fan_out: Final = next(
(
processor
for processor in _attached_processors(provider if provider is not None else trace.get_tracer_provider())
if isinstance(processor, TenantFanOutSpanProcessor)
),
None,
)
return fan_out.deliverable(destinations) if fan_out is not None else ()
def operator_sink_keys(*configs: OpenTelemetryV2Config) -> frozenset[_SinkKey]:
"""The accounts the operator's own exporters write to, in destination terms.
Every v2 logger's config counts, since each logger exports through its own
provider. An exporter with no endpoint of its own resolves one from the
environment at export time, so it has no comparable identity and is left out,
and so is one that never reaches the wire: a console kind ignores the endpoint,
and a header-gated spec with no credentials is skipped when the provider is built.
"""
return frozenset(
key
for config in configs
for spec in config.exporters
if _exports_to_the_wire(spec) and (key := _sink_key(spec.endpoint, parse_headers(spec.headers))) is not None
)
def _exports_to_the_wire(spec: ExporterSpec) -> bool:
"""Whether ``build_tracer_provider`` gives ``spec`` an exporter that sends OTLP."""
return exporter_transport(spec.kind) != "headerless" and not (spec.requires_headers and not spec.headers)
def _sink_key(endpoint: str | None, headers: Mapping[str, str]) -> "_SinkKey | None":
"""The account an exporter writes to, or ``None`` when it has no fixed one.
Normalized on the three counts that make one account look like two: the operator's
spec carries the signal path a tenant destination leaves for the exporter to
append, header names survive one round trip lowercased and the other not, and one
credential answers to more than one name (see :data:`_CREDENTIAL_ALIASES`).
"""
normalized: Final = _otlp_traces_endpoint(endpoint)
if normalized is None:
return None
return (normalized, tuple(sorted((_credential_name(name), value) for name, value in headers.items())))
def _credential_name(header: str) -> str:
"""The credential a header carries, under whichever name the backend spells it."""
normalized: Final = header.strip().lower().replace("-", "_")
return _CREDENTIAL_ALIASES.get(normalized, normalized)
def _attached_processors(provider: trace.TracerProvider) -> "tuple[SpanProcessor, ...]":
"""The processors already on ``provider``, or empty when the SDK hides them."""
multi: Final = getattr(provider, "_active_span_processor", None)
return tuple(getattr(multi, "_span_processors", ()))
def get_tracer(provider: TracerProvider, name: str = "litellm") -> Tracer:
# Stamp the instrumentation scope with the LiteLLM package version so every
# emitted span carries a deterministic ``scope.version`` (the standard OTel

View file

@ -25,6 +25,7 @@ from opentelemetry.trace import Tracer
from litellm._logging import verbose_logger
from litellm.constants import OTEL_SERVICE_NAME_METADATA_KEYS
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.integrations.otel.plumbing.context import destination_backends
from litellm.integrations.otel.plumbing.providers import (
build_tracer_provider,
exporter_transport,
@ -231,10 +232,21 @@ class TenantTracerCache:
concurrent overflow eviction can't shut it down between selection and
the caller's span start. The caller must ``release`` it exactly once.
"""
# A backend with a destination is delivered by the fan-out processor, which
# carries the whole trace and already carries this tenant's credentials and
# service name. Routing here too would detach this span onto a second provider,
# so the tenant would get the request tree plus a stray one-span trace.
if self._callback_name is not None and self._callback_name in destination_backends():
return TenantRoute(tracer=default, detached=False)
credential_headers: Final = self._credential_headers(dynamic_params)
project_headers: Final = self._project_headers(auth_metadata)
service_name: Final = tenant_service_name(auth_metadata)
if not credential_headers and not project_headers and service_name is None:
tenant_account: Final = bool(credential_headers) or bool(project_headers)
# A service name on its own only relabels the operator's own backend, so moving
# the span to a second provider for it while some other backend has a
# destination would drop the model call out of the trace the fan-out delivers.
# The destination stamps the same service name itself.
if not tenant_account and (service_name is None or destination_backends()):
return TenantRoute(tracer=default, detached=False)
# A fixed per-integration region endpoint (New Relic us/eu), never a
# caller-supplied host; ``None`` keeps the preset's own endpoint.
@ -255,7 +267,7 @@ class TenantTracerCache:
_shutdown_provider(evicted)
return TenantRoute(
tracer=get_tracer(provider, self._tracer_name),
detached=bool(project_headers) or bool(credential_headers),
detached=tenant_account,
provider=provider,
)

View file

@ -39,6 +39,7 @@ class _AgentOpsSettings(BaseSettings):
def agentops_preset(
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config:
"""Build the AgentOps config without any network I/O.

View file

@ -26,10 +26,12 @@ class _ArizeSettings(BaseSettings):
def arize_preset(
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config:
base: Final = config_overrides or OpenTelemetryV2Config()
mappers: Final = ensure_mappers(base.mapper_names, "openinference")
arize_cfg: Final = _V1ArizeLogger.get_arize_config()
headers: Final = _arize_headers(arize_cfg)
base: Final = config_overrides or OpenTelemetryV2Config()
return base.model_copy(
update={
"exporters": [
@ -41,7 +43,7 @@ def arize_preset(
owner=ExporterOwner.ARIZE_AX,
),
],
"mapper_names": ensure_mappers(base.mapper_names, "openinference"),
"mapper_names": mappers,
"resource_attributes": {
**base.resource_attributes,
**({"model_id": arize_cfg.project_name} if arize_cfg.project_name else {}),

View file

@ -18,6 +18,18 @@ class Preset(Protocol):
``config_overrides`` lets one preset layer onto another's config (or onto
test-supplied defaults); the factory calls presets with no arguments.
``allow_missing_credentials`` lets a credential-mandatory backend (langfuse and
weave) degrade to an exporter-less, mapper-only config instead of raising when the
operator set no env credentials of their own. That is a real
deployment: every team brings its own account and the operator keeps none, and
without it the whole V2 path silently falls back to the legacy integration, so
no team destination is ever reached. Credential-optional backends ignore it.
"""
def __call__(self, *, config_overrides: OpenTelemetryV2Config | None = None) -> OpenTelemetryV2Config: ...
def __call__(
self,
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config: ...

View file

@ -0,0 +1,152 @@
"""Map a key's or team's callback vars to the OTLP destination its traces export to.
Header building is delegated to each preset's existing ``*_dynamic_headers`` builder,
so a destination authenticates exactly the way the per-request tracer route already
did; only the endpoint and transport need a per-backend rule.
"""
import os
from collections.abc import Callable, Mapping
from functools import lru_cache
from types import MappingProxyType
from typing import Final
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.otel.model.destination import OtelDestination
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
from litellm.types.utils import StandardCallbackDynamicParams
#: An endpoint plus the OTLP transport to reach it with, or ``None`` when the backend
#: names no destination. The transport is ``None`` where the backend has only one.
_Destination = tuple[str, str | None]
@lru_cache(maxsize=128)
def _warn_host_not_allowlisted(host: str) -> None:
"""Cached so one misconfigured team logs once rather than once per request."""
verbose_logger.warning(
"OTel V2: not exporting to key/team Langfuse host '%s'. Add it to "
"litellm_settings.provider_url_destination_allowed_hosts to permit it",
host,
)
def _langfuse_destination(params: StandardCallbackDynamicParams) -> "_Destination | None":
"""The tenant's own Langfuse host, else the operator's, else Langfuse US cloud.
A host the tenant named has to be allowlisted by the operator, the same way a
URL-valued ``model`` is: anyone who can mint a key can write it, and it becomes an
endpoint the proxy posts the request's whole trace to, carrying the tenant's own
credentials. The operator's own ``LANGFUSE_HOST`` is not checked, since an internal
collector there is a deployment choice.
"""
from litellm.integrations.langfuse.langfuse_otel import (
LANGFUSE_CLOUD_US_ENDPOINT,
LangfuseOtelLogger,
)
tenant_host: Final = params.get("langfuse_host") or None
host: Final = tenant_host or LangfuseOtelLogger._get_langfuse_otel_host() # pyright: ignore[reportPrivateUsage] # reuse the backend's own env host resolver rather than duplicating it
if not host:
return (LANGFUSE_CLOUD_US_ENDPOINT, None)
normalized: Final = host if host.startswith("http") else f"https://{host}"
endpoint: Final = f"{normalized.rstrip('/')}/api/public/otel"
if tenant_host is None:
return (endpoint, None)
if not is_url_destination_allowed_by_host(endpoint, litellm.provider_url_destination_allowed_hosts):
_warn_host_not_allowlisted(host)
return None
return (endpoint, None)
def _arize_destination(params: StandardCallbackDynamicParams) -> "_Destination | None":
from litellm.integrations.arize.arize import ArizeLogger
config: Final = ArizeLogger.get_arize_config()
return (config.endpoint, config.protocol)
def _weave_destination(params: StandardCallbackDynamicParams) -> "_Destination | None":
from litellm.integrations.weave.weave_otel import weave_otel_endpoint
return (weave_otel_endpoint(os.environ.get("WANDB_HOST")), None)
def _newrelic_destination(params: StandardCallbackDynamicParams) -> "_Destination | None":
from litellm.integrations.otel.presets.newrelic import newrelic_dynamic_endpoint
endpoint: Final = newrelic_dynamic_endpoint(params)
return (endpoint, None) if endpoint else None
#: Callback name -> destination resolver. A backend is destination-capable exactly
#: when it appears here AND in ``DYNAMIC_HEADERS_BY_CALLBACK``: without a header
#: builder the destination would carry no tenant credentials, and the exporter
#: would post the tenant's traffic to the operator's account.
_DESTINATION_BY_CALLBACK: Final[Mapping[str, Callable[[StandardCallbackDynamicParams], "_Destination | None"]]] = (
MappingProxyType(
{
"langfuse_otel": _langfuse_destination,
"arize": _arize_destination,
"weave_otel": _weave_destination,
"newrelic": _newrelic_destination,
}
)
)
#: Headers a destination must carry to authenticate. Several dynamic-header builders
#: gate each credential independently, so a half-configured backend yields a non-empty
#: but unusable header set; accepting it would suppress the operator's own exporter and
#: send the request's whole trace where it cannot be stored.
_REQUIRED_HEADERS_BY_CALLBACK: Final[Mapping[str, frozenset[str]]] = MappingProxyType(
{
"langfuse_otel": frozenset({"Authorization"}),
"arize": frozenset({"arize-space-id", "api_key"}),
"weave_otel": frozenset({"Authorization", "project_id"}),
"newrelic": frozenset({"api-key"}),
}
)
_NO_ATTRS: Final[Mapping[str, str]] = MappingProxyType({})
def destination_capable_backends() -> frozenset[str]:
"""Backends a key or team can point at its own account."""
from litellm.integrations.otel.presets import DYNAMIC_HEADERS_BY_CALLBACK
return frozenset(_DESTINATION_BY_CALLBACK) & frozenset(DYNAMIC_HEADERS_BY_CALLBACK)
def destination_for(
callback_name: str,
params: StandardCallbackDynamicParams,
service_name: str | None = None,
) -> OtelDestination | None:
"""The destination ``params`` names for ``callback_name``, or ``None``.
``None`` means the caller configured nothing usable for this backend, so the
request keeps the operator's global exporters. ``service_name`` is the key's or
team's ``otel_service_name``, which the per-request tracer route applies when the
backend is not overridden and the destination has to apply once it is.
"""
from litellm.integrations.otel.presets import DYNAMIC_HEADERS_BY_CALLBACK
header_builder: Final = DYNAMIC_HEADERS_BY_CALLBACK.get(callback_name)
destination_builder: Final = _DESTINATION_BY_CALLBACK.get(callback_name)
if header_builder is None or destination_builder is None:
return None
headers: Final = header_builder(params)
if not headers or not _REQUIRED_HEADERS_BY_CALLBACK[callback_name] <= frozenset(headers):
return None
resolved: Final = destination_builder(params)
if resolved is None:
return None
endpoint, protocol = resolved
return OtelDestination(
endpoint=endpoint,
headers=MappingProxyType(dict(headers)), # mutable-ok: MappingProxyType needs a concrete mapping to wrap
resource_attributes=MappingProxyType({"service.name": service_name}) if service_name else _NO_ATTRS,
callback_name=callback_name,
protocol=protocol,
)

View file

@ -10,17 +10,32 @@ from litellm.integrations.otel.model.config import (
ExporterSpec,
OpenTelemetryV2Config,
)
from litellm.integrations.otel.presets.utils import ensure_mappers
from litellm.integrations.otel.presets.utils import (
credential_gated_exporters,
ensure_mappers,
)
from litellm.types.utils import StandardCallbackDynamicParams
def langfuse_preset(
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config:
cfg: Final = _V1Langfuse.get_langfuse_otel_config()
kind: Final = cfg.exporter if isinstance(cfg.exporter, str) else "otlp_http"
base: Final = config_overrides or OpenTelemetryV2Config()
mappers: Final = ensure_mappers(base.mapper_names, "langfuse")
try:
cfg: Final = _V1Langfuse.get_langfuse_otel_config()
except Exception:
if not allow_missing_credentials:
raise
return base.model_copy(
update={ # mutable-ok: pydantic model_copy takes a plain update mapping
"exporters": credential_gated_exporters(base.exporters, ExporterOwner.LANGFUSE_OTEL),
"mapper_names": mappers,
}
)
kind: Final = cfg.exporter if isinstance(cfg.exporter, str) else "otlp_http"
return base.model_copy(
update={
"exporters": [
@ -32,7 +47,7 @@ def langfuse_preset(
owner=ExporterOwner.LANGFUSE_OTEL,
),
],
"mapper_names": ensure_mappers(base.mapper_names, "langfuse"),
"mapper_names": mappers,
}
)

View file

@ -9,6 +9,7 @@ from litellm.integrations.otel.presets.utils import ensure_mappers
def langtrace_preset(
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config:
"""Compose the Langtrace mapper on top of the customer's OTLP destination.

View file

@ -13,6 +13,7 @@ from litellm.integrations.otel.model.config import (
def levo_preset(
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config:
cfg: Final = _V1Levo.get_levo_config()
base: Final = config_overrides or OpenTelemetryV2Config()

View file

@ -44,6 +44,7 @@ class _NewRelicSettings(BaseSettings):
def newrelic_preset(
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config:
settings: Final = _NewRelicSettings()
base: Final = config_overrides or OpenTelemetryV2Config()

View file

@ -60,6 +60,7 @@ def phoenix_project_headers(auth_metadata: Mapping[str, str] | None) -> Mapping[
def phoenix_preset(
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config:
cfg: Final = _V1Phoenix.get_arize_phoenix_config()
headers: Final = cfg.otlp_auth_headers if hasattr(cfg, "otlp_auth_headers") else None

View file

@ -3,6 +3,8 @@
from collections.abc import Iterable
from typing import Final
from litellm.integrations.otel.model.config import ExporterOwner, ExporterSpec
def ensure_mappers(mapper_names: Iterable[str], *names: str) -> list[str]:
"""Return ``mapper_names`` with each of ``names`` appended if not already present.
@ -15,3 +17,32 @@ def ensure_mappers(mapper_names: Iterable[str], *names: str) -> list[str]:
if name not in result:
result.append(name)
return result
def credential_gated_exporters(
exporters: "Iterable[ExporterSpec]", owner: "ExporterOwner"
) -> "tuple[ExporterSpec, ...]":
"""``exporters`` with the operator's destination replaced by a header-gated one.
Used when a credential-mandatory backend is asked to build without the operator's
own credentials, so only key/team destinations receive spans. Two things have to
happen for that to mean "export nowhere": the placeholder console spec that
``OpenTelemetryV2Config`` folds in for an empty exporter list is dropped, or every
span would be printed to stdout, and the gated spec keeps the owner so the
override filter still recognises which backend this provider speaks for.
"""
return (
*(spec for spec in exporters if not is_unconfigured_placeholder(spec)),
ExporterSpec(owner=owner, requires_headers=True),
)
def is_unconfigured_placeholder(spec: "ExporterSpec") -> bool:
"""Whether ``spec`` is the one ``_normalize`` folds in when nothing was configured.
No field set is what says the operator asked for nothing: an exporter they did
configure survives, even ``OTEL_EXPORTER=console`` whose value matches the default,
and so does the gated spec this module appends, which would otherwise eat itself
when one preset layers onto another.
"""
return not spec.model_fields_set

View file

@ -7,7 +7,10 @@ from litellm.integrations.otel.model.config import (
ExporterSpec,
OpenTelemetryV2Config,
)
from litellm.integrations.otel.presets.utils import ensure_mappers
from litellm.integrations.otel.presets.utils import (
credential_gated_exporters,
ensure_mappers,
)
from litellm.integrations.weave.weave_otel import (
_get_weave_authorization_header,
get_weave_otel_config,
@ -18,9 +21,21 @@ from litellm.types.utils import StandardCallbackDynamicParams
def weave_preset(
*,
config_overrides: OpenTelemetryV2Config | None = None,
allow_missing_credentials: bool = False,
) -> OpenTelemetryV2Config:
weave_cfg: Final = get_weave_otel_config()
base: Final = config_overrides or OpenTelemetryV2Config()
mappers: Final = ensure_mappers(base.mapper_names, "openinference", "weave")
try:
weave_cfg: Final = get_weave_otel_config()
except Exception:
if not allow_missing_credentials:
raise
return base.model_copy(
update={ # mutable-ok: pydantic model_copy takes a plain update mapping
"exporters": credential_gated_exporters(base.exporters, ExporterOwner.WEAVE_OTEL),
"mapper_names": mappers,
}
)
return base.model_copy(
update={
"exporters": [
@ -33,7 +48,7 @@ def weave_preset(
),
],
# Weave consumes OpenInference + a small Weave-specific overlay.
"mapper_names": ensure_mappers(base.mapper_names, "openinference", "weave"),
"mapper_names": mappers,
}
)

View file

@ -117,6 +117,14 @@ def _get_weave_authorization_header(api_key: str) -> str:
return f"Basic {auth_header}"
def weave_otel_endpoint(host: str | None) -> str:
"""The OTLP traces endpoint for a self-managed ``host``, else Weave cloud."""
if not host:
return WEAVE_BASE_URL + WEAVE_OTEL_ENDPOINT
normalized: Final = host if host.startswith("http") else f"https://{host}"
return normalized.rstrip("/") + WEAVE_OTEL_ENDPOINT
def get_weave_otel_config() -> WeaveOtelConfig:
"""
Retrieves the Weave OpenTelemetry configuration based on environment variables.
@ -134,7 +142,6 @@ def get_weave_otel_config() -> WeaveOtelConfig:
"""
api_key: Final = os.getenv("WANDB_API_KEY")
project_id: Final = os.getenv("WANDB_PROJECT_ID")
host = os.getenv("WANDB_HOST")
if not api_key:
raise ValueError("WANDB_API_KEY must be set for Weave OpenTelemetry integration.")
@ -144,15 +151,8 @@ def get_weave_otel_config() -> WeaveOtelConfig:
"WANDB_PROJECT_ID must be set for Weave OpenTelemetry integration. Format: <entity>/<project_name>"
)
if host:
if not host.startswith("http"):
host = "https://" + host
# Self-managed instances use a different path
endpoint = host.rstrip("/") + WEAVE_OTEL_ENDPOINT
verbose_logger.debug("Using Weave OTEL endpoint from host: %s", endpoint)
else:
endpoint = WEAVE_BASE_URL + WEAVE_OTEL_ENDPOINT
verbose_logger.debug("Using Weave cloud endpoint: %s", endpoint)
endpoint: Final = weave_otel_endpoint(os.getenv("WANDB_HOST"))
verbose_logger.debug("Using Weave OTEL endpoint: %s", endpoint)
# Weave uses Basic auth with format: api:<WANDB_API_KEY>
auth_header: Final = _get_weave_authorization_header(api_key=api_key)

View file

@ -2,6 +2,7 @@
Pulls the cost + context window + provider route for known models from https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json
This can be disabled by setting the LITELLM_LOCAL_MODEL_COST_MAP environment variable to True.
The ``lite`` and ``litellm-proxy`` CLI entry points also use the bundled map without fetching.
```
export LITELLM_LOCAL_MODEL_COST_MAP=True
@ -9,17 +10,22 @@ export LITELLM_LOCAL_MODEL_COST_MAP=True
"""
import asyncio
import hashlib
import json
import os
import random
import sys
import threading
import time
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from dataclasses import dataclass, replace
from datetime import datetime, timezone
from importlib.resources import files
from pathlib import Path
from typing import Final, Protocol
import httpx
from typing_extensions import ReadOnly, TypedDict
from litellm import verbose_logger
from litellm.constants import (
@ -31,6 +37,12 @@ from litellm.litellm_core_utils.fallback_generalizations import (
)
FALLBACK_GENERALIZATIONS_KEY: Final = "fallback_generalizations"
_CLI_ENTRYPOINT_NAMES: Final = frozenset({"lite", "litellm-proxy"})
def _is_cli_process() -> bool:
return Path(sys.argv[0]).stem in _CLI_ENTRYPOINT_NAMES
# Reserved top-level keys that are not model entries. They must be excluded
# from the model-count integrity check so a real upstream shrink can't be masked.
@ -42,6 +54,10 @@ def _count_model_entries(model_cost: dict) -> int:
return sum(1 for key in model_cost if key not in RESERVED_TOP_LEVEL_KEYS)
def git_blob_id(body: bytes) -> str:
return hashlib.sha1(b"blob %d\0" % len(body) + body, usedforsecurity=False).hexdigest()
class GetModelCostMap:
"""
Handles fetching, validating, and loading the model cost map.
@ -53,15 +69,24 @@ class GetModelCostMap:
_backup_model_count: int = -1 # -1 = not yet loaded
@staticmethod
def read_local_model_cost_map_bytes() -> bytes:
return files("litellm").joinpath("model_prices_and_context_window_backup.json").read_bytes()
@staticmethod
def read_local_model_cost_map_text() -> str:
return files("litellm").joinpath("model_prices_and_context_window_backup.json").read_text(encoding="utf-8")
return GetModelCostMap.read_local_model_cost_map_bytes().decode("utf-8")
@staticmethod
def load_local_model_cost_map_with_revision() -> "ModelCostMapReloaded":
body: Final = GetModelCostMap.read_local_model_cost_map_bytes()
content: Final = json.loads(body)
return ModelCostMapReloaded(model_cost_map=content, revision=git_blob_id(body))
@staticmethod
def load_local_model_cost_map() -> dict:
"""Load the local backup model cost map bundled with the package."""
content: Final = json.loads(GetModelCostMap.read_local_model_cost_map_text())
return content
return GetModelCostMap.load_local_model_cost_map_with_revision().model_cost_map
@classmethod
def _get_backup_model_count(cls) -> int:
@ -161,11 +186,18 @@ class GetModelCostMap:
RETRYABLE_FETCH_STATUS_CODES: Final = frozenset({429, 500, 502, 503, 504})
MODEL_COST_MAP_FETCH_MAX_ATTEMPTS: Final = 3
MODEL_COST_MAP_FETCH_MAX_WAIT_SECONDS: Final = 30.0
_litellm_import_complete = threading.Event()
def mark_litellm_import_complete() -> None:
_litellm_import_complete.set()
@dataclass(frozen=True, slots=True)
class ModelCostMapReloaded:
model_cost_map: dict # mutable-ok: adopted as litellm.model_cost, whose consumer contract is a plain mutable dict
revision: str | None = None
etag: str | None = None
@dataclass(frozen=True, slots=True)
@ -254,7 +286,9 @@ def _classify_fetch_response(response: httpx.Response, url: str) -> _FetchAttemp
return ModelCostMapReloadUnavailable(reason=f"invalid JSON from {url}: {e}")
if not isinstance(parsed, dict):
return ModelCostMapReloadUnavailable(reason=f"expected a JSON object from {url}, got {type(parsed).__name__}")
return ModelCostMapReloaded(model_cost_map=parsed)
return ModelCostMapReloaded(
model_cost_map=parsed, revision=git_blob_id(response.content), etag=response.headers.get("etag")
)
def _next_retry_wait(
@ -295,12 +329,13 @@ async def _fetch_remote_model_cost_map_with_retry(
def _fetch_remote_model_cost_map_with_retry_sync(
url: str,
timeout: int,
max_attempts: int,
attempts: range,
sleep: Callable[[float], None],
rng: random.Random,
client: _SyncGetClient,
) -> ModelCostMapReloadResult:
for attempt in range(1, max_attempts + 1):
max_attempts: Final = attempts.stop - 1
for attempt in attempts:
outcome = _attempt_fetch_sync(client=client, url=url, timeout=timeout)
if not isinstance(outcome, _FetchAttemptRetryable):
return outcome
@ -328,13 +363,12 @@ async def refetch_model_cost_map(
map they already have.
"""
if os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", "").lower() == "true":
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
_cost_map_source_info.source = "local"
_cost_map_source_info.url = None
_cost_map_source_info.is_env_forced = True
_cost_map_source_info.fallback_reason = None
return ModelCostMapReloaded(
model_cost_map=_finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
)
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision())
result: Final = await _fetch_remote_model_cost_map_with_retry(
url=url,
@ -355,11 +389,12 @@ async def refetch_model_cost_map(
backup_model_count=GetModelCostMap._get_backup_model_count(),
):
return ModelCostMapReloadUnavailable(reason=f"model cost map from {url} failed integrity validation")
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
_cost_map_source_info.source = "remote"
_cost_map_source_info.url = url
_cost_map_source_info.is_env_forced = False
_cost_map_source_info.fallback_reason = None
return ModelCostMapReloaded(model_cost_map=_finalize_model_cost_map(result.model_cost_map))
return _finalize_loaded_model_cost_map(result)
class ModelCostMapSourceInfo:
@ -370,13 +405,35 @@ class ModelCostMapSourceInfo:
is_env_forced: bool = False
fallback_reason: str | None = None
loaded_at: "datetime | None" = None
source_revision: str | None = None
etag: str | None = None
# Module-level singleton tracking the source of the current cost map
_cost_map_source_info: Final = ModelCostMapSourceInfo()
def get_model_cost_map_source_info() -> dict:
class CostMapProvenance(TypedDict):
source_revision: ReadOnly[str | None]
etag: ReadOnly[str | None]
class CostMapSourceInfo(CostMapProvenance):
source: ReadOnly[str]
url: ReadOnly[str | None]
is_env_forced: ReadOnly[bool]
fallback_reason: ReadOnly[str | None]
loaded_at: ReadOnly[str | None]
def get_model_cost_map_provenance() -> CostMapProvenance:
return {
"source_revision": _cost_map_source_info.source_revision,
"etag": _cost_map_source_info.etag,
}
def get_model_cost_map_source_info() -> CostMapSourceInfo:
"""
Return metadata about where the current model cost map was loaded from.
@ -385,12 +442,19 @@ def get_model_cost_map_source_info() -> dict:
- url: the remote URL attempted (or None for local-only)
- is_env_forced: True if LITELLM_LOCAL_MODEL_COST_MAP=True forced local usage
- fallback_reason: human-readable reason if remote failed and local was used
- loaded_at: ISO 8601 time this process last loaded the map
- source_revision: git blob id of the loaded file's bytes
- etag: the ETag of the remote fetch (None for the bundled backup)
"""
loaded_at: Final = _cost_map_source_info.loaded_at
return {
"source": _cost_map_source_info.source,
"url": _cost_map_source_info.url,
"is_env_forced": _cost_map_source_info.is_env_forced,
"fallback_reason": _cost_map_source_info.fallback_reason,
"loaded_at": loaded_at.isoformat() if loaded_at is not None else None,
"source_revision": _cost_map_source_info.source_revision,
"etag": _cost_map_source_info.etag,
}
@ -466,6 +530,74 @@ def _finalize_model_cost_map(model_cost: dict) -> dict:
return _expand_model_aliases(model_cost)
def _finalize_loaded_model_cost_map(loaded: ModelCostMapReloaded) -> ModelCostMapReloaded:
_cost_map_source_info.source_revision = loaded.revision
_cost_map_source_info.etag = loaded.etag
return replace(loaded, model_cost_map=_finalize_model_cost_map(loaded.model_cost_map))
def adopt_model_cost_map(
new_model_cost_map: dict, # mutable-ok: public API preserves the mutable cost-map contract
) -> int:
import litellm
from litellm import utils
litellm.model_cost = new_model_cost_map
utils._invalidate_model_cost_lowercase_map() # pyright: ignore[reportPrivateUsage] # required cache invalidation
litellm.add_known_models(model_cost_map=new_model_cost_map)
fetched_model_count: Final = len(new_model_cost_map) if new_model_cost_map else 0
utils.reapply_runtime_model_cost_registrations()
return fetched_model_count
def _retry_remote_fetch_in_background(
url: str,
timeout: int,
max_attempts: int,
sleep: Callable[[float], None],
rng: random.Random,
client: _SyncGetClient,
first_outcome: _FetchAttemptRetryable,
) -> None:
try:
first_wait: Final = _next_retry_wait(outcome=first_outcome, attempt=1, max_attempts=max_attempts, rng=rng)
if isinstance(first_wait, ModelCostMapReloadUnavailable):
return
sleep(first_wait)
result: Final = _fetch_remote_model_cost_map_with_retry_sync(
url=url,
timeout=timeout,
attempts=range(2, max_attempts + 1),
sleep=sleep,
rng=rng,
client=client,
)
if isinstance(result, ModelCostMapReloadUnavailable):
verbose_logger.warning(
"LiteLLM: Failed to fetch remote model cost map from %s after %d attempts; keeping local backup",
url,
max_attempts,
)
return
_litellm_import_complete.wait()
if not GetModelCostMap.validate_model_cost_map(
fetched_map=result.model_cost_map,
backup_model_count=GetModelCostMap._get_backup_model_count(), # pyright: ignore[reportPrivateUsage] # integrity cache
):
verbose_logger.warning(
"LiteLLM: Fetched model cost map failed integrity check. Using local backup instead. url=%s",
url,
)
return
finalized: Final = _finalize_loaded_model_cost_map(result).model_cost_map
_cost_map_source_info.source = "remote"
_cost_map_source_info.fallback_reason = None
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
adopt_model_cost_map(finalized)
except Exception as e: # noqa: BLE001 # a failed background retry must not kill the task; the backup stays
verbose_logger.warning("LiteLLM: Background model cost map retry failed: %s", e)
def get_model_cost_map(
url: str,
timeout: int = 5,
@ -477,10 +609,12 @@ def get_model_cost_map(
"""
Public entry point returns the model cost map dict.
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set, uses the local backup only.
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set or this is a ``lite`` /
``litellm-proxy`` CLI process, uses the local backup only.
2. Otherwise fetches from ``url``, retrying transient HTTP errors
(429/5xx/transport) with Retry-After-aware backoff, validates
integrity, and falls back to the local backup on any failure.
(429/5xx/transport) with Retry-After-aware backoff in a background
thread, validates integrity, and falls back to the local backup on any
failure.
Only the backup model count is cached (a single int) for validation.
The full backup dict is only parsed when it must be *returned* as a
@ -489,34 +623,44 @@ def get_model_cost_map(
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
# Note: can't use get_secret_bool here — this runs during litellm.__init__
# before litellm._key_management_settings is set.
if os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", "").lower() == "true":
if os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", "").lower() == "true" or _is_cli_process():
_cost_map_source_info.source = "local"
_cost_map_source_info.url = None
_cost_map_source_info.is_env_forced = True
_cost_map_source_info.fallback_reason = None
return _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision()).model_cost_map
_cost_map_source_info.url = url
_cost_map_source_info.is_env_forced = False
result: Final = _fetch_remote_model_cost_map_with_retry_sync(
url=url,
timeout=timeout,
max_attempts=max_attempts,
sleep=sleep,
rng=rng if rng is not None else random.Random(),
client=client if client is not None else httpx,
)
if isinstance(result, ModelCostMapReloadUnavailable):
fetch_client: Final = client if client is not None else httpx
fetch_rng: Final = rng if rng is not None else random.Random()
outcome: Final = _attempt_fetch_sync(client=fetch_client, url=url, timeout=timeout)
if isinstance(outcome, _FetchAttemptRetryable) and max_attempts > 1:
threading.Thread(
target=_retry_remote_fetch_in_background,
kwargs={ # mutable-ok: threading requires a mutable keyword-arguments mapping
"url": url,
"timeout": timeout,
"max_attempts": max_attempts,
"sleep": sleep,
"rng": fetch_rng,
"client": fetch_client,
"first_outcome": outcome,
},
name="litellm-model-cost-map-retry",
daemon=True,
).start()
if not isinstance(outcome, ModelCostMapReloaded):
verbose_logger.warning(
"LiteLLM: Failed to fetch remote model cost map from %s: %s. Falling back to local backup.",
url,
result.reason,
outcome.reason,
)
_cost_map_source_info.source = "local"
_cost_map_source_info.fallback_reason = f"Remote fetch failed: {result.reason}"
return _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
content: Final = result.model_cost_map
_cost_map_source_info.fallback_reason = f"Remote fetch failed: {outcome.reason}"
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision()).model_cost_map
content: Final = outcome.model_cost_map
# Validate using cached count (cheap int comparison, no file I/O)
if not GetModelCostMap.validate_model_cost_map(
@ -529,8 +673,8 @@ def get_model_cost_map(
)
_cost_map_source_info.source = "local"
_cost_map_source_info.fallback_reason = "Remote data failed integrity validation"
return _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision()).model_cost_map
_cost_map_source_info.source = "remote"
_cost_map_source_info.fallback_reason = None
return _finalize_model_cost_map(content)
return _finalize_loaded_model_cost_map(outcome).model_cost_map

View file

@ -202,6 +202,8 @@ if TYPE_CHECKING:
from mcp.types import EmbeddedResource, ImageContent, TextContent
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
@ -589,6 +591,7 @@ class Logging(LiteLLMLoggingBaseClass):
# Initialize cost breakdown field
self.cost_breakdown: CostBreakdown | None = None
self.billed_token_rates: BilledTokenRates | None = None
# Init Caching related details
self.caching_details: CachingDetails | None = None
@ -1586,6 +1589,7 @@ class Logging(LiteLLMLoggingBaseClass):
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
billed_token_rates: "BilledTokenRates | None" = None,
) -> None:
"""
Helper method to store cost breakdown in the logging object.
@ -1605,8 +1609,10 @@ class Logging(LiteLLMLoggingBaseClass):
service_tier: Tier the costs above were priced on, already resolved
data_residency: Region uplift the costs above were priced on, already resolved
vertex_location: Vertex AI location the costs above were priced on, already resolved
billed_token_rates: Per-token rates the costs above were billed at, already resolved
"""
self.billed_token_rates = billed_token_rates
self.cost_breakdown = CostBreakdown(
input_cost=input_cost,
output_cost=output_cost,
@ -4850,31 +4856,83 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[Custom
Returns ``None`` when V2 is off OR when there's no preset registered for
``callback_name`` callers should then fall through to the legacy path.
A preset that needs operator credentials it cannot find is allowed to build
only when this request has a key/team destination for that backend and another
V2 logger is already registered to carry the fan-out. The resulting logger keeps
only its credential-gated exporter, while the registered logger owns operator
delivery. Without that carrier, a preset that raises or that ends up with nothing
but its gated exporter and the default console placeholder returns ``None``, so the
caller falls through to the legacy path exactly as before V2 landed.
"""
from litellm.integrations.otel.model.config import is_otel_v2_enabled
if not is_otel_v2_enabled():
return None
from litellm.integrations.otel.logger import OpenTelemetryV2, build_otel_v2_logger
from litellm.integrations.otel.plumbing.context import destination_backends
from litellm.integrations.otel.presets import PRESET_BY_CALLBACK
preset_fn: Final = PRESET_BY_CALLBACK.get(callback_name)
if preset_fn is None:
return None
serves_a_destination: Final = callback_name in destination_backends()
has_v2_logger: Final = any(isinstance(callback, OpenTelemetryV2) for callback in _in_memory_loggers)
carried: Final = serves_a_destination and has_v2_logger
for callback in _in_memory_loggers:
if isinstance(callback, OpenTelemetryV2) and getattr(callback, "callback_name", None) == callback_name:
if (
isinstance(callback, OpenTelemetryV2)
and getattr(callback, "callback_name", None) == callback_name
and (serves_a_destination or not _exports_nowhere(callback.config))
):
return callback
try:
config: Final = preset_fn()
built: Final = preset_fn(allow_missing_credentials=carried)
except Exception:
# If env vars are missing or the preset raises, defer to the legacy path
# so customers get the same error story they had before V2 landed.
return None
gated: Final = _is_credential_gated(built)
if gated and not carried and not _has_operator_exporter(built):
return None
config: Final = _only_the_gated_exporter(built) if gated and carried else built
if _exports_nowhere(config):
verbose_logger.warning(
"OTel V2: no operator credentials for '%s'; only key/team destinations will receive its traces",
callback_name,
)
v2_logger: Final = build_otel_v2_logger(config=config, callback_name=callback_name)
_in_memory_loggers.append(v2_logger)
return v2_logger
def _exports_nowhere(config: "OpenTelemetryV2Config") -> bool:
"""Whether every exporter in ``config`` is waiting on credentials it never got."""
return all(_is_gated(spec) for spec in config.exporters)
def _is_credential_gated(config: "OpenTelemetryV2Config") -> bool:
"""Whether the preset built without the operator's own credentials for its backend."""
return any(_is_gated(spec) for spec in config.exporters)
def _has_operator_exporter(config: "OpenTelemetryV2Config") -> bool:
"""Whether the operator configured somewhere real to export, beyond the default console placeholder."""
from litellm.integrations.otel.presets.utils import is_unconfigured_placeholder
return any(not _is_gated(spec) and not is_unconfigured_placeholder(spec) for spec in config.exporters)
def _only_the_gated_exporter(config: "OpenTelemetryV2Config") -> "OpenTelemetryV2Config":
return config.model_copy(
update={"exporters": [spec for spec in config.exporters if _is_gated(spec)]} # mutable-ok: model_copy update
)
def _is_gated(spec: "ExporterSpec") -> bool:
return spec.requires_headers and not spec.headers
def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list[CustomLogger]) -> None:
"""
Auto-initialize ArizePhoenixLogger when Phoenix env vars are detected.

View file

@ -10,6 +10,7 @@ from typing import Any, Final, Literal, TypedDict, cast
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
import litellm
from litellm._internal_context import current_billing_time
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import (
select_tier_for_input,
@ -19,6 +20,7 @@ from litellm.types.utils import (
CacheCreationTokenDetails,
CallTypes,
CompletionTokensDetailsWrapper,
CostPerToken,
DataResidency,
ImageResponse,
ModelInfo,
@ -305,7 +307,7 @@ def _is_within_off_peak_window(off_peak_hours_utc: str | Sequence[str], current_
than being localised, so callers must pass datetime.now(timezone.utc), never datetime.now(),
or every window shifts by the host's offset.
"""
reference: Final = current_time if current_time is not None else datetime.now(timezone.utc)
reference: Final = current_time if current_time is not None else current_billing_time()
now: Final = (reference.astimezone(timezone.utc) if reference.tzinfo is not None else reference).time()
windows: Final = (off_peak_hours_utc,) if isinstance(off_peak_hours_utc, str) else off_peak_hours_utc
for window in windows:
@ -392,7 +394,7 @@ def _is_off_peak(off_peak: Mapping[str, object], current_time: datetime | None =
rules: the flat hours_utc windows, which apply every day, or any entry in windows, whose
hours apply only on its weekdays.
"""
reference: Final = current_time if current_time is not None else datetime.now(timezone.utc)
reference: Final = current_time if current_time is not None else current_billing_time()
reference_utc: Final = (
reference.astimezone(timezone.utc) if reference.tzinfo is not None else reference.replace(tzinfo=timezone.utc)
)
@ -780,6 +782,7 @@ class PromptTokensDetailsResult(TypedDict):
image_count: int
video_length_seconds: float
audio_length_seconds: float
query_count: int
def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
@ -828,6 +831,7 @@ def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
)
or 0.0
)
query_count: Final = _coerce_token_count(getattr(usage.prompt_tokens_details, "query_count", 0))
return PromptTokensDetailsResult(
cache_hit_tokens=cache_hit_tokens,
@ -841,6 +845,7 @@ def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
image_count=image_count,
video_length_seconds=float(video_length_seconds),
audio_length_seconds=float(audio_length_seconds),
query_count=query_count,
)
@ -978,6 +983,11 @@ def _calculate_input_cost(
prompt_tokens_details["audio_length_seconds"],
)
if prompt_tokens_details["query_count"]:
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_query", prompt_tokens_details["query_count"]
)
return prompt_cost
@ -1149,6 +1159,7 @@ def generic_cost_per_token(
image_count=0,
video_length_seconds=0.0,
audio_length_seconds=0.0,
query_count=0,
)
if usage.prompt_tokens_details:
prompt_tokens_details = parse_prompt_tokens_details(usage)
@ -1186,7 +1197,7 @@ def generic_cost_per_token(
usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens, 0
)
billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
billing_time: Final = current_time if current_time is not None else current_billing_time()
(
prompt_base_cost,
completion_base_cost,
@ -1300,42 +1311,90 @@ def _coerce_token_count(value: object) -> int:
return value if isinstance(value, int) and value > 0 else 0
@dataclass(frozen=True, slots=True)
class BilledTokenRates:
"""Per-token rates one request's usage bills at, after token tiers, off-peak windows and the
regional multipliers the totals apply, so each cost line equals its token count times its rate."""
input_cost_per_token: float
output_cost_per_token: float
cache_read_input_token_cost: float
cache_creation_input_token_cost: float
cache_creation_input_token_cost_above_1hr: float
output_cost_per_reasoning_token: float
def scaled(self, multiplier: float) -> "BilledTokenRates":
if multiplier == 1.0:
return self
return BilledTokenRates(
input_cost_per_token=self.input_cost_per_token * multiplier,
output_cost_per_token=self.output_cost_per_token * multiplier,
cache_read_input_token_cost=self.cache_read_input_token_cost * multiplier,
cache_creation_input_token_cost=self.cache_creation_input_token_cost * multiplier,
cache_creation_input_token_cost_above_1hr=self.cache_creation_input_token_cost_above_1hr * multiplier,
output_cost_per_reasoning_token=self.output_cost_per_reasoning_token * multiplier,
)
@dataclass(frozen=True, slots=True)
class TokenTypeCostBreakdown:
reasoning_cost: float
cache_read_cost: float
cache_creation_cost: float
rates: BilledTokenRates | None = None
"""Rates these lines were billed at, so a caller reporting both cannot resolve them a second,
differently-argued way. None when the model's pricing could not be resolved."""
def get_token_type_cost_breakdown(
model: str,
custom_llm_provider: str | None,
def _reasoning_token_count(usage: Usage) -> int:
parsed: Final = (
parse_completion_tokens_details(usage)["reasoning_tokens"] if usage.completion_tokens_details is not None else 0
)
return parsed or _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
def _cache_token_counts(usage: Usage) -> tuple[int, int, CacheCreationTokenDetails | None]:
"""(cache read tokens, cache creation tokens, cache creation details): read from prompt_tokens_details
first, then the private top-level counters the Usage constructor mirrors cache tokens onto for
providers/callers that bypass the details."""
parsed: Final = parse_prompt_tokens_details(usage) if usage.prompt_tokens_details is not None else None
parsed_read: Final = parsed["cache_hit_tokens"] if parsed is not None else 0
parsed_creation: Final = parsed["cache_creation_tokens"] if parsed is not None else 0
return (
parsed_read or _coerce_token_count(getattr(usage, "_cache_read_input_tokens", 0)),
parsed_creation or _coerce_token_count(getattr(usage, "_cache_creation_input_tokens", 0)),
parsed["cache_creation_token_details"] if parsed is not None else None,
)
def _custom_pricing_rates(custom_cost_per_token: CostPerToken) -> BilledTokenRates:
"""Flat custom pricing has no tiers, uplifts or reasoning rate: cache tokens bill at the configured
cache rates (else the input rate) and reasoning at the output rate, as _cost_per_token_custom_pricing_helper does."""
input_rate: Final = custom_cost_per_token["input_cost_per_token"]
output_rate: Final = custom_cost_per_token["output_cost_per_token"]
cache_creation_rate: Final = custom_cost_per_token.get("cache_creation_input_token_cost", input_rate)
return BilledTokenRates(
input_cost_per_token=input_rate,
output_cost_per_token=output_rate,
cache_read_input_token_cost=custom_cost_per_token.get("cache_read_input_token_cost", input_rate),
cache_creation_input_token_cost=cache_creation_rate,
cache_creation_input_token_cost_above_1hr=cache_creation_rate,
output_cost_per_reasoning_token=output_rate,
)
def _cost_map_billed_rates(
model_info: ModelInfo,
usage: Usage,
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
current_time: datetime | None = None,
) -> TokenTypeCostBreakdown:
"""
Provider-agnostic cost of reasoning and cache tokens, derived from the usage
object and model pricing alone.
This works for every provider, including Perplexity/Cerebras/Dashscope whose
cost calculators bypass ``generic_cost_per_token``, because cache tokens always
land on ``prompt_tokens_details`` (via the Usage constructor and provider
transformations) and reasoning tokens on ``completion_tokens_details``. It reuses
the same rate-resolution primitives as the total-cost path so the breakdown can
never drift from the totals. Returns zeros (never raises) when the model or its
pricing cannot be resolved.
"""
try:
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception:
return TokenTypeCostBreakdown(0.0, 0.0, 0.0)
billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
custom_llm_provider: str | None,
service_tier: str | None,
data_residency: str | None,
vertex_location: str | None,
current_time: datetime | None,
) -> BilledTokenRates:
billing_time: Final = current_time if current_time is not None else current_billing_time()
(
_prompt_base_cost,
prompt_base_cost,
completion_base_cost,
cache_creation_cost_rate,
cache_creation_cost_above_1hr_rate,
@ -1347,13 +1406,6 @@ def get_token_type_cost_breakdown(
current_time=billing_time,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
)
reasoning_tokens = (
parse_completion_tokens_details(usage)["reasoning_tokens"] if usage.completion_tokens_details is not None else 0
)
if not reasoning_tokens:
reasoning_tokens = _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
reasoning_rate: Final = _resolve_billed_reasoning_rate(
model_info=model_info,
usage=usage,
@ -1361,57 +1413,103 @@ def get_token_type_cost_breakdown(
completion_base_cost=completion_base_cost,
current_time=billing_time,
)
reasoning_cost = float(reasoning_tokens) * reasoning_rate
multiplier: Final = (
_get_regional_uplift_multiplier(model_info, data_residency)
* get_vertex_regional_endpoint_uplift(model_info, vertex_location)
* get_provider_specific_geo_multiplier(model_info=model_info, usage=usage)
)
return BilledTokenRates(
input_cost_per_token=prompt_base_cost,
output_cost_per_token=completion_base_cost,
cache_read_input_token_cost=cache_read_cost_rate,
cache_creation_input_token_cost=cache_creation_cost_rate,
cache_creation_input_token_cost_above_1hr=cache_creation_cost_above_1hr_rate,
output_cost_per_reasoning_token=reasoning_rate,
).scaled(multiplier)
cache_read_tokens = 0
cache_creation_tokens = 0
cache_creation_token_details: CacheCreationTokenDetails | None = None
if usage.prompt_tokens_details is not None:
prompt_tokens_details: Final = parse_prompt_tokens_details(usage)
cache_read_tokens = prompt_tokens_details["cache_hit_tokens"]
cache_creation_tokens = prompt_tokens_details["cache_creation_tokens"]
cache_creation_token_details = prompt_tokens_details["cache_creation_token_details"]
# Fall back to the private top-level counters the Usage constructor mirrors cache
# tokens onto, so providers/callers that bypass prompt_tokens_details are covered.
if not cache_read_tokens:
cache_read_tokens = _coerce_token_count(getattr(usage, "_cache_read_input_tokens", 0))
if not cache_creation_tokens:
cache_creation_tokens = _coerce_token_count(getattr(usage, "_cache_creation_input_tokens", 0))
cache_read_cost = float(cache_read_tokens) * cache_read_cost_rate
cache_creation_cost = calculate_cache_writing_cost(
cache_creation_tokens=cache_creation_tokens,
cache_creation_token_details=cache_creation_token_details,
cache_creation_cost_above_1hr=cache_creation_cost_above_1hr_rate,
cache_creation_cost=cache_creation_cost_rate,
def get_billed_token_rates(
model: str,
custom_llm_provider: str | None,
usage: Usage,
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
current_time: datetime | None = None,
custom_cost_per_token: CostPerToken | None = None,
) -> BilledTokenRates | None:
"""Rates the cost calculator bills ``usage`` at, resolved exactly as the totals and the token-type
breakdown resolve them. None when the model's pricing cannot be resolved."""
if custom_cost_per_token is not None:
return _custom_pricing_rates(custom_cost_per_token)
try:
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception: # noqa: BLE001 # get_model_info raises a bare Exception for an unmapped model: no rates
return None
return _cost_map_billed_rates(
model_info=model_info,
usage=usage,
custom_llm_provider=custom_llm_provider,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
current_time=current_time,
)
# Apply the same flat regional-processing uplift the totals get, so per-type
# costs stay reconciled with input_cost/output_cost for regionalized OpenAI hosts.
uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency)
if uplift != 1.0:
reasoning_cost *= uplift
cache_read_cost *= uplift
cache_creation_cost *= uplift
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
if vertex_uplift != 1.0:
reasoning_cost *= vertex_uplift
cache_read_cost *= vertex_uplift
cache_creation_cost *= vertex_uplift
def get_token_type_cost_breakdown(
model: str,
custom_llm_provider: str | None,
usage: Usage,
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
current_time: datetime | None = None,
custom_cost_per_token: CostPerToken | None = None,
) -> TokenTypeCostBreakdown:
"""
Provider-agnostic cost of reasoning and cache tokens, derived from the usage
object and model pricing alone.
# Mirror the provider-specific geo uplift (e.g. Anthropic us: 1.1) the totals
# apply, so cache and reasoning line items stay reconciled with them.
geo_multiplier: Final = get_provider_specific_geo_multiplier(model_info=model_info, usage=usage)
if geo_multiplier != 1.0:
reasoning_cost *= geo_multiplier
cache_read_cost *= geo_multiplier
cache_creation_cost *= geo_multiplier
This works for every provider, including Perplexity/Cerebras/Dashscope whose
cost calculators bypass ``generic_cost_per_token``, because cache tokens always
land on ``prompt_tokens_details`` (via the Usage constructor and provider
transformations) and reasoning tokens on ``completion_tokens_details``. It reuses
the same rate resolution as the total-cost path (``get_billed_token_rates``) so the
breakdown can never drift from the totals. A deployment billed by
``custom_cost_per_token`` is priced from those flat rates instead of the cost map and,
like its totals, bills cache writes flat rather than by their 5m/1h split.
Returns zeros (never raises) when the model or its pricing cannot be resolved.
"""
rates: Final = get_billed_token_rates(
model=model,
custom_llm_provider=custom_llm_provider,
usage=usage,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
current_time=current_time,
custom_cost_per_token=custom_cost_per_token,
)
if rates is None:
return TokenTypeCostBreakdown(0.0, 0.0, 0.0)
cache_read_tokens, cache_creation_tokens, cache_creation_token_details = _cache_token_counts(usage)
cache_creation_cost: Final = (
float(cache_creation_tokens) * rates.cache_creation_input_token_cost
if custom_cost_per_token is not None
else calculate_cache_writing_cost(
cache_creation_tokens=cache_creation_tokens,
cache_creation_token_details=cache_creation_token_details,
cache_creation_cost_above_1hr=rates.cache_creation_input_token_cost_above_1hr,
cache_creation_cost=rates.cache_creation_input_token_cost,
)
)
return TokenTypeCostBreakdown(
reasoning_cost=reasoning_cost,
cache_read_cost=cache_read_cost,
reasoning_cost=float(_reasoning_token_count(usage)) * rates.output_cost_per_reasoning_token,
cache_read_cost=float(cache_read_tokens) * rates.cache_read_input_token_cost,
cache_creation_cost=cache_creation_cost,
rates=rates,
)

View file

@ -183,6 +183,7 @@ class _RemoteSource:
class RemoteMedia:
url: str
fields: Mapping[str, object]
part_type: str
_NO_FIELDS: Final[Mapping[str, object]] = MappingProxyType({})
@ -192,6 +193,10 @@ def inline_every_remote_url(_media: RemoteMedia) -> bool:
return True
def inline_remote_image_urls(media: RemoteMedia) -> bool:
return media.part_type == "image_url"
def _parse_remote_image(fields: Mapping[str, object]) -> _RemoteImage | None:
if fields.get("type") != "image_url":
return None
@ -223,11 +228,11 @@ def _parse_remote_part(part: object) -> _RemoteImage | _RemoteFile | _RemoteSour
def _remote_media(remote: _RemoteImage | _RemoteFile | _RemoteSource) -> RemoteMedia:
match remote:
case _RemoteImage(_, image_url, url):
return RemoteMedia(url, image_url if image_url is not None else _NO_FIELDS)
return RemoteMedia(url, image_url if image_url is not None else _NO_FIELDS, "image_url")
case _RemoteFile(_, file, url):
return RemoteMedia(url, file)
case _RemoteSource(_, source, url):
return RemoteMedia(url, source)
return RemoteMedia(url, file, "file")
case _RemoteSource(part, source, url):
return RemoteMedia(url, source, str(part.get("type")))
_PDF_FORMAT: Final = MappingProxyType({"format": "application/pdf"})

View file

@ -24,6 +24,7 @@ from litellm.types.utils import (
Choices,
CompletionTokensDetails,
CompletionTokensDetailsWrapper,
Delta,
Function,
FunctionCall,
ModelResponse,
@ -326,6 +327,18 @@ class ChunkProcessor:
return chunk_id
return ""
@staticmethod
def _get_role_from_chunks(chunks: Sequence["_BaseChunk"]) -> str:
return ChunkProcessor._role_of_choice(next((c["choices"][0] for c in chunks if c.get("choices")), None))
@staticmethod
def _role_of_choice(choice: object) -> str:
match choice:
case StreamingChoices(delta=Delta(role=str() as role)) | {"delta": {"role": str() as role}} if role:
return role
case _:
return "assistant"
@staticmethod
def _get_model_from_chunks(chunks: Sequence["_BaseChunk"], first_chunk_model: str) -> str:
"""
@ -353,8 +366,7 @@ class ChunkProcessor:
model: Final = ChunkProcessor._get_model_from_chunks(chunks, first_chunk_model)
system_fingerprint: Final = chunk.get("system_fingerprint", None)
first_chunk_with_choices: Final = next((c for c in chunks if c.get("choices")), chunk)
role: Final = first_chunk_with_choices["choices"][0]["delta"]["role"]
role: Final = ChunkProcessor._get_role_from_chunks(chunks)
finish_reason = "stop"
for chunk in chunks:
if "choices" in chunk and len(chunk["choices"]) > 0:

View file

@ -368,27 +368,92 @@ class AnthropicChatCompletion(BaseLLM):
if config is None:
raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}")
def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream
"""Translate the request the Python way, returning `(headers, data)`.
transform_params: Final = {**optional_params, "is_vertex_request": is_vertex_request}
def finish_request(request_data: dict) -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream
"""Filter beta headers and emit pre_call, returning `(headers, data)`.
The pair stays mutable because the streaming path rewrites it in
place (`data["stream"] = True`) before sending.
Shared by the normal path and by the Rust path's fallback, which
builds it only when the Rust call did not serve the request.
place (`data["stream"] = True`) before sending. A Rust attempt that
declined already emitted pre_call for this request, so skip it there.
"""
request_data: Final = config.transform_request(
model=model,
messages=messages,
optional_params={**optional_params, "is_vertex_request": is_vertex_request},
litellm_params=litellm_params,
headers=headers,
)
return update_request_with_filtered_beta(
request_headers, data = update_request_with_filtered_beta(
headers=headers,
request_data=request_data,
provider=custom_llm_provider,
)
if not serves_via_rust:
logging_obj.pre_call(
input=messages,
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": request_headers,
},
)
print_verbose(f"_is_function_call: {_is_function_call}")
return request_headers, data
async def acompletion_dispatch() -> "ModelResponse | CustomStreamWrapper":
"""Translate then send, so the provider config can inline remote media off the event loop."""
request_headers, data = finish_request(
await config.async_transform_request(
model=model,
messages=messages,
optional_params=transform_params,
litellm_params=litellm_params,
headers=headers,
)
)
if (
stream is True
): # if function call - fake the streaming (need complete blocks for output parsing in openai format)
print_verbose("makes async anthropic streaming POST request")
data["stream"] = stream
return await self.acompletion_stream_function(
model=model,
messages=messages,
data=data,
api_base=api_base,
custom_prompt_dict=custom_prompt_dict,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
api_key=api_key,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
_is_function_call=_is_function_call,
json_mode=json_mode,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=request_headers,
timeout=timeout,
client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None),
)
return await self.acompletion_function(
model=model,
messages=messages,
data=data,
api_base=api_base,
custom_prompt_dict=custom_prompt_dict,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
api_key=api_key,
provider_config=config,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
_is_function_call=_is_function_call,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=request_headers,
client=client,
json_mode=json_mode,
timeout=timeout,
)
# The Rust core owns the whole call for the subset it accepts, so ask
# before transforming: whichever path runs emits pre_call exactly once.
@ -424,35 +489,6 @@ class AnthropicChatCompletion(BaseLLM):
additional_args=rust_logging_args,
)
if acompletion is True:
async def python_fallback() -> "ModelResponse | CustomStreamWrapper":
# pre_call already fired for this request above. The Rust
# path only declines before the provider is called, so this
# is the same attempt continuing, not a second one.
fallback_headers, fallback_data = build_request()
return await self.acompletion_function(
model=model,
messages=messages,
data=fallback_data,
api_base=api_base,
custom_prompt_dict=custom_prompt_dict,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
api_key=api_key,
provider_config=config,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
_is_function_call=_is_function_call,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=fallback_headers,
client=client,
json_mode=json_mode,
timeout=timeout,
)
return rust_chat_completions_bridge.achat_completions_or_fallback(
model=model,
messages=messages,
@ -464,7 +500,7 @@ class AnthropicChatCompletion(BaseLLM):
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
python_fallback=python_fallback,
python_fallback=acompletion_dispatch,
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
model=model,
@ -481,74 +517,18 @@ class AnthropicChatCompletion(BaseLLM):
if rust_response is not None:
return rust_response
headers, data = build_request()
## LOGGING
# Reaching here with `serves_via_rust` set means the Rust attempt
# declined at call time, before the provider was called, and already
# logged this request. That is the same attempt continuing.
if not serves_via_rust:
logging_obj.pre_call(
input=messages,
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": headers,
},
)
print_verbose(f"_is_function_call: {_is_function_call}")
if acompletion is True:
if (
stream is True
): # if function call - fake the streaming (need complete blocks for output parsing in openai format)
print_verbose("makes async anthropic streaming POST request")
data["stream"] = stream
return self.acompletion_stream_function(
model=model,
messages=messages,
data=data,
api_base=api_base,
custom_prompt_dict=custom_prompt_dict,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
api_key=api_key,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
_is_function_call=_is_function_call,
json_mode=json_mode,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=headers,
timeout=timeout,
client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None),
)
else:
return self.acompletion_function(
model=model,
messages=messages,
data=data,
api_base=api_base,
custom_prompt_dict=custom_prompt_dict,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
api_key=api_key,
provider_config=config,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
_is_function_call=_is_function_call,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=headers,
client=client,
json_mode=json_mode,
timeout=timeout,
)
return acompletion_dispatch()
else:
headers, data = finish_request(
config.transform_request(
model=model,
messages=messages,
optional_params=transform_params,
litellm_params=litellm_params,
headers=headers,
)
)
## COMPLETION CALL
if (
stream is True

View file

@ -26,6 +26,11 @@ from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.prompt_templates.common_utils import (
sanitize_input_schema_for_anthropic,
)
from litellm.litellm_core_utils.prompt_templates.image_handling import (
RemoteMedia,
async_inline_remote_media,
inline_remote_image_urls,
)
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.anthropic import (
@ -1840,6 +1845,25 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
break
return headers
def inlines_remote_media(self, media: RemoteMedia) -> bool:
return inline_remote_image_urls(media) and media.url.startswith("http://")
async def async_transform_request(
self,
model: str,
messages: list[AllMessageValues], # mutable-ok: BaseConfig signature
optional_params: dict[str, object], # mutable-ok: BaseConfig signature
litellm_params: dict[str, object], # mutable-ok: BaseConfig signature
headers: dict[str, object], # mutable-ok: BaseConfig signature
) -> dict[str, object]: # mutable-ok: BaseConfig signature
return self.transform_request(
model=model,
messages=await async_inline_remote_media(messages, should_inline=self.inlines_remote_media),
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
def transform_request(
self,
model: str,

View file

@ -3,6 +3,8 @@ from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from pydantic import BaseModel, ConfigDict, ValidationError
import litellm
from litellm.types.utils import ModelInfo
@ -21,10 +23,27 @@ _EFFORT_DEGRADATION_CHAIN: Final[Mapping[str, tuple[str, ...]]] = MappingProxyTy
_THINKING_OFF: Final = "none"
class _ClaudeCodeUserId(BaseModel):
"""The JSON Claude Code packs into ``metadata.user_id``; only ``session_id`` is per conversation."""
model_config = ConfigDict(frozen=True)
session_id: str
def prompt_cache_key_from_user_id(user_id: object) -> str | None:
if user_id is None:
"""The per-session key Claude Code carries inside ``metadata.user_id``, or nothing.
Anthropic defines ``user_id`` as an opaque end-user id, so a plain string names a person, not
a conversation. Keying the provider cache on it pins every parallel session and subagent of that
person to one slot, which caches worse than the provider's own prompt-prefix hashing does.
"""
if not isinstance(user_id, str):
return None
try:
return _ClaudeCodeUserId.model_validate_json(user_id).session_id[:OPENAI_MAX_PROMPT_CACHE_KEY_LENGTH] or None
except ValidationError:
return None
return str(user_id)[:OPENAI_MAX_PROMPT_CACHE_KEY_LENGTH] or None
def litellm_logging_obj_from_kwargs(kwargs: Mapping[str, object]) -> "LiteLLMLoggingObject | None":

View file

@ -37,6 +37,7 @@ class AzureAudioTranscription(AzureChatCompletion):
azure_ad_token: str | None = None,
atranscription: bool = False,
litellm_params: dict | None = None,
custom_llm_provider: str = "azure",
) -> TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse]:
data: Final = {"model": model, "file": audio_file, **optional_params}
@ -53,6 +54,7 @@ class AzureAudioTranscription(AzureChatCompletion):
logging_obj=logging_obj,
model=model,
litellm_params=litellm_params,
custom_llm_provider=custom_llm_provider,
)
azure_client: Final = self.get_azure_openai_client(
@ -99,7 +101,7 @@ class AzureAudioTranscription(AzureChatCompletion):
additional_args={"complete_input_dict": data},
original_response=stringified_response,
)
hidden_params: Final = {"model": model, "custom_llm_provider": "azure"}
hidden_params: Final = {"model": model, "custom_llm_provider": custom_llm_provider}
final_response: Final[TranscriptionResponse] = convert_to_model_response_object(
response_object=stringified_response,
model_response_object=model_response,
@ -122,6 +124,7 @@ class AzureAudioTranscription(AzureChatCompletion):
client=None,
max_retries=None,
litellm_params: dict | None = None,
custom_llm_provider: str = "azure",
) -> TranscriptionResponse:
response = None
try:
@ -178,7 +181,7 @@ class AzureAudioTranscription(AzureChatCompletion):
},
original_response=stringified_response,
)
hidden_params: Final = {"model": model, "custom_llm_provider": "azure"}
hidden_params: Final = {"model": model, "custom_llm_provider": custom_llm_provider}
response = convert_to_model_response_object(
_response_headers=headers,
response_object=stringified_response,

View file

@ -11,7 +11,7 @@ from litellm.types.utils import Usage
from litellm.utils import get_model_info
def _is_azure_model_router(model: str) -> bool:
def is_azure_model_router(model: str) -> bool:
"""
Check if the model is Azure AI Foundry Model Router.
@ -31,6 +31,18 @@ def _is_azure_model_router(model: str) -> bool:
return "model-router" in model_lower or "model_router" in model_lower or model_lower == "azure-model-router"
ROUTER_FEE_ENTRY_NAMES: Final = frozenset({"model-router", "model_router"})
def is_router_fee_entry(model: str) -> bool:
return model.lower().removeprefix("azure_ai/") in ROUTER_FEE_ENTRY_NAMES
def _router_fee_entry_name(model: str) -> str:
entry_name: Final = model.lower().removeprefix("azure_ai/")
return entry_name if entry_name in ROUTER_FEE_ENTRY_NAMES else "model_router"
def calculate_azure_model_router_flat_cost(model: str, prompt_tokens: int) -> float:
"""
Calculate the flat cost for Azure AI Foundry Model Router.
@ -42,20 +54,39 @@ def calculate_azure_model_router_flat_cost(model: str, prompt_tokens: int) -> fl
Returns:
float: The flat cost in USD, or 0.0 if not applicable
"""
if not _is_azure_model_router(model):
if not is_azure_model_router(model):
return 0.0
# Get the model router pricing from model_prices_and_context_window.json
# Use "model_router" as the key (without actual model name suffix)
model_info: Final = get_model_info(model="model_router", custom_llm_provider="azure_ai")
model_info: Final = get_model_info(model=_router_fee_entry_name(model), custom_llm_provider="azure_ai")
router_flat_cost_per_token: Final = model_info.get("input_cost_per_token", 0)
if router_flat_cost_per_token and router_flat_cost_per_token > 0:
return prompt_tokens * router_flat_cost_per_token
return 0.0
def _response_model_cost(model: str, usage: Usage, service_tier: str | None) -> tuple[float, float]:
try:
return generic_cost_per_token(
model=model, usage=usage, custom_llm_provider="azure_ai", service_tier=service_tier
)
except Exception as e:
if not is_azure_model_router(model):
raise
verbose_logger.debug(
"Azure AI Model Router: model '%s' not in cost map, only the routing fee applies. Error: %s", model, e
)
return 0.0, 0.0
def _router_fee_name(model: str, request_model: str | None) -> str | None:
if is_router_fee_entry(model):
return None
if is_azure_model_router(model):
return model
if request_model is not None and is_azure_model_router(request_model):
return request_model
return None
def cost_per_token(
model: str,
usage: Usage,
@ -64,68 +95,31 @@ def cost_per_token(
service_tier: str | None = None,
) -> tuple[float, float]:
"""
Calculate the cost per token for Azure AI models.
Price the response model's own tokens for Azure AI, plus the Model Router fee exactly once when either the
priced name or request_model is a Model Router name.
For Azure AI Foundry Model Router:
- Adds a flat cost of $0.14 per million input tokens (from model_prices_and_context_window.json)
- Plus the cost of the actual model used (handled by generic_cost_per_token)
A response priced as the router entry itself already carries the fee, so nothing is added on top of it. A
router deployment name that is missing from the cost map prices at the fee alone.
completion_cost passes only the priced name: when that name is a routed model it adds the fee itself through
AzureModelRouterConfig.calculate_additional_costs as the "Azure Model Router Flat Cost" line of the cost
breakdown, and when the name is router-shaped the fee is already in the prompt cost returned here.
Args:
model: str, the model name without provider prefix (from response)
usage: LiteLLM Usage block
response_time_ms: Optional response time in milliseconds
request_model: Optional[str], the original request model name (to detect router usage)
request_model: Optional[str], the original request model name; a Model Router name adds the routing fee
service_tier: Optional service tier the request was priced on
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
Raises:
ValueError: If the model is not found in the cost map and cost cannot be calculated
(except for Model Router models where we return just the routing flat cost)
ValueError: If a model that is not a Model Router name is missing from the cost map
"""
prompt_cost = 0.0
completion_cost = 0.0
# Determine if this was a model router request
# Check both the response model and the request model
is_router_request: Final = _is_azure_model_router(model) or (
request_model is not None and _is_azure_model_router(request_model)
)
# Calculate base cost using generic cost calculator
# This may raise an exception if the model is not in the cost map
try:
prompt_cost, completion_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="azure_ai",
service_tier=service_tier,
)
except Exception as e:
# For Model Router, the model name (e.g., "azure-model-router") may not be in the cost map
# because it's a routing service, not an actual model. In this case, we continue
# to calculate just the routing flat cost.
if not _is_azure_model_router(model):
# Re-raise for non-router models - they should have pricing defined
raise
verbose_logger.debug(
"Azure AI Model Router: model '%s' not in cost map, calculating routing flat cost only. Error: %s", model, e
)
# Add flat cost for Azure Model Router
# The flat cost is defined in model_prices_and_context_window.json for azure_ai/model_router
if is_router_request:
# Use the request model for flat cost calculation if available, otherwise use response model
router_model_for_calc: Final = request_model if request_model else model
router_flat_cost: Final = calculate_azure_model_router_flat_cost(router_model_for_calc, usage.prompt_tokens)
if router_flat_cost > 0:
verbose_logger.debug(
f"Azure AI Model Router flat cost: ${router_flat_cost:.6f} "
f"({usage.prompt_tokens} tokens × ${router_flat_cost / usage.prompt_tokens:.9f}/token)"
)
# Add flat cost to prompt cost
prompt_cost += router_flat_cost
return prompt_cost, completion_cost
prompt_cost, completion_cost = _response_model_cost(model=model, usage=usage, service_tier=service_tier)
fee_name: Final = _router_fee_name(model=model, request_model=request_model)
if fee_name is None:
return prompt_cost, completion_cost
return prompt_cost + calculate_azure_model_router_flat_cost(fee_name, usage.prompt_tokens), completion_cost

View file

@ -121,6 +121,9 @@ class RouterVectorStoreEmbeddingExecutor:
class BaseVectorStoreConfig:
def validate_create_vector_store(self) -> None:
return None
def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]:
return []

View file

@ -35,7 +35,7 @@ from .amazon_titan_multimodal_transformation import (
)
from .amazon_titan_v2_transformation import AmazonTitanV2Config
from .cohere_transformation import BedrockCohereEmbeddingConfig
from .twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig
from .twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig, drop_params_enabled
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -239,7 +239,7 @@ class BedrockEmbedding(BaseAWSLLM):
returned_response = AmazonTitanG1Config()._transform_response(response_list=response_list, model=model)
elif provider == "twelvelabs":
returned_response = TwelveLabsMarengoEmbeddingConfig()._transform_response(
response_list=response_list, model=model
response_list=response_list, model=model, batch_data=batch_data
)
elif provider == "nova":
returned_response = AmazonNovaEmbeddingConfig()._transform_response(
@ -484,12 +484,13 @@ class BedrockEmbedding(BaseAWSLLM):
elif provider == "twelvelabs":
batch_data = []
for i in input:
twelvelabs_request = TwelveLabsMarengoEmbeddingConfig()._transform_request(
twelvelabs_request = TwelveLabsMarengoEmbeddingConfig(model=model)._transform_request(
input=i,
inference_params=inference_params,
async_invoke_route=has_async_invoke,
model_id=modelId,
output_s3_uri=inference_params.get("output_s3_uri"),
drop_params=drop_params_enabled(litellm_params),
)
batch_data.append(twelvelabs_request)
elif provider == "nova":

View file

@ -0,0 +1,239 @@
"""
Request builder for Bedrock TwelveLabs Marengo Embed 3.0, whose payload nests the input under a key named after
``inputType`` instead of the flat 2.7 layout.
Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo-3.html
"""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.types.llms.bedrock import (
TWELVELABS_MARENGO_3_EMBEDDING_OPTIONS,
TWELVELABS_MARENGO_3_EMBEDDING_SCOPES,
TWELVELABS_MARENGO_3_EMBEDDING_TYPES,
TWELVELABS_MARENGO_3_INPUT_TYPES,
TwelveLabsMarengo3AudioRequest,
TwelveLabsMarengo3EmbeddingRequest,
TwelveLabsMarengo3ImageRequest,
TwelveLabsMarengo3MultiInputRequest,
TwelveLabsMarengo3NamedMediaSource,
TwelveLabsMarengo3RequestBase,
TwelveLabsMarengo3Segmentation,
TwelveLabsMarengo3TextImageRequest,
TwelveLabsMarengo3TextRequest,
TwelveLabsMarengo3TimedMediaInput,
TwelveLabsMarengo3TimedMediaOptions,
TwelveLabsMarengo3VideoRequest,
TwelveLabsMediaSource,
TwelveLabsS3Location,
)
from litellm.utils import get_base64_str
MARENGO_3_MODEL_MARKER: Final = "marengo-embed-3-"
S3_URI_PREFIX: Final = "s3://"
TIMED_MEDIA_OPTION_FIELDS: Final = MappingProxyType(
{
"startSec": True,
"endSec": True,
"segmentation": True,
"embeddingOption": True,
"embeddingType": True,
"embeddingScope": True,
}
)
TIMED_MEDIA_OPTIONS: Final = TypeAdapter(TwelveLabsMarengo3TimedMediaOptions)
TIMED_INPUT_TYPES: Final = frozenset({"video", "audio"})
MARENGO_2_7_ONLY_PARAMS: Final = ("textTruncate", "lengthSec", "useFixedLengthSec", "minClipSec")
MARENGO_2_7_ONLY_FIELDS: Final = MappingProxyType({name: True for name in MARENGO_2_7_ONLY_PARAMS})
def is_marengo_3_model(model: str | None) -> bool:
return MARENGO_3_MODEL_MARKER in (model or "")
class Marengo3Params(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
inputType: TWELVELABS_MARENGO_3_INPUT_TYPES | None = None
input_type: TWELVELABS_MARENGO_3_INPUT_TYPES | None = None
media_source: str | None = None
media_sources: Mapping[str, str] | None = None
bucketOwner: str | None = None
startSec: float | None = None
endSec: float | None = None
segmentation: TwelveLabsMarengo3Segmentation | None = None
embeddingOption: tuple[TWELVELABS_MARENGO_3_EMBEDDING_OPTIONS, ...] | None = None
embeddingType: tuple[TWELVELABS_MARENGO_3_EMBEDDING_TYPES, ...] | None = None
embeddingScope: tuple[TWELVELABS_MARENGO_3_EMBEDDING_SCOPES, ...] | None = None
inferenceId: str | None = None
textTruncate: object = None
lengthSec: object = None
useFixedLengthSec: object = None
minClipSec: object = None
@property
def resolved_input_type(self) -> TWELVELABS_MARENGO_3_INPUT_TYPES:
return self.inputType or self.input_type or "text"
def timed_media_options(self) -> TwelveLabsMarengo3TimedMediaOptions:
return TIMED_MEDIA_OPTIONS.validate_python(self.given_timed_media_options())
def given_timed_media_options(self) -> dict[str, object]:
return self.model_dump(include=TIMED_MEDIA_OPTION_FIELDS, exclude_none=True)
def given_2_7_only_params(self) -> dict[str, object]:
return self.model_dump(include=MARENGO_2_7_ONLY_FIELDS, exclude_none=True)
def _require_bucket_owner(bucket_owner: str | None) -> str:
if bucket_owner is None:
raise BedrockError(
status_code=400,
message="s3:// media requires the 'bucketOwner' parameter, the account id that owns the bucket",
)
return bucket_owner
def _media_source(media: str, bucket_owner: str | None) -> TwelveLabsMediaSource:
if not media.startswith(S3_URI_PREFIX):
inline: Final[TwelveLabsMediaSource] = {"base64String": get_base64_str(media)}
return inline
s3_location: Final[TwelveLabsS3Location] = {"uri": media, "bucketOwner": _require_bucket_owner(bucket_owner)}
remote: Final[TwelveLabsMediaSource] = {"s3Location": s3_location}
return remote
def _named_media_source(name: str, media: str, bucket_owner: str | None) -> TwelveLabsMarengo3NamedMediaSource:
named: Final[TwelveLabsMarengo3NamedMediaSource] = {
"name": name,
"mediaType": "image",
**_media_source(media, bucket_owner),
}
return named
def _timed_media_input(media: str, params: Marengo3Params) -> TwelveLabsMarengo3TimedMediaInput:
timed: Final[TwelveLabsMarengo3TimedMediaInput] = {
"mediaSource": _media_source(media, params.bucketOwner),
**params.timed_media_options(),
}
return timed
def _describe(error: ValidationError) -> str:
return "; ".join(
f"{'.'.join(str(part) for part in problem['loc'])}: {problem['msg']}" for problem in error.errors()
)
def _validated_params(inference_params: Mapping[str, object]) -> Marengo3Params:
try:
return Marengo3Params.model_validate(inference_params)
except ValidationError as error:
raise BedrockError(status_code=400, message=f"Invalid Marengo 3.0 parameters: {_describe(error)}") from error
def _reject_unless_dropped(given: Mapping[str, object], drop_params: bool, reason: str) -> None:
if not given or drop_params:
return
raise BedrockError(status_code=400, message=f"{reason} {', '.join(given)}; set drop_params to drop them")
def _require(value: str | None, input_type: str, param_name: str) -> str:
if value is None:
raise BedrockError(status_code=400, message=f"Input type '{input_type}' requires the '{param_name}' parameter")
return value
def _require_media_sources(value: Mapping[str, str] | None) -> Mapping[str, str]:
if not value:
raise BedrockError(
status_code=400,
message="Input type 'multi_input' requires a non-empty 'media_sources' mapping of name to media",
)
return value
def _request_base(inference_id: str | None) -> TwelveLabsMarengo3RequestBase:
if inference_id is None:
anonymous: Final[TwelveLabsMarengo3RequestBase] = {}
return anonymous
identified: Final[TwelveLabsMarengo3RequestBase] = {"inferenceId": inference_id}
return identified
def build_marengo_3_request(
input: str, inference_params: Mapping[str, object], drop_params: bool = False
) -> TwelveLabsMarengo3EmbeddingRequest:
params: Final = _validated_params(inference_params)
base: Final = _request_base(params.inferenceId)
input_type: Final = params.resolved_input_type
_reject_unless_dropped(
params.given_2_7_only_params(), drop_params, "Marengo 3.0 does not accept the Marengo 2.7 parameters"
)
if input_type not in TIMED_INPUT_TYPES:
_reject_unless_dropped(
params.given_timed_media_options(), drop_params, f"Input type '{input_type}' does not accept"
)
match input_type:
case "text":
text_request: Final[TwelveLabsMarengo3TextRequest] = {
**base,
"inputType": "text",
"text": {"inputText": input},
}
return text_request
case "image":
image_request: Final[TwelveLabsMarengo3ImageRequest] = {
**base,
"inputType": "image",
"image": {"mediaSource": _media_source(input, params.bucketOwner)},
}
return image_request
case "video":
video_request: Final[TwelveLabsMarengo3VideoRequest] = {
**base,
"inputType": "video",
"video": _timed_media_input(input, params),
}
return video_request
case "audio":
audio_request: Final[TwelveLabsMarengo3AudioRequest] = {
**base,
"inputType": "audio",
"audio": _timed_media_input(input, params),
}
return audio_request
case "text_image":
text_image_request: Final[TwelveLabsMarengo3TextImageRequest] = {
**base,
"inputType": "text_image",
"text_image": {
"inputText": input,
"mediaSource": _media_source(
_require(params.media_source, input_type, "media_source"), params.bucketOwner
),
},
}
return text_image_request
case "multi_input":
media_sources: Final = tuple(
_named_media_source(name, media, params.bucketOwner)
for name, media in _require_media_sources(params.media_sources).items()
)
multi_input_request: Final[TwelveLabsMarengo3MultiInputRequest] = {
**base,
"inputType": "multi_input",
"multi_input": {"inputText": input, "mediaSources": media_sources}
if input
else {"mediaSources": media_sources},
}
return multi_input_request
case _:
assert_never(input_type)

View file

@ -4,19 +4,120 @@ Transformation logic from OpenAI /v1/embeddings format to Bedrock TwelveLabs Mar
Why separate file? Make it easy to see how transformation works
Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo.html
Marengo 3.0 docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo-3.html
"""
from collections.abc import Mapping
from typing import Final, cast
from pydantic import BaseModel, ConfigDict, TypeAdapter
from typing_extensions import assert_never
import litellm
from litellm.llms.bedrock.embed.twelvelabs_marengo_3_transformation import (
MARENGO_2_7_ONLY_PARAMS,
build_marengo_3_request,
is_marengo_3_model,
)
from litellm.types.llms.bedrock import (
TWELVELABS_EMBEDDING_INPUT_TYPES,
TWELVELABS_MARENGO_3_INPUT_TYPES,
TwelveLabsAsyncInvokeRequest,
TwelveLabsMarengo3EmbeddingRequest,
TwelveLabsMarengoEmbeddingRequest,
TwelveLabsOutputDataConfig,
TwelveLabsS3Location,
TwelveLabsS3OutputDataConfig,
)
from litellm.types.utils import Embedding, EmbeddingResponse, Usage
from litellm.types.utils import Embedding, EmbeddingResponse, PromptTokensDetailsWrapper, Usage
class MarengoEmbeddingItem(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
embedding: tuple[float, ...] | None = None
class MarengoInvokeResponse(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
data: tuple[MarengoEmbeddingItem, ...] = ()
embedding: tuple[float, ...] | None = None
embeddings: tuple[MarengoEmbeddingItem, ...] = ()
def vectors(self) -> tuple[tuple[float, ...], ...]:
if self.data:
return tuple(item.embedding for item in self.data if item.embedding is not None)
if self.embedding is not None:
return (self.embedding,)
return tuple(item.embedding for item in self.embeddings if item.embedding is not None)
class MarengoBilledMultiInput(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
inputText: str | None = None
mediaSources: tuple[Mapping[str, object], ...] = ()
class MarengoBilledRequest(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
inputType: TWELVELABS_MARENGO_3_INPUT_TYPES | None = None
multi_input: MarengoBilledMultiInput | None = None
INVOKE_RESPONSES: Final = TypeAdapter(tuple[MarengoInvokeResponse, ...])
BILLED_REQUESTS: Final = TypeAdapter(tuple[MarengoBilledRequest, ...])
def _billed_units(request: MarengoBilledRequest) -> tuple[int, int]:
input_type: Final = request.inputType
match input_type:
case "text":
return (1, 0)
case "image":
return (0, 1)
case "text_image":
return (1, 1)
case "multi_input":
multi_input: Final = request.multi_input or MarengoBilledMultiInput()
return (1 if multi_input.inputText else 0, len(multi_input.mediaSources))
case "video" | "audio" | None:
return (0, 0)
case _:
assert_never(input_type)
def _billed_usage(batch_data: list[dict] | None) -> Usage:
units: Final = tuple(_billed_units(request) for request in BILLED_REQUESTS.validate_python(batch_data or ()))
query_count: Final = sum(text_requests for text_requests, _ in units)
image_count: Final = sum(images for _, images in units)
details: Final = (
PromptTokensDetailsWrapper(query_count=query_count or None, image_count=image_count or None)
if query_count or image_count
else None
)
return Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0, prompt_tokens_details=details)
MARENGO_SHARED_PARAMS: Final = (
"encoding_format",
"embeddingOption",
"startSec",
"input_type",
"endSec",
"segmentation",
"embeddingType",
"embeddingScope",
"inferenceId",
"media_source",
"media_sources",
)
def drop_params_enabled(litellm_params: Mapping[str, object]) -> bool:
return litellm.drop_params is True or litellm_params.get("drop_params") is True
class TwelveLabsMarengoEmbeddingConfig:
@ -26,28 +127,24 @@ class TwelveLabsMarengoEmbeddingConfig:
Supports text, image, video, and audio inputs.
- InvokeModel: text and image inputs
- StartAsyncInvoke: video, audio, image, and text inputs
Marengo 3.0 (model ids containing "marengo-embed-3") nests the input under a key named after inputType and
adds the text_image and multi_input input types; that payload is built by build_marengo_3_request.
"""
def __init__(self) -> None:
pass
def __init__(self, model: str | None = None) -> None:
self.is_marengo_3: Final = is_marengo_3_model(model)
def get_supported_openai_params(self) -> list[str]:
return [
"encoding_format",
"textTruncate",
"embeddingOption",
"startSec",
"lengthSec",
"useFixedLengthSec",
"minClipSec",
"input_type",
]
if self.is_marengo_3:
return list(MARENGO_SHARED_PARAMS)
return [*MARENGO_SHARED_PARAMS, *MARENGO_2_7_ONLY_PARAMS]
def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict:
for k, v in non_default_params.items():
if k == "encoding_format":
# TwelveLabs doesn't have encoding_format, but we can map it to embeddingOption
if v == "float":
if v == "float" and not self.is_marengo_3:
optional_params["embeddingOption"] = ["visual-text", "visual-image"]
elif k == "textTruncate":
optional_params["textTruncate"] = v
@ -56,7 +153,19 @@ class TwelveLabsMarengoEmbeddingConfig:
elif k == "input_type":
# Map input_type to inputType for Bedrock
optional_params["inputType"] = v
elif k in ["startSec", "lengthSec", "useFixedLengthSec", "minClipSec"]:
elif k in (
"startSec",
"lengthSec",
"useFixedLengthSec",
"minClipSec",
"endSec",
"segmentation",
"embeddingType",
"embeddingScope",
"inferenceId",
"media_source",
"media_sources",
):
optional_params[k] = v
return optional_params
@ -77,7 +186,8 @@ class TwelveLabsMarengoEmbeddingConfig:
async_invoke_route: bool = False,
model_id: str | None = None,
output_s3_uri: str | None = None,
) -> TwelveLabsMarengoEmbeddingRequest | TwelveLabsAsyncInvokeRequest:
drop_params: bool = False,
) -> TwelveLabsMarengoEmbeddingRequest | TwelveLabsMarengo3EmbeddingRequest | TwelveLabsAsyncInvokeRequest:
"""
Transform OpenAI-style input to TwelveLabs Marengo format/async-invoke format.
@ -87,20 +197,29 @@ class TwelveLabsMarengoEmbeddingConfig:
- Video inputs (async-invoke only)
- Audio inputs (async-invoke only)
- S3 URLs for all media types (async-invoke only)
- Marengo 3.0 only: text_image and multi_input inputs (nested payload)
"""
# Get input_type or default to "text"
input_type: Final = cast(
TWELVELABS_EMBEDDING_INPUT_TYPES,
inference_params.get("inputType") or inference_params.get("input_type") or "text",
)
# Validate that async-invoke is used for video/audio
if input_type in ["video", "audio"] and not async_invoke_route:
raise ValueError(
f"Input type '{input_type}' requires async_invoke route. "
f"Use model format: 'bedrock/async_invoke/model_id'"
)
if self.is_marengo_3:
marengo_3_request: Final = build_marengo_3_request(
input=input, inference_params=inference_params, drop_params=drop_params
)
if async_invoke_route and model_id:
return self._wrap_async_invoke_request(
model_input=marengo_3_request, model_id=model_id, output_s3_uri=output_s3_uri
)
return marengo_3_request
transformed_request: Final[TwelveLabsMarengoEmbeddingRequest] = {"inputType": input_type}
if input_type == "text":
@ -154,7 +273,7 @@ class TwelveLabsMarengoEmbeddingConfig:
def _wrap_async_invoke_request(
self,
model_input: TwelveLabsMarengoEmbeddingRequest,
model_input: TwelveLabsMarengoEmbeddingRequest | TwelveLabsMarengo3EmbeddingRequest,
model_id: str,
output_s3_uri: str | None = None,
) -> TwelveLabsAsyncInvokeRequest:
@ -188,62 +307,16 @@ class TwelveLabsMarengoEmbeddingConfig:
),
)
def _transform_response(self, response_list: list[dict], model: str) -> EmbeddingResponse:
"""
Transform TwelveLabs response to OpenAI format.
Handles the actual TwelveLabs response format: {"data": [{"embedding": [...]}]}
"""
embeddings: Final[list[Embedding]] = []
total_tokens = 0
for response in response_list:
# TwelveLabs response format has a "data" field containing the embeddings
if "data" in response and isinstance(response["data"], list):
for item in response["data"]:
if "embedding" in item:
# Single embedding response
embedding = Embedding(
embedding=item["embedding"],
index=len(embeddings),
object="embedding",
)
embeddings.append(embedding)
# Estimate token count (rough approximation)
if "inputTextTokenCount" in item:
total_tokens += item["inputTextTokenCount"]
else:
# Rough estimate: 1 token per 4 characters for text, or use embedding size
total_tokens += len(item["embedding"]) // 4
elif "embedding" in response:
# Direct embedding response (fallback for other formats)
embedding = Embedding(
embedding=response["embedding"],
index=len(embeddings),
object="embedding",
)
embeddings.append(embedding)
# Estimate token count (rough approximation)
if "inputTextTokenCount" in response:
total_tokens += response["inputTextTokenCount"]
else:
# Rough estimate: 1 token per 4 characters for text
total_tokens += len(response.get("inputText", "")) // 4
elif "embeddings" in response:
# Multiple embeddings response (from video/audio)
for i, emb in enumerate(response["embeddings"]):
embedding = Embedding(
embedding=emb["embedding"],
index=len(embeddings),
object="embedding",
)
embeddings.append(embedding)
total_tokens += len(emb["embedding"]) // 4 # Rough estimate
usage: Final = Usage(prompt_tokens=total_tokens, total_tokens=total_tokens)
return EmbeddingResponse(data=embeddings, model=model, usage=usage)
def _transform_response(
self, response_list: list[dict], model: str, batch_data: list[dict] | None = None
) -> EmbeddingResponse:
vectors: Final = tuple(
vector for response in INVOKE_RESPONSES.validate_python(response_list) for vector in response.vectors()
)
embeddings: Final = [
Embedding(embedding=list(vector), index=index, object="embedding") for index, vector in enumerate(vectors)
]
return EmbeddingResponse(data=embeddings, model=model, usage=_billed_usage(batch_data))
def _transform_async_invoke_response(self, response: dict, model: str) -> EmbeddingResponse:
"""

View file

@ -7,13 +7,13 @@ from contextlib import suppress
from functools import cache
from itertools import chain
from types import MappingProxyType
from typing import Any, Final, TypeAlias, TypedDict
from typing import Any, Final, Literal, TypeAlias, TypedDict
from urllib.parse import unquote
import httpx
from httpx import Headers, Response
from openai.types.file_deleted import FileDeleted
from pydantic import BaseModel, ConfigDict, TypeAdapter
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
@ -60,11 +60,12 @@ from litellm.utils import get_llm_provider
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError, merge_bedrock_aws_request_params, resolve_s3_encryption_key_id
# litellm_params key used to hand the SigV4-signed GET headers from
# `transform_file_content_request` to `validate_environment` (the only hook
# the shared file-content HTTP handler exposes for setting request headers).
# Same pattern as the `upload_url` handoff in `transform_create_file_request`.
S3_SIGNED_GET_HEADERS_PARAM: Final = "_s3_signed_get_headers"
S3_SIGNED_REQUEST_HEADERS_PARAM: Final = "_s3_signed_request_headers"
class _S3DeleteContext(BaseModel):
file_id: str = Field(min_length=1)
# litellm_params key carrying the size of the body uploaded to S3, handed from
# `transform_create_file_request` to `transform_create_file_response`.
@ -291,7 +292,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
) -> dict:
result: Final[dict[str, object]] = {}
result.update(headers)
signed_headers: Final = litellm_params.pop(S3_SIGNED_GET_HEADERS_PARAM, None)
signed_headers: Final = litellm_params.pop(S3_SIGNED_REQUEST_HEADERS_PARAM, None)
if isinstance(signed_headers, Mapping):
result.update(signed_headers) # any-ok: untyped handoff headers
# otherwise no extra headers - AWS credentials are handled by BaseAWSLLM
@ -1187,18 +1188,27 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
def transform_delete_file_request(
self,
file_id: str,
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
raise NotImplementedError("BedrockFilesConfig does not support file deletion")
optional_params: Mapping[str, object],
litellm_params: MutableMapping[str, object],
) -> tuple[str, dict[str, str]]:
return self._transform_s3_file_request(
file_id=file_id, method="DELETE", optional_params=optional_params, litellm_params=litellm_params
)
def transform_delete_file_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
litellm_params: Mapping[str, object],
) -> FileDeleted:
raise NotImplementedError("BedrockFilesConfig does not support file deletion")
if raw_response.status_code != 204:
raise BedrockError(
status_code=raw_response.status_code if raw_response.status_code >= 400 else 502,
message=raw_response.text or f"S3 file deletion returned HTTP {raw_response.status_code}",
headers=raw_response.headers,
)
context: Final = _S3DeleteContext.model_validate(logging_obj.model_call_details.get("additional_args"))
return FileDeleted(id=context.file_id, deleted=True, object="file")
def transform_list_files_request(
self,
@ -1233,6 +1243,18 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
if not file_id:
raise ValueError("file_id is required for Bedrock file content retrieval")
return self._transform_s3_file_request(
file_id=file_id, method="GET", optional_params=optional_params, litellm_params=litellm_params
)
def _transform_s3_file_request(
self,
*,
file_id: str,
method: Literal["GET", "DELETE"],
optional_params: Mapping[str, object],
litellm_params: MutableMapping[str, object],
) -> tuple[str, dict[str, str]]:
s3_uri: Final = extract_s3_uri_from_file_id(file_id)
bucket_name, object_key = _validate_file_id_against_configured_buckets(
s3_uri=s3_uri,
@ -1240,40 +1262,32 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(litellm_params),
)
# The shared file-content handler passes optional_params={}, so AWS
# credentials/region arrive via litellm_params here (unlike the upload
# path). s3_region_name wins over aws_region_name, same priority as
# get_complete_file_url above.
merged_params: Final[dict[str, object]] = {}
merged_params.update(litellm_params)
merged_params.update(optional_params)
request_params: Final = _BedrockS3RequestParams.model_validate(merged_params)
request_params: Final = _BedrockS3RequestParams.model_validate({**litellm_params, **optional_params})
region_preference: Final = request_params.s3_region_name or request_params.aws_region_name
region_params: Final[dict[str, str | None]] = {"aws_region_name": region_preference}
aws_region_name: Final = self._get_aws_region_name(optional_params=region_params, model="")
s3_endpoint_url = (
s3_endpoint_url: Final = (
request_params.s3_endpoint_url or f"https://s3.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}"
).rstrip("/")
url: Final = f"{s3_endpoint_url}/{bucket_name}/{encode_s3_object_key_for_url(object_key)}"
litellm_params[S3_SIGNED_GET_HEADERS_PARAM] = self._sign_s3_get_request(
litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM] = self._sign_s3_request_without_body(
api_base=url,
aws_region_name=aws_region_name,
request_params=request_params,
method=method,
)
return url, {}
def _sign_s3_get_request(
def _sign_s3_request_without_body(
self,
api_base: str,
aws_region_name: str,
request_params: _BedrockS3RequestParams,
method: Literal["GET", "DELETE"] = "GET",
) -> dict[str, str]:
"""
SigV4-sign an S3 GetObject request, mirroring `_sign_s3_request` (PUT).
"""
try:
import hashlib
@ -1297,7 +1311,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
empty_body_hash: Final = hashlib.sha256(b"").hexdigest()
aws_request: Final = AWSRequest( # any-ok: botocore AWSRequest is untyped
method="GET",
method=method,
url=api_base,
headers={"x-amz-content-sha256": empty_body_hash},
)

View file

@ -579,7 +579,7 @@ class BaseLLMHTTPHandler:
data: dict[str, object], # mutable-ok: async_completion takes dict
signed_headers: dict[str, object], # mutable-ok: async_completion takes dict
signed_json_body: bytes | None,
):
) -> Coroutine[object, object, ModelResponse | CustomStreamWrapper]:
async_client: Final = client if isinstance(client, AsyncHTTPHandler) else None
if stream is True:
return self.acompletion_stream_function(
@ -626,7 +626,7 @@ class BaseLLMHTTPHandler:
if acompletion is True and provider_config.uses_async_transform_request:
async def transform_then_dispatch():
async def transform_then_dispatch() -> ModelResponse | CustomStreamWrapper:
transformed: Final = cast( # cast-ok: async_transform_request is declared as a bare dict
"dict[str, object]",
await provider_config.async_transform_request(
@ -9826,7 +9826,7 @@ class BaseLLMHTTPHandler:
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
litellm_params=MappingProxyType(dict(litellm_params, timeout=timeout)),
extra_body=extra_body,
embedding_executor=embedding_executor,
)
@ -9871,6 +9871,12 @@ class BaseLLMHTTPHandler:
data=request_data,
timeout=timeout,
)
except httpx.TimeoutException:
raise vector_store_provider_config.get_error_class(
error_message="Vector store search exceeded the caller timeout.",
status_code=408,
headers=httpx.Headers(),
) from None
except Exception as e:
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
@ -9955,7 +9961,7 @@ class BaseLLMHTTPHandler:
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
litellm_params=MappingProxyType(dict(litellm_params, timeout=timeout)),
extra_body=extra_body,
embedding_executor=embedding_executor,
)
@ -10000,7 +10006,14 @@ class BaseLLMHTTPHandler:
url=url,
headers=headers,
data=request_data,
timeout=timeout,
)
except httpx.TimeoutException:
raise vector_store_provider_config.get_error_class(
error_message="Vector store search exceeded the caller timeout.",
status_code=408,
headers=httpx.Headers(),
) from None
except Exception as e:
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
@ -10030,6 +10043,8 @@ class BaseLLMHTTPHandler:
else:
async_httpx_client = client
vector_store_provider_config.validate_create_vector_store()
headers: Final = vector_store_provider_config.validate_environment(
headers=extra_headers or {}, litellm_params=litellm_params
)
@ -10100,6 +10115,8 @@ class BaseLLMHTTPHandler:
else:
sync_httpx_client = client
vector_store_provider_config.validate_create_vector_store()
headers: Final = vector_store_provider_config.validate_environment(
headers=extra_headers or {}, litellm_params=litellm_params
)

View file

@ -1,10 +1,10 @@
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from urllib.parse import unquote
import httpx
from openai.types.responses import EasyInputMessageParam, ResponseInputItemParam
from openai.types.responses import EasyInputMessageParam, ResponseInputContentParam, ResponseInputItemParam
from litellm.llms.fireworks_ai.common_utils import (
resolve_fireworks_api_key,
@ -31,6 +31,17 @@ def _session_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object
)
_INSTRUCTION_ROLES: Final = frozenset({"system", "developer"})
def _role(item: ResponseInputItemParam) -> str | None:
match item:
case {"role": str(role)}:
return role
case _:
return None
def _developer_item_as_system(item: ResponseInputItemParam) -> ResponseInputItemParam:
if "role" not in item or item["role"] != "developer":
return item
@ -43,6 +54,69 @@ def _developer_items_as_system(input: str | ResponseInputParam) -> str | Respons
return [_developer_item_as_system(item) for item in input]
def _text_part(part: ResponseInputContentParam) -> str | None:
match part:
case {"type": "input_text", "text": str(text)}:
return text
case _:
return None
def _text_only_content(item: ResponseInputItemParam) -> str | None:
match item:
case {"role": "system" | "developer", "content": str(text)}:
return text
case {"role": "system" | "developer", "content": [*parts]}:
texts: Final = tuple(map(_text_part, parts))
return None if any(text is None for text in texts) else "\n\n".join(text for text in texts if text)
case _:
return None
def _leading_instruction_block_length(roles: Sequence[str | None]) -> int:
return next((index for index, role in enumerate(roles) if role not in _INSTRUCTION_ROLES), len(roles))
def _closing_instruction_block_start(roles: Sequence[str | None], leading_length: int) -> int:
last_conversation_index: Final = next(
(index for index in range(len(roles) - 1, leading_length - 1, -1) if roles[index] not in _INSTRUCTION_ROLES),
None,
)
if last_conversation_index is None or roles[last_conversation_index] != "assistant":
return len(roles)
return last_conversation_index + 1
def _hoisted_indices(roles: Sequence[str | None]) -> tuple[int, ...]:
leading_length: Final = _leading_instruction_block_length(roles)
closing_start: Final = _closing_instruction_block_start(roles, leading_length)
return tuple(
index for index, role in enumerate(roles[:closing_start]) if index < leading_length or role == "developer"
)
def _with_instruction_items_folded(
input: str | ResponseInputParam, instructions: str | None
) -> tuple[str | None, str | ResponseInputParam]:
if isinstance(input, str):
return instructions, input
items: Final = tuple(input)
folded: Final = MappingProxyType(
{
index: text
for index in _hoisted_indices(tuple(map(_role, items)))
if (text := _text_only_content(items[index])) is not None
}
)
joined: Final = "\n\n".join(chunk for chunk in (instructions, *folded.values()) if chunk)
return (
instructions if not folded else joined or None,
[ # mutable-ok: the base class takes the input items as a list
_developer_item_as_system(item) for index, item in enumerate(items) if index not in folded
],
)
class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
@property
def custom_llm_provider(self) -> LlmProviders:
@ -68,9 +142,6 @@ class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
base: Final = (api_base or get_secret_str("FIREWORKS_API_BASE") or FIREWORKS_AI_DEFAULT_API_BASE).rstrip("/")
return f"{base}/responses"
def _validate_input_param(self, input: str | ResponseInputParam) -> str | ResponseInputParam:
return _developer_items_as_system(super()._validate_input_param(input))
def transform_responses_api_request(
self,
model: str,
@ -79,10 +150,25 @@ class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
litellm_params: GenericLiteLLMParams,
headers: dict, # mutable-ok: overrides the base class signature
) -> dict: # mutable-ok: overrides the base class signature
instructions_param: Final[object] = response_api_optional_request_params.get("instructions")
validated_input: Final = self._validate_input_param(input)
instructions, folded_input = (
_with_instruction_items_folded(validated_input, instructions_param)
if isinstance(instructions_param, str | None)
else (instructions_param, _developer_items_as_system(validated_input))
)
instruction_entries: Final = () if instructions is None else (("instructions", instructions),)
folded_params: Final = { # mutable-ok: the base class takes the optional params as a dict
key: value
for key, value in (
*((key, value) for key, value in response_api_optional_request_params.items() if key != "instructions"),
*instruction_entries,
)
}
return super().transform_responses_api_request(
model=resolve_fireworks_resource_name(model),
input=input,
response_api_optional_request_params=response_api_optional_request_params,
input=folded_input,
response_api_optional_request_params=folded_params,
litellm_params=litellm_params,
headers=headers,
)

View file

@ -1,303 +0,0 @@
"""Shared helpers for the MongoDB integrations. pymongo lives in the optional ``mongodb`` extra,
so every import of it is deferred to call time."""
import asyncio
import threading
import weakref
from asyncio import AbstractEventLoop
from collections import OrderedDict
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias, TypeVar
from litellm.exceptions import BadRequestError, ServiceUnavailableError, Timeout
if TYPE_CHECKING:
from pymongo import AsyncMongoClient, MongoClient
PYMONGO_INSTALL_HINT: Final = (
"The MongoDB vector store requires the 'pymongo' package. "
"Run 'pip install litellm[mongodb]' (or 'pip install pymongo') to install it."
)
MONGODB_PROVIDER: Final = "mongodb"
def config_error(message: str) -> BadRequestError:
"""400 rather than the 500 a bare ValueError becomes once litellm.exception_type wraps it."""
return BadRequestError(message=message, model=None, llm_provider=MONGODB_PROVIDER)
def timeout_error(message: str) -> Timeout:
return Timeout(message=message, model=None, llm_provider=MONGODB_PROVIDER)
def unavailable_error(message: str) -> ServiceUnavailableError:
"""litellm only retries 408, 409, 429 and 5xx, so a 400 here would make a failover permanent."""
return ServiceUnavailableError(message=message, model=None, llm_provider=MONGODB_PROVIDER)
DEFAULT_CONNECT_TIMEOUT_MS: Final = 10_000
DEFAULT_SOCKET_TIMEOUT_MS: Final = 30_000
DEFAULT_SERVER_SELECTION_TIMEOUT_MS: Final = 10_000
_MAX_CACHED_CLIENTS: Final = 32
_APP_NAME: Final = "litellm"
@dataclass(frozen=True, slots=True)
class MongoClientKey:
connection_string: str
connect_timeout_ms: int
socket_timeout_ms: int
server_selection_timeout_ms: int
SyncClientFactory: TypeAlias = Callable[..., "MongoClient"]
AsyncClientFactory: TypeAlias = Callable[..., "AsyncMongoClient"]
_K = TypeVar("_K")
_V = TypeVar("_V")
_AsyncClientCacheKey: TypeAlias = tuple[MongoClientKey, int]
# CPython recycles id() aggressively, so the id alone would hand a new loop a closed loop's client
_AsyncClientEntry: TypeAlias = tuple["weakref.ref[AbstractEventLoop]", "AsyncMongoClient"]
_SyncClientCache: TypeAlias = "OrderedDict[MongoClientKey, MongoClient]"
_AsyncClientCache: TypeAlias = "OrderedDict[_AsyncClientCacheKey, _AsyncClientEntry]"
_sync_clients: Final[_SyncClientCache] = OrderedDict() # mutable-ok: process-level client cache
_async_clients: Final[_AsyncClientCache] = OrderedDict() # mutable-ok: same cache, per loop
# async searches reach the sync client through executor threads, so both caches are shared state
_cache_lock: Final = threading.Lock()
def _store_bounded(cache: "OrderedDict[_K, _V]", cache_key: "_K", value: "_V") -> None:
"""Eviction only drops this cache's reference; an in-flight search keeps its client alive."""
with _cache_lock:
cache[cache_key] = value # mutable-ok: an LRU cache is mutable state by definition
cache.move_to_end(cache_key)
while len(cache) > _MAX_CACHED_CLIENTS:
cache.popitem(last=False)
def _mark_used(cache: "OrderedDict[_K, _V]", cache_key: "_K") -> None:
with _cache_lock:
if cache_key in cache:
cache.move_to_end(cache_key)
def import_sync_mongo_client() -> "type[MongoClient]":
try:
from pymongo import MongoClient as SyncMongoClient
except ImportError as e:
raise config_error(PYMONGO_INSTALL_HINT) from e
return SyncMongoClient
def import_async_mongo_client() -> "type[AsyncMongoClient]":
try:
from pymongo import AsyncMongoClient as AsyncMongoClientClass
except ImportError as e:
raise config_error(PYMONGO_INSTALL_HINT) from e
return AsyncMongoClientClass
def _client_kwargs(key: MongoClientKey) -> Mapping[str, object]:
return MappingProxyType(
{
"connectTimeoutMS": key.connect_timeout_ms,
"socketTimeoutMS": key.socket_timeout_ms,
"serverSelectionTimeoutMS": key.server_selection_timeout_ms,
"appname": _APP_NAME,
}
)
def get_sync_client(key: MongoClientKey, client_class: SyncClientFactory | None = None) -> "MongoClient":
cached: Final = _sync_clients.get(key)
if cached is not None:
_mark_used(_sync_clients, key)
return cached
build: Final = client_class if client_class is not None else import_sync_mongo_client()
client: Final = build(key.connection_string, **_client_kwargs(key))
_store_bounded(_sync_clients, key, client)
return client
def _purge_dead_loops() -> None:
"""A cached client holds its loop alive, so a closed loop's entry would pin that client and its
sockets for the life of the process."""
with _cache_lock:
for stale in tuple(
cache_key
for cache_key, (loop_ref, _) in _async_clients.items()
if (cached_loop := loop_ref()) is None or cached_loop.is_closed()
):
del _async_clients[stale]
def get_async_client(key: MongoClientKey, client_class: AsyncClientFactory | None = None) -> "AsyncMongoClient":
"""Async clients bind to the loop that created them, so the cache is keyed per loop."""
loop: Final = asyncio.get_running_loop()
loop_key: Final = (key, id(loop))
cached: Final = _async_clients.get(loop_key)
if cached is not None and cached[0]() is loop:
_mark_used(_async_clients, loop_key)
return cached[1]
_purge_dead_loops()
build: Final = client_class if client_class is not None else import_async_mongo_client()
client: Final = build(key.connection_string, **_client_kwargs(key))
_store_bounded(_async_clients, loop_key, (weakref.ref(loop), client))
return client
def reset_client_cache() -> None:
with _cache_lock:
_sync_clients.clear()
_async_clients.clear()
_AUTHENTICATION_FAILED_CODE: Final = 18
_UNAUTHORIZED_CODE: Final = 13
# Atlas reports a rejected user as code 8000 "AtlasError" where a self-managed mongod reports 18
_AUTHENTICATION_MESSAGE_MARKERS: Final = ("bad auth", "authentication failed", "not authorized")
_RESOLUTION_TIMEOUT_MARKERS: Final = ("resolution lifetime expired", "dns operation timed out")
_UNKNOWN_HOSTNAME_MARKERS: Final = ("dns query name does not exist", "name or service not known")
_CREDENTIAL_ESCAPING_MARKERS: Final = ("must be escaped according to rfc 3986", "bad database name")
def _index_hint(index_name: str, database: str, collection: str) -> str:
return (
f"No queryable MongoDB Vector Search index named '{index_name}' was found on "
f"'{database}.{collection}'. Confirm the index exists on that exact collection, that its "
"status is READY rather than still building, and that the vector store id matches the index name."
)
def missing_index_error(index_name: str, database: str, collection: str) -> BadRequestError:
"""$vectorSearch against a missing index, database or collection returns zero documents rather
than failing, so an empty result set is checked against the catalogue and reported as this."""
return config_error(
f"{_index_hint(index_name, database, collection)} A vector search against a database, "
"collection or index that does not exist returns no results rather than an error, so this "
"was reported as an empty result set by MongoDB."
)
def index_not_ready_error(index_name: str, database: str, collection: str, status: str) -> BadRequestError:
return config_error(
f"The MongoDB Vector Search index '{index_name}' on '{database}.{collection}' is not queryable "
f"yet; its status is {status}. Searches against it return no results until the build finishes."
)
def translate_mongo_error(error: Exception, index_name: str, database: str, collection: str) -> Exception:
"""Returns the exception to raise, so callers keep the driver error as ``__cause__``."""
try:
from pymongo.errors import (
ConfigurationError,
ConnectionFailure,
ExecutionTimeout,
InvalidOperation,
NetworkTimeout,
OperationFailure,
ServerSelectionTimeoutError,
)
except ImportError:
return error
if isinstance(error, ServerSelectionTimeoutError):
return timeout_error(
"Could not reach the MongoDB deployment before the timeout. On Atlas this is usually the "
"project's IP access list not containing this host, or a paused cluster. On a self-managed "
"deployment it is usually the host or port in the URI, or a firewall between this process "
f"and mongod. Either way it can also be an unresolvable hostname. Driver detail: {error}"
)
# ExecutionTimeout subclasses OperationFailure, so it has to be matched before it
if isinstance(error, (NetworkTimeout, ExecutionTimeout)):
return timeout_error(
f"The MongoDB vector search against '{database}.{collection}' timed out before returning. "
f"Driver detail: {error}"
)
# ServerSelectionTimeoutError and NetworkTimeout also subclass ConnectionFailure, so this only
# sees what those branches left
if isinstance(error, ConnectionFailure):
return unavailable_error(
f"The connection to '{database}.{collection}' was dropped or refused. That is usually a "
"replica set failover or a restarted node, so the search is worth retrying. If it keeps "
"happening: on Atlas the usual cause is a connection string with no username and password, "
"or a TLS failure, so confirm the URI is the one Atlas shows under Connect, Drivers; on a "
"self-managed deployment, check that mongod is listening on the host and port in the URI. "
f"Driver detail: {error}"
)
if isinstance(error, OperationFailure):
code: Final = error.code
detail: Final = str(error).lower()
if code in (_AUTHENTICATION_FAILED_CODE, _UNAUTHORIZED_CODE) or any(
marker in detail for marker in _AUTHENTICATION_MESSAGE_MARKERS
):
return config_error(
"MongoDB rejected the credentials in mongodb_connection_string, or the database user "
f"lacks read access to '{database}.{collection}'. Driver detail: {error.details}"
)
if "dimension" in detail:
return config_error(
"The query embedding does not match the vector dimensions the index was built for. "
"litellm_embedding_model must be the same model that produced the stored vectors. "
f"Driver detail: {error}"
)
if "is not indexed as vector" in detail:
return config_error(
"mongodb_embedding_field names a field the MongoDB Vector Search index does not cover. "
f"It must match the 'path' the index '{index_name}' was created on. Driver detail: {error}"
)
if "index" in detail and ("not found" in detail or "does not exist" in detail or "unknown" in detail):
return config_error(f"{_index_hint(index_name, database, collection)} Driver detail: {error}")
return config_error(
f"MongoDB rejected the vector search against '{database}.{collection}' using index "
f"'{index_name}'. Driver detail: {error}"
)
if isinstance(error, ConfigurationError):
configuration_detail: Final = str(error).lower()
if any(marker in configuration_detail for marker in _RESOLUTION_TIMEOUT_MARKERS):
return timeout_error(
"The DNS lookup for the cluster in mongodb_connection_string did not finish in time. "
"A mongodb+srv:// URI needs an SRV lookup before any connection is attempted, so this "
f"is DNS or the configured timeout, not MongoDB. Driver detail: {error}"
)
if any(marker in configuration_detail for marker in _UNKNOWN_HOSTNAME_MARKERS):
return config_error(
"The hostname in mongodb_connection_string does not exist in DNS. On Atlas, check the "
"cluster name against the URI shown under Connect, Drivers. On a self-managed deployment, "
f"check that the hostname resolves from this process. Driver detail: {error}"
)
if any(marker in configuration_detail for marker in _CREDENTIAL_ESCAPING_MARKERS):
return config_error(
"mongodb_connection_string could not be parsed. A username or password containing "
"'@', '/', ':' or '%' has to be percent-encoded per RFC 3986, so 'p@ss/word' becomes "
"'p%40ss%2Fword'. If the credentials are already encoded, check the database name in "
f"the URI path instead. Driver detail: {error}"
)
return config_error(
f"mongodb_connection_string is not a usable MongoDB connection string. Driver detail: {error}"
)
if isinstance(error, InvalidOperation):
return config_error(f"The MongoDB client was already closed or is unusable. Driver detail: {error}")
# An unreadable tlsCAFile or tlsCertificateKeyFile raises OSError, not a PyMongoError
if isinstance(error, OSError) and error.filename:
return config_error(
f"'{error.filename}', named by a TLS option in mongodb_connection_string, could not be read. "
"Check that tlsCAFile and tlsCertificateKeyFile point at files this process can open; inside "
f"a container that is the path in the container, not on the host. Driver detail: {error}"
)
# pymongo raises a plain ValueError, not a PyMongoError, for an unusable port
if isinstance(error, ValueError):
return config_error(
"The host and port in mongodb_connection_string could not be parsed. If the port is a "
"number between 0 and 65535, the cause is usually an unescaped ':' in the password, which "
f"has to be percent-encoded per RFC 3986 as '%3A'. Driver detail: {error}"
)
return error

View file

@ -1,37 +1,29 @@
"""MongoDB Vector Search has no HTTP query API, so this is a direct provider that runs the
``$vectorSearch`` aggregation through pymongo. ``vector_store_id`` is the search index name."""
from collections.abc import Callable, Mapping, Sequence
from collections.abc import Mapping, Sequence
from ipaddress import ip_address
from math import isfinite
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, NoReturn
from typing import TYPE_CHECKING, Final, Literal, NoReturn
from urllib.parse import quote, urlsplit
import httpx
from pydantic import BaseModel, ConfigDict
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from litellm.exceptions import AuthenticationError, BadRequestError, ServiceUnavailableError, Timeout
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.vector_store.transformation import (
BaseDirectVectorStoreConfig,
BaseQueryEmbeddingVectorStoreConfig,
LiteLLMVectorStoreEmbeddingExecutor,
VectorStoreEmbeddingExecutor,
)
from litellm.llms.mongodb.common_utils import (
DEFAULT_CONNECT_TIMEOUT_MS,
DEFAULT_SERVER_SELECTION_TIMEOUT_MS,
DEFAULT_SOCKET_TIMEOUT_MS,
MongoClientKey,
config_error,
get_async_client,
get_sync_client,
index_not_ready_error,
missing_index_error,
translate_mongo_error,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import EmbeddingResponse
from litellm.types.vector_stores import (
BaseVectorStoreAuthCredentials,
VectorStoreCreateOptionalRequestParams,
VectorStoreResultContent,
VectorStoreIndexEndpoints,
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
VectorStoreSearchResult,
)
if TYPE_CHECKING:
@ -39,26 +31,45 @@ if TYPE_CHECKING:
DEFAULT_EMBEDDING_FIELD_NAME: Final = "embedding"
DEFAULT_TEXT_FIELD_NAME: Final = "text"
SCORE_FIELD_NAME: Final = "score"
DEFAULT_MAX_NUM_RESULTS: Final = 10
MIN_MAX_NUM_RESULTS: Final = 1
MAX_MAX_NUM_RESULTS: Final = 50
NUM_CANDIDATES_MULTIPLIER: Final = 10
MIN_NUM_CANDIDATES: Final = 100
MAX_NUM_CANDIDATES: Final = 10_000
MAX_QUERY_CHARACTERS: Final = 32_000
_EMPTY_EMBEDDING_CONFIG: Final = MappingProxyType({})
_SEARCH_ONLY_MESSAGE: Final = (
"MongoDB vector store is search-only. Create the collection and its MongoDB Vector Search "
"index in MongoDB directly, then register it here by index name."
)
def config_error(message: str) -> BadRequestError:
return BadRequestError(message=message, model=None, llm_provider="mongodb")
class _Content(BaseModel):
model_config = ConfigDict(frozen=True, strict=True)
type: Literal["text"]
text: str
class _Result(BaseModel):
model_config = ConfigDict(frozen=True, strict=True, allow_inf_nan=False)
score: float | None
content: Sequence[_Content]
file_id: str | None
filename: str | None
class _SearchResponse(BaseModel):
model_config = ConfigDict(frozen=True, strict=True)
object: Literal["vector_store.search_results.page"]
search_query: str
data: Sequence[_Result]
class _MongoDBSearchParams(BaseModel):
"""Typed view over the vector store's litellm_params; unrelated keys are ignored."""
@ -66,7 +77,6 @@ class _MongoDBSearchParams(BaseModel):
litellm_embedding_model: str | None = None
litellm_embedding_config: Mapping[str, object] | None = None
mongodb_connection_string: str | None = None
mongodb_database: str | None = None
mongodb_collection: str | None = None
mongodb_text_field: str | None = None
@ -91,21 +101,6 @@ class _MongoDBSearchParams(BaseModel):
)
return self.litellm_embedding_model
def require_connection_string(self) -> str:
if not self.mongodb_connection_string:
raise config_error(
"mongodb_connection_string is required in litellm_params for the MongoDB vector store. "
"Example: mongodb+srv://<user>:<password>@<cluster>.mongodb.net for Atlas, or "
"mongodb://<user>:<password>@<host>:27017 for a self-managed deployment"
)
scheme: Final = self.mongodb_connection_string.split("://", 1)[0].lower()
if scheme not in ("mongodb", "mongodb+srv"):
raise config_error(
"mongodb_connection_string must start with 'mongodb://' or 'mongodb+srv://', "
f"got '{self.mongodb_connection_string.split('://', 1)[0]}://'"
)
return self.mongodb_connection_string
def require_database(self) -> str:
if not self.mongodb_database:
raise config_error(
@ -127,30 +122,28 @@ _MONGODB_PARAM_PREFIX: Final = "mongodb_"
_KNOWN_MONGODB_PARAMS: Final = frozenset(
name for name in _MongoDBSearchParams.model_fields if name.startswith(_MONGODB_PARAM_PREFIX)
)
_RESPONSE_ADAPTER: Final = TypeAdapter(VectorStoreSearchResponse)
class MongoDBVectorStoreConfig(BaseDirectVectorStoreConfig):
def __init__(
self,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
sync_client_factory: Callable[[MongoClientKey], object] | None = None,
async_client_factory: Callable[[MongoClientKey], object] | None = None,
) -> None:
super().__init__()
self.embedding_executor: Final[VectorStoreEmbeddingExecutor] = (
embedding_executor if embedding_executor is not None else LiteLLMVectorStoreEmbeddingExecutor()
)
self.sync_client_factory: Final[Callable[[MongoClientKey], object]] = (
sync_client_factory if sync_client_factory is not None else get_sync_client
)
self.async_client_factory: Final[Callable[[MongoClientKey], object]] = (
async_client_factory if async_client_factory is not None else get_async_client
)
class MongoDBVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
def __init__(self, embedding_executor: VectorStoreEmbeddingExecutor | None = None) -> None:
self.embedding_executor: Final = embedding_executor or LiteLLMVectorStoreEmbeddingExecutor()
def get_auth_credentials(self, litellm_params: Mapping[str, object]) -> BaseVectorStoreAuthCredentials:
return BaseVectorStoreAuthCredentials()
def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints:
return VectorStoreIndexEndpoints(read=[], write=[]) # mutable-ok: the TypedDict declares list fields
@staticmethod
def _reject_unknown_params(litellm_params: Mapping[str, object]) -> None:
"""Without this a mistyped mongodb_collection reads as 'mongodb_collection is required',
naming a key the reader can see they have set."""
if litellm_params.get("mongodb_connection_string") is not None:
raise config_error(
"MongoDB vector stores now use the BETA sidecar. Move mongodb_connection_string to "
"MONGODB_CONNECTION_STRING in the sidecar, remove it from LiteLLM, and configure api_base and api_key."
)
unknown: Final = sorted(
key for key in litellm_params if key.startswith(_MONGODB_PARAM_PREFIX) and key not in _KNOWN_MONGODB_PARAMS
)
@ -191,239 +184,203 @@ class MongoDBVectorStoreConfig(BaseDirectVectorStoreConfig):
return configured
return min(max(limit * NUM_CANDIDATES_MULTIPLIER, MIN_NUM_CANDIDATES), MAX_NUM_CANDIDATES)
@staticmethod
def _timeout_ms(timeout: float | httpx.Timeout | None) -> tuple[int, int]:
"""The connect and socket budgets pymongo is built with, in that order."""
if isinstance(timeout, httpx.Timeout):
return (
int((timeout.connect or DEFAULT_CONNECT_TIMEOUT_MS / 1000) * 1000),
int((timeout.read or DEFAULT_SOCKET_TIMEOUT_MS / 1000) * 1000),
def validate_environment(
self, headers: Mapping[str, object], litellm_params: GenericLiteLLMParams | None
) -> dict[str, object]: # mutable-ok: the shared HTTP handler requires writable headers
if litellm_params is None:
raise config_error("Configure api_base and api_key for the MongoDB BETA sidecar.")
self._reject_unknown_params(MappingProxyType(dict(litellm_params)))
api_key: Final = litellm_params.api_key or get_secret_str("MONGODB_SIDECAR_API_KEY")
if not api_key:
raise config_error("MongoDB sidecar api_key is required. Set api_key or MONGODB_SIDECAR_API_KEY.")
return {
**headers,
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
} # mutable-ok: writable HTTP headers
def get_complete_url(self, api_base: str | None, litellm_params: Mapping[str, object]) -> str:
if not api_base:
raise config_error("MongoDB sidecar api_base is required, for example http://127.0.0.1:8080.")
try:
parsed: Final = urlsplit(api_base)
valid: Final = parsed.scheme in ("http", "https") and bool(parsed.hostname) and parsed.port != 0
except ValueError:
raise config_error("MongoDB sidecar api_base must be a valid HTTP or HTTPS URL.") from None
if not valid or parsed.username or parsed.password or parsed.query or parsed.fragment:
raise config_error(
"MongoDB sidecar api_base must be an HTTP or HTTPS URL without credentials, query, or fragment."
)
if timeout is None:
return DEFAULT_CONNECT_TIMEOUT_MS, DEFAULT_SOCKET_TIMEOUT_MS
return min(int(float(timeout) * 1000), DEFAULT_CONNECT_TIMEOUT_MS), int(float(timeout) * 1000)
if parsed.scheme == "http":
try:
loopback: Final = ip_address(parsed.hostname or "").is_loopback
except ValueError:
raise config_error(
"MongoDB sidecar requires HTTPS. HTTP is supported only for a loopback IP such as 127.0.0.1."
) from None
if not loopback:
raise config_error(
"MongoDB sidecar requires HTTPS. HTTP is supported only for a loopback IP such as 127.0.0.1."
)
return api_base.rstrip("/")
@staticmethod
def _timeout_ms(value: object) -> int:
seconds: Final = value.read if isinstance(value, httpx.Timeout) else value
if seconds is None:
return 30_000
if not isinstance(seconds, (int, float)) or not isfinite(seconds) or seconds <= 0:
raise config_error("MongoDB search timeout must be a positive finite number.")
try:
return max(1, int(seconds * 1000))
except (ValueError, OverflowError):
raise config_error("MongoDB search timeout must be a positive finite number.") from None
@classmethod
def _client_key(cls, params: _MongoDBSearchParams, timeout: float | httpx.Timeout | None) -> MongoClientKey:
connect_ms, socket_ms = cls._timeout_ms(timeout)
return MongoClientKey(
connection_string=params.require_connection_string(),
connect_timeout_ms=connect_ms,
socket_timeout_ms=socket_ms,
server_selection_timeout_ms=min(connect_ms, DEFAULT_SERVER_SELECTION_TIMEOUT_MS),
)
def _params(
cls,
litellm_params: Mapping[str, object],
optional_params: VectorStoreSearchOptionalRequestParams,
extra_body: Mapping[str, object] | None,
) -> _MongoDBSearchParams:
cls._reject_unknown_params(litellm_params)
if extra_body:
raise config_error("MongoDB vector store does not support extra_body overrides.")
for unsupported in ("filters", "ranking_options", "rewrite_query"):
if optional_params.get(unsupported) is not None:
raise config_error(f"MongoDB vector store does not support the {unsupported} parameter.")
try:
params: Final = _MongoDBSearchParams.model_validate(litellm_params)
except ValidationError:
raise config_error(
"Invalid MongoDB vector-store configuration. Check the database, collection, fields, and candidate count."
) from None
params.require_database()
params.require_collection()
params.require_embedding_model()
cls._num_candidates(cls._limit(optional_params), params.mongodb_num_candidates)
cls._timeout_ms(litellm_params.get("timeout"))
return params
@classmethod
def _pipeline(
def _request(
cls,
vector_store_id: str,
query_vector: Sequence[float],
query_text: str,
params: _MongoDBSearchParams,
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
) -> Sequence[Mapping[str, object]]:
if vector_store_search_optional_params.get("filters") is not None:
optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
embedding_response: EmbeddingResponse,
timeout: object,
) -> tuple[str, dict[str, object]]: # mutable-ok: the provider contract returns a writable JSON request body
if not embedding_response.data:
raise config_error(
"MongoDB vector store does not support the filters parameter yet. "
"Restrict the collection or the MongoDB Vector Search index definition instead."
"The embedding model returned no embedding for the search query. Check litellm_embedding_model."
)
if vector_store_search_optional_params.get("ranking_options") is not None:
raise config_error(
"MongoDB vector store does not support the ranking_options parameter yet. "
"Every result already carries the vectorSearchScore, so filter or re-rank "
"on that rather than having the threshold silently ignored."
)
if vector_store_search_optional_params.get("rewrite_query") is not None:
raise config_error(
"MongoDB vector store does not support the rewrite_query parameter. The query is "
"embedded exactly as sent; rewrite it before calling if you need that."
)
limit: Final = cls._limit(vector_store_search_optional_params)
search: Final = MappingProxyType(
{
"index": vector_store_id,
"path": params.embedding_field,
"queryVector": tuple(query_vector),
"numCandidates": cls._num_candidates(limit, params.mongodb_num_candidates),
"limit": limit,
}
)
projection: Final = MappingProxyType(
{params.text_field: 1, SCORE_FIELD_NAME: MappingProxyType({"$meta": "vectorSearchScore"})}
)
return [ # mutable-ok: pymongo rejects any non-list pipeline in common.validate_list
MappingProxyType({"$vectorSearch": search}),
MappingProxyType({"$project": projection}),
]
@classmethod
def _field_value(cls, document: Mapping[str, object], dotted_path: str) -> str | None:
"""None means absent, which is what separates a mistyped field from genuinely empty text."""
head, _, rest = dotted_path.partition(".")
if head not in document:
return None
value: Final = document[head]
if not rest:
return None if value is None else str(value)
return cls._field_value(value, rest) if isinstance(value, Mapping) else None
@classmethod
def _to_result(cls, document: Mapping[str, object], text_field: str) -> VectorStoreSearchResult:
document_id: Final = document.get("_id")
identifier: Final = None if document_id is None else str(document_id)
content: Final = [ # mutable-ok: VectorStoreSearchResult declares a list of content parts
VectorStoreResultContent(text=cls._field_value(document, text_field) or "", type="text")
]
raw_score: Final = document.get(SCORE_FIELD_NAME)
return VectorStoreSearchResult(
score=float(raw_score) if isinstance(raw_score, (int, float)) else None,
content=content,
file_id=identifier,
filename=identifier,
vector: Final = embedding_response.data[0]["embedding"]
if not vector or any(not isinstance(value, (float, int)) or not isfinite(value) for value in vector):
raise config_error("The embedding model must return a non-empty, finite query vector.")
limit: Final = cls._limit(optional_params)
return (
f"{api_base}/v1/vector_stores/{quote(vector_store_id, safe='')}/search",
{ # mutable-ok: JSON transport requires a dict
"query": query_text,
"query_vector": tuple(vector),
"mongodb_database": params.require_database(),
"mongodb_collection": params.require_collection(),
"mongodb_embedding_field": params.embedding_field,
"mongodb_text_field": params.text_field,
"mongodb_num_candidates": cls._num_candidates(limit, params.mongodb_num_candidates),
"max_num_results": limit,
"timeout_ms": cls._timeout_ms(timeout),
},
)
@classmethod
def _raise_for_missing_text_field(
cls, documents: Sequence[Mapping[str, object]], text_field: str, database: str, collection: str
) -> None:
"""$vectorSearch matches documents carrying no text, so a mistyped mongodb_text_field
returns well-scored results with empty content instead of failing."""
if documents and all(cls._field_value(document, text_field) is None for document in documents):
raise config_error(
f"None of the {len(documents)} matched documents in '{database}.{collection}' has a "
f"'{text_field}' field, so every result would carry empty text. Set mongodb_text_field "
"to the field holding the readable text; it accepts a dotted path such as metadata.body."
)
@classmethod
def _to_response(
cls, documents: Sequence[Mapping[str, object]], query_text: str, text_field: str
) -> VectorStoreSearchResponse:
return VectorStoreSearchResponse(
object="vector_store.search_results.page",
search_query=query_text,
data=[ # mutable-ok: VectorStoreSearchResponse declares data as a list
cls._to_result(document, text_field) for document in documents
],
)
@staticmethod
def _raise_for_unusable_index(
catalogue: Sequence[Mapping[str, object]], index_name: str, database: str, collection: str
) -> None:
"""mongod returns zero documents both for a query that matched nothing and for a missing
database, collection or index, so the catalogue decides which one happened."""
if not catalogue:
raise missing_index_error(index_name, database, collection)
entry: Final = catalogue[0]
if not entry.get("queryable"):
raise index_not_ready_error(index_name, database, collection, str(entry.get("status") or "unknown"))
@staticmethod
def _embedding_vector(embedding_response: EmbeddingResponse) -> Sequence[float]:
data: Final = embedding_response.data
if not data:
raise config_error(
"The embedding model returned no embedding for the search query, so there is nothing "
"to search MongoDB with. Check the embedding deployment named by litellm_embedding_model."
)
return data[0]["embedding"]
def execute_search_vector_store_request(
def transform_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
litellm_logging_obj: "LiteLLMLoggingObj",
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
timeout: float | httpx.Timeout | None = None,
) -> VectorStoreSearchResponse:
self._reject_unknown_params(litellm_params)
params: Final = _MongoDBSearchParams.model_validate(litellm_params)
) -> tuple[str, dict[str, object]]: # mutable-ok: the provider contract returns a writable JSON request body
params: Final = self._params(litellm_params, vector_store_search_optional_params, extra_body)
query_text: Final = self._query_text(query)
key: Final = self._client_key(params, timeout)
database: Final = params.require_database()
collection: Final = params.require_collection()
embedding_response: Final = (embedding_executor or self.embedding_executor).embed(
params.require_embedding_model(),
response: Final = (embedding_executor or self.embedding_executor).embed(
params.require_embedding_model(), query_text, params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG
)
return self._request(
vector_store_id,
query_text,
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
)
pipeline: Final = self._pipeline(
vector_store_id, self._embedding_vector(embedding_response), params, vector_store_search_optional_params
params,
vector_store_search_optional_params,
api_base,
response,
litellm_params.get("timeout"),
)
try:
client: Final = self.sync_client_factory(key)
target: Final = client[database][collection] # pyright: ignore[reportIndexIssue] # factory is typed as returning object so injected doubles are accepted
documents: Final = tuple(target.aggregate(pipeline))
except Exception as e:
raise translate_mongo_error(e, index_name=vector_store_id, database=database, collection=collection) from e
if not documents:
try:
catalogue: Final = tuple(target.list_search_indexes(vector_store_id))
except Exception as e:
raise translate_mongo_error(
e, index_name=vector_store_id, database=database, collection=collection
) from e
self._raise_for_unusable_index(catalogue, vector_store_id, database, collection)
self._raise_for_missing_text_field(documents, params.text_field, database, collection)
return self._to_response(documents, query_text, params.text_field)
async def aexecute_search_vector_store_request(
async def atransform_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
litellm_logging_obj: "LiteLLMLoggingObj",
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
timeout: float | httpx.Timeout | None = None,
) -> VectorStoreSearchResponse:
self._reject_unknown_params(litellm_params)
params: Final = _MongoDBSearchParams.model_validate(litellm_params)
) -> tuple[str, dict[str, object]]: # mutable-ok: the provider contract returns a writable JSON request body
params: Final = self._params(litellm_params, vector_store_search_optional_params, extra_body)
query_text: Final = self._query_text(query)
key: Final = self._client_key(params, timeout)
database: Final = params.require_database()
collection: Final = params.require_collection()
embedding_response: Final = await (embedding_executor or self.embedding_executor).aembed(
params.require_embedding_model(),
response: Final = await (embedding_executor or self.embedding_executor).aembed(
params.require_embedding_model(), query_text, params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG
)
return self._request(
vector_store_id,
query_text,
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
)
pipeline: Final = self._pipeline(
vector_store_id, self._embedding_vector(embedding_response), params, vector_store_search_optional_params
params,
vector_store_search_optional_params,
api_base,
response,
litellm_params.get("timeout"),
)
def transform_search_vector_store_response(
self, response: httpx.Response, litellm_logging_obj: "LiteLLMLoggingObj"
) -> VectorStoreSearchResponse:
try:
client: Final = self.async_client_factory(key)
target: Final = client[database][collection] # pyright: ignore[reportIndexIssue] # factory is typed as returning object so injected doubles are accepted
cursor: Final = await target.aggregate(pipeline)
documents: Final = [ # mutable-ok: an async comprehension cannot build a tuple directly
document async for document in cursor
]
except Exception as e:
raise translate_mongo_error(e, index_name=vector_store_id, database=database, collection=collection) from e
if not documents:
try:
index_cursor: Final = await target.list_search_indexes(vector_store_id)
catalogue: Final = [ # mutable-ok: an async comprehension cannot build a tuple directly
entry async for entry in index_cursor
]
except Exception as e:
raise translate_mongo_error(
e, index_name=vector_store_id, database=database, collection=collection
) from e
self._raise_for_unusable_index(catalogue, vector_store_id, database, collection)
self._raise_for_missing_text_field(documents, params.text_field, database, collection)
return self._to_response(documents, query_text, params.text_field)
validated: Final = _SearchResponse.model_validate_json(response.content)
return _RESPONSE_ADAPTER.validate_python(validated.model_dump())
except ValidationError:
raise ServiceUnavailableError(
message="MongoDB sidecar returned an invalid search response. Check the sidecar version and deployment.",
model=None,
llm_provider="mongodb",
) from None
def get_error_class(
self, error_message: str, status_code: int, headers: Mapping[str, object] | httpx.Headers
) -> BaseLLMException:
if status_code == 400:
raise config_error(error_message)
if status_code == 401:
raise AuthenticationError(message="MongoDB sidecar rejected api_key.", model=None, llm_provider="mongodb")
if status_code == 408:
raise Timeout(message=error_message, model=None, llm_provider="mongodb")
raise ServiceUnavailableError(
message="MongoDB sidecar is unavailable. Check its address, health, and logs.",
model=None,
llm_provider="mongodb",
)
def validate_create_vector_store(self) -> NoReturn:
raise config_error(_SEARCH_ONLY_MESSAGE)
def transform_create_vector_store_request(
self,
vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams,
api_base: str,
self, vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, api_base: str
) -> NoReturn:
raise config_error(_SEARCH_ONLY_MESSAGE)

View file

@ -16,6 +16,10 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
custom_prompt,
ollama_pt,
)
from litellm.litellm_core_utils.prompt_templates.image_handling import (
async_inline_remote_media,
inline_remote_image_urls,
)
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUsageBlock
@ -344,6 +348,26 @@ class OllamaConfig(BaseConfig):
)
return model_response
@property
def uses_async_transform_request(self) -> bool:
return True
async def async_transform_request(
self,
model: str,
messages: list[AllMessageValues], # mutable-ok: BaseConfig signature
optional_params: dict[str, object], # mutable-ok: BaseConfig signature
litellm_params: dict[str, object], # mutable-ok: BaseConfig signature
headers: dict[str, object], # mutable-ok: BaseConfig signature
) -> dict[str, object]: # mutable-ok: BaseConfig signature
return self.transform_request(
model=model,
messages=await async_inline_remote_media(messages, should_inline=inline_remote_image_urls),
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
def transform_request(
self,
model: str,

View file

@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Final
import httpx
import litellm
from litellm.litellm_core_utils.prompt_templates.image_handling import RemoteMedia, inline_remote_image_urls
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
@ -51,6 +52,9 @@ class VertexAIAnthropicConfig(AnthropicConfig):
def custom_llm_provider(self) -> str | None:
return "vertex_ai"
def inlines_remote_media(self, media: RemoteMedia) -> bool:
return inline_remote_image_urls(media)
def should_strip_billing_metadata(self) -> bool:
return True

View file

@ -157,11 +157,9 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig):
@staticmethod
async def aapply_prompt_template(model: str, messages: list[dict[str, str]]) -> str | None:
"""Apply prompt template (async version)"""
import litellm
from litellm.litellm_core_utils.prompt_templates.factory import (
ahf_chat_template,
custom_prompt,
hf_chat_template,
ibm_granite_pt,
mistral_instruct_pt,
)
@ -179,11 +177,7 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig):
else:
hf_model = model
try:
# Use sync if cached, async if not
if hf_model in litellm.known_tokenizer_config:
result = hf_chat_template(model=hf_model, messages=messages)
else:
result = await ahf_chat_template(model=hf_model, messages=messages)
result = await ahf_chat_template(model=hf_model, messages=messages)
# Return result if it's truthy (not None and not empty string)
# The caller (_aconvert_watsonx_messages_core) will handle None/empty by falling back to default
if result:

View file

@ -16,6 +16,7 @@ from ..common_utils import (
IBMWatsonXMixin,
WatsonXAIError,
_get_api_params,
aconvert_watsonx_messages_to_prompt,
convert_watsonx_messages_to_prompt,
)
@ -236,7 +237,11 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig):
**watsonx_auth_payload,
}
async def atransform_request(
@property
def uses_async_transform_request(self) -> bool:
return True
async def async_transform_request(
self,
model: str,
messages: list[AllMessageValues],
@ -244,11 +249,6 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig):
litellm_params: dict,
headers: dict,
) -> dict:
"""Async version of transform_request"""
from litellm.llms.watsonx.common_utils import (
aconvert_watsonx_messages_to_prompt,
)
provider: Final = model.split("/")[0]
prompt: Final = await aconvert_watsonx_messages_to_prompt(
model=model, messages=messages, provider=provider, custom_prompt_dict={}

View file

@ -6545,7 +6545,7 @@ def embedding(
client=client,
timeout=timeout,
aembedding=aembedding,
litellm_params={},
litellm_params=litellm_params_dict,
api_base=api_base,
print_verbose=print_verbose,
extra_headers=headers,
@ -7805,6 +7805,7 @@ def transcription(
azure_ad_token=azure_ad_token,
max_retries=max_retries,
litellm_params=litellm_params_dict,
custom_llm_provider=custom_llm_provider,
)
elif custom_llm_provider == "openai" or (custom_llm_provider in litellm.openai_compatible_providers):
api_base = (

View file

@ -650,7 +650,10 @@
},
"twelvelabs.marengo-embed-2-7-v1:0": {
"deprecation_date": "2026-11-30",
"input_cost_per_token": 7e-05,
"input_cost_per_query": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
@ -662,7 +665,7 @@
},
"us.twelvelabs.marengo-embed-2-7-v1:0": {
"deprecation_date": "2026-11-30",
"input_cost_per_token": 7e-05,
"input_cost_per_query": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
@ -677,7 +680,7 @@
},
"eu.twelvelabs.marengo-embed-2-7-v1:0": {
"deprecation_date": "2026-11-30",
"input_cost_per_token": 7e-05,
"input_cost_per_query": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
@ -690,6 +693,48 @@
"supports_embedding_image_input": true,
"supports_image_input": true
},
"twelvelabs.marengo-embed-3-0-v1:0": {
"input_cost_per_query": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 500,
"max_tokens": 500,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 512,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"us.twelvelabs.marengo-embed-3-0-v1:0": {
"input_cost_per_query": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 500,
"max_tokens": 500,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 512,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"eu.twelvelabs.marengo-embed-3-0-v1:0": {
"input_cost_per_query": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 500,
"max_tokens": 500,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 512,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
@ -3581,6 +3626,79 @@
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure_ai/gpt-chat-latest": {
"cache_read_input_token_cost": 5e-07,
"deprecation_date": "2026-12-02",
"input_cost_per_token": 5e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3e-05,
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
},
"azure_ai/codex-mini": {
"cache_read_input_token_cost": 3.75e-07,
"deprecation_date": "2026-11-15",
"input_cost_per_token": 1.5e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 200000,
"max_output_tokens": 100000,
"max_tokens": 100000,
"mode": "responses",
"output_cost_per_token": 6e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
"supported_endpoints": [
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true
},
"azure_ai/whisper": {
"deprecation_date": "2026-12-15",
"input_cost_per_second": 0.0001,
"litellm_provider": "azure_ai",
"mode": "audio_transcription",
"output_cost_per_second": 0.0001,
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/"
},
"azure_ai/gpt-5.5-2026-04-23": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
@ -3984,13 +4102,29 @@
"supports_minimal_reasoning_effort": false
},
"azure_ai/model_router": {
"deprecation_date": "2027-05-20",
"input_cost_per_token": 1.4e-07,
"output_cost_per_token": 0,
"litellm_provider": "azure_ai",
"max_input_tokens": 200000,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/aoai/",
"comment": "Flat cost of $0.14 per M input tokens for Azure AI Foundry Model Router infrastructure. Use pattern: azure_ai/model_router/<deployment-name> where deployment-name is your Azure deployment (e.g., azure-model-router)"
},
"azure_ai/model-router": {
"deprecation_date": "2027-05-20",
"input_cost_per_token": 1.4e-07,
"output_cost_per_token": 0,
"litellm_provider": "azure_ai",
"max_input_tokens": 200000,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/aoai/",
"comment": "Catalog-name twin of azure_ai/model_router: the flat $0.14 per M input tokens is the router's own fee, the routed model is priced on top of it"
},
"azure/eu/gpt-4o-2024-08-06": {
"deprecation_date": "2027-04-14",
"cache_read_input_token_cost": 1.375e-06,
@ -10302,6 +10436,18 @@
"/v1/ocr"
]
},
"azure_ai/cohere-command-a": {
"input_cost_per_token": 2.5e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 131072,
"max_output_tokens": 8182,
"max_tokens": 8182,
"mode": "chat",
"output_cost_per_token": 1e-05,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/cohere/",
"supports_function_calling": true,
"supports_tool_choice": true
},
"azure_ai/doc-intelligence/prebuilt-read": {
"litellm_provider": "azure_ai",
"ocr_cost_per_page": 0.0015,
@ -10653,6 +10799,41 @@
"supports_vision": true,
"supports_web_search": true
},
"azure_ai/grok-4-20-reasoning": {
"cache_read_input_token_cost": 1.25e-06,
"deprecation_date": "2027-04-06",
"input_cost_per_token": 1.25e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 262000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/grok/",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_reasoning": true
},
"azure_ai/grok-4-20-non-reasoning": {
"cache_read_input_token_cost": 1.25e-06,
"deprecation_date": "2027-04-06",
"input_cost_per_token": 1.25e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 262000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/grok/",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
},
"azure_ai/grok-4-fast-non-reasoning": {
"deprecation_date": "2026-05-01",
"input_cost_per_token": 2e-07,

View file

@ -21,3 +21,6 @@ _mcp_gateway_initialize_instructions: Final[ContextVar[str | None]] = ContextVar
# Per-request scoped server name; set in MCP HTTP/SSE handlers when the path
# identifies exactly one upstream server. Never populated from client-supplied headers.
_mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gateway_server_name", default=None)
# Set server-side by the /mcp/proxy route. Never populated from client-supplied headers.
_mcp_proxy_mode: Final[ContextVar[bool]] = ContextVar("_mcp_proxy_mode", default=False)

View file

@ -3,12 +3,15 @@ import importlib
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from traceback import walk_tb
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
from uuid import uuid4
import anyio
import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from pydantic import ValidationError
from starlette.datastructures import Headers
from litellm._logging import verbose_logger
@ -30,6 +33,8 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
list_fault_http_status,
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree
from litellm.proxy._experimental.mcp_server.oauth_utils import _redact_mcp_resource_url
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
acting_user_auth,
build_effective_auth_contexts,
@ -78,11 +83,39 @@ _MCP_GUARDRAIL_REJECTIONS: Final = (
def _connection_error_message(exc: BaseException, url: str | None, timeout_seconds: float) -> str:
reference: Final = uuid4().hex
verbose_logger.error(
"MCP connection test failed (reference=%s): %s",
reference,
tuple(
(
type(cause).__name__,
tuple(
(frame.f_code.co_filename, lineno, frame.f_code.co_name)
for frame, lineno in walk_tb(cause.__traceback__)
),
)
for cause in iter_exception_tree(exc)
),
)
return next(
(
message
for cause in iter_exception_tree(exc)
if (message := _known_connection_error_message(cause, url, timeout_seconds)) is not None
),
"An unexpected error occurred while testing the MCP connection. "
f"Retry; if it persists, share reference {reference} with your gateway administrator.",
)
def _known_connection_error_message(exc: BaseException, url: str | None, timeout_seconds: float) -> str | None:
if isinstance(exc, MCPServerURLCredentialsError):
return str(exc.detail)
if isinstance(exc, TimeoutError):
return (
f"Failed to connect to MCP server: no response from {url or 'the server'} "
"Failed to connect to MCP server: no valid MCP response received from "
f"{_redact_mcp_resource_url(url) or 'the server'} "
f"within {timeout_seconds:.0f}s. Check that the LiteLLM proxy can reach this URL "
"from its network (DNS, egress rules, firewalls) and that the server answers MCP requests."
)
@ -99,13 +132,45 @@ def _connection_error_message(exc: BaseException, url: str | None, timeout_secon
return "Failed to connect to MCP server: the connection timed out."
if isinstance(exc, httpx.HTTPStatusError):
return f"Failed to connect to MCP server: it returned HTTP {exc.response.status_code}."
return "Failed to connect to MCP server. Check proxy logs for details."
if isinstance(exc, (httpx.NetworkError, httpx.RemoteProtocolError, ConnectionError)):
return (
"Failed to connect to MCP server: the connection was interrupted. "
"Check the server and network connection, then retry."
)
if isinstance(exc, ValueError) and str(exc).startswith("Unexpected content type:"):
return (
"Failed to connect to MCP server: the endpoint returned an unsupported content type. "
"Check that the URL is an MCP endpoint, not a web page, and matches the selected transport."
)
if isinstance(exc, ValidationError) and exc.title in ("JSONRPCMessage", "InitializeResult", "ListToolsResult"):
return (
"Failed to connect to MCP server: the endpoint returned invalid JSON or an invalid MCP response. "
"Check the MCP endpoint URL and the server's protocol implementation."
)
if MCP_AVAILABLE and isinstance(exc, McpError):
if exc.error.code == -32000 and exc.error.message == "Connection closed":
return (
"Failed to connect to MCP server: the connection was closed before the request completed. "
"Check that the server stays running and returns a complete MCP response, then retry."
)
if exc.error.code == 32600 and exc.error.message == "Session terminated":
return (
"Failed to connect to MCP server: the MCP session was terminated. "
"Check that the URL points to an MCP endpoint and matches the selected transport, "
"then retry to start a new session."
)
return (
f"Failed to connect to MCP server: the MCP request failed (JSON-RPC code {exc.error.code}). "
"Check that the endpoint supports MCP initialization and tool listing, and check the upstream server logs."
)
return None
if MCP_AVAILABLE:
from mcp.shared.exceptions import McpError
from mcp.types import Tool as MCPTool
from litellm.experimental_mcp_client.client import MCPClient
from litellm.experimental_mcp_client.client import MCPClient, as_mcp_read_timeout
from litellm.llms.litellm_proxy.skills.skill_search import (
DEFAULT_SKILL_SEARCH_TOP_K,
)
@ -1342,11 +1407,18 @@ if MCP_AVAILABLE:
except (KeyboardInterrupt, SystemExit, asyncio.CancelledError):
raise
except BaseException as e:
verbose_logger.error("Error in MCP operation: %s", e, exc_info=True)
effective_timeout: Final = (
min(request.timeout if request.timeout is not None else MCP_CLIENT_TIMEOUT, timeout_seconds)
if any(
isinstance(cause, McpError) and as_mcp_read_timeout(cause) is not None
for cause in iter_exception_tree(e)
)
else timeout_seconds
)
return {
"status": "error",
"error": True,
"message": _connection_error_message(e, request.url, timeout_seconds),
"message": _connection_error_message(e, request.url, effective_timeout),
}
async def _preview_openapi_tools(spec_path: str) -> dict:

View file

@ -15,11 +15,11 @@ import types
import uuid
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Protocol
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol
import httpx
from fastapi import FastAPI, HTTPException
from pydantic import AnyUrl, ConfigDict
from pydantic import AnyUrl, ConfigDict, TypeAdapter, ValidationError
from starlette.requests import Request as StarletteRequest
from starlette.responses import JSONResponse
from starlette.types import Message, Receive, Scope, Send
@ -47,6 +47,7 @@ from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
_mcp_gateway_initialize_instructions,
_mcp_gateway_server_name,
_mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # server-owned request mode
)
from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug
from litellm.proxy._experimental.mcp_server.oauth_utils import (
@ -108,9 +109,9 @@ _MAX_STATEFUL_SESSIONS_PER_OWNER: Final = 100
# prevents an authenticated client from forcing the proxy to buffer an
# arbitrarily large body just to make a routing decision.
_MCP_ROUTING_PEEK_MAX_BYTES: Final = 4096
# ASGI scope key holding the tracing span of the request carrying an MCP
# message, written on the request task and read back by the message handler.
# ASGI scope keys carrying OTel request state into a stateful MCP message handler.
_MCP_TRANSPORT_SPAN_SCOPE_KEY: Final = "litellm_otel_transport_span"
_MCP_DESTINATIONS_SCOPE_KEY: Final = "litellm_otel_request_destinations"
def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None:
@ -328,18 +329,17 @@ def _otel_publish_transport_span_on_scope(scope: Scope) -> None:
scope[_MCP_TRANSPORT_SPAN_SCOPE_KEY] = span
def _otel_transport_span_from_message(req_ctx: object) -> object:
"""The tracing span of the HTTP request that carried this MCP message.
Read off that request's ASGI scope, reached through the ``Request`` the
streamable-HTTP transport attaches to each message, so it is this message's
transport and not whichever request happens to have touched the session last.
Returns whatever the scope holds; the otel plumbing validates it."""
def _otel_value_from_message_scope(req_ctx: object, key: str) -> object:
request: Final = getattr(req_ctx, "request", None)
scope: Final = getattr(request, "scope", None)
if not isinstance(scope, Mapping):
return None
return scope.get(_MCP_TRANSPORT_SPAN_SCOPE_KEY)
return scope.get(key)
def _otel_transport_span_from_message(req_ctx: object) -> object:
"""The tracing span of the HTTP request that carried this MCP message."""
return _otel_value_from_message_scope(req_ctx, _MCP_TRANSPORT_SPAN_SCOPE_KEY)
def _otel_set_mcp_transport_span(span: object) -> object:
@ -372,6 +372,44 @@ def _otel_reset_mcp_transport_span(token: object) -> None:
return
def _otel_publish_request_destinations_on_scope(scope: Scope) -> None:
try:
from litellm.integrations.otel.plumbing.context import request_destinations
scope[_MCP_DESTINATIONS_SCOPE_KEY] = request_destinations()
except ImportError:
return
def _otel_set_mcp_request_destinations(req_ctx: object) -> object:
destinations: Final = _otel_value_from_message_scope(req_ctx, _MCP_DESTINATIONS_SCOPE_KEY)
if not isinstance(destinations, tuple):
return None
try:
from litellm.integrations.otel.model.destination import OtelDestination
from litellm.integrations.otel.plumbing.context import set_request_destinations
destination_adapter: Final[TypeAdapter[tuple[OtelDestination, ...]]] = TypeAdapter(
tuple[OtelDestination, ...],
config=ConfigDict(revalidate_instances="always"),
)
validated_destinations: Final = destination_adapter.validate_python(destinations, strict=True)
return set_request_destinations(validated_destinations)
except (ImportError, ValidationError):
return None
def _otel_reset_mcp_request_destinations(token: object) -> None:
if token is None:
return
try:
from litellm.integrations.otel.plumbing.context import reset_request_destinations
reset_request_destinations(token)
except ImportError:
return
def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException:
"""Map a ``ProxyException`` to an ``HTTPException`` that preserves its real
status code and headers.
@ -500,11 +538,22 @@ if MCP_AVAILABLE:
notification_options: NotificationOptions | None = None,
experimental_capabilities: dict[str, dict[str, object]] | None = None,
) -> InitializationOptions:
opts: Final = Server.create_initialization_options(
base_options: Final = Server.create_initialization_options(
self,
notification_options=notification_options,
experimental_capabilities=experimental_capabilities or {},
)
opts: Final = (
base_options.model_copy(
update={ # mutable-ok: Pydantic update payload
"capabilities": base_options.capabilities.model_copy(
update={"prompts": None, "resources": None} # mutable-ok: Pydantic update payload
)
}
)
if _mcp_proxy_mode.get()
else base_options
)
updates: Final[dict[str, str]] = {}
merged: Final = _mcp_gateway_initialize_instructions.get()
if merged is not None:
@ -718,12 +767,12 @@ if MCP_AVAILABLE:
_stateful_auth_context_cleanup_task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await _stateful_auth_context_cleanup_task
if _session_manager_cm:
await _session_manager_cm.__aexit__(None, None, None)
if _session_manager_stateful_cm:
await _session_manager_stateful_cm.__aexit__(None, None, None)
if _sse_session_manager_cm:
await _sse_session_manager_cm.__aexit__(None, None, None)
if _session_manager_stateful_cm:
await _session_manager_stateful_cm.__aexit__(None, None, None)
if _session_manager_cm:
await _session_manager_cm.__aexit__(None, None, None)
except Exception as e:
verbose_logger.exception("Error during session manager shutdown: %s", e)
@ -763,10 +812,12 @@ if MCP_AVAILABLE:
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
_trace_token = None
_transport_token = None
_destinations_token = None
try:
_trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx))
_transport_token = _otel_set_mcp_transport_span(_otel_transport_span_from_message(req_ctx))
_destinations_token = _otel_set_mcp_request_destinations(req_ctx)
# Get user authentication from context variable
(
user_api_key_auth,
@ -783,17 +834,20 @@ if MCP_AVAILABLE:
"MCP list_tools - MCP server auth headers: %s",
list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None,
)
from mcp.types import Tool
from litellm.proxy._experimental.mcp_server.tool_search import (
get_mcp_proxy_tool_definitions,
get_virtual_tool_definitions,
)
if _mcp_proxy_mode.get():
return [Tool.model_validate(d) for d in get_mcp_proxy_tool_definitions()] # mutable-ok: MCP SDK list
if getattr(
getattr(user_api_key_auth, "object_permission", None),
"mcp_tool_search_enabled",
False,
):
from mcp.types import Tool
from litellm.proxy._experimental.mcp_server.tool_search import (
get_virtual_tool_definitions,
)
return [Tool.model_validate(d) for d in get_virtual_tool_definitions()]
# Get mcp_servers from context variable
@ -828,6 +882,7 @@ if MCP_AVAILABLE:
# This prevents the HTTP stream from failing and allows the client to get a response
return []
finally:
_otel_reset_mcp_request_destinations(_destinations_token)
_otel_reset_mcp_transport_span(_transport_token)
_otel_reset_mcp_trace_carrier(_trace_token)
if _session_reset_token is not None:
@ -866,6 +921,12 @@ if MCP_AVAILABLE:
verbose_logger.debug("Host progressToken captured: %s...", str(host_token)[:8])
return forward_progress
def _reject_mcp_proxy_operation() -> NoReturn:
from mcp.shared.exceptions import McpError
from mcp.types import METHOD_NOT_FOUND, ErrorData
raise McpError(ErrorData(code=METHOD_NOT_FOUND, message="Operation unavailable on /mcp/proxy"))
async def _build_virtual_call_logging_obj(
name: str,
arguments: dict[str, object],
@ -921,16 +982,91 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.tool_search import (
AGENT_SEARCH_TOOL_NAME,
DEFAULT_AGENT_SEARCH_TOP_K,
MCP_PROXY_CALL_TOOL_NAME,
MCP_PROXY_TOOL_NAMES,
MCP_TOOL_SEARCH_TOOL_NAME,
SKILL_SEARCH_TOOL_NAME,
VIRTUAL_TOOL_NAMES,
coerce_top_k,
handle_agent_search,
handle_mcp_proxy_tool,
handle_mcp_tool_call,
handle_mcp_tool_search,
handle_skill_search,
)
if _mcp_proxy_mode.get() and name not in MCP_PROXY_TOOL_NAMES:
return CallToolResult(
content=[ # mutable-ok: MCP result content
TextContent(type="text", text=f"Tool {name} is unavailable on /mcp/proxy")
],
isError=True,
)
if _mcp_proxy_mode.get() and name in MCP_PROXY_TOOL_NAMES:
assert user_api_key_auth is not None
proxy_call_start: Final = datetime.now() # noqa: DTZ005 # logging pipeline uses naive datetimes
proxy_logging_obj: Final = (
await _build_virtual_call_logging_obj(
name=name,
arguments=arguments or {}, # mutable-ok: logging pipeline payload
user_api_key_auth=user_api_key_auth,
raw_headers=raw_headers,
client_ip=client_ip,
)
if name == MCP_PROXY_CALL_TOOL_NAME
else None
)
try:
proxy_result: Final = await handle_mcp_proxy_tool(
name=name,
arguments=arguments or {}, # mutable-ok: proxy handler payload
user_api_key_dict=user_api_key_auth,
client_ip=client_ip,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
litellm_logging_obj=proxy_logging_obj,
)
except Exception as exc:
if proxy_logging_obj is not None:
from litellm.proxy.proxy_server import proxy_logging_obj as request_logging_obj
failure_end: Final = datetime.now() # noqa: DTZ005 # matches the logging pipeline start time
failure_traceback: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG)
try:
proxy_logging_obj.failure_handler(exc, failure_traceback, proxy_call_start, failure_end)
await proxy_logging_obj.async_failure_handler(
exc, failure_traceback, proxy_call_start, failure_end
)
if not isinstance(exc, MCPUpstreamAuthError):
await request_logging_obj.post_call_failure_hook(
request_data={ # mutable-ok: failure hook mutates its request payload
"name": name,
"arguments": arguments,
"litellm_logging_obj": proxy_logging_obj,
},
original_exception=exc,
user_api_key_dict=user_api_key_auth,
route="/mcp/call_tool",
traceback_str=failure_traceback,
)
except Exception: # noqa: BLE001 # a failing failure hook must not mask the tool call's own error
verbose_logger.exception("Error logging failed MCP proxy tool call")
raise
if proxy_logging_obj is not None:
return await _fire_mcp_tool_call_logging(
logging_obj=proxy_logging_obj,
result=proxy_result,
start_time=proxy_call_start,
end_time=datetime.now(), # noqa: DTZ005 # matches the logging pipeline start time
user_api_key_auth=user_api_key_auth,
request_data=types.MappingProxyType({"name": name, "arguments": arguments}),
)
return proxy_result
if name not in VIRTUAL_TOOL_NAMES:
return None
@ -1021,10 +1157,12 @@ if MCP_AVAILABLE:
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
_trace_token = None
_transport_token = None
_destinations_token = None
try:
_trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx))
_transport_token = _otel_set_mcp_transport_span(_otel_transport_span_from_message(req_ctx))
_destinations_token = _otel_set_mcp_request_destinations(req_ctx)
# Validate arguments
(
user_api_key_auth,
@ -1163,6 +1301,7 @@ if MCP_AVAILABLE:
return response
finally:
_otel_reset_mcp_request_destinations(_destinations_token)
_otel_reset_mcp_transport_span(_transport_token)
_otel_reset_mcp_trace_carrier(_trace_token)
if _session_reset_token is not None:
@ -1173,6 +1312,8 @@ if MCP_AVAILABLE:
"""
List all available prompts
"""
if _mcp_proxy_mode.get():
_reject_mcp_proxy_operation()
from mcp.server.lowlevel.server import request_ctx
req_ctx: Final = request_ctx.get(None)
@ -1230,8 +1371,8 @@ if MCP_AVAILABLE:
Returns:
GetPromptResult: Getting prompt execution results
"""
# Validate arguments
if _mcp_proxy_mode.get():
_reject_mcp_proxy_operation()
from mcp.server.lowlevel.server import request_ctx
req_ctx: Final = request_ctx.get(None)
@ -1268,6 +1409,8 @@ if MCP_AVAILABLE:
@server.list_resources()
async def list_resources() -> list[Resource]:
"""List all available resources."""
if _mcp_proxy_mode.get():
_reject_mcp_proxy_operation()
from mcp.server.lowlevel.server import request_ctx
req_ctx: Final = request_ctx.get(None)
@ -1312,6 +1455,8 @@ if MCP_AVAILABLE:
@server.list_resource_templates()
async def list_resource_templates() -> list[ResourceTemplate]:
"""List all available resource templates."""
if _mcp_proxy_mode.get():
_reject_mcp_proxy_operation()
from mcp.server.lowlevel.server import request_ctx
req_ctx: Final = request_ctx.get(None)
@ -1357,6 +1502,8 @@ if MCP_AVAILABLE:
@server.read_resource()
async def read_resource(url: AnyUrl) -> list[ReadResourceContents]:
if _mcp_proxy_mode.get():
_reject_mcp_proxy_operation()
from mcp.server.lowlevel.server import request_ctx
req_ctx: Final = request_ctx.get(None)
@ -1955,6 +2102,7 @@ if MCP_AVAILABLE:
litellm_trace_id: str | None = None,
request_tags: list[str] | None = None,
client_ip: str | None = None,
mcp_proxy_mode: bool = False,
) -> AggregateToolListing:
"""
Helper method to fetch tools from MCP servers based on server filtering criteria.
@ -2134,9 +2282,14 @@ if MCP_AVAILABLE:
user_api_key_auth=user_api_key_auth,
)
# Apply display-name/description overrides last so that
# permission filtering always works against original names.
filtered_tools = apply_tool_overrides(filtered_tools, server)
if mcp_proxy_mode:
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
filtered_tools = [ # mutable-ok: MCP tool pipeline
with_mcp_proxy_identity(tool, server.server_id) for tool in filtered_tools
]
else:
filtered_tools = apply_tool_overrides(filtered_tools, server)
verbose_logger.debug(
"Successfully fetched %s tools from server %s, %s after filtering",
@ -2448,6 +2601,7 @@ if MCP_AVAILABLE:
log_list_tools_to_spendlogs: bool = False,
list_tools_log_source: str | None = None,
client_ip: str | None = None,
mcp_proxy_mode: bool = False,
) -> AggregateToolListing:
"""
List all available MCP tools.
@ -2477,6 +2631,7 @@ if MCP_AVAILABLE:
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
list_tools_log_source=list_tools_log_source,
client_ip=client_ip,
mcp_proxy_mode=mcp_proxy_mode,
)
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
return listing
@ -3376,7 +3531,9 @@ if MCP_AVAILABLE:
server_name: str | None,
session_id: str | None = None,
) -> StandardLoggingMCPToolCall:
mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name(
add_server_prefix_to_name(name, server_name) if server_name else name
)
namespaced_tool_name: Final = f"{server_name}/{name}" if server_name else name
if mcp_server:
mcp_info: Final = mcp_server.mcp_info or {}
@ -4493,6 +4650,7 @@ if MCP_AVAILABLE:
async def _dispatch() -> None:
_otel_publish_transport_span_on_scope(scope)
_otel_publish_request_destinations_on_scope(scope)
auth_user: Final = _set_or_update_auth_context(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,

View file

@ -1,5 +1,6 @@
from __future__ import annotations
import hashlib
import json
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
@ -30,6 +31,12 @@ if TYPE_CHECKING:
MCP_TOOL_SEARCH_SETTINGS_KEY: Final[str] = "mcp_tool_search"
MCP_TOOL_SEARCH_TOOL_NAME: Final[str] = "mcp_tool_search"
MCP_TOOL_CALL_TOOL_NAME: Final[str] = "mcp_tool_call"
MCP_PROXY_SEARCH_TOOL_NAME: Final[str] = "search_tools"
MCP_PROXY_SCHEMA_TOOL_NAME: Final[str] = "get_tool_schema"
MCP_PROXY_CALL_TOOL_NAME: Final[str] = "call_tool"
MCP_PROXY_TOOL_NAMES: Final = frozenset(
(MCP_PROXY_SEARCH_TOOL_NAME, MCP_PROXY_SCHEMA_TOOL_NAME, MCP_PROXY_CALL_TOOL_NAME)
)
AGENT_SEARCH_TOOL_NAME: Final[str] = "agent_search"
SKILL_SEARCH_TOOL_NAME: Final[str] = "skill_search"
VIRTUAL_TOOL_NAMES: Final = frozenset(
@ -51,6 +58,29 @@ class ToolSearchResult(TypedDict, total=False):
score: ReadOnly[float]
class MCPProxySearchResult(TypedDict, total=False):
tool_id: Required[ReadOnly[str]]
name: Required[ReadOnly[str]]
description: Required[ReadOnly[str]]
score: ReadOnly[float]
class MCPProxySchemaResult(MCPProxySearchResult, total=False):
inputSchema: Required[ReadOnly[Mapping[str, object]]]
outputSchema: ReadOnly[Mapping[str, object]]
class MCPProxyToolIdentity(TypedDict):
server_id: ReadOnly[str]
tool_name: ReadOnly[str]
@dataclass(frozen=True, slots=True)
class MCPToolSearchHit:
tool: Tool
score: float | None = None
@dataclass(frozen=True, slots=True)
class SemanticToolRanker:
embed: Embedder
@ -76,6 +106,55 @@ def _scored_result(tool: Tool, score: float) -> ToolSearchResult:
return {"name": tool.name, "description": tool.description or "", "inputSchema": tool.inputSchema, "score": score}
_MCP_PROXY_IDENTITY_META_KEY: Final[str] = "litellm.ai/proxy_tool_identity"
def with_mcp_proxy_identity(tool: Tool, server_id: str) -> Tool:
identity: Final[MCPProxyToolIdentity] = {"server_id": server_id, "tool_name": tool.name}
return tool.model_copy( # mutable-ok: Pydantic requires mutable update and metadata mappings
update={ # mutable-ok: Pydantic update payload
"meta": {**(tool.meta or {}), _MCP_PROXY_IDENTITY_META_KEY: identity} # mutable-ok: metadata mapping
}
)
def _mcp_proxy_identity(tool: Tool) -> MCPProxyToolIdentity:
identity: Final = (tool.meta or {}).get(_MCP_PROXY_IDENTITY_META_KEY) # mutable-ok: absent metadata default
if not isinstance(identity, Mapping):
raise TypeError("MCP proxy tool identity is missing")
server_id: Final = identity.get("server_id")
tool_name: Final = identity.get("tool_name")
if not isinstance(server_id, str) or not isinstance(tool_name, str):
raise TypeError("MCP proxy tool identity is invalid")
return {"server_id": server_id, "tool_name": tool_name} # mutable-ok: TypedDict identity payload
def mcp_proxy_tool_id(tool: Tool) -> str:
identity: Final = _mcp_proxy_identity(tool)
return hashlib.sha256(f"{identity['server_id']}\0{identity['tool_name']}".encode()).hexdigest()[:32]
def _proxy_search_result(hit: MCPToolSearchHit) -> MCPProxySearchResult:
base: Final[MCPProxySearchResult] = {
"tool_id": mcp_proxy_tool_id(hit.tool),
"name": hit.tool.name,
"description": hit.tool.description or "",
}
return {**base, "score": hit.score} if hit.score is not None else base # mutable-ok: wire result payload
def _proxy_schema_result(tool: Tool) -> MCPProxySchemaResult:
base: Final[MCPProxySchemaResult] = {
"tool_id": mcp_proxy_tool_id(tool),
"name": tool.name,
"description": tool.description or "",
"inputSchema": tool.inputSchema,
}
if tool.outputSchema is None:
return base
return {**base, "outputSchema": tool.outputSchema} # mutable-ok: wire schema payload
def _tool_text(tool: Tool) -> str:
return "\n".join(part for part in (tool.name, tool.description or "") if part)
@ -107,6 +186,38 @@ def search_tools(query: str, tools: Sequence[Tool], top_k: int = 5) -> tuple[Too
return tuple(_tool_result(tool) for _, tool in _top_hits(tools, scores, minimum=1.0, limit=top_k))
async def rank_mcp_tools(
query: str,
tools: Sequence[Tool],
top_k: int,
settings: MCPToolSearchSettings,
ranker: SemanticToolRanker | None,
) -> tuple[MCPToolSearchHit, ...] | EmbeddingFailed:
core, rest = _split_core_tools(tools, settings.core_tools)
core_hits: Final = tuple(MCPToolSearchHit(tool) for tool in core)
if not query:
return core_hits
limit: Final = min(top_k, settings.top_k)
if ranker is None:
scores: Final = tuple(_keyword_score(query, tool) for tool in rest)
return (
*core_hits,
*(MCPToolSearchHit(tool) for _, tool in _top_hits(rest, scores, minimum=1.0, limit=limit)),
)
semantic_scores: Final = await ranker.index.scores(
query, tuple(_tool_text(tool) for tool in rest), ranker.embed, ranker.embedding_model
)
if isinstance(semantic_scores, EmbeddingFailed):
return semantic_scores
return (
*core_hits,
*(
MCPToolSearchHit(tool, score)
for score, tool in _top_hits(rest, semantic_scores, settings.similarity_threshold, limit)
),
)
async def search_mcp_tools(
query: str,
tools: Sequence[Tool],
@ -114,21 +225,12 @@ async def search_mcp_tools(
settings: MCPToolSearchSettings,
ranker: SemanticToolRanker | None,
) -> tuple[ToolSearchResult, ...] | EmbeddingFailed:
"""Core tools the caller can access come first, then up to `top_k` ranked matches from the remaining tools."""
core, rest = _split_core_tools(tools, settings.core_tools)
limit: Final = min(top_k, settings.top_k)
core_results: Final = tuple(_tool_result(tool) for tool in core)
if ranker is None:
return (*core_results, *search_tools(query, rest, limit))
if not query:
return core_results
scores: Final = await ranker.index.scores(
query, tuple(_tool_text(tool) for tool in rest), ranker.embed, ranker.embedding_model
hits: Final = await rank_mcp_tools(query, tools, top_k, settings, ranker)
if isinstance(hits, EmbeddingFailed):
return hits
return tuple(
_scored_result(hit.tool, hit.score) if hit.score is not None else _tool_result(hit.tool) for hit in hits
)
if isinstance(scores, EmbeddingFailed):
return scores
hits: Final = _top_hits(rest, scores, minimum=settings.similarity_threshold, limit=limit)
return (*core_results, *(_scored_result(tool, score) for score, tool in hits))
class _ToolParamSchema(TypedDict, total=False):
@ -223,10 +325,48 @@ _SKILL_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
}
_MCP_PROXY_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
"name": MCP_PROXY_SEARCH_TOOL_NAME,
"description": "Search accessible MCP tools by describing what you need. Returns opaque tool IDs.",
"inputSchema": {
"type": "object",
"properties": {"query": {"type": "string", "description": "What the tool should do."}},
"required": _json_array("query"),
},
}
_MCP_PROXY_SCHEMA_DEFINITION: Final[VirtualToolDefinition] = {
"name": MCP_PROXY_SCHEMA_TOOL_NAME,
"description": "Return the complete schema for an accessible MCP tool ID.",
"inputSchema": {
"type": "object",
"properties": {"tool_id": {"type": "string", "description": "Opaque ID from search_tools."}},
"required": _json_array("tool_id"),
},
}
_MCP_PROXY_CALL_DEFINITION: Final[VirtualToolDefinition] = {
"name": MCP_PROXY_CALL_TOOL_NAME,
"description": "Call an accessible MCP tool by opaque ID with schema-valid arguments.",
"inputSchema": {
"type": "object",
"properties": {
"tool_id": {"type": "string", "description": "Opaque ID from search_tools."},
"arguments": {"type": "object", "description": "Arguments validated against the selected tool schema."},
},
"required": _json_array("tool_id"),
},
}
def get_virtual_tool_definitions() -> tuple[VirtualToolDefinition, ...]:
return (_MCP_TOOL_SEARCH_DEFINITION, _MCP_TOOL_CALL_DEFINITION, _AGENT_SEARCH_DEFINITION, _SKILL_SEARCH_DEFINITION)
def get_mcp_proxy_tool_definitions() -> tuple[VirtualToolDefinition, ...]:
return (_MCP_PROXY_SEARCH_DEFINITION, _MCP_PROXY_SCHEMA_DEFINITION, _MCP_PROXY_CALL_DEFINITION)
def _text_tool_result(text: str, is_error: bool) -> CallToolResult:
from mcp.types import CallToolResult, TextContent
@ -314,7 +454,9 @@ async def handle_mcp_tool_search(
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
) -> CallToolResult:
from litellm.proxy._experimental.mcp_server.server import _list_mcp_tools
from litellm.proxy._experimental.mcp_server.server import (
_list_mcp_tools, # pyright: ignore[reportPrivateUsage] # shared catalog owner
)
from litellm.proxy.proxy_server import llm_router, proxy_logging_obj
settings: Final = mcp_tool_search_settings()
@ -351,6 +493,97 @@ async def handle_mcp_tool_search(
return _text_tool_result(json.dumps(results), is_error=False)
async def handle_mcp_proxy_tool(
name: str,
arguments: dict[str, object], # mutable-ok: MCP dispatcher passes mutable call arguments
user_api_key_dict: UserAPIKeyAuth,
client_ip: str | None = None,
mcp_servers: list[str] | None = None, # mutable-ok: preserve MCP scope container for existing resolver
mcp_auth_header: str | None = None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, # mutable-ok: preserve forwarded headers
oauth2_headers: dict[str, str] | None = None, # mutable-ok: preserve forwarded headers
raw_headers: dict[str, str] | None = None, # mutable-ok: preserve request headers
litellm_logging_obj: LiteLLMLoggingObj | None = None,
) -> CallToolResult:
from fastapi import HTTPException
from jsonschema import ValidationError as JsonSchemaValidationError
from jsonschema import validate
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.server import ( # pyright: ignore[reportPrivateUsage] # shared catalog owner
_list_mcp_tools, # pyright: ignore[reportPrivateUsage] # shared catalog owner
)
listing: Final = await _list_mcp_tools(
user_api_key_auth=user_api_key_dict,
mcp_servers=mcp_servers,
client_ip=client_ip,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
mcp_proxy_mode=True,
)
tools_by_id: Final = {mcp_proxy_tool_id(tool): tool for tool in listing.tools} # mutable-ok: lookup index
if name == MCP_PROXY_SEARCH_TOOL_NAME:
llm_router: Final = proxy_server.llm_router
proxy_logging_obj: Final = proxy_server.proxy_logging_obj
settings: Final = mcp_tool_search_settings()
if isinstance(settings, ValidationError):
return _text_tool_result(str(settings), is_error=True)
if settings.embedding_model is not None and llm_router is None:
return _text_tool_result(
f"litellm_settings.{MCP_TOOL_SEARCH_SETTINGS_KEY}.embedding_model needs a model_list so it can be called",
is_error=True,
)
ranker: Final = (
SemanticToolRanker(
embed=router_embedder(llm_router, settings.embedding_model, user_api_key_dict, proxy_logging_obj),
embedding_model=settings.embedding_model,
index=global_mcp_tool_search_index,
)
if settings.embedding_model is not None and llm_router is not None
else None
)
results: Final = await rank_mcp_tools(str(arguments.get("query", "")), listing.tools, 5, settings, ranker)
if isinstance(results, EmbeddingFailed):
return _text_tool_result(results.reason, is_error=True)
return _text_tool_result(json.dumps(tuple(_proxy_search_result(hit) for hit in results)), is_error=False)
tool_id: Final = arguments.get("tool_id")
tool: Final = tools_by_id.get(tool_id) if isinstance(tool_id, str) else None
if tool is None:
return _text_tool_result("Unknown or unauthorized tool_id", is_error=True)
if name == MCP_PROXY_SCHEMA_TOOL_NAME:
return _text_tool_result(json.dumps(_proxy_schema_result(tool)), is_error=False)
if name != MCP_PROXY_CALL_TOOL_NAME:
raise HTTPException(status_code=400, detail=f"Unknown MCP proxy tool: {name}")
tool_arguments: Final = arguments.get("arguments", {}) # mutable-ok: JSON Schema validator consumes mapping
if not isinstance(tool_arguments, dict):
return _text_tool_result("arguments must be an object", is_error=True)
try:
validate(instance=tool_arguments, schema=tool.inputSchema)
except JsonSchemaValidationError as exc:
return _text_tool_result(f"Invalid arguments: {exc.message}", is_error=True)
return await handle_mcp_tool_call(
tool_name=_mcp_proxy_identity(tool)["tool_name"],
arguments=tool_arguments,
user_api_key_dict=user_api_key_dict,
requested_server_id=_mcp_proxy_identity(tool)["server_id"],
client_ip=client_ip,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
)
async def handle_mcp_tool_call(
tool_name: str,
arguments: dict[str, Any],
@ -362,6 +595,7 @@ async def handle_mcp_tool_call(
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
litellm_logging_obj: LiteLLMLoggingObj | None = None,
requested_server_id: str | None = None,
) -> CallToolResult:
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers,
@ -400,4 +634,5 @@ async def handle_mcp_tool_call(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
requested_server_id=requested_server_id,
)

View file

@ -132,10 +132,18 @@ async def update_mcp_toolset(
data: UpdateMCPToolsetRequest,
touched_by: str,
) -> MCPToolset | None:
data_dict: Final = data.model_dump(exclude_none=True, exclude={"toolset_id"})
if "tools" in data_dict:
data_dict["tools"] = json.dumps(data_dict["tools"])
data_dict["updated_by"] = touched_by
"""A partial update: absent keeps, null clears. A toolset always has a name and a
tool list, so a null ``toolset_name`` or ``tools`` is a no-op rather than a clear;
emptying the tool selection is an explicit ``[]``, which cannot be mistaken for a
caller that left the field out."""
data_dict: Final = dict( # mutable-ok: Prisma requires a plain dict for JSON query serialization
(
(field, json.dumps(value) if field == "tools" else value)
for field, value in data.model_dump(exclude_unset=True).items()
if field != "toolset_id" and (field not in ("toolset_name", "tools") or value is not None)
),
updated_by=touched_by,
)
try:
row: Final = await _toolset_table(prisma_client).update(
where={"toolset_id": data.toolset_id},

View file

@ -17027,6 +17027,134 @@
"mcp_app"
]
}
},
"/mcp/proxy": {
"delete": {
"description": "Serve the fixed three-tool MCP proxy surface.",
"operationId": "proxy_mcp_route_mcp_proxy_delete",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Mcp Route",
"tags": [
"mcp_app"
]
},
"get": {
"description": "Serve the fixed three-tool MCP proxy surface.",
"operationId": "proxy_mcp_route_mcp_proxy_get",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Mcp Route",
"tags": [
"mcp_app"
]
},
"head": {
"description": "Serve the fixed three-tool MCP proxy surface.",
"operationId": "proxy_mcp_route_mcp_proxy_head",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Mcp Route",
"tags": [
"mcp_app"
]
},
"options": {
"description": "Serve the fixed three-tool MCP proxy surface.",
"operationId": "proxy_mcp_route_mcp_proxy_options",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Mcp Route",
"tags": [
"mcp_app"
]
},
"patch": {
"description": "Serve the fixed three-tool MCP proxy surface.",
"operationId": "proxy_mcp_route_mcp_proxy_patch",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Mcp Route",
"tags": [
"mcp_app"
]
},
"post": {
"description": "Serve the fixed three-tool MCP proxy surface.",
"operationId": "proxy_mcp_route_mcp_proxy_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Mcp Route",
"tags": [
"mcp_app"
]
},
"put": {
"description": "Serve the fixed three-tool MCP proxy surface.",
"operationId": "proxy_mcp_route_mcp_proxy_put",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Mcp Route",
"tags": [
"mcp_app"
]
}
}
}
},

View file

@ -3,6 +3,7 @@ import json
import os
from collections.abc import Callable, Mapping
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple
import httpx
@ -502,6 +503,7 @@ class LiteLLMRoutes(enum.Enum):
mcp_inference_routes = [
"/mcp",
"/mcp/",
"/mcp/proxy",
"/mcp/{subpath}",
"/mcp/tools",
"/mcp/tools/list",
@ -835,6 +837,14 @@ class LiteLLMRoutes(enum.Enum):
"/team/daily/activity/aggregated",
"/team/spend/by_user",
"/team/{team_id}/members/me",
# POST/GET the team's logging callbacks, and DELETE one of them. Every
# handler calls _verify_team_access, which admits only a proxy admin, an
# org admin for the team, or an admin of this team.
#
# team_id is a free-form string, so it spells these with the same path
# converter the router uses; the gate matches that converter.
"/team/{team_id:path}/callback",
"/team/{team_id:path}/callback/{callback_name}",
"/model/new",
"/model/update",
"/model/delete",
@ -1285,6 +1295,13 @@ class UpdateKeyRequest(KeyRequestBase):
rotation_interval: str | None = None
organization_id: str | None = None
@model_validator(mode="before")
@classmethod
def drop_blank_team_id(cls, values: object) -> object:
if isinstance(values, Mapping) and values.get("team_id") == "":
return MappingProxyType({k: v for k, v in values.items() if k != "team_id"})
return values
@field_validator("organization_id", mode="before")
@classmethod
def treat_cleared_organization_id_as_unset(cls, v: object) -> object:
@ -3584,6 +3601,7 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
ui_callback_name="OpenTelemetry",
litellm_callback_params=[
"OTEL_EXPORTER",
"OTEL_EXPORTER_OTLP_PROTOCOL",
"OTEL_ENDPOINT",
"OTEL_TRACES_ENDPOINT",
"OTEL_HEADERS",
@ -5138,9 +5156,26 @@ class CostEstimateRequest(LiteLLMPydanticObjectBase):
model: str = Field(description="Model name (from /model_group/info)")
input_tokens: int = Field(description="Expected input tokens per request", ge=0)
output_tokens: int = Field(description="Expected output tokens per request", ge=0)
cache_read_input_tokens: int = Field(
default=0, description="Input tokens read from the prompt cache; counted within input_tokens", ge=0
)
cache_creation_input_tokens: int = Field(
default=0, description="Input tokens written to the prompt cache; counted within input_tokens", ge=0
)
reasoning_tokens: int = Field(
default=0, description="Reasoning tokens the model emits; counted within output_tokens", ge=0
)
num_requests_per_day: int | None = Field(default=None, description="Number of requests per day", ge=0)
num_requests_per_month: int | None = Field(default=None, description="Number of requests per month", ge=0)
@model_validator(mode="after")
def validate_token_subsets(self) -> "CostEstimateRequest":
if self.cache_read_input_tokens + self.cache_creation_input_tokens > self.input_tokens:
raise ValueError("cache_read_input_tokens plus cache_creation_input_tokens cannot exceed input_tokens")
if self.reasoning_tokens > self.output_tokens:
raise ValueError("reasoning_tokens cannot exceed output_tokens")
return self
class CostEstimateResponse(LiteLLMPydanticObjectBase):
"""Response body for /cost/estimate endpoint."""
@ -5148,6 +5183,9 @@ class CostEstimateResponse(LiteLLMPydanticObjectBase):
model: str
input_tokens: int
output_tokens: int
cache_read_input_tokens: int = 0
cache_creation_input_tokens: int = 0
reasoning_tokens: int = 0
num_requests_per_day: int | None = None
num_requests_per_month: int | None = None
# Per-request costs
@ -5155,17 +5193,33 @@ class CostEstimateResponse(LiteLLMPydanticObjectBase):
input_cost_per_request: float = Field(description="Input token cost per request (before margin)")
output_cost_per_request: float = Field(description="Output token cost per request (before margin)")
margin_cost_per_request: float = Field(default=0.0, description="Margin/fee added per request")
cache_read_cost_per_request: float = Field(default=0.0, description="Cache-read share of input_cost_per_request")
cache_creation_cost_per_request: float = Field(
default=0.0, description="Cache-write share of input_cost_per_request"
)
reasoning_cost_per_request: float = Field(default=0.0, description="Reasoning share of output_cost_per_request")
# Daily costs (if num_requests_per_day provided)
daily_cost: float | None = Field(default=None, description="Total daily cost (includes margin)")
daily_input_cost: float | None = Field(default=None, description="Daily input token cost")
daily_output_cost: float | None = Field(default=None, description="Daily output token cost")
daily_margin_cost: float | None = Field(default=None, description="Daily margin/fee")
daily_cache_read_cost: float | None = Field(default=None, description="Cache-read share of daily_input_cost")
daily_cache_creation_cost: float | None = Field(default=None, description="Cache-write share of daily_input_cost")
daily_reasoning_cost: float | None = Field(default=None, description="Reasoning share of daily_output_cost")
# Monthly costs (if num_requests_per_month provided)
monthly_cost: float | None = Field(default=None, description="Total monthly cost (includes margin)")
monthly_input_cost: float | None = Field(default=None, description="Monthly input token cost")
monthly_output_cost: float | None = Field(default=None, description="Monthly output token cost")
monthly_margin_cost: float | None = Field(default=None, description="Monthly margin/fee")
# Pricing info
input_cost_per_token: float | None = None
output_cost_per_token: float | None = None
monthly_cache_read_cost: float | None = Field(default=None, description="Cache-read share of monthly_input_cost")
monthly_cache_creation_cost: float | None = Field(
default=None, description="Cache-write share of monthly_input_cost"
)
monthly_reasoning_cost: float | None = Field(default=None, description="Reasoning share of monthly_output_cost")
# Pricing info: the rates this request's usage bills at, after token tiers and regional multipliers
input_cost_per_token: float | None = Field(default=None, description="Rate billed per input token")
output_cost_per_token: float | None = Field(default=None, description="Rate billed per output token")
cache_read_input_token_cost: float | None = Field(default=None, description="Rate billed per cache-read token")
cache_creation_input_token_cost: float | None = Field(default=None, description="Rate billed per cache-write token")
output_cost_per_reasoning_token: float | None = Field(default=None, description="Rate billed per reasoning token")
provider: str | None = None

View file

@ -144,31 +144,32 @@ def _validate_push_notification_url(url: str) -> None:
raise HTTPException(status_code=400, detail=str(e)) from e
def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> dict[str, str]:
headers: Final[dict[str, str]] = {}
if user_api_key_dict.user_id:
headers["X-LiteLLM-User-Id"] = user_api_key_dict.user_id
if user_api_key_dict.team_id:
headers["X-LiteLLM-Team-Id"] = user_api_key_dict.team_id
return headers
def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, str]:
return MappingProxyType(
{
name: value
for name, value in (
("X-LiteLLM-User-Id", user_api_key_dict.user_id),
("X-LiteLLM-Team-Id", user_api_key_dict.team_id),
)
if value
}
)
def _forwarding_headers(
user_api_key_dict: UserAPIKeyAuth,
caller_identity: Mapping[str, str],
request_data: Mapping[str, object],
agent_extra_headers: Mapping[str, str] | None,
) -> Mapping[str, str] | None:
sanitized: Final = (
{k: v for k, v in agent_extra_headers.items() if not k.lower().startswith("x-litellm-")}
if agent_extra_headers
else None
) -> dict[str, str] | None:
passthrough: Final = tuple(
(name, value)
for name, value in (agent_extra_headers.items() if agent_extra_headers else ())
if not name.lower().startswith("x-litellm-")
)
merged: Final = merge_agent_headers(dynamic_headers=sanitized, static_headers=None) or {}
identity: Final = _caller_identity_headers(user_api_key_dict)
trace_id: Final = request_data.get("litellm_trace_id")
if trace_id:
identity["X-LiteLLM-Trace-Id"] = str(trace_id)
merged.update(identity)
trace: Final = (("X-LiteLLM-Trace-Id", str(trace_id)),) if trace_id else ()
merged: Final = dict((*passthrough, *caller_identity.items(), *trace))
return merged or None
@ -755,6 +756,7 @@ async def invoke_agent_a2a(
ProxyBaseLLMRequestProcessing,
)
caller_identity: Final = _caller_identity_headers(user_api_key_dict)
processor: Final = ProxyBaseLLMRequestProcessing(data=body)
data, logging_obj = await processor.common_processing_pre_call_logic(
request=request,
@ -793,9 +795,13 @@ async def invoke_agent_a2a(
if header_name:
dynamic_headers[header_name] = val
agent_extra_headers = merge_agent_headers(
dynamic_headers=dynamic_headers or None,
static_headers=static_headers or None,
agent_extra_headers = _forwarding_headers(
caller_identity=caller_identity,
request_data=data,
agent_extra_headers=merge_agent_headers(
dynamic_headers=dynamic_headers or None,
static_headers=static_headers or None,
),
)
# Databricks App endpoints require a short-lived OAuth M2M token rather
@ -942,12 +948,7 @@ async def invoke_agent_a2a(
"method": method,
"params": params,
}
caller_headers: Final = _forwarding_headers(
user_api_key_dict=user_api_key_dict,
request_data=data,
agent_extra_headers=agent_extra_headers,
)
result = await _forward_jsonrpc(agent_url, forward_body, extra_headers=caller_headers)
result = await _forward_jsonrpc(agent_url, forward_body, extra_headers=agent_extra_headers)
if method == "agent/getAuthenticatedExtendedCard":
card: Final = result.get("result")
if isinstance(card, dict):
@ -988,16 +989,11 @@ async def invoke_agent_a2a(
"method": method,
"params": params,
}
sse_caller_headers: Final = _forwarding_headers(
user_api_key_dict=user_api_key_dict,
request_data=data,
agent_extra_headers=agent_extra_headers,
)
return await _forward_jsonrpc_sse(
agent_url,
forward_body,
request_id=request_id,
extra_headers=sse_caller_headers,
extra_headers=agent_extra_headers,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=data,

View file

@ -25,6 +25,11 @@ from litellm.proxy.common_request_processing import (
proxy_exception_from_http_exception,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.openai_error_payload import (
error_status_code,
openai_error_param,
openai_error_type,
)
from litellm.types.utils import TokenCountResponse
router: Final = APIRouter()
@ -243,9 +248,9 @@ async def anthropic_response(
return _anthropic_error_json_response(
ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
code=error_status_code(e, 500),
headers=headers,
),
request,

View file

@ -475,6 +475,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None
_NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({})
_TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True})
def _has_ptu_flat_cost(model: str, llm_router: "Router") -> bool:
@ -2858,7 +2859,9 @@ class TeamNotFoundError(HTTPException):
async def _get_team_db_check(
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
) -> "_PrismaTeamRow | None":
response = await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id})
response = await _team_table(TeamRepository(prisma_client)).find_unique(
where={"team_id": team_id}, include=_TEAM_GRANT_RELATIONS
)
if response is None and team_id_upsert:
from litellm.proxy.management_endpoints.team_endpoints import new_team
@ -3158,7 +3161,9 @@ async def get_team_object_by_alias(
# Query database by team_alias
try:
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(where={"team_alias": team_alias})
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(
where={"team_alias": team_alias}, include=_TEAM_GRANT_RELATIONS
)
if not teams:
raise HTTPException(

View file

@ -53,6 +53,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.team_grants import team_model_aliases
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
get_management_object_ttl,
@ -1595,7 +1596,7 @@ class JWTAuthManager:
model=requested_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=None,
team_model_aliases=team_model_aliases(team_object),
)
):
is_allowed = allowed_routes_check(
@ -2132,7 +2133,7 @@ class JWTAuthManager:
model=requested_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=None,
team_model_aliases=team_model_aliases(team_object),
)
except ProxyException:
continue

View file

@ -497,10 +497,22 @@ class RouteChecks:
def _placeholder_to_regex(match: re.Match) -> str:
placeholder: Final = match.group(0).strip("{}")
if placeholder.endswith(":path"):
# allow "/" in the placeholder value, but don't eat the route suffix after ":"
return r"[^:]+"
return r"[^/]+"
if not placeholder.endswith(":path"):
return r"[^/]+"
# A ":path" placeholder takes whatever the router's own path
# converter takes, slashes and colons alike, so an id spelled with
# either (or both) still matches the template it was mounted under.
#
# Unless the template puts a ":" literal of its own after the
# placeholder: the Google routes end in ":generateContent" and
# friends, and there the value has to stop before that suffix
# rather than swallow it and match a different verb.
#
# "[\s\S]" rather than ".", because "." stops at a newline and the
# path converter does not: a %0A anywhere in the value would leave
# the route unmatched here while still reaching the handler, which
# turns this gate into a bypass for the lists built on it.
return r"[^:]+" if ":" in match.string[match.end() :] else r"[\s\S]+"
pattern = re.sub(r"\{[^}]+\}", _placeholder_to_regex, pattern)
# Anchor the pattern to match the entire string

View file

@ -0,0 +1,122 @@
"""Project a team row (plus the caller's membership in it) onto the ``team_*`` fields of ``UserAPIKeyAuth``.
The virtual-key path gets these fields for free from the combined-view SQL join. Every other auth path
starts from a ``LiteLLM_TeamTable`` object instead and has to copy them over by hand, which is how JWT
callers kept losing grants (aliases, permissions, limits) one field at a time. Build the badge through
``team_grants`` and the two paths cannot drift.
"""
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Annotated, Final
from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError
from pydantic.main import IncEx
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy._types import (
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
Member,
)
_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
_JSON_COLUMNS: Final[Mapping[str, IncEx | bool]] = MappingProxyType(
{"metadata": True, "litellm_model_table": MappingProxyType({"model_aliases": True})}
)
def _decode_model_aliases(value: object) -> object:
"""``LiteLLM_ModelTable.model_aliases`` is typed ``str | dict``; writers hand Prisma ``json.dumps(...)``, so take both."""
if not isinstance(value, str):
return value
try:
return _MODEL_ALIASES_ADAPTER.validate_json(value)
except ValidationError:
return None
class TeamModelAliasTable(BaseModel):
model_config = ConfigDict(protected_namespaces=())
model_aliases: Annotated[Mapping[str, str] | None, BeforeValidator(_decode_model_aliases)] = None
class _TeamJsonColumns(BaseModel):
"""The two loosely typed columns on ``LiteLLM_TeamTable``, re-read with the shape the badge needs."""
metadata: Mapping[str, object] | None = None
litellm_model_table: TeamModelAliasTable | None = None
class TeamGrants(TypedDict, total=False):
"""Keyword arguments for ``UserAPIKeyAuth``. Empty when the caller has no team, so the model's own defaults apply."""
team_alias: ReadOnly[str | None]
team_tpm_limit: ReadOnly[int | None]
team_rpm_limit: ReadOnly[int | None]
team_max_budget: ReadOnly[float | None]
team_soft_budget: ReadOnly[float | None]
team_spend: ReadOnly[float | None]
team_models: ReadOnly[Sequence[str]]
team_blocked: ReadOnly[bool]
team_metadata: ReadOnly[Mapping[str, object] | None]
team_model_aliases: ReadOnly[Mapping[str, str] | None]
team_object_permission_id: ReadOnly[str | None]
team_object_permission: ReadOnly[LiteLLM_ObjectPermissionTable | None]
team_member: ReadOnly[Member | None]
team_member_spend: ReadOnly[float | None]
team_member_tpm_limit: ReadOnly[int | None]
team_member_rpm_limit: ReadOnly[int | None]
def _json_columns(team_object: LiteLLM_TeamTable) -> _TeamJsonColumns:
try:
return _TeamJsonColumns.model_validate(team_object.model_dump(include=_JSON_COLUMNS))
except ValidationError:
return _TeamJsonColumns()
def team_model_aliases(team_object: LiteLLM_TeamTable | None) -> Mapping[str, str] | None:
if team_object is None:
return None
alias_table: Final = _json_columns(team_object).litellm_model_table
return alias_table.model_aliases if alias_table is not None else None
def team_grants(
team_object: LiteLLM_TeamTable | None,
team_membership: LiteLLM_TeamMembership | None,
user_id: str | None,
) -> TeamGrants:
if team_object is None:
return TeamGrants()
json_columns: Final = _json_columns(team_object)
return TeamGrants(
team_alias=team_object.team_alias,
team_tpm_limit=team_object.tpm_limit,
team_rpm_limit=team_object.rpm_limit,
team_max_budget=team_object.max_budget,
team_soft_budget=team_object.soft_budget,
team_spend=team_object.spend,
team_models=tuple(team_object.models),
team_blocked=team_object.blocked,
team_metadata=json_columns.metadata,
team_model_aliases=(
json_columns.litellm_model_table.model_aliases if json_columns.litellm_model_table is not None else None
),
team_object_permission_id=team_object.object_permission_id,
team_object_permission=team_object.object_permission,
team_member=next(
(m for m in team_object.members_with_roles if user_id is not None and m.user_id == user_id),
None,
),
team_member_spend=team_membership.spend if team_membership is not None else None,
team_member_tpm_limit=(
team_membership.safe_get_team_member_tpm_limit() if team_membership is not None else None
),
team_member_rpm_limit=(
team_membership.safe_get_team_member_rpm_limit() if team_membership is not None else None
),
)

View file

@ -82,6 +82,7 @@ from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
from litellm.proxy.auth.resolvers import CredentialRef, Principal
from litellm.proxy.auth.resolvers.store import IdentityStore
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.team_grants import team_grants
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
from litellm.proxy.common_utils.http_parsing_utils import (
@ -1476,24 +1477,16 @@ async def _user_api_key_auth_builder(
user_id=user_id,
user_email=user_email,
team_id=team_id,
team_alias=(team_object.team_alias if team_object is not None else None),
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
team_models=(team_object.models if team_object is not None else []),
team_metadata=(team_object.metadata if team_object is not None else None),
org_id=org_id,
end_user_id=end_user_id,
parent_otel_span=parent_otel_span,
jwt_claims=jwt_claims,
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
valid_token = UserAPIKeyAuth(
api_key=None,
team_id=team_id,
team_alias=(team_object.team_alias if team_object is not None else None),
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
team_models=(team_object.models if team_object is not None else []),
user_role=(
LitellmUserRoles(user_object.user_role)
if user_object is not None and user_object.user_role is not None
@ -1507,17 +1500,8 @@ async def _user_api_key_auth_builder(
user_tpm_limit=(user_object.tpm_limit if user_object is not None else None),
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
team_member_rpm_limit=(
team_membership.safe_get_team_member_rpm_limit() if team_membership is not None else None
),
team_member_tpm_limit=(
team_membership.safe_get_team_member_tpm_limit() if team_membership is not None else None
),
team_metadata=(team_object.metadata if team_object is not None else None),
jwt_claims=jwt_claims,
)
valid_token.team_object_permission = (
team_object.object_permission if team_object is not None else None
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
@ -2836,6 +2820,43 @@ async def _authorize_authenticated_request(
return None
def _seed_request_destinations(user_api_key_dict: UserAPIKeyAuth, request: Request | None = None) -> None:
"""Anchor the OTLP destinations this key or team overrides its traces to.
Called inside the ``auth`` phase span so that span reaches the tenant's account
as well, and on the request task so the ``ContextVar`` is inherited by the logging
tasks that close the LLM span. Best-effort: trace routing must never fail auth.
``request`` carries the headers, so a backend this request disabled with
``x-litellm-disable-callbacks`` resolves to no destination.
Only destinations the published fan-out can build are anchored. Anchoring one is
what tells the operator's exporter to hold that backend's spans back under
``override``, so an unbuildable one would leave the span with nowhere to go.
The ``postgres`` spans under ``auth`` close before this runs, because they are the
reads that resolve the identity being read here. They never reach the tenant's
account, and they are never withheld from the operator's backend, whichever mode
is set.
"""
try:
from litellm.integrations.otel.logger import fan_out_provider
from litellm.integrations.otel.plumbing.context import set_request_destinations
from litellm.integrations.otel.plumbing.providers import deliverable_destinations
from litellm.proxy.litellm_pre_call_utils import (
resolve_tenant_otel_destinations,
)
set_request_destinations(
deliverable_destinations(
resolve_tenant_otel_destinations(user_api_key_dict, _safe_get_request_headers(request)),
fan_out_provider(),
)
)
except Exception as exc: # noqa: BLE001 # telemetry routing is best-effort and must never break authentication
verbose_proxy_logger.debug("OTel V2: tenant destination resolution failed: %s", exc)
@tracer.wrap()
async def user_api_key_auth(
request: Request,
@ -2883,6 +2904,7 @@ async def user_api_key_auth(
raise body_parse_exception
raise
user_api_key_auth_obj.budget_reservation = None
_seed_request_destinations(user_api_key_auth_obj, request)
# A body that never parsed is authenticated (so the trace carries identity
# and this ``auth`` span) but not authorized: there is no model to check it

View file

@ -54,6 +54,12 @@ from litellm.proxy.common_utils.callback_utils import (
get_logging_caching_headers,
get_remaining_tokens_and_requests_from_request_data,
)
from litellm.proxy.common_utils.openai_error_payload import (
attribute_of,
error_status_code,
openai_error_param,
openai_error_type,
)
from litellm.proxy.common_utils.sse_keepalive import (
SSE_COMMENT_PING_BYTES,
coerce_keepalive_interval,
@ -464,46 +470,6 @@ def _stream_usage_tracking_updates(
}
def _getattr_object(value: object, name: str, default: object = None) -> object:
return getattr(value, name, default)
_OPENAI_ERROR_TYPE_BY_STATUS: Final[Mapping[int, str]] = MappingProxyType(
{
status.HTTP_401_UNAUTHORIZED: "authentication_error",
status.HTTP_403_FORBIDDEN: "permission_error",
status.HTTP_429_TOO_MANY_REQUESTS: "rate_limit_error",
}
)
def _error_status_code(exc: object, default: int) -> int:
"""The HTTP status an exception carries, or ``default`` when it carries none."""
carried: Final = _getattr_object(exc, "status_code")
return carried if isinstance(carried, int) and not isinstance(carried, bool) else default
def _openai_error_type(exc: object, status_code: int) -> str:
"""OpenAI types ``error.type`` as a required string, so an exception carrying none
falls back to the type its status code stands for."""
carried: Final = _getattr_object(exc, "type")
if isinstance(carried, str):
return carried
mapped: Final = _OPENAI_ERROR_TYPE_BY_STATUS.get(status_code)
if mapped is not None:
return mapped
if status_code < status.HTTP_500_INTERNAL_SERVER_ERROR:
return "invalid_request_error"
return "internal_server_error"
def _openai_error_param(exc: object) -> str | None:
"""OpenAI types ``error.param`` as nullable, so an exception carrying none
serializes as JSON ``null``."""
carried: Final = _getattr_object(exc, "param")
return carried if isinstance(carried, str) else None
class _UpstreamHttpResponse(Protocol):
@property
def status_code(self) -> int: ...
@ -573,15 +539,15 @@ def serialize_http_exception_detail(
def proxy_exception_from_http_exception(exc: HTTPException, headers: dict[str, str]) -> ProxyException:
raw_detail: Final = _getattr_object(exc, "detail", str(exc))
raw_detail: Final = attribute_of(exc, "detail", str(exc))
message, structured_fields = serialize_http_exception_detail(raw_detail)
existing_fields: Final = getattr(exc, "provider_specific_fields", None) or {}
merged_fields: Final = {**existing_fields, **structured_fields} if structured_fields else (existing_fields or None)
error_status: Final = _error_status_code(exc, status.HTTP_400_BAD_REQUEST)
error_status: Final = error_status_code(exc, status.HTTP_400_BAD_REQUEST)
return ProxyException(
message=message,
type=_openai_error_type(exc, error_status),
param=_openai_error_param(exc),
type=openai_error_type(exc, error_status),
param=openai_error_param(exc),
code=error_status,
provider_specific_fields=merged_fields,
headers=headers,
@ -865,8 +831,8 @@ def sse_error_payload(exc: BaseException) -> tuple[int, Mapping[str, object]]:
are byte-identical.
"""
# Preserve status code from HTTPException (e.g. guardrail blocks)
error_status: Final = _error_status_code(exc, status.HTTP_500_INTERNAL_SERVER_ERROR)
raw_detail: Final = _getattr_object(exc, "detail", "Error processing stream start")
error_status: Final = error_status_code(exc, status.HTTP_500_INTERNAL_SERVER_ERROR)
raw_detail: Final = attribute_of(exc, "detail", "Error processing stream start")
message, structured_fields = serialize_http_exception_detail(raw_detail)
existing_fields: Final = getattr(exc, "provider_specific_fields", None) or {}
@ -874,8 +840,8 @@ def sse_error_payload(exc: BaseException) -> tuple[int, Mapping[str, object]]:
error_obj: Final = {
"message": message,
"type": _openai_error_type(exc, error_status),
"param": _openai_error_param(exc),
"type": openai_error_type(exc, error_status),
"param": openai_error_param(exc),
"code": str(error_status),
}
if not merged_fields:
@ -2815,10 +2781,10 @@ class ProxyBaseLLMRequestProcessing:
``ResponsesAPIResponse`` directly. Handle both shapes so the
container-ownership recording path can walk ``.output`` either way.
"""
completed: Final = _getattr_object(stream_response, "completed_response")
completed: Final = attribute_of(stream_response, "completed_response")
if completed is None:
return None
response_obj: Final = _getattr_object(completed, "response")
response_obj: Final = attribute_of(completed, "response")
if response_obj is not None:
return response_obj
return completed
@ -3475,7 +3441,7 @@ class ProxyBaseLLMRequestProcessing:
headers = getattr(e, "headers", None) or {}
if not headers:
# Try to get headers from e.response.headers (httpx.Response)
_response: Final = _getattr_object(e, "response")
_response: Final = attribute_of(e, "response")
if _response is not None:
_response_headers: Final = getattr(_response, "headers", None)
if _response_headers:
@ -3550,8 +3516,8 @@ class ProxyBaseLLMRequestProcessing:
_code = status.HTTP_500_INTERNAL_SERVER_ERROR
raise ProxyException(
message=redact_internal_details_from_client_message(getattr(e, "message", error_msg)),
type=_openai_error_type(e, _code),
param=_openai_error_param(e),
type=openai_error_type(e, _code),
param=openai_error_param(e),
openai_code=getattr(e, "code", None),
code=_code,
provider_specific_fields=getattr(e, "provider_specific_fields", None),
@ -3761,11 +3727,11 @@ class ProxyBaseLLMRequestProcessing:
if isinstance(e, HTTPException):
raise e
stream_error_status: Final = _error_status_code(e, status.HTTP_500_INTERNAL_SERVER_ERROR)
stream_error_status: Final = error_status_code(e, status.HTTP_500_INTERNAL_SERVER_ERROR)
proxy_exception: Final = ProxyException(
message=redact_internal_details_from_client_message(getattr(e, "message", str(e))),
type=_openai_error_type(e, stream_error_status),
param=_openai_error_param(e),
type=openai_error_type(e, stream_error_status),
param=openai_error_param(e),
code=stream_error_status,
)
stream_completed = True

View file

@ -44,6 +44,91 @@ def _langfuse_environment_error(callback_vars: Mapping[str, str]) -> str | None:
return None
# Which credential family a dynamic variable belongs to. The families are the
# integrations that share one account: every langfuse_* variable configures the
# same Langfuse project whether it rides the classic callback or the OTel one,
# and every dd_* variable configures the same Datadog account.
_VAR_FAMILIES: Final[Mapping[str, str]] = MappingProxyType(
{
"arize_": "Arize",
"dd_": "Datadog",
"gcs_": "GCS",
"humanloop_": "Humanloop",
"langfuse_": "Langfuse",
"langsmith_": "LangSmith",
"newrelic_": "New Relic",
"posthog_": "PostHog",
"wandb_": "Weights & Biases",
"weave_": "Weights & Biases",
}
)
def _family_of(var: str) -> str | None:
"""The credential family ``var`` configures, or ``None`` if it configures none.
``turn_off_message_logging`` and friends belong to no backend, so they carry
no credentials anyone could redirect.
"""
return next((family for prefix, family in _VAR_FAMILIES.items() if var.startswith(prefix)), None)
def cross_entry_family_error(
callback_vars: Mapping[str, str] | None,
stored_vars_by_entry: Sequence[Mapping[str, str]],
) -> str | None:
"""Reject an entry that changes what a family another entry holds resolves to.
Every stored entry's variables are flattened into one dict before a request
reads them, and the flattened dict is what the exporter authenticates and
addresses with. So an entry naming only a destination is enough to redirect
credentials that were written somewhere else: a host on a second entry pairs
with the key from the first, and the request carries that key to the new
host.
Two rules together keep the flattened dict out of the caller's hands. A
variable the family already configures has to keep the value it has, so
nothing already in use can be moved. A variable the family does not yet
configure may only carry a value the family already holds, which is what lets
the same credential go in under its other spelling (``langfuse_secret`` and
``langfuse_secret_key`` are one key) without anything here having to list the
spellings. Between them, no value the caller chose can enter the family, and
repeating the family as it stands is still allowed -- that is how one
integration gets registered for both the success and the failure event.
A team admin who does want to move a family deletes the entry holding it
first, which reveals nothing.
Only the writers this endpoint newly admits are held to this, because a proxy
admin already holds every credential the proxy has.
``stored_vars_by_entry`` has to arrive decrypted; the credential values are
encrypted at rest and ciphertext never equals the plaintext coming in.
"""
if not callback_vars:
return None
stored_by_var: Final = {
var: value for entry in stored_vars_by_entry for var, value in entry.items() if _family_of(var) is not None
}
family_values: Final = frozenset(
(family, value)
for entry in stored_vars_by_entry
for var, value in entry.items()
if (family := _family_of(var)) is not None
)
held_families: Final = frozenset(family for family, _ in family_values)
return next(
(
f"{family} is already configured by another callback entry on this team. "
f"Remove that entry before setting {var} here."
for var, value, family in ((v, callback_vars[v], _family_of(v)) for v in callback_vars)
if family in held_families
and (stored_by_var[var] != value if var in stored_by_var else (family, value) not in family_values)
),
None,
)
def logging_metadata_config_error(metadata: Mapping[str, object] | None) -> str | None:
"""Validate every ``logging`` entry of a team/key metadata payload."""
if not metadata:

View file

@ -1,10 +1,12 @@
import copy
import json
import os
from collections.abc import Callable, Iterable, Mapping
from dataclasses import dataclass
from itertools import accumulate
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias
from typing_extensions import assert_never
from typing_extensions import ReadOnly, TypedDict, assert_never
import litellm
from litellm import get_secret
@ -12,6 +14,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import (
CLIENT_OUTPUT_CEILING_METADATA_KEY,
CONSUMED_REQUEST_TAGS_METADATA_KEY,
MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH,
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
ROUTING_REQUEST_TAGS_METADATA_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
@ -28,6 +31,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
from litellm.proxy.types_utils.utils import get_instance_fn
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
StandardLoggingGuardrailInformation,
StandardLoggingPayload,
@ -52,6 +56,15 @@ reset_color_code: Final = "\033[0m"
TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY: Final = "_pillar_response_headers_trusted"
GUARDRAIL_SCAN_IDS_METADATA_KEY: Final = "guardrail_scan_ids"
GUARDRAIL_SCAN_METADATA_METADATA_KEY: Final = "guardrail_scan_metadata"
class GuardrailScanMetadata(TypedDict):
guardrail: ReadOnly[str | None]
stage: ReadOnly[str]
provider: ReadOnly[str]
scan_id: ReadOnly[str]
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
@ -450,6 +463,16 @@ def get_remaining_tokens_and_requests_from_request_data(data: dict) -> dict[str,
return headers
def _serialize_scan_metadata_header(entries: Iterable[object], *, max_length: int) -> str | None:
"""Compact JSON list of scan metadata entries, dropping trailing entries so the header fits in max_length."""
encoded: Final = tuple(json.dumps(entry, separators=(",", ":")) for entry in entries)
lengths: Final = tuple(accumulate(len(item) + 1 for item in encoded))
kept: Final = sum(1 for length in lengths if length + 1 <= max_length)
if kept == 0:
return None
return f"[{','.join(encoded[:kept])}]"
def get_logging_caching_headers(request_data: dict) -> dict | None:
_metadata: Final[dict] = {}
metadata_bucket: Final = request_data.get("metadata")
@ -468,6 +491,15 @@ def get_logging_caching_headers(request_data: dict) -> dict | None:
if scan_ids:
headers["x-litellm-guardrail-scan-id"] = ",".join(scan_ids)
scan_metadata: Final = _metadata.get(GUARDRAIL_SCAN_METADATA_METADATA_KEY)
scan_metadata_header: Final = (
_serialize_scan_metadata_header(scan_metadata, max_length=MAX_GUARDRAIL_SCAN_METADATA_HEADER_LENGTH)
if isinstance(scan_metadata, (list, tuple))
else None
)
if scan_metadata_header:
headers["x-litellm-guardrail-scan-metadata"] = scan_metadata_header
if "applied_policies" in _metadata:
headers["x-litellm-applied-policies"] = ",".join(_metadata["applied_policies"])
@ -501,6 +533,7 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset(
"applied_policies",
"applied_guardrails",
GUARDRAIL_SCAN_IDS_METADATA_KEY,
GUARDRAIL_SCAN_METADATA_METADATA_KEY,
"policy_sources",
"guardrails",
"guardrail_config",
@ -565,21 +598,40 @@ def add_guardrail_to_applied_guardrails_header(request_data: dict, guardrail_nam
_metadata["applied_guardrails"] = [guardrail_name]
def add_guardrail_scan_id(request_data: dict, scan_id: str | None) -> None:
def add_guardrail_scan_id(
request_data: dict[str, object],
scan_id: str | None,
*,
guardrail_name: str | None,
provider: str,
stage: GuardrailEventHooks,
) -> None:
"""
Record a provider scan id so it can be surfaced to the caller.
Record a provider scan id, keyed to the guardrail execution that produced it, so it can be surfaced to the caller.
Guardrails only return scan details to the client when they block, so allowed requests carry no
audit trail. Ids recorded here become the x-litellm-guardrail-scan-id response header.
audit trail. Ids recorded here become the x-litellm-guardrail-scan-id response header, and the
(guardrail, stage, provider, scan_id) entries become the x-litellm-guardrail-scan-metadata header.
"""
if not scan_id:
return
_, _metadata = get_or_create_metadata_bucket(request_data)
existing: Final = _metadata.get(GUARDRAIL_SCAN_IDS_METADATA_KEY)
scan_ids: Final = tuple(existing) if isinstance(existing, (list, tuple)) else ()
scan_ids: Final[tuple[object, ...]] = tuple(existing) if isinstance(existing, (list, tuple)) else ()
if scan_id not in scan_ids:
_metadata[GUARDRAIL_SCAN_IDS_METADATA_KEY] = (*scan_ids, scan_id)
entry: Final[GuardrailScanMetadata] = {
"guardrail": guardrail_name,
"stage": stage.value,
"provider": provider,
"scan_id": scan_id,
}
existing_entries: Final = _metadata.get(GUARDRAIL_SCAN_METADATA_METADATA_KEY)
entries: Final[tuple[object, ...]] = tuple(existing_entries) if isinstance(existing_entries, (list, tuple)) else ()
if entry not in entries:
_metadata[GUARDRAIL_SCAN_METADATA_METADATA_KEY] = (*entries, entry)
def add_policy_to_applied_policies_header(request_data: dict, policy_name: str | None):
"""

View file

@ -0,0 +1,52 @@
"""Shapes the ``error`` object the proxy answers with so it matches OpenAI's contract:
``type`` is a required string and ``param`` is nullable, neither of which the literal
string ``"None"`` satisfies."""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from fastapi import status
_OPENAI_ERROR_TYPE_BY_STATUS: Final[Mapping[int, str]] = MappingProxyType(
{
status.HTTP_401_UNAUTHORIZED: "authentication_error",
status.HTTP_403_FORBIDDEN: "permission_error",
status.HTTP_429_TOO_MANY_REQUESTS: "rate_limit_error",
}
)
def attribute_of(value: object, name: str, default: object = None) -> object:
return getattr(value, name, default)
def error_status_code(exc: object, default: int) -> int:
"""The HTTP status an exception carries as ``status_code`` or, the way ``ProxyException``
stores it, as a stringified ``code``; ``default`` when it carries neither."""
carried: Final = attribute_of(exc, "status_code")
if isinstance(carried, int) and not isinstance(carried, bool):
return carried
stringified: Final = attribute_of(exc, "code")
return int(stringified) if isinstance(stringified, str) and stringified.isdecimal() else default
def openai_error_type(exc: object, status_code: int) -> str:
"""OpenAI types ``error.type`` as a required string, so an exception carrying none
falls back to the type its status code stands for."""
carried: Final = attribute_of(exc, "type")
if isinstance(carried, str):
return carried
mapped: Final = _OPENAI_ERROR_TYPE_BY_STATUS.get(status_code)
if mapped is not None:
return mapped
if status_code < status.HTTP_500_INTERNAL_SERVER_ERROR:
return "invalid_request_error"
return "internal_server_error"
def openai_error_param(exc: object) -> str | None:
"""OpenAI types ``error.param`` as nullable, so an exception carrying none
serializes as JSON ``null``."""
carried: Final = attribute_of(exc, "param")
return carried if isinstance(carried, str) else None

View file

@ -1063,7 +1063,7 @@ class DBSpendUpdateWriter:
await enqueue_spend_logs(prisma_client, (payload,))
if payload.get("call_type") in RESPONSES_SESSION_CALL_TYPES:
request_spend_log_flush()
request_spend_log_flush(prisma_client)
else:
verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.")

View file

@ -34,12 +34,20 @@ writer's connection params (pool size, timeouts, pgbouncer mode) for the
ones the reader URL does not pin itself.
"""
import _ssl
import hashlib
import os
import socket
import ssl
import struct
import sys
import tempfile
import urllib.parse
from collections.abc import Mapping
from collections.abc import Callable, Mapping, Sequence
from functools import partial
from pathlib import Path
from types import MappingProxyType
from typing import Annotated, Final, cast
from typing import Annotated, Final, Protocol, TypeAlias, cast
from pydantic import AliasChoices, BeforeValidator, Field
from pydantic_settings import BaseSettings, SettingsConfigDict
@ -64,6 +72,7 @@ DISABLE_PREPARED_STATEMENTS_ENV_VAR: Final = "DATABASE_DISABLE_PREPARED_STATEMEN
DisablePreparedStatementsFlag = Annotated[
bool, BeforeValidator(partial(token_auth_flag_enabled, env_var=DISABLE_PREPARED_STATEMENTS_ENV_VAR))
]
MAX_IDLE_CONNECTION_LIFETIME_ENV_VAR: Final = "DATABASE_MAX_IDLE_CONNECTION_LIFETIME"
# schema.prisma pins `provider = "postgresql"`, so these are the only schemes
# Prisma can actually connect with.
@ -125,21 +134,100 @@ def add_missing_query_params(url: str, params: Mapping[str, str | int | float])
LIBPQ_VERIFY_SSLMODES: Final[frozenset[str]] = frozenset({"verify-ca", "verify-full"})
PEM_CERT_HEADER: Final = b"-----BEGIN CERTIFICATE-----"
PG_SSL_REQUEST: Final = struct.pack("!ii", 8, 80877103)
TLS_PROBE_TIMEOUT_SECONDS: Final = 10.0
RootCertResolver: TypeAlias = Callable[[str, str, int], str] # mutable-ok: Callable parameter syntax
def translate_libpq_ssl_params(url: str) -> str:
class _VerifiedChainSource(Protocol):
def get_verified_chain(self) -> Sequence[_ssl.Certificate] | None: ...
def _verified_chain_der(tls: ssl.SSLSocket) -> tuple[bytes, ...]:
if sys.version_info >= (3, 13):
return tuple(tls.get_verified_chain())
legacy: Final = cast( # cast-ok: the stub omits _sslobj, the C object has get_verified_chain since 3.10
"_VerifiedChainSource | None",
tls._sslobj, # pyright: ignore[reportAttributeAccessIssue, reportUnknownMemberType] # public API only from 3.13
)
chain: Final = () if legacy is None else legacy.get_verified_chain() or ()
return tuple(cert.public_bytes(_ssl.ENCODING_DER) for cert in chain)
def _server_trust_anchor(cafile: str, host: str, port: int) -> bytes | None:
try:
context: Final = ssl.create_default_context(cafile=cafile)
with socket.create_connection((host, port), timeout=TLS_PROBE_TIMEOUT_SECONDS) as raw:
raw.sendall(PG_SSL_REQUEST)
if raw.recv(1) != b"S":
return None
with context.wrap_socket(raw, server_hostname=host) as tls:
chain: Final = _verified_chain_der(tls)
except (OSError, ValueError):
return None
return chain[-1] if chain else None
def pin_bundle_root(cert_path: str, host: str, port: int) -> str:
"""Reduce a multi-root CA bundle to the one root that verifies ``host``.
Prisma's ``sslcert`` loads a single PEM certificate (native-tls
``Certificate::from_pem``), so pointing it at a bundle such as the AWS RDS
global bundle trusts only the first of its 108 regional roots and the
handshake fails with "unable to get local issuer certificate" for every
other region. A single-certificate file is returned as is. For a bundle,
one verifying handshake (chain and hostname, whole bundle as trust store)
identifies the trust anchor the server actually chains to, which is
written to a single-certificate file for Prisma. If the probe fails the
bundle path is returned unchanged, so Prisma fails closed exactly as
before rather than trusting anything the bundle would not.
"""
try:
if Path(cert_path).read_bytes().count(PEM_CERT_HEADER) < 2:
return cert_path
except OSError:
return cert_path
root: Final = _server_trust_anchor(cert_path, host, port)
if root is None:
return cert_path
pinned: Final = Path(tempfile.gettempdir()) / f"litellm-sslcert-{hashlib.sha256(root).hexdigest()[:16]}.pem"
return str(pinned) if _replace_file(pinned, ssl.DER_cert_to_PEM_cert(root)) else cert_path
def _replace_file(target: Path, content: str) -> bool:
"""Write ``content`` to a private temp file and rename it over ``target``, so
readers never see a partial file and a symlink planted at ``target`` is
replaced rather than followed."""
try:
fd, staged = tempfile.mkstemp(dir=target.parent, prefix=f"{target.name}.")
except OSError:
return False
try:
with os.fdopen(fd, "w") as handle:
handle.write(content)
os.replace(staged, target)
except OSError:
Path(staged).unlink(missing_ok=True)
return False
return True
def translate_libpq_ssl_params(url: str, resolve_root_cert: RootCertResolver = pin_bundle_root) -> str:
"""Rewrite libpq's certificate-verification params into Prisma's dialect.
Prisma's engine only knows ``sslmode=disable|prefer|require``, ``sslcert``
(the CA bundle) and ``sslaccept=strict``. It silently discards
(a single CA certificate) and ``sslaccept=strict``. It silently discards
``sslrootcert`` and downgrades ``sslmode=verify-ca`` / ``verify-full`` to
``prefer``, so a URL copied from libpq / RDS docs connects over TLS with no
certificate check at all. ``verify-ca`` and ``verify-full`` both become
``require`` (Prisma has no CA-only mode), ``sslrootcert`` becomes
``sslcert``, and either one turns on ``sslaccept=strict`` (chain and
hostname), matching libpq where a root cert makes ``require`` verify.
Prisma params the operator pinned themselves win; anything else is left
untouched.
``sslcert`` (run through ``resolve_root_cert``, which pins a multi-root
bundle down to the server's root), and either one turns on
``sslaccept=strict`` (chain and hostname), matching libpq where a root
cert makes ``require`` verify. Prisma params the operator pinned
themselves win; anything else is left untouched.
"""
parsed: Final = urllib.parse.urlsplit(url)
pairs: Final = tuple(urllib.parse.parse_qsl(parsed.query, keep_blank_values=True))
@ -153,7 +241,9 @@ def translate_libpq_ssl_params(url: str) -> str:
if key != "sslrootcert"
)
root_cert: Final = tuple(
("sslcert", value) for key, value in pairs if key == "sslrootcert" and "sslcert" not in keys
("sslcert", resolve_root_cert(value, parsed.hostname or "", parsed.port or int(DEFAULT_POSTGRES_PORT)))
for key, value in pairs
if key == "sslrootcert" and "sslcert" not in keys
)
strict: Final = () if "sslaccept" in keys else (("sslaccept", "strict"),)
query: Final = urllib.parse.urlencode(translated + root_cert + strict)
@ -217,6 +307,9 @@ class DatabaseURLSettings(BaseSettings):
disable_prepared_statements: DisablePreparedStatementsFlag = Field(
default=False, validation_alias=DISABLE_PREPARED_STATEMENTS_ENV_VAR
)
max_idle_connection_lifetime: int | None = Field(
default=None, validation_alias=MAX_IDLE_CONNECTION_LIFETIME_ENV_VAR
)
# Writer
database_url: str | None = Field(default=None, validation_alias="DATABASE_URL")
@ -453,6 +546,12 @@ class DatabaseURLSettings(BaseSettings):
if url:
os.environ[env_var] = add_missing_query_params(url, MappingProxyType({"pgbouncer": "true"}))
lifetime_params: Final = idle_lifetime_params(self.max_idle_connection_lifetime)
for env_var in ("DATABASE_URL", "DIRECT_URL"):
url = os.environ.get(env_var)
if url:
os.environ[env_var] = add_missing_query_params(url, lifetime_params)
# The reader inherits the writer's connection params (pool size, timeouts,
# pgbouncer mode). Without this the reader pool ignores the configured cap
# and falls back to Prisma's `num_physical_cpus * 2 + 1` default.

View file

@ -10,13 +10,28 @@ strings rather than passing the raw path through. Nothing a caller sends can
add a key, so the fold and the table it commits to are bounded by (days x
routes) however much traffic arrives, and the response path carries no
unbounded queue that would block once full.
A flush commits its whole snapshot as one multi-row ``INSERT ... ON CONFLICT DO
UPDATE`` rather than one upsert per key, so a worker costs the primary one
statement per interval however many routes it served. With
``use_redis_transaction_buffer`` on, workers instead push their snapshot to a
Redis list and one lock-holding pod folds every entry and writes the table, so
the deployment as a whole costs the primary one statement per interval.
"""
from dataclasses import asdict
import json
from collections.abc import AsyncIterator, Iterable
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Final
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.caching import RedisCache
from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_GATEWAY_REQUESTS_BUFFER_KEY
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.middleware.billable_request_metrics_middleware import BillableCategory
from litellm.types.proxy.gateway_requests import (
GatewayRequestCounts,
@ -28,6 +43,15 @@ if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
_EMPTY: Final = GatewayRequestCounts(successful_requests=0, failed_requests=0)
_TABLE: Final = '"LiteLLM_DailyGatewayRequests"'
_COLUMNS_PER_ROW: Final = 5
_UTC_NOW: Final = "(NOW() AT TIME ZONE 'UTC')"
GATEWAY_REQUESTS_JOB_NAME: Final = "update_gateway_requests_job"
_BufferedRows: TypeAlias = tuple[tuple[str, str, str, int, int], ...]
_BUFFERED_ROWS: Final = TypeAdapter(_BufferedRows)
_BUFFERED_ENTRIES: Final = TypeAdapter(tuple[str | bytes, ...])
_NO_COUNTS: Final[GatewayRequestSnapshot] = MappingProxyType({})
def _utc_date() -> str:
@ -59,20 +83,54 @@ class GatewayRequestAccumulator:
route) however long the database is unreachable.
This buys at-least-once, not exactly-once, and the cost is worth stating.
The batch commits inside its context manager's ``__aexit__``, so a failure
raised after the transaction committed (a connection dropped while reading
the acknowledgement) restores counts that are already persisted, and the
next flush increments them a second time. Exactly-once would need a dedup
key the upserts could ignore on replay. For a traffic-volume metric a rare
The statement commits on the server before its acknowledgement is read, so
a failure raised after the commit (a connection dropped while reading the
acknowledgement) restores counts that are already persisted, and the next
flush increments them a second time. Exactly-once would need a dedup key
the upsert could ignore on replay. For a traffic-volume metric a rare
overcount on a dropped acknowledgement beats losing a whole interval to
every database blip, so the trade is deliberate.
"""
for key, counts in snapshot.items():
existing = self._counts.get(key, _EMPTY)
self._counts[key] = GatewayRequestCounts(
successful_requests=existing.successful_requests + counts.successful_requests,
failed_requests=existing.failed_requests + counts.failed_requests,
)
self._counts = dict(fold_counts(chain(self._counts.items(), snapshot.items()))) # mutable-ok: fold replaced
def fold_counts(items: Iterable[tuple[GatewayRequestKey, GatewayRequestCounts]]) -> GatewayRequestSnapshot:
"""Sum counts key-wise; the result stays bounded by (date x category x route)."""
folded: Final[dict[GatewayRequestKey, GatewayRequestCounts]] = {} # mutable-ok: local fold returned once
for key, counts in items:
existing = folded.get(key, _EMPTY)
folded[key] = GatewayRequestCounts(
successful_requests=existing.successful_requests + counts.successful_requests,
failed_requests=existing.failed_requests + counts.failed_requests,
)
return folded
def build_gateway_requests_upsert(snapshot: GatewayRequestSnapshot) -> tuple[str, tuple[str | int, ...]]:
"""
One ``INSERT ... ON CONFLICT DO UPDATE`` that increments every (date, category,
route) in the snapshot. Rows are ordered by the conflict key so concurrent
writers lock rows in the same order and cannot deadlock.
"""
ordered: Final = sorted(snapshot.items(), key=lambda item: (item[0].date, item[0].category, item[0].route))
rows: Final = ", ".join(
f"(${base + 1}::text, ${base + 2}::text, ${base + 3}::text, ${base + 4}::bigint, ${base + 5}::bigint, {_UTC_NOW})"
for base in range(0, len(ordered) * _COLUMNS_PER_ROW, _COLUMNS_PER_ROW)
)
sql: Final = (
f'INSERT INTO {_TABLE} ("date", "category", "route", "successful_requests", "failed_requests", "updated_at")\n'
f"VALUES {rows}\n"
'ON CONFLICT ("date", "category", "route") DO UPDATE SET\n'
f' "successful_requests" = {_TABLE}."successful_requests" + EXCLUDED."successful_requests",\n'
f' "failed_requests" = {_TABLE}."failed_requests" + EXCLUDED."failed_requests",\n'
f' "updated_at" = {_UTC_NOW}'
)
params: Final[tuple[str | int, ...]] = tuple(
value
for key, counts in ordered
for value in (key.date, key.category, key.route, counts.successful_requests, counts.failed_requests)
)
return sql, params
async def commit_gateway_requests_to_db(
@ -80,50 +138,130 @@ async def commit_gateway_requests_to_db(
prisma_client: "PrismaClient",
snapshot: GatewayRequestSnapshot,
) -> None:
"""Upsert one incrementing row per (date, category, route)."""
"""Increment every (date, category, route) in the snapshot with a single statement."""
if not snapshot:
return
ordered: Final = sorted(snapshot.items(), key=lambda item: (item[0].date, item[0].category, item[0].route))
sql, params = build_gateway_requests_upsert(snapshot)
await prisma_client.db.execute_raw(sql, *params) # pyright: ignore[reportAny] # untyped prisma client
# pyright: ignore[reportAny] on both lines -- prisma's generated client is untyped,
# so .db and every table action off it resolve to Any at this boundary. The dict
# literals below are the shape prisma's generated inputs require.
async with prisma_client.db.batch_() as batcher: # pyright: ignore[reportAny] # untyped prisma client
for key, counts in ordered:
columns = asdict(key)
batcher.litellm_dailygatewayrequests.upsert( # pyright: ignore[reportAny] # untyped prisma client
where={"date_category_route": columns}, # mutable-ok: prisma input is dict-shaped
data={ # mutable-ok: prisma input is dict-shaped
"create": { # mutable-ok: prisma input is dict-shaped
**columns,
"successful_requests": counts.successful_requests,
"failed_requests": counts.failed_requests,
},
"update": { # mutable-ok: prisma input is dict-shaped
"successful_requests": {"increment": counts.successful_requests}, # mutable-ok: as above
"failed_requests": {"increment": counts.failed_requests}, # mutable-ok: as above
},
},
verbose_proxy_logger.debug(
"Gateway request tracking - committed %d aggregated rows in one statement", len(snapshot)
)
class GatewayRequestRedisBuffer:
"""
Folds every worker's snapshot through one Redis list so a single pod per
interval writes the table, mirroring the spend writer's transaction buffer.
Each entry is one worker's snapshot as JSON rows; the lock holder pops them,
sums them, and commits one statement. A commit failure pushes the summed
rows back so the next holder retries, keeping the at-least-once guarantee.
If that push fails too, the rows go back to the holder's own accumulator so
they ride along with its next flush instead of vanishing with the pop.
"""
def __init__(self, *, redis_cache: RedisCache, pod_lock_manager: PodLockManager) -> None:
self._redis_cache: Final = redis_cache
self._pod_lock_manager: Final = pod_lock_manager
async def push(self, snapshot: GatewayRequestSnapshot) -> None:
if not snapshot:
return
rows: Final[_BufferedRows] = tuple(
(key.date, key.category, key.route, counts.successful_requests, counts.failed_requests)
for key, counts in snapshot.items()
)
await self._redis_cache.async_rpush(key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, values=(json.dumps(rows),))
async def _pop_batch(self) -> tuple[str | bytes, ...]:
popped: Final[object] = await self._redis_cache.async_lpop( # pyright: ignore[reportAny] # redis returns Any
key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT
)
if not popped:
return ()
return _BUFFERED_ENTRIES.validate_python(popped if isinstance(popped, list) else (popped,))
async def _pop_all(self) -> AsyncIterator[str | bytes]:
while True:
batch = await self._pop_batch()
for entry in batch:
yield entry
if len(batch) < MAX_REDIS_BUFFER_DEQUEUE_COUNT:
return
async def pop(self) -> GatewayRequestSnapshot:
entries: Final = tuple([entry async for entry in self._pop_all()])
return fold_counts(
(
GatewayRequestKey(date=date, category=category, route=route),
GatewayRequestCounts(successful_requests=succeeded, failed_requests=failed),
)
for entry in entries
for date, category, route, succeeded, failed in _BUFFERED_ROWS.validate_json(entry)
)
verbose_proxy_logger.debug("Gateway request tracking - committed %d aggregated rows", len(ordered))
async def commit_if_leader(self, prisma_client: "PrismaClient") -> GatewayRequestSnapshot:
"""
Drain the list and write it as one statement, but only on the pod holding the job lock.
The lock is a lease, never released: the holder re-enters it on every flush and
keeps committing alone until the TTL lapses, so the primary sees one statement
per flush interval deployment-wide instead of one per worker.
Returns the popped rows that could be neither committed nor re-queued, for the
caller to keep in memory. Empty on success.
"""
if not await self._pod_lock_manager.acquire_lock(cronjob_id=GATEWAY_REQUESTS_JOB_NAME):
return _NO_COUNTS
buffered: Final = await self.pop()
try:
await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=buffered)
except Exception: # noqa: BLE001 -- a failed commit must not stop the scheduler
verbose_proxy_logger.warning(
"Gateway request tracking - failed to commit %d buffered rows, re-queuing to Redis for the next flush",
len(buffered),
exc_info=True,
)
return await self._requeue(buffered)
return _NO_COUNTS
async def _requeue(self, snapshot: GatewayRequestSnapshot) -> GatewayRequestSnapshot:
try:
await self.push(snapshot)
except Exception: # noqa: BLE001 -- the rows go back to the caller's accumulator instead
verbose_proxy_logger.warning(
"Gateway request tracking - Redis re-queue failed, keeping %d rows in memory for the next flush",
len(snapshot),
exc_info=True,
)
return snapshot
return _NO_COUNTS
async def flush_gateway_requests(
prisma_client: "PrismaClient",
accumulator: GatewayRequestAccumulator,
redis_buffer: GatewayRequestRedisBuffer | None = None,
) -> None:
"""
Scheduler entrypoint. Never raises: a metering failure must not kill the job.
With ``redis_buffer`` the snapshot goes to Redis and only the lease holder
writes to Postgres. Shutdown passes no buffer so a departing worker writes its
own counts directly instead of parking them behind a lease it may not hold.
``CancelledError`` is deliberately not caught, so a flush cancelled during
shutdown drops its snapshot rather than restoring counts onto an accumulator
the process is about to discard.
"""
snapshot: Final = accumulator.drain()
try:
await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=snapshot)
if redis_buffer is None:
await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=snapshot)
else:
await redis_buffer.push(snapshot)
except Exception: # noqa: BLE001 -- a failed flush must not stop the scheduler
accumulator.restore(snapshot)
verbose_proxy_logger.warning(
@ -131,3 +269,13 @@ async def flush_gateway_requests(
len(snapshot),
exc_info=True,
)
return
if redis_buffer is None:
return
try:
accumulator.restore(await redis_buffer.commit_if_leader(prisma_client))
except Exception: # noqa: BLE001 -- entries still in Redis are drained by the next flush
verbose_proxy_logger.warning(
"Gateway request tracking - leader drain failed, buffered rows stay in Redis for the next flush",
exc_info=True,
)

View file

@ -16,6 +16,7 @@ from datetime import datetime, timedelta
from typing import Any, Final, Protocol
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.db_url_settings import add_missing_query_params, connection_params_from_url
from litellm.proxy.db.token_auth import (
DEFAULT_POSTGRES_PORT,
DatabaseTokenAuth,
@ -438,7 +439,10 @@ class PrismaWrapper:
return None
endpoint: Final = self._iam_endpoint if self._iam_endpoint is not None else self._endpoint_from_env()
db_url: Final = endpoint.build_url(mint_database_token(auth, endpoint))
db_url: Final = add_missing_query_params(
endpoint.build_url(mint_database_token(auth, endpoint)),
connection_params_from_url(os.environ.get(self._db_url_env_var, "")),
)
os.environ[self._db_url_env_var] = db_url
return db_url

View file

@ -17,7 +17,8 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.proxy.common_utils.callback_utils import add_guardrail_scan_id
from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailIntegrations
from litellm.types.utils import (
GenericGuardrailAPIInputs,
GuardrailStatus,
@ -218,6 +219,13 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
metadata: Final = request_data.get("metadata") or {}
request_data["metadata"] = metadata
metadata["_openai_moderation_response"] = moderation_response.model_dump()
add_guardrail_scan_id(
request_data=request_data,
scan_id=moderation_response.id,
guardrail_name=self.guardrail_name,
provider=SupportedGuardrailIntegrations.OPENAI_MODERATION.value,
stage=GuardrailEventHooks.post_call if input_type == "response" else GuardrailEventHooks.pre_call,
)
# Check if content is flagged and raise exception if needed
self._check_moderation_result(moderation_response)

View file

@ -721,10 +721,18 @@ class PanwPrismaAirsHandler(CustomGuardrail):
}
}
def _record_scan_id(self, request_data: dict[str, object], scan_result: Mapping[str, object]) -> None:
def _record_scan_id(
self, request_data: dict[str, object], scan_result: Mapping[str, object], stage: GuardrailEventHooks
) -> None:
"""Surface the AIRS scan id on the response, so allowed calls are auditable too."""
scan_id: Final = scan_result.get("scan_id")
add_guardrail_scan_id(request_data=request_data, scan_id=str(scan_id) if scan_id else None)
add_guardrail_scan_id(
request_data=request_data,
scan_id=str(scan_id) if scan_id else None,
guardrail_name=self.guardrail_name,
provider=self._PROVIDER_NAME,
stage=stage,
)
def _handle_api_error_with_logging(
self,
@ -948,7 +956,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
event_type=GuardrailEventHooks.post_call,
)
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
self._record_scan_id(request_data, scan_result)
self._record_scan_id(request_data, scan_result, GuardrailEventHooks.post_call)
def _check_and_mark_scanned(self, data: dict, scan_type: str) -> bool:
"""
@ -1078,7 +1086,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.pre_call,
)
self._record_scan_id(data, scan_result)
self._record_scan_id(data, scan_result, GuardrailEventHooks.pre_call)
action: Final = scan_result.get("action", "block")
category: Final = scan_result.get("category", "unknown")
@ -1199,7 +1207,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.post_call,
)
self._record_scan_id(data, scan_result)
self._record_scan_id(data, scan_result, GuardrailEventHooks.post_call)
action: Final = scan_result.get("action", "block")
category: Final = scan_result.get("category", "unknown")
@ -1401,7 +1409,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
duration=(end_time - start_time).total_seconds(),
event_type=GuardrailEventHooks.post_call,
)
self._record_scan_id(request_data, scan_result)
self._record_scan_id(request_data, scan_result, GuardrailEventHooks.post_call)
# Add guardrail to applied guardrails header for observability
add_guardrail_to_applied_guardrails_header(
@ -1475,7 +1483,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
)
continue
self._record_scan_id(request_data, scan_result)
self._record_scan_id(
request_data,
scan_result,
GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call,
)
action = scan_result.get("action", "block")
masked_args = self._masked_tool_call_arguments(
@ -1829,7 +1841,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
new_texts.append(text)
continue
self._record_scan_id(request_data, scan_result)
self._record_scan_id(
request_data,
scan_result,
GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call,
)
action = scan_result.get("action", "block")
masked_text = self._get_masked_text(scan_result, is_response=is_response)
@ -1901,7 +1917,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
)
# If we reach here, fallback_on_error="allow"
else:
self._record_scan_id(request_data, mcp_scan_result)
self._record_scan_id(request_data, mcp_scan_result, GuardrailEventHooks.pre_call)
action = mcp_scan_result.get("action", "block")
masked_text = self._get_masked_text(mcp_scan_result, is_response=False)
if action == "allow":

View file

@ -20,6 +20,11 @@ from litellm.proxy.common_utils.http_parsing_utils import (
coerce_numeric_form_fields,
numeric_form_fields,
)
from litellm.proxy.common_utils.openai_error_payload import (
error_status_code,
openai_error_param,
openai_error_type,
)
from litellm.proxy.route_llm_request import route_request
from litellm.types.images.main import ImageEditRequestParams
from litellm.types.llms.openai import ChatCompletionUserMessage
@ -200,18 +205,18 @@ async def image_generation(
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
param=openai_error_param(e),
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
)
else:
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
openai_code=getattr(e, "code", None),
code=getattr(e, "status_code", 500),
code=error_status_code(e, 500),
)

View file

@ -10,6 +10,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, cast
from fastapi import HTTPException, Request
from pydantic import TypeAdapter
from pydantic import ValidationError as PydanticValidationError
from starlette.datastructures import Headers
@ -28,6 +29,7 @@ from litellm.constants import (
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
SESSION_ID_GENERATED_METADATA_KEY,
SESSION_ID_OMITTED_METADATA_KEY,
X_LITELLM_DISABLE_CALLBACKS,
)
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
@ -159,6 +161,7 @@ from litellm.types.utils import (
CustomPricingLiteLLMParams,
LlmProviders,
ProviderSpecificHeader,
StandardCallbackDynamicParams,
StandardLoggingUserAPIKeyMetadata,
SupportedCacheControls,
)
@ -172,6 +175,7 @@ _ENABLE_TEAM_STALE_ALIAS_BYPASS: bool | None = None
if TYPE_CHECKING:
from litellm.integrations.otel.model.destination import OtelDestination
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
from litellm.types.proxy.policy_engine import Policy, PolicyMatchContext
@ -231,6 +235,7 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
"applied_policies",
"policy_sources",
"guardrail_scan_ids",
"guardrail_scan_metadata",
"routing_decision",
GATEWAY_INJECTED_CACHE_METADATA_KEY,
"pillar_response_headers",
@ -287,6 +292,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS: Final = (
"applied_policies",
"policy_sources",
"guardrail_scan_ids",
"guardrail_scan_metadata",
"routing_decision",
GATEWAY_INJECTED_CACHE_METADATA_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
@ -975,6 +981,145 @@ def _get_dynamic_logging_metadata(
return callback_settings_obj
_TENANT_OTEL_PARAMS: Final = TypeAdapter(StandardCallbackDynamicParams)
def _tenant_otel_params(callback_vars: Mapping[str, str]) -> StandardCallbackDynamicParams:
try:
return _TENANT_OTEL_PARAMS.validate_python(callback_vars)
except PydanticValidationError:
return StandardCallbackDynamicParams()
_NO_REQUEST_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
def _dynamically_disabled_backends(
user_api_key_dict: UserAPIKeyAuth,
request_headers: Mapping[str, str] | None,
) -> frozenset[str]:
"""The callbacks this request turned off, read the way dispatch reads them.
Same sources, precedence, and premium gate ``EnterpriseCallbackControls`` applies
before it skips a callback: the ``x-litellm-disable-callbacks`` header wins over the
key's stored list, team settings are not a source, and a non-premium proxy honours
neither. A destination has to agree with that decision, or a backend the key turned
off would still be exported to, now through the fan-out instead of the callback.
"""
from litellm.proxy.proxy_server import premium_user
if litellm.allow_dynamic_callback_disabling is not True or not premium_user:
return frozenset()
header: Final = (request_headers if request_headers is not None else _NO_REQUEST_HEADERS).get(
X_LITELLM_DISABLE_CALLBACKS
)
if header is not None:
return frozenset(name.strip().lower() for name in header.split(","))
metadata: Final = user_api_key_dict.metadata
disabled: Final = metadata.get("litellm_disabled_callbacks") if metadata else None
if not isinstance(disabled, list):
return frozenset()
return frozenset(name.lower() for name in disabled if isinstance(name, str))
def resolve_tenant_otel_destinations(
user_api_key_dict: UserAPIKeyAuth,
request_headers: Mapping[str, str] | None = None,
) -> "tuple[OtelDestination, ...]":
"""The OTLP destinations this request's key or team config overrides its traces to.
Key settings win over team settings outright, the same precedence
``_get_dynamic_logging_metadata`` applies, so one caller never exports the same
backend to two accounts. An empty key-level list counts as configured, since that
is what disabling a key's callbacks writes. Returns empty when OTEL V2 is off, when
neither level named a destination-capable backend, or when the config is
incomplete, and the request then keeps the operator's own exporters.
Two entries naming the same backend merge their ``callback_vars`` last-wins, the
way ``convert_key_logging_metadata_to_callback`` merges them, so the destination
and the per-request tracer routing cannot read one config two ways.
A ``failure``-only entry is skipped: a destination is resolved during auth, before
the request has an outcome, so honouring the filter would mean holding every span
back until the call finishes. Those entries keep today's behaviour instead, where
the tenant's credentials reach the backend through per-request tracer routing and
the operator's exporter is left alone. Its ``callback_vars`` still take part in the
merge for a backend another entry made eligible, so the destination carries the
same credentials the runtime parser resolves for that request.
A backend the request disabled dynamically, through the key's
``litellm_disabled_callbacks`` or the ``x-litellm-disable-callbacks`` header in
``request_headers``, resolves to no destination, so the fan-out never carries the
request tree to that account and the operator's exporter is never suppressed for
it. That leaves the request exactly where it stood before destinations existed:
the OTel V2 logger itself is not on the disable list's class registry, so its own
span still routes to the tenant's credentials the way it did then.
"""
from litellm.integrations.otel.model.config import is_otel_v2_enabled
from litellm.integrations.otel.presets.destinations import destination_for
if not is_otel_v2_enabled():
return ()
key_entries: Final = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict)
entries: Final = (
key_entries
if key_entries is not None
else KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict)
)
if not entries:
return ()
disabled: Final = _dynamically_disabled_backends(user_api_key_dict, request_headers)
callbacks: Final = tuple(
callback
for item in entries
if (callback := _get_validated_callback_metadata(item=item, source="otel-destination")) is not None
if callback.callback_name.lower() not in disabled
)
return tuple(
destination
for name in dict.fromkeys(
callback.callback_name for callback in callbacks if callback.callback_type != "failure"
)
if (
destination := destination_for(
name,
_tenant_otel_params(
MappingProxyType(
{
var: value
for callback in callbacks
if callback.callback_name == name
for var, value in callback.callback_vars.items()
}
)
),
_tenant_service_name(user_api_key_dict),
)
)
is not None
)
def _tenant_service_name(user_api_key_dict: UserAPIKeyAuth) -> str | None:
"""The ``service.name`` this key or team configured, the key winning over its team.
Same fields and same precedence the request-metadata build applies, read straight
off the auth object because destinations resolve during auth, before that metadata
is assembled.
"""
sources: Final = (user_api_key_dict.metadata, user_api_key_dict.team_metadata)
return next(
(
stripped
for source in sources
if source
for field in OTEL_SERVICE_NAME_METADATA_KEYS
if isinstance(value := source.get(field), str) and (stripped := value.strip())
),
None,
)
def clean_headers(
headers: Headers,
litellm_key_header_name: str | None = None,

View file

@ -14,6 +14,7 @@ from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.spend_tracking.key_metadata_recovery import (
attach_user_emails,
recover_double_hashed_key_metadata,
recover_key_metadata_from_spend_logs,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
from litellm.proxy.utils import PrismaClient
@ -433,9 +434,29 @@ def update_breakdown_metrics(
return breakdown
def _spend_logs_window(dates: AbstractSet[str | None]) -> tuple[datetime, datetime] | None:
parsed: Final = sorted(day for day in (_parse_spend_date(raw) for raw in dates) if day is not None)
if not parsed:
return None
return (parsed[0] - timedelta(days=1), parsed[-1] + timedelta(days=2))
def _parse_spend_date(raw: str | None) -> datetime | None:
if not isinstance(raw, str):
return None
try:
return datetime.fromisoformat(raw)
except ValueError:
return None
_EMPTY_KEY_METADATA: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType({})
async def get_api_key_metadata(
prisma_client: PrismaClient,
api_keys: AbstractSet[str],
spend_logs_window: tuple[datetime, datetime] | None = None,
) -> Mapping[str, _KeyMetadataDict]:
"""Get api key metadata, falling back to deleted keys table for keys not found in active table.
@ -481,11 +502,17 @@ async def get_api_key_metadata(
)
still_missing: Final = api_keys - frozenset(result)
combined: Final = (
result
if not still_missing
else MappingProxyType({**result, **(await recover_double_hashed_key_metadata(prisma_client, still_missing))})
from_reverse_hash: Final = (
await recover_double_hashed_key_metadata(prisma_client, still_missing) if still_missing else _EMPTY_KEY_METADATA
)
after_token_recovery: Final = MappingProxyType({**result, **from_reverse_hash})
unresolved: Final = api_keys - frozenset(after_token_recovery)
from_spend_logs: Final = (
await recover_key_metadata_from_spend_logs(prisma_client, unresolved, spend_logs_window)
if unresolved and spend_logs_window is not None
else _EMPTY_KEY_METADATA
)
combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs})
return await attach_user_emails(prisma_client, combined)
@ -898,7 +925,9 @@ async def _aggregate_spend_records(
api_key_metadata: dict[str, _KeyMetadataDict] = {}
if api_keys:
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
api_key_metadata = await get_api_key_metadata(
prisma_client, api_keys, _spend_logs_window(frozenset(record.date for record in records))
)
return await asyncio.to_thread(
_aggregate_spend_records_sync,
@ -1094,7 +1123,9 @@ async def _aggregate_grouping_sets_records(
api_key_metadata: dict[str, _KeyMetadataDict] = {}
if api_keys:
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
api_key_metadata = await get_api_key_metadata(
prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records))
)
return await asyncio.to_thread(
_aggregate_grouping_sets_records_sync,
@ -1357,7 +1388,9 @@ async def get_daily_activity_aggregated(
r.api_key for r in entity_records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY
)
entity_key_metadata: Final = (
await get_api_key_metadata(prisma_client, entity_api_keys)
await get_api_key_metadata(
prisma_client, entity_api_keys, _spend_logs_window(frozenset(r.date for r in entity_records))
)
if entity_api_keys
else {} # mutable-ok: matches the helper's dict return
)

View file

@ -18,6 +18,7 @@ from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
import litellm
from litellm._internal_context import current_billing_time, pinned_billing_time
from litellm._logging import verbose_proxy_logger
from litellm.cost_calculator import completion_cost
from litellm.proxy._types import (
@ -27,7 +28,15 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.utils import CostPerToken, LlmProvidersSet, ModelInfo
from litellm.types.utils import (
CostBreakdown,
CostPerToken,
LlmProvidersSet,
ModelInfo,
ModelResponse,
PromptTokensDetailsWrapper,
Usage,
)
router: Final = APIRouter()
@ -46,13 +55,15 @@ def _configured_price(key: str, sources: tuple[Mapping[str, object], ...]) -> fl
def _extract_custom_pricing(
litellm_params: Mapping[str, object], model_info: Mapping[str, object]
litellm_params: Mapping[str, object], model_info: Mapping[str, object], builtin: ModelInfo | None
) -> CostPerToken | None:
"""
Pull per-token pricing configured on a deployment so on-prem / self-hosted
models (absent from the public cost map) still estimate a real cost.
Pricing may live on ``litellm_params`` or ``model_info``; ``litellm_params``
wins, matching the router's cost-map registration precedence.
wins, matching the router's cost-map registration precedence. Cache rates the
deployment leaves unset come from the backend model's built-in entry, then its
own input rate, again matching what the router registers for live billing.
"""
sources: Final = (litellm_params, model_info)
input_price: Final = _configured_price("input_cost_per_token", sources)
@ -61,15 +72,21 @@ def _extract_custom_pricing(
if input_price is None and output_price is None:
return None
input_rate: Final = input_price or 0.0
cache_sources: Final = sources if builtin is None else (*sources, builtin)
cache_read_price: Final = _configured_price("cache_read_input_token_cost", cache_sources)
cache_creation_price: Final = _configured_price("cache_creation_input_token_cost", cache_sources)
return CostPerToken(
input_cost_per_token=input_price or 0.0,
input_cost_per_token=input_rate,
output_cost_per_token=output_price or 0.0,
cache_read_input_token_cost=input_rate if cache_read_price is None else cache_read_price,
cache_creation_input_token_cost=input_rate if cache_creation_price is None else cache_creation_price,
)
def _lookup_model_info(model: str) -> ModelInfo | None:
def _lookup_model_info(model: str, custom_llm_provider: str | None = None) -> ModelInfo | None:
try:
return litellm.get_model_info(model=model)
return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception:
return None
@ -98,17 +115,14 @@ def _resolve_model_for_cost_lookup(model: str) -> ResolvedCostModel:
model_info: Final = first_deployment.get("model_info", {})
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
provider: Final = str(custom_llm_provider) if custom_llm_provider is not None else None
custom_cost_per_token: Final = _extract_custom_pricing(litellm_params, model_info)
# Check base_model first (needed for Azure custom deployment names)
# base_model wins (needed for Azure custom deployment names)
base_model: Final = model_info.get("base_model") or litellm_params.get("base_model")
if base_model:
verbose_proxy_logger.debug("Resolved model '%s' to base_model '%s' from router", model, base_model)
return ResolvedCostModel(str(base_model), provider, custom_cost_per_token)
resolved_model: Final = litellm_params.get("model")
resolved_model: Final = base_model or litellm_params.get("model")
if resolved_model:
verbose_proxy_logger.debug("Resolved model '%s' to '%s' from router", model, resolved_model)
custom_cost_per_token: Final = _extract_custom_pricing(
litellm_params, model_info, _lookup_model_info(str(resolved_model), provider)
)
return ResolvedCostModel(str(resolved_model), provider, custom_cost_per_token)
except Exception as e:
verbose_proxy_logger.debug("Could not resolve model '%s' from router: %s", model, e)
@ -117,19 +131,59 @@ def _resolve_model_for_cost_lookup(model: str) -> ResolvedCostModel:
return ResolvedCostModel(model, None, None)
def _calculate_period_costs(num_requests, cost_per_request, input_cost, output_cost, margin_cost):
"""
Calculate costs for a given number of requests.
@dataclass(frozen=True, slots=True)
class CostLines:
"""Cost of one request split the way the spend logs split it: the cache lines are
shares of input_cost and the reasoning line is a share of output_cost."""
Returns tuple of (total_cost, input_cost, output_cost, margin_cost) or all None if num_requests is None/0.
"""
if not num_requests:
return None, None, None, None
return (
cost_per_request * num_requests,
input_cost * num_requests,
output_cost * num_requests,
margin_cost * num_requests,
total_cost: float
input_cost: float
output_cost: float
margin_cost: float
cache_read_cost: float
cache_creation_cost: float
reasoning_cost: float
def times(self, num_requests: int | None) -> "CostLines | None":
if not num_requests:
return None
return CostLines(
total_cost=self.total_cost * num_requests,
input_cost=self.input_cost * num_requests,
output_cost=self.output_cost * num_requests,
margin_cost=self.margin_cost * num_requests,
cache_read_cost=self.cache_read_cost * num_requests,
cache_creation_cost=self.cache_creation_cost * num_requests,
reasoning_cost=self.reasoning_cost * num_requests,
)
def _cost_lines(cost_per_request: float, cost_breakdown: CostBreakdown | None) -> CostLines:
breakdown: Final = cost_breakdown if cost_breakdown is not None else CostBreakdown()
return CostLines(
total_cost=cost_per_request,
input_cost=breakdown.get("input_cost", 0.0),
output_cost=breakdown.get("output_cost", 0.0),
margin_cost=breakdown.get("margin_total_amount", 0.0),
cache_read_cost=breakdown.get("cache_read_cost", 0.0),
cache_creation_cost=breakdown.get("cache_creation_cost", 0.0),
reasoning_cost=breakdown.get("reasoning_cost", 0.0),
)
def _usage_for_estimate(request: CostEstimateRequest) -> Usage:
cache_tokens: Final = request.cache_read_input_tokens + request.cache_creation_input_tokens
return Usage(
prompt_tokens=request.input_tokens,
completion_tokens=request.output_tokens,
total_tokens=request.input_tokens + request.output_tokens,
reasoning_tokens=request.reasoning_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=request.cache_read_input_tokens,
cache_creation_tokens=request.cache_creation_input_tokens,
)
if cache_tokens
else None,
)
@ -530,11 +584,14 @@ async def estimate_cost(
- model: Model name (e.g., "gpt-4", "claude-3-opus")
- input_tokens: Expected input tokens per request
- output_tokens: Expected output tokens per request
- cache_read_input_tokens: Cache-read tokens per request, counted within input_tokens (optional)
- cache_creation_input_tokens: Cache-write tokens per request, counted within input_tokens (optional)
- reasoning_tokens: Reasoning tokens per request, counted within output_tokens (optional)
- num_requests_per_day: Number of requests per day (optional)
- num_requests_per_month: Number of requests per month (optional)
Returns cost breakdown including:
- Per-request costs (input, output, margin)
- Per-request costs (input, output, margin, plus the cache-read, cache-write and reasoning shares)
- Daily costs (if num_requests_per_day provided)
- Monthly costs (if num_requests_per_month provided)
@ -543,14 +600,15 @@ async def estimate_cost(
{
"model": "gpt-4",
"input_tokens": 1000,
"cache_read_input_tokens": 800,
"output_tokens": 500,
"reasoning_tokens": 200,
"num_requests_per_day": 100,
"num_requests_per_month": 3000
}
```
"""
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import ModelResponse, Usage
# Resolve model name (handles router aliases like 'e-model-router' -> 'azure_ai/gpt-4')
resolved: Final = _resolve_model_for_cost_lookup(request.model)
@ -559,15 +617,8 @@ async def estimate_cost(
verbose_proxy_logger.debug("Cost estimate: request.model='%s' resolved to '%s'", request.model, resolved_model)
# Create a mock response with usage for completion_cost
mock_response: Final = ModelResponse(
model=resolved_model,
usage=Usage(
prompt_tokens=request.input_tokens,
completion_tokens=request.output_tokens,
total_tokens=request.input_tokens + request.output_tokens,
),
)
usage: Final = _usage_for_estimate(request)
mock_response: Final = ModelResponse(model=resolved_model, usage=usage)
# Create a logging object to capture cost breakdown
litellm_logging_obj: Final = LiteLLMLoggingObj(
@ -580,92 +631,73 @@ async def estimate_cost(
function_id="cost-estimate",
)
# Use completion_cost which handles all the logic including margins/discounts
try:
cost_per_request: Final = completion_cost(
completion_response=mock_response,
model=resolved_model,
custom_llm_provider=resolved_provider,
custom_cost_per_token=resolved.custom_cost_per_token,
litellm_logging_obj=litellm_logging_obj,
)
except Exception as e:
raise HTTPException(
status_code=404,
detail={
"error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e}"
},
)
# Pinning one moment keeps an off-peak window that opens mid-quote from pricing the totals on
# one side of it and the reported rates on the other.
with pinned_billing_time(current_billing_time()):
# Use completion_cost which handles all the logic including margins/discounts
try:
cost_per_request: Final = completion_cost(
completion_response=mock_response,
model=resolved_model,
custom_llm_provider=resolved_provider,
custom_cost_per_token=resolved.custom_cost_per_token,
litellm_logging_obj=litellm_logging_obj,
)
except Exception as e: # noqa: BLE001 # completion_cost raises a bare Exception for an unpriceable model
raise HTTPException(
status_code=404,
detail={
"error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e}"
},
)
# Get cost breakdown from the logging object
cost_breakdown: Final = litellm_logging_obj.cost_breakdown
# The rates come back from the pricing call itself rather than a second lookup, so they are the
# ones the cost lines above billed at even when completion_cost infers a provider this endpoint
# never resolved (an unrouted "xai/grok-4" prices on xai's inclusive tier thresholds; a lookup
# here without that provider would report the sub-200k rate for a line billed above it).
rates: Final = litellm_logging_obj.billed_token_rates
per_request: Final = _cost_lines(cost_per_request, litellm_logging_obj.cost_breakdown)
daily: Final = per_request.times(request.num_requests_per_day)
monthly: Final = per_request.times(request.num_requests_per_month)
input_cost: Final = cost_breakdown.get("input_cost", 0.0) if cost_breakdown else 0.0
output_cost: Final = cost_breakdown.get("output_cost", 0.0) if cost_breakdown else 0.0
margin_cost: Final = cost_breakdown.get("margin_total_amount", 0.0) if cost_breakdown else 0.0
model_info: Final = _lookup_model_info(resolved_model)
mapped_input_price: Final = model_info.get("input_cost_per_token") if model_info is not None else None
mapped_output_price: Final = model_info.get("output_cost_per_token") if model_info is not None else None
model_info: Final = _lookup_model_info(resolved_model, resolved_provider)
mapped_provider: Final = model_info.get("litellm_provider") if model_info is not None else None
input_cost_per_token: Final = (
resolved.custom_cost_per_token["input_cost_per_token"]
if resolved.custom_cost_per_token is not None
else mapped_input_price
)
output_cost_per_token: Final = (
resolved.custom_cost_per_token["output_cost_per_token"]
if resolved.custom_cost_per_token is not None
else mapped_output_price
)
custom_llm_provider: Final = mapped_provider if mapped_provider is not None else resolved_provider
# Calculate daily and monthly costs
(
daily_cost,
daily_input_cost,
daily_output_cost,
daily_margin_cost,
) = _calculate_period_costs(
num_requests=request.num_requests_per_day,
cost_per_request=cost_per_request,
input_cost=input_cost,
output_cost=output_cost,
margin_cost=margin_cost,
)
(
monthly_cost,
monthly_input_cost,
monthly_output_cost,
monthly_margin_cost,
) = _calculate_period_costs(
num_requests=request.num_requests_per_month,
cost_per_request=cost_per_request,
input_cost=input_cost,
output_cost=output_cost,
margin_cost=margin_cost,
)
return CostEstimateResponse(
model=request.model,
input_tokens=request.input_tokens,
output_tokens=request.output_tokens,
cache_read_input_tokens=request.cache_read_input_tokens,
cache_creation_input_tokens=request.cache_creation_input_tokens,
reasoning_tokens=request.reasoning_tokens,
num_requests_per_day=request.num_requests_per_day,
num_requests_per_month=request.num_requests_per_month,
cost_per_request=cost_per_request,
input_cost_per_request=input_cost,
output_cost_per_request=output_cost,
margin_cost_per_request=margin_cost,
daily_cost=daily_cost,
daily_input_cost=daily_input_cost,
daily_output_cost=daily_output_cost,
daily_margin_cost=daily_margin_cost,
monthly_cost=monthly_cost,
monthly_input_cost=monthly_input_cost,
monthly_output_cost=monthly_output_cost,
monthly_margin_cost=monthly_margin_cost,
input_cost_per_token=input_cost_per_token,
output_cost_per_token=output_cost_per_token,
cost_per_request=per_request.total_cost,
input_cost_per_request=per_request.input_cost,
output_cost_per_request=per_request.output_cost,
margin_cost_per_request=per_request.margin_cost,
cache_read_cost_per_request=per_request.cache_read_cost,
cache_creation_cost_per_request=per_request.cache_creation_cost,
reasoning_cost_per_request=per_request.reasoning_cost,
daily_cost=daily.total_cost if daily is not None else None,
daily_input_cost=daily.input_cost if daily is not None else None,
daily_output_cost=daily.output_cost if daily is not None else None,
daily_margin_cost=daily.margin_cost if daily is not None else None,
daily_cache_read_cost=daily.cache_read_cost if daily is not None else None,
daily_cache_creation_cost=daily.cache_creation_cost if daily is not None else None,
daily_reasoning_cost=daily.reasoning_cost if daily is not None else None,
monthly_cost=monthly.total_cost if monthly is not None else None,
monthly_input_cost=monthly.input_cost if monthly is not None else None,
monthly_output_cost=monthly.output_cost if monthly is not None else None,
monthly_margin_cost=monthly.margin_cost if monthly is not None else None,
monthly_cache_read_cost=monthly.cache_read_cost if monthly is not None else None,
monthly_cache_creation_cost=monthly.cache_creation_cost if monthly is not None else None,
monthly_reasoning_cost=monthly.reasoning_cost if monthly is not None else None,
input_cost_per_token=rates.input_cost_per_token if rates is not None else None,
output_cost_per_token=rates.output_cost_per_token if rates is not None else None,
cache_read_input_token_cost=rates.cache_read_input_token_cost if rates is not None else None,
cache_creation_input_token_cost=rates.cache_creation_input_token_cost if rates is not None else None,
output_cost_per_reasoning_token=rates.output_cost_per_reasoning_token if rates is not None else None,
provider=custom_llm_provider,
)

View file

@ -2673,6 +2673,8 @@ if MCP_AVAILABLE:
"""
Updates the MCP Server in the db.
Partial update: a field left out of the payload keeps its stored value, and a field sent as null is cleared.
Parameters:
- payload: UpdateMCPServerRequest - Required. The updated mcp server data.
```
@ -3098,6 +3100,8 @@ if MCP_AVAILABLE:
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: str | None = Header(None),
):
"""Partial update: a field left out keeps its stored value, and a field sent as null is cleared, except
``toolset_name`` and ``tools``, which a toolset always has; empty the tool selection with an explicit []."""
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
raise HTTPException(

View file

@ -20,6 +20,7 @@ from litellm.proxy._types import (
LiteLLM_AuditLogs,
LiteLLM_TeamTable,
LitellmTableNames,
LitellmUserRoles,
ProxyErrorTypes,
ProxyException,
TeamCallbackDeleteResponse,
@ -28,7 +29,10 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.callback_config_validation import callback_config_error
from litellm.proxy.common_utils.callback_config_validation import (
callback_config_error,
cross_entry_family_error,
)
from litellm.proxy.common_utils.callback_utils import (
_CALLBACK_VAR_ENCRYPTED_PREFIX,
decrypt_callback_vars,
@ -230,6 +234,22 @@ def _callback_error(status_code: int, message: str) -> HTTPException:
)
def _unknown_team_error(team_id: str, user_api_key_dict: UserAPIKeyAuth, status_code: int) -> HTTPException:
"""Report an unknown team without telling an unauthorized caller that it is unknown.
These routes are reachable by any authenticated caller so that a team admin can
get as far as _verify_team_access. A distinct "does not exist" would therefore let
any valid key probe which team ids exist, so a caller who could not have managed
the team either way gets the same 403 body _verify_team_access raises.
"""
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
return _callback_error(status_code, f"Team id = {team_id} does not exist.")
return HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="You do not have access to this team",
)
@router.post(
"/team/{team_id:path}/callback",
tags=["team management"],
@ -304,10 +324,7 @@ async def add_team_callbacks(
# Check if team_id exists already
_existing_team = await prisma_client.get_data(team_id=team_id, table_name="team", query_type="find_unique")
if _existing_team is None:
raise HTTPException(
status_code=400,
detail={"error": f"Team id = {team_id} does not exist. Please use a different team id."},
)
raise _unknown_team_error(team_id, user_api_key_dict, status.HTTP_400_BAD_REQUEST)
# IDOR guard: only proxy admins / org admins / team admins of THIS
# team may write callback credentials. Without this, any
@ -326,6 +343,28 @@ async def add_team_callbacks(
if team_callback_settings is None or not isinstance(team_callback_settings, list):
team_callback_settings = []
# One entry has to own a credential family end to end. The entries are
# flattened into one dict before a request reads them, so an entry
# naming only a destination would pair with a key written on another
# entry and carry it to that destination -- a key a team admin can read
# back nowhere. Repeating a value the owning entry already stores is
# fine, which is how one integration covers both events. Proxy admins
# are exempt: they already hold every credential the proxy has.
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
# Decrypted, because the check compares the incoming values against
# the stored ones and the credentials are encrypted at rest.
decrypted_logging: Final = decrypt_callback_vars(team_metadata).get("logging")
stored_entries: Final = decrypted_logging if isinstance(decrypted_logging, list) else ()
stored_entry_vars: Final = [ # mutable-ok: read-only input to the check, never stored
entry.get("callback_vars") or {} for entry in stored_entries
]
family_error: Final = cross_entry_family_error(data.callback_vars, stored_entry_vars)
if family_error is not None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=family_error,
)
## check if it already exists, for the same callback event
for callback in team_callback_settings:
if (
@ -452,7 +491,7 @@ async def delete_team_callback(
team_id=team_id, table_name="team", query_type="find_unique"
)
if _existing_team is None:
raise _callback_error(404, f"Team id = {team_id} does not exist.")
raise _unknown_team_error(team_id, user_api_key_dict, status.HTTP_404_NOT_FOUND)
# IDOR guard: only proxy admins / org admins / team admins of THIS team may
# deregister its callbacks, otherwise any authenticated key holder could
@ -726,10 +765,7 @@ async def get_team_callbacks(
# Check if team_id exists
_existing_team = await prisma_client.get_data(team_id=team_id, table_name="team", query_type="find_unique")
if _existing_team is None:
raise HTTPException(
status_code=404,
detail={"error": f"Team id = {team_id} does not exist."},
)
raise _unknown_team_error(team_id, user_api_key_dict, status.HTTP_404_NOT_FOUND)
# IDOR guard: callback metadata holds third-party API credentials
# (Langfuse / Langsmith / GCS). Only proxy admins / org admins /

View file

@ -5601,7 +5601,7 @@ async def team_model_add(
updated_team: Final = await _team_db(prisma_client).update(
where={"team_id": data.team_id},
data={"updated_at": datetime.now(timezone.utc)},
include={"object_permission": True},
include={"litellm_model_table": True, "object_permission": True},
)
if updated_team is None:
raise HTTPException(
@ -5688,7 +5688,7 @@ async def team_model_delete(
updated_team: Final = await _team_db(prisma_client).update(
where={"team_id": data.team_id},
data={"models": updated_models},
include={"object_permission": True},
include={"litellm_model_table": True, "object_permission": True},
)
if updated_team is None:
raise HTTPException(

View file

@ -22,7 +22,6 @@ from html import escape
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
Annotated,
Any,
Final,
Literal,
@ -42,7 +41,7 @@ if TYPE_CHECKING:
import jwt
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
from fastapi.responses import RedirectResponse
from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError
from pydantic import BaseModel, TypeAdapter, ValidationError
import litellm
from litellm._logging import verbose_proxy_logger
@ -95,6 +94,7 @@ from litellm.proxy.auth.auth_utils import (
)
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.team_grants import TeamModelAliasTable
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.admin_ui_utils import (
admin_ui_disabled,
@ -209,31 +209,14 @@ def _team_detail_db(repo: TeamRepository) -> "TableActions[_TeamDetailRow]":
return repo.table
_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def _decode_model_aliases(value: object) -> object:
"""``/team/new`` stores team model aliases as a JSON-encoded string in the Json column."""
if not isinstance(value, str):
return value
try:
return _MODEL_ALIASES_ADAPTER.validate_json(value)
except ValidationError:
return None
class _TeamModelAliasTable(BaseModel):
model_config = ConfigDict(protected_namespaces=())
model_aliases: Annotated[Mapping[str, str] | None, BeforeValidator(_decode_model_aliases)] = None
class _TeamRowGrants(BaseModel):
team_id: str
team_alias: str | None = None
models: tuple[str, ...] = ()
litellm_model_table: _TeamModelAliasTable | None = None
litellm_model_table: TeamModelAliasTable | None = None
class CliSsoTeamDetail(BaseModel):

View file

@ -45,6 +45,11 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_headers,
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.common_utils.openai_error_payload import (
error_status_code,
openai_error_param,
openai_error_type,
)
from litellm.proxy.openai_files_endpoints.batch_file_validation import (
check_batch_file_upload,
raise_batch_file_validation_failure,
@ -296,22 +301,22 @@ async def route_create_file(
if managed_files_obj is None:
raise ProxyException(
message="Managed files hook not found",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
if llm_router is None:
raise ProxyException(
message="LLM Router not found",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
if not isinstance(managed_files_obj, BaseFileEndpoints):
raise ProxyException(
message="Managed files hook is not a BaseFileEndpoints",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
# Managed files internally calls llm_router.acreate_file() which includes loadbalancing
@ -713,17 +718,17 @@ async def create_file(
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e.detail)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
param=openai_error_param(e),
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
)
else:
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
code=error_status_code(e, 500),
)
finally:
for spool in spools:
@ -812,22 +817,22 @@ async def get_file_content(
if managed_files_obj is None:
raise ProxyException(
message="Managed files hook not found",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
if llm_router is None:
raise ProxyException(
message="LLM Router not found",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
if not isinstance(managed_files_obj, BaseFileEndpoints):
raise ProxyException(
message="Managed files hook is not a BaseFileEndpoints",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
@ -1021,17 +1026,17 @@ async def get_file_content(
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e.detail)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
param=openai_error_param(e),
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
)
else:
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
code=error_status_code(e, 500),
)
@ -1151,15 +1156,15 @@ async def get_file(
if managed_files_obj is None:
raise ProxyException(
message="Managed files hook not found",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
if not isinstance(managed_files_obj, BaseFileEndpoints):
raise ProxyException(
message="Managed files hook is not a BaseFileEndpoints",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
response = await managed_files_obj.afile_retrieve(
@ -1215,17 +1220,17 @@ async def get_file(
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e.detail)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
param=openai_error_param(e),
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
)
else:
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
code=error_status_code(e, 500),
)
@ -1355,22 +1360,22 @@ async def delete_file(
if managed_files_obj is None:
raise ProxyException(
message="Managed files hook not found",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
if llm_router is None:
raise ProxyException(
message="LLM Router not found",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
if not isinstance(managed_files_obj, BaseFileEndpoints):
raise ProxyException(
message="Managed files hook is not a BaseFileEndpoints",
type="None",
param="None",
type=ProxyErrorTypes.internal_server_error.value,
param=None,
code=500,
)
@ -1427,17 +1432,17 @@ async def delete_file(
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e.detail)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
param=openai_error_param(e),
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
)
else:
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
code=error_status_code(e, 500),
)
@ -1629,15 +1634,15 @@ async def list_files(
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e.detail)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
param=openai_error_param(e),
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
)
else:
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
code=error_status_code(e, 500),
)

View file

@ -78,6 +78,11 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
)
from litellm.proxy.common_utils.openai_error_payload import (
error_status_code,
openai_error_param,
openai_error_type,
)
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
)
@ -311,9 +316,9 @@ async def chat_completion_pass_through_endpoint(
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
code=error_status_code(e, 500),
)
@ -1728,18 +1733,18 @@ async def pass_through_request(
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(getattr(e, "detail", str(e)))),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
param=openai_error_param(e),
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
headers=custom_headers,
)
else:
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=openai_error_type(e, error_status_code(e, 500)),
param=openai_error_param(e),
code=error_status_code(e, 500),
headers=custom_headers,
)

View file

@ -17,6 +17,7 @@ import time
import traceback
import warnings
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Collection, Mapping, MutableMapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import MappingProxyType, UnionType
from typing import (
@ -131,6 +132,7 @@ from litellm.router_utils.auto_router_tuning_baseline import (
snapshot_tuning_baselines,
tuning_limit_violation,
)
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
@ -138,11 +140,7 @@ from litellm.types.utils import (
TextCompletionResponse,
TokenCountResponse,
)
from litellm.utils import (
_invalidate_model_cost_lowercase_map,
load_credentials_from_list,
reapply_runtime_model_cost_registrations,
)
from litellm.utils import load_credentials_from_list
if TYPE_CHECKING:
from aiohttp import ClientSession
@ -426,6 +424,7 @@ from litellm.proxy.db.exception_handler import (
)
from litellm.proxy.db.gateway_request_tracking import (
GatewayRequestAccumulator,
GatewayRequestRedisBuffer,
flush_gateway_requests,
)
from litellm.proxy.db.proxy_worker_heartbeat import (
@ -2357,6 +2356,17 @@ open_telemetry_logger: OpenTelemetry | None = None
gateway_request_accumulator: Final = GatewayRequestAccumulator()
### INITIALIZE GLOBAL LOGGING OBJECT ###
proxy_logging_obj: ProxyLogging = ProxyLogging(user_api_key_cache=user_api_key_cache, premium_user=premium_user)
def _gateway_request_redis_buffer() -> GatewayRequestRedisBuffer | None:
"""Shares the spend writer's transaction-buffer Redis and pod lock when use_redis_transaction_buffer is on."""
writer: Final = proxy_logging_obj.db_spend_update_writer
redis_cache: Final = writer.redis_update_buffer.redis_cache
if redis_cache is None or not writer.redis_update_buffer._should_commit_spend_updates_to_redis():
return None
return GatewayRequestRedisBuffer(redis_cache=redis_cache, pod_lock_manager=writer.pod_lock_manager)
### REDIS QUEUE ###
async_result: Final = None
celery_app_conn: Final = None
@ -2707,6 +2717,12 @@ async def _read_spend_counter_estimate(counter_key: str, fallback_spend: float)
return fallback_spend, False
@dataclass(frozen=True, slots=True)
class _PendingSpendIncrement:
counter_key: str
increment: float
async def increment_spend_counters(
token: str | None,
team_id: str | None,
@ -2741,7 +2757,7 @@ async def increment_spend_counters(
cost: Final[float] = response_cost
async def _key_scope(key_token: str) -> None:
async def _key_scope(key_token: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
# key_token arrives pre-hashed from metadata["user_api_key"] (auth flow
# hashes raw "sk-..." keys before they reach the callback). The
# startswith("sk-") check is a safety net matching update_cache —
@ -2752,30 +2768,29 @@ async def increment_spend_counters(
hash_token(token=key_token) if isinstance(key_token, str) and key_token.startswith("sk-") else key_token
)
key_counter_key: Final = f"spend:key:{hashed_token}"
if key_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=key_counter_key,
source_cache_key=hashed_token,
increment=cost,
key_pending: Final[tuple[_PendingSpendIncrement, ...]] = (
()
if key_counter_key in reserved_counter_keys
else (
await _prepare_spend_counter_increment(
counter_key=key_counter_key,
source_cache_key=hashed_token,
increment=cost,
),
)
key_obj: Final[object] = await user_api_key_cache.async_get_cache(key=hashed_token)
if key_obj is None:
return
key_budget_limits = getattr(key_obj, "budget_limits", None) or (
key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
)
if isinstance(key_budget_limits, str):
key_budget_limits = json.loads(key_budget_limits)
if not isinstance(key_budget_limits, list):
return
for window in key_budget_limits:
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
key_window_reset_at = window.get("reset_at") if isinstance(window, dict) else window.reset_at
key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
async def _key_window_increment(window: object) -> _PendingSpendIncrement | None:
duration = (
window["budget_duration"] if isinstance(window, dict) else getattr(window, "budget_duration", None)
)
key_window_reset_at = (
window.get("reset_at") if isinstance(window, dict) else getattr(window, "reset_at", None)
)
key_window_counter: Final = f"spend:key:{hashed_token}:window:{duration}"
key_window_start = get_budget_window_start(window)
if key_window_counter not in reserved_counter_keys:
await _init_and_increment_window_spend_counter(
pending_window: Final = (
await _prepare_window_spend_counter_increment(
counter_key=key_window_counter,
entity_type="Key",
entity_id=hashed_token,
@ -2783,6 +2798,9 @@ async def increment_spend_counters(
window_start=key_window_start,
increment=cost,
)
if key_window_counter not in reserved_counter_keys
else None
)
await _enqueue_window_spend_row_update(
entity_type=Litellm_EntityType.KEY,
entity_id=hashed_token,
@ -2792,33 +2810,48 @@ async def increment_spend_counters(
increment=cost,
request_started_at=request_started_at,
)
return pending_window
async def _team_scope(scope_team_id: str) -> None:
team_counter_key: Final = f"spend:team:{scope_team_id}"
if team_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=team_counter_key,
source_cache_key=f"team_id:{scope_team_id}",
increment=cost,
)
team_obj: Final[object] = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
if team_obj is None:
return
team_budget_limits = getattr(team_obj, "budget_limits", None) or (
team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
key_obj: Final[object] = await user_api_key_cache.async_get_cache(key=hashed_token)
if key_obj is None:
return key_pending
key_budget_limits = getattr(key_obj, "budget_limits", None) or (
key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
)
if isinstance(team_budget_limits, str):
team_budget_limits = json.loads(team_budget_limits)
if not isinstance(team_budget_limits, list):
return
for window in team_budget_limits:
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
team_window_reset_at = window.get("reset_at") if isinstance(window, dict) else window.reset_at
team_window_counter = f"spend:team:{scope_team_id}:window:{duration}"
if isinstance(key_budget_limits, str):
key_budget_limits = json.loads(key_budget_limits)
if not isinstance(key_budget_limits, list):
return key_pending
window_pending: Final = await asyncio.gather(
*(_key_window_increment(window) for window in key_budget_limits), return_exceptions=True
)
return key_pending + tuple(item for item in window_pending if item is not None)
async def _team_scope(scope_team_id: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
team_counter_key: Final = f"spend:team:{scope_team_id}"
team_pending: Final[tuple[_PendingSpendIncrement, ...]] = (
()
if team_counter_key in reserved_counter_keys
else (
await _prepare_spend_counter_increment(
counter_key=team_counter_key,
source_cache_key=f"team_id:{scope_team_id}",
increment=cost,
),
)
)
async def _team_window_increment(window: object) -> _PendingSpendIncrement | None:
duration = (
window["budget_duration"] if isinstance(window, dict) else getattr(window, "budget_duration", None)
)
team_window_reset_at = (
window.get("reset_at") if isinstance(window, dict) else getattr(window, "reset_at", None)
)
team_window_counter: Final = f"spend:team:{scope_team_id}:window:{duration}"
team_window_start = get_budget_window_start(window)
if team_window_counter not in reserved_counter_keys:
await _init_and_increment_window_spend_counter(
pending_window: Final = (
await _prepare_window_spend_counter_increment(
counter_key=team_window_counter,
entity_type="Team",
entity_id=scope_team_id,
@ -2826,6 +2859,9 @@ async def increment_spend_counters(
window_start=team_window_start,
increment=cost,
)
if team_window_counter not in reserved_counter_keys
else None
)
await _enqueue_window_spend_row_update(
entity_type=Litellm_EntityType.TEAM,
entity_id=scope_team_id,
@ -2835,25 +2871,47 @@ async def increment_spend_counters(
increment=cost,
request_started_at=request_started_at,
)
return pending_window
async def _team_member_scope(scope_user_id: str, scope_team_id: str) -> None:
team_obj: Final[object] = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
if team_obj is None:
return team_pending
team_budget_limits = getattr(team_obj, "budget_limits", None) or (
team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
)
if isinstance(team_budget_limits, str):
team_budget_limits = json.loads(team_budget_limits)
if not isinstance(team_budget_limits, list):
return team_pending
window_pending: Final = await asyncio.gather(
*(_team_window_increment(window) for window in team_budget_limits), return_exceptions=True
)
return team_pending + tuple(item for item in window_pending if item is not None)
async def _team_member_scope(
scope_user_id: str, scope_team_id: str
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
team_member_counter_key: Final = f"spend:team_member:{scope_user_id}:{scope_team_id}"
if team_member_counter_key in reserved_counter_keys:
return
await _init_and_increment_spend_counter(
counter_key=team_member_counter_key,
source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
increment=cost,
return ()
return (
await _prepare_spend_counter_increment(
counter_key=team_member_counter_key,
source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
increment=cost,
),
)
async def _user_scope(scope_user_id: str) -> None:
async def _user_scope(scope_user_id: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
user_counter_key: Final = f"spend:user:{scope_user_id}"
if user_counter_key in reserved_counter_keys:
return
await _init_and_increment_spend_counter(
counter_key=user_counter_key,
source_cache_key=scope_user_id,
increment=cost,
return ()
return (
await _prepare_spend_counter_increment(
counter_key=user_counter_key,
source_cache_key=scope_user_id,
increment=cost,
),
)
scope_coros: Final = tuple(
@ -2863,7 +2921,7 @@ async def increment_spend_counters(
_team_scope(team_id) if team_id is not None else None,
_team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None,
_user_scope(user_id) if user_id is not None else None,
_increment_end_user_and_tag_spend_counters(
_prepare_end_user_and_tag_spend_increments(
end_user_id=end_user_id,
tags=tags,
response_cost=cost,
@ -2871,14 +2929,14 @@ async def increment_spend_counters(
)
if end_user_id is not None or tags is not None
else None,
_increment_model_access_group_spend_counters(
_prepare_model_access_group_spend_increments(
model_access_groups=model_access_groups,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
)
if model_access_groups
else None,
_increment_org_spend_counter(
_prepare_org_spend_increment(
org_id=org_id,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
@ -2893,7 +2951,20 @@ async def increment_spend_counters(
# as orphaned tasks that race the caller's reservation-counter invalidation;
# all scopes settle, then the first error propagates as before.
scope_results: Final = await asyncio.gather(*scope_coros, return_exceptions=True)
scope_errors: Final = [r for r in scope_results if isinstance(r, BaseException)]
scope_errors: Final = tuple(
item
for scope in scope_results
for item in (scope if isinstance(scope, tuple) else (scope,))
if isinstance(item, BaseException)
)
pending: Final = tuple(
item
for scope in scope_results
if not isinstance(scope, BaseException)
for item in scope
if not isinstance(item, BaseException)
)
await _apply_spend_counter_increments(pending=pending)
if scope_errors:
raise scope_errors[0]
@ -2936,41 +3007,49 @@ async def _reconcile_budget_reservation_for_counter_update(
return reserved_counter_keys
async def _increment_end_user_and_tag_spend_counters(
async def _prepare_end_user_and_tag_spend_increments(
end_user_id: str | None,
tags: list[str] | None,
response_cost: float,
reserved_counter_keys: set[str],
) -> None:
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=end_user_cache_key(end_user_id),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
if tags is None:
return
seen_tags: Final[set[str]] = set()
for tag_name in tags:
if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags:
continue
seen_tags.add(tag_name)
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:tag:{tag_name}",
source_cache_key=tag_cache_key(tag_name),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
unique_tags: Final = (
tuple(dict.fromkeys(tag for tag in tags if tag and isinstance(tag, str))) if tags is not None else ()
)
results: Final = await asyncio.gather(
*(
coro
for coro in (
_prepare_unreserved_spend_counter_increment(
counter_key=f"spend:end_user:{end_user_id}",
source_cache_key=end_user_cache_key(end_user_id),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
if end_user_id is not None
else None,
*(
_prepare_unreserved_spend_counter_increment(
counter_key=f"spend:tag:{tag_name}",
source_cache_key=tag_cache_key(tag_name),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
for tag_name in unique_tags
),
)
if coro is not None
),
return_exceptions=True,
)
return tuple(item for item in results if item is not None)
async def _increment_model_access_group_spend_counters(
async def _prepare_model_access_group_spend_increments(
model_access_groups: Sequence[object],
response_cost: float,
reserved_counter_keys: set[str],
) -> None:
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
"""Charge the model access groups that authorized this request.
Without this the counter auth reads is written only by the reservation path, so
@ -2984,55 +3063,63 @@ async def _increment_model_access_group_spend_counters(
unique_groups: Final = tuple(
dict.fromkeys(group for group in model_access_groups if group and isinstance(group, str))
)
for group in unique_groups:
await _init_and_increment_unreserved_spend_counter(
counter_key=model_access_group_spend_counter_key(group),
source_cache_key=model_access_group_cache_key(group),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
results: Final = await asyncio.gather(
*(
_prepare_unreserved_spend_counter_increment(
counter_key=model_access_group_spend_counter_key(group),
source_cache_key=model_access_group_cache_key(group),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
for group in unique_groups
),
return_exceptions=True,
)
return tuple(item for item in results if item is not None)
async def _increment_org_spend_counter(
async def _prepare_org_spend_increment(
org_id: str | None,
response_cost: float,
reserved_counter_keys: set[str],
) -> None:
) -> tuple[_PendingSpendIncrement, ...]:
if org_id is None:
return
return ()
await _init_and_increment_unreserved_spend_counter(
pending: Final = await _prepare_unreserved_spend_counter_increment(
counter_key=f"spend:org:{org_id}",
source_cache_key=[f"org_id:{org_id}:with_budget", f"org_id:{org_id}"],
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
return (pending,) if pending is not None else ()
async def _init_and_increment_unreserved_spend_counter(
async def _prepare_unreserved_spend_counter_increment(
counter_key: str,
source_cache_key: str | list[str],
increment: float,
reserved_counter_keys: set[str],
) -> None:
) -> _PendingSpendIncrement | None:
if counter_key in reserved_counter_keys:
return
return None
await _init_and_increment_spend_counter(
return await _prepare_spend_counter_increment(
counter_key=counter_key,
source_cache_key=source_cache_key,
increment=increment,
)
async def _init_and_increment_spend_counter(
async def _prepare_spend_counter_increment(
counter_key: str,
source_cache_key: str | list[str],
increment: float,
):
) -> _PendingSpendIncrement:
"""
Initialize counter from the authoritative DB spend value if not yet
set, then atomically increment in both in-memory and Redis.
set, then return the pending increment for the caller to apply in one
pipelined Redis call.
On first access per pod:
1. Check spend_counter_cache (in-memory -> Redis via DualCache)
@ -3044,13 +3131,13 @@ async def _init_and_increment_spend_counter(
the counter as absent and seed it. Using increment means the worst case
is over-counting (conservative, blocks slightly early) rather than
under-counting (would allow overspend).
4. Increment atomically (both in-memory + Redis)
4. Increment is returned for the caller to apply via pipeline
"""
await _ensure_spend_counter_initialized(
counter_key=counter_key,
source_cache_key=source_cache_key,
)
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
return _PendingSpendIncrement(counter_key=counter_key, increment=increment)
async def _enqueue_window_spend_row_update(
@ -3102,20 +3189,20 @@ async def _enqueue_window_spend_row_update(
)
async def _init_and_increment_window_spend_counter(
async def _prepare_window_spend_counter_increment(
counter_key: str,
entity_type: str,
entity_id: str,
window_duration: str | None,
window_start: datetime | None,
increment: float,
):
) -> _PendingSpendIncrement | None:
if window_start is None:
verbose_proxy_logger.warning(
"Skipping spend counter increment for invalid budget window %s",
counter_key,
)
return
return None
initialized: Final = await _ensure_window_spend_counter_initialized(
counter_key=counter_key,
@ -3125,8 +3212,8 @@ async def _init_and_increment_window_spend_counter(
window_start=window_start,
)
if initialized is False:
return
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
return None
return _PendingSpendIncrement(counter_key=counter_key, increment=increment)
async def _ensure_spend_counter_initialized(
@ -3259,6 +3346,32 @@ async def _invalidate_spend_counter(counter_key: str):
)
async def _apply_spend_counter_increments(pending: Sequence[_PendingSpendIncrement]) -> None:
if not pending:
return
redis_cache: Final = spend_counter_cache.redis_cache
if redis_cache is None:
for item in pending:
await spend_counter_cache.async_increment_cache(
key=item.counter_key,
value=item.increment,
refresh_ttl=True,
)
return
ttl: Final = redis_cache.get_ttl()
increment_list: Final = [ # mutable-ok: async_increment_pipeline signature requires list[RedisPipelineIncrementOperation]
RedisPipelineIncrementOperation(key=item.counter_key, increment_value=item.increment, ttl=ttl)
for item in pending
]
try:
results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list)
except Exception:
await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending))
raise
for item, current_value in zip(pending, results or ()):
spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value)
async def update_cache(
token: str | None,
user_id: str | None,
@ -4436,20 +4549,9 @@ def resolve_classifier_plugin(
def _swap_in_model_cost_map(new_model_cost_map: dict) -> int:
"""Adopt a freshly fetched cost map into this process's litellm state, return the model count"""
litellm.model_cost = new_model_cost_map
# Invalidate case-insensitive lookup map since model_cost was replaced
_invalidate_model_cost_lowercase_map()
# Repopulate provider model sets (e.g. litellm.anthropic_models) so that
# wildcard patterns like "anthropic/*" include any newly added models.
litellm.add_known_models(model_cost_map=new_model_cost_map)
# Counted before the re-apply below, which writes into this same dict, so the
# number reported describes the fetched price data alone.
fetched_model_count: Final = len(new_model_cost_map) if new_model_cost_map else 0
# The swap discards everything registered at runtime (deployment model_info,
# register_model overrides), so put it back on top of the fresh catalog.
reapply_runtime_model_cost_registrations()
return fetched_model_count
from litellm.litellm_core_utils.get_model_cost_map import adopt_model_cost_map
return adopt_model_cost_map(new_model_cost_map)
def should_load_db_object(object_type: str | SupportedDBObjectType) -> bool:
@ -9543,7 +9645,7 @@ class ProxyStartupEvent:
flush_gateway_requests,
"interval",
seconds=batch_writing_interval,
args=(prisma_client, gateway_request_accumulator),
args=(prisma_client, gateway_request_accumulator, _gateway_request_redis_buffer()),
id="update_gateway_requests_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
@ -17750,6 +17852,7 @@ async def reload_model_cost_map(
# Immediately reload the model cost map in the current pod
from litellm.litellm_core_utils.get_model_cost_map import (
ModelCostMapReloadUnavailable,
get_model_cost_map_provenance,
refetch_model_cost_map,
)
@ -17763,6 +17866,7 @@ async def reload_model_cost_map(
models_count = _swap_in_model_cost_map(reload_result.model_cost_map)
current_time = utc_now()
proxy_config.model_cost_map_loaded_at = current_time
provenance: Final = get_model_cost_map_provenance()
# Publish a new revision so every other pod reloads on its next poll; this pod has
# already served it, so adopt it here rather than reloading again a tick later
@ -17777,6 +17881,7 @@ async def reload_model_cost_map(
"status": "success",
"models_count": models_count,
"timestamp": current_time.isoformat(),
**provenance,
}
except HTTPException:
raise
@ -17897,12 +18002,17 @@ async def get_model_cost_map_reload_status(
try:
global prisma_client
from litellm.litellm_core_utils.get_model_cost_map import (
get_model_cost_map_provenance,
)
provenance: Final = get_model_cost_map_provenance()
if prisma_client is None:
verbose_proxy_logger.info("No database connection, returning not scheduled")
return reload_schedule_status(None)
return {**reload_schedule_status(None), **provenance}
return reload_schedule_status(await read_reload_schedule(prisma_client, MODEL_COST_MAP_RELOAD_PARAM_NAME))
schedule: Final = await read_reload_schedule(prisma_client, MODEL_COST_MAP_RELOAD_PARAM_NAME)
return {**reload_schedule_status(schedule), **provenance}
except Exception as e:
verbose_proxy_logger.exception("Failed to get model cost map reload status: %s", e)
raise HTTPException(
@ -17930,6 +18040,9 @@ async def get_model_cost_map_source(
- url: the remote URL that was attempted (null when env-forced local)
- is_env_forced: true if LITELLM_LOCAL_MODEL_COST_MAP=True forced local usage
- fallback_reason: human-readable reason why remote failed (null on success)
- loaded_at: when this pod last loaded the map
- source_revision: git blob id of the loaded file, what git rev-parse <commit>:<path> prints for it
- etag: the ETag of the remote fetch (null for the bundled backup)
- model_count: number of models in the currently loaded cost map
"""
# Read-only source info — admin viewers can read.
@ -18484,6 +18597,31 @@ async def _stream_mcp_asgi_response(handle_fn, scope: dict, receive) -> "Streami
########################################################
@app.api_route(
"/mcp/proxy",
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], # mutable-ok: FastAPI route methods
)
async def proxy_mcp_route(request: Request) -> Response:
"""Serve the fixed three-tool MCP proxy surface."""
from litellm.proxy._experimental.mcp_server.mcp_context import ( # pyright: ignore[reportPrivateUsage] # route-owned mode
_mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # route-owned mode
)
from litellm.proxy._experimental.mcp_server.server import handle_streamable_http_mcp
from litellm.proxy._experimental.mcp_server.utils import is_mcp_available
if not is_mcp_available():
raise HTTPException(status_code=404, detail="Not Found")
token: Final = _mcp_proxy_mode.set(True)
try:
scope: Final = dict(request.scope) # mutable-ok: ASGI scope rewrite
scope["_original_path"] = scope.get("path", "")
scope["path"] = BASE_MCP_ROUTE
return await _stream_mcp_asgi_response(handle_streamable_http_mcp, scope, request.receive)
finally:
_mcp_proxy_mode.reset(token)
@app.api_route(
BASE_MCP_ROUTE,
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"],

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