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

# Conflicts:
#	tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py
This commit is contained in:
mateo-berri 2026-09-09 15:43:48 -07:00
commit d644e4970f
416 changed files with 33970 additions and 6287 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

@ -1,5 +1,11 @@
#!/bin/sh
# stale samples from a previous container incarnation would be summed into the aggregate
if [ -n "$PROMETHEUS_MULTIPROC_DIR" ]; then
mkdir -p "$PROMETHEUS_MULTIPROC_DIR"
rm -f "$PROMETHEUS_MULTIPROC_DIR"/*.db
fi
case "$USE_DDTRACE" in
[Tt][Rr][Uu][Ee])
export DD_TRACE_OPENAI_ENABLED="False"

View file

@ -142,6 +142,40 @@ class CheckBatchCost:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up team alias for team {team_id}: {e}")
return None
async def _get_org_id(self, job: "LiteLLM_ManagedObjectTable", batch_id: str) -> str | None:
org_id = getattr(job, "org_id", None)
if org_id:
return org_id
api_key = getattr(job, "api_key", None)
team_id = getattr(job, "team_id", None)
if api_key:
try:
key_row: prisma_models.LiteLLM_VerificationToken | None = (
await self.prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": api_key}
)
)
key_org_id = getattr(key_row, "organization_id", None) if key_row is not None else None
if key_org_id:
return key_org_id
except Exception as e:
verbose_proxy_logger.error(
f"CheckBatchCost: could not resolve the key's org for batch {batch_id}, "
f"still trying the team's: {e}"
)
if not team_id:
return None
try:
team_row: prisma_models.LiteLLM_TeamTable | None = (
await self.prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
)
)
return getattr(team_row, "organization_id", None) if team_row is not None else None
except Exception as e:
verbose_proxy_logger.error(f"CheckBatchCost: could not resolve the team's org for batch {batch_id}: {e}")
return None
async def _build_creator_attribution_metadata(
self, job: "LiteLLM_ManagedObjectTable", batch_id: str
) -> dict[str, object]:
@ -153,6 +187,10 @@ class CheckBatchCost:
user_api_key_alias; when it has no alias, or the key has since been rotated or
deleted, the field keeps the creating user's alias that _get_user_info filled in,
because a resolvable name is more useful on the spend row than a null.
user_api_key_org_id must be resolved here too: the spend update writer reads it
off this metadata to increment organization spend, so leaving it out silently
drops batch cost from org accounting for keys and teams that belong to one.
"""
api_key = getattr(job, "api_key", None)
team_id = getattr(job, "team_id", None)
@ -172,6 +210,9 @@ class CheckBatchCost:
team_alias = await self._get_team_alias(team_id)
if team_alias is not None:
metadata["user_api_key_team_alias"] = team_alias
org_id: Final = await self._get_org_id(job, batch_id)
if org_id is not None:
metadata["user_api_key_org_id"] = org_id
if isinstance(request_tags, list) and request_tags:
metadata["tags"] = [tag for tag in request_tags if isinstance(tag, str)]
@ -641,7 +682,7 @@ class CheckBatchCost:
from litellm.files.main import afile_content
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info
from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info, mask_api_base_credentials
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
@ -805,6 +846,7 @@ class CheckBatchCost:
function_id=str(uuid.uuid4()),
)
deployment_api_base: Final = deployment_info.litellm_params.api_base
logging_obj.update_environment_variables(
litellm_params={
# set the user-agent header so that S3 callback consumers can easily identify CheckBatchCost callbacks
@ -813,9 +855,17 @@ class CheckBatchCost:
"user-agent": CHECK_BATCH_COST_USER_AGENT,
}
},
"metadata": await self._build_creator_attribution_metadata(job, batch_id),
**({"api_base": mask_api_base_credentials(deployment_api_base)} if deployment_api_base else {}),
"metadata": {
**(await self._build_creator_attribution_metadata(job, batch_id)),
# spend logs read the deployment identity off these metadata keys, so
# without them the batch cost row carries no model_id or model_group
"model_info": {"id": model_id},
"model_group": deployment_info.model_name,
},
},
optional_params={},
custom_llm_provider=str(llm_provider) if llm_provider else None,
)
if not await self._claim_job_for_costing(job):
@ -833,6 +883,8 @@ class CheckBatchCost:
batch_models=batch_result.models,
batch_successful_requests=batch_result.successful_requests,
batch_failed_requests=batch_result.failed_requests,
batch_prompt_cost=batch_result.prompt_cost,
batch_completion_cost=batch_result.completion_cost,
)
except Exception:
await self._release_job_claim(job)

View file

@ -280,6 +280,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
verbose_logger.debug(f"LiteLLM Managed File object with id={file_id} stored in db: {result}")
async def _resolve_creator_org_id(self, user_api_key_dict: UserAPIKeyAuth) -> Optional[str]:
if user_api_key_dict.org_id:
return user_api_key_dict.org_id
if not user_api_key_dict.team_id:
return None
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
try:
team: Final = await get_team_object(
team_id=user_api_key_dict.team_id,
prisma_client=self.prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_dict.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
return team.organization_id
except Exception as e:
verbose_logger.warning(f"could not resolve org for managed object attribution: {e}")
return None
async def store_unified_object_id(
self,
unified_object_id: str,
@ -352,6 +373,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"file_purpose": file_purpose,
"created_by": resolve_resource_owner_id(user_api_key_dict),
"team_id": user_api_key_dict.team_id,
"org_id": await self._resolve_creator_org_id(user_api_key_dict),
"updated_by": user_api_key_dict.user_id,
"status": file_object.status,
**attribution_columns,
@ -1779,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)
@ -1790,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

@ -152,6 +152,13 @@ spec:
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsEnv" . | nindent 12 }}
{{- end }}
{{- if .Values.metricsServer.enabled }}
{{- if eq (int .Values.metricsServer.port) (int .Values.service.port) }}
{{- fail "metricsServer.port must differ from service.port" }}
{{- end }}
- name: PROMETHEUS_METRICS_PORT
value: {{ .Values.metricsServer.port | quote }}
{{- end }}
{{- if .Values.migrationJob.enabled }}
# Schema updates are owned by the dedicated migrations Job; skip
# the proxy's startup `prisma db push` so N replicas don't race
@ -189,6 +196,11 @@ spec:
- name: http
containerPort: {{ .Values.service.port }}
protocol: TCP
{{- if .Values.metricsServer.enabled }}
- name: metrics
containerPort: {{ .Values.metricsServer.port }}
protocol: TCP
{{- end }}
livenessProbe:
httpGet:
path: {{ .Values.livenessProbe.path | quote }}

View file

@ -0,0 +1,17 @@
{{- if .Values.metricsServer.enabled }}
apiVersion: v1
kind: Service
metadata:
name: {{ include "litellm.fullname" . }}-metrics
labels:
{{- include "litellm.labels" . | nindent 4 }}
spec:
type: ClusterIP
ports:
- port: {{ .Values.metricsServer.port }}
targetPort: metrics
protocol: TCP
name: metrics
selector:
{{- include "litellm.selectorLabels" . | nindent 4 }}
{{- end }}

View file

@ -26,7 +26,7 @@ spec:
{{- toYaml .namespaceSelector.matchNames | nindent 4 }}
{{- end }}
endpoints:
- port: http
- port: {{ ternary "metrics" "http" $.Values.metricsServer.enabled }}
path: /metrics/
interval: {{ .interval }}
scrapeTimeout: {{ .scrapeTimeout }}

View file

@ -0,0 +1,106 @@
suite: separate metrics server
templates:
- configmap-litellm.yaml
- deployment.yaml
- service.yaml
- service-metrics.yaml
- servicemonitor.yaml
tests:
- it: should not expose a metrics port or PROMETHEUS_METRICS_PORT by default
asserts:
- notContains:
path: spec.template.spec.containers[0].ports
content:
name: metrics
any: true
template: deployment.yaml
- notContains:
path: spec.template.spec.containers[0].env
content:
name: PROMETHEUS_METRICS_PORT
any: true
template: deployment.yaml
- lengthEqual:
path: spec.ports
count: 1
template: service.yaml
- hasDocuments:
count: 0
template: service-metrics.yaml
- it: should scrape the proxy port when the metrics server is disabled
template: servicemonitor.yaml
set:
serviceMonitor.enabled: true
asserts:
- equal:
path: spec.endpoints[0].port
value: http
- it: should wire the separate metrics server through container, a ClusterIP metrics service and servicemonitor
set:
metricsServer.enabled: true
metricsServer.port: 4101
serviceMonitor.enabled: true
service.type: LoadBalancer
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: PROMETHEUS_METRICS_PORT
value: "4101"
template: deployment.yaml
- contains:
path: spec.template.spec.containers[0].ports
content:
name: metrics
containerPort: 4101
protocol: TCP
template: deployment.yaml
- lengthEqual:
path: spec.ports
count: 1
template: service.yaml
- equal:
path: spec.type
value: LoadBalancer
template: service.yaml
- equal:
path: metadata.name
value: RELEASE-NAME-litellm-metrics
template: service-metrics.yaml
- equal:
path: spec.type
value: ClusterIP
template: service-metrics.yaml
- equal:
path: spec.ports
value:
- port: 4101
targetPort: metrics
protocol: TCP
name: metrics
template: service-metrics.yaml
- equal:
path: spec.selector
value:
app.kubernetes.io/name: litellm
app.kubernetes.io/instance: RELEASE-NAME
template: service-metrics.yaml
- equal:
path: spec.endpoints[0].port
value: metrics
template: servicemonitor.yaml
- equal:
path: spec.endpoints[0].path
value: /metrics/
template: servicemonitor.yaml
- it: should reject a metrics port equal to the proxy port
template: deployment.yaml
set:
metricsServer.enabled: true
metricsServer.port: 4000
asserts:
- failedTemplate:
errorMessage: metricsServer.port must differ from service.port

View file

@ -180,6 +180,16 @@ proxy_config:
general_settings:
master_key: os.environ/PROXY_MASTER_KEY
# Serve Prometheus /metrics from a separate process (PROMETHEUS_METRICS_PORT)
# so a scrape never runs on an inference worker. Adds a `metrics` port to the
# container and a dedicated ClusterIP `<release>-metrics` Service, and the
# ServiceMonitor scrapes it instead of the proxy port. The separate port has
# no virtual-key auth: keep it off public ingress. Needs the proxy image
# v1.101.0 or newer.
metricsServer:
enabled: false
port: 4001
resources:
{}
# Unset by default so the chart installs on small clusters such as Minikube, and so an

View file

@ -441,3 +441,5 @@ ImplementationSpecific
{{- .pathType -}}
{{- end -}}
{{- end -}}
{{- define "litellm.gateway.prometheusMultiprocDir" -}}/tmp/litellm_prometheus_multiproc{{- end -}}

View file

@ -64,14 +64,25 @@ spec:
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsEnv" . | nindent 12 }}
{{- end }}
{{- if .Values.gateway.metricsServer.enabled }}
{{- if eq (int .Values.gateway.metricsServer.port) 4000 }}
{{- fail "gateway.metricsServer.port must differ from the gateway port 4000" }}
{{- end }}
- name: PROMETHEUS_MULTIPROC_DIR
value: {{ include "litellm.gateway.prometheusMultiprocDir" . }}
{{- end }}
{{- include "litellm.envFrom" .Values.gateway | nindent 10 }}
{{- if or .Values.gateway.config.create .Values.gateway.volumeMounts .Values.billingMetrics.enabled }}
{{- if or .Values.gateway.config.create .Values.gateway.volumeMounts .Values.billingMetrics.enabled .Values.gateway.metricsServer.enabled }}
volumeMounts:
{{- if .Values.gateway.config.create }}
- name: gateway-config
mountPath: /app/config/config.yaml
subPath: config.yaml
{{- end }}
{{- if .Values.gateway.metricsServer.enabled }}
- name: prometheus-multiproc
mountPath: {{ include "litellm.gateway.prometheusMultiprocDir" . }}
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsVolumeMounts" . | nindent 12 }}
{{- end }}
@ -97,16 +108,54 @@ spec:
{{- end }}
resources:
{{- toYaml .Values.gateway.resources | nindent 12 }}
{{- if .Values.gateway.metricsServer.enabled }}
- name: metrics
image: "{{ .Values.gateway.image.repository }}:{{ .Values.gateway.image.tag | default .Chart.AppVersion }}"
imagePullPolicy: {{ .Values.gateway.image.pullPolicy }}
{{- with .Values.gateway.securityContext }}
securityContext:
{{- toYaml . | nindent 12 }}
{{- end }}
command:
- python
- -m
- litellm.proxy.prometheus_metrics_server
- --port
- {{ .Values.gateway.metricsServer.port | quote }}
env:
- name: PROMETHEUS_MULTIPROC_DIR
value: {{ include "litellm.gateway.prometheusMultiprocDir" . }}
ports:
- name: metrics
containerPort: {{ .Values.gateway.metricsServer.port }}
protocol: TCP
volumeMounts:
- name: prometheus-multiproc
mountPath: {{ include "litellm.gateway.prometheusMultiprocDir" . }}
readinessProbe:
tcpSocket: { port: metrics }
periodSeconds: 10
livenessProbe:
tcpSocket: { port: metrics }
periodSeconds: 15
failureThreshold: 6
resources:
{{- toYaml .Values.gateway.metricsServer.resources | nindent 12 }}
{{- end }}
{{- with .Values.gateway.extraContainers }}
{{- tpl (toYaml .) $ | nindent 8 }}
{{- end }}
{{- if or .Values.gateway.config.create .Values.gateway.volumes .Values.billingMetrics.enabled }}
{{- if or .Values.gateway.config.create .Values.gateway.volumes .Values.billingMetrics.enabled .Values.gateway.metricsServer.enabled }}
volumes:
{{- if .Values.gateway.config.create }}
- name: gateway-config
configMap:
name: {{ include "litellm.gateway.fullname" . }}-config
{{- end }}
{{- if .Values.gateway.metricsServer.enabled }}
- name: prometheus-multiproc
emptyDir: {}
{{- end }}
{{- if .Values.billingMetrics.enabled }}
{{- include "litellm.billingMetricsVolumes" . | nindent 8 }}
{{- end }}

View file

@ -0,0 +1,18 @@
{{- if and .Values.gateway.enabled .Values.gateway.metricsServer.enabled }}
apiVersion: v1
kind: Service
metadata:
name: {{ include "litellm.gateway.fullname" . }}-metrics
labels:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: gateway
spec:
type: ClusterIP
ports:
- port: {{ .Values.gateway.metricsServer.port }}
targetPort: metrics
protocol: TCP
name: metrics
selector:
{{- include "litellm.gateway.selectorLabels" . | nindent 4 }}
{{- end }}

View file

@ -0,0 +1,148 @@
suite: test gateway metrics sidecar
templates:
- gateway/configmap.yaml
- gateway/deployment.yaml
- gateway/service.yaml
- gateway/service-metrics.yaml
values:
- ./values/required.yaml
tests:
- it: adds no sidecar, volume, env or service port when the metrics server is off
asserts:
- lengthEqual:
path: spec.template.spec.containers
count: 1
template: gateway/deployment.yaml
- notContains:
path: spec.template.spec.containers[0].env
content:
name: PROMETHEUS_MULTIPROC_DIR
any: true
template: gateway/deployment.yaml
- notContains:
path: spec.template.spec.volumes
content:
name: prometheus-multiproc
any: true
template: gateway/deployment.yaml
- lengthEqual:
path: spec.ports
count: 1
template: gateway/service.yaml
- hasDocuments:
count: 0
template: gateway/service-metrics.yaml
- it: runs the metrics server as a sidecar over a shared multiproc dir and exposes it on a ClusterIP metrics service
set:
gateway.metricsServer.enabled: true
gateway.metricsServer.port: 4101
gateway.service.type: LoadBalancer
gateway.image.tag: v1.101.0
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: PROMETHEUS_MULTIPROC_DIR
value: /tmp/litellm_prometheus_multiproc
template: gateway/deployment.yaml
- contains:
path: spec.template.spec.containers[0].volumeMounts
content:
name: prometheus-multiproc
mountPath: /tmp/litellm_prometheus_multiproc
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].name
value: metrics
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].image
value: ghcr.io/berriai/litellm-gateway:v1.101.0
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].command
value:
- python
- -m
- litellm.proxy.prometheus_metrics_server
- --port
- "4101"
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].env
value:
- name: PROMETHEUS_MULTIPROC_DIR
value: /tmp/litellm_prometheus_multiproc
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].ports
value:
- name: metrics
containerPort: 4101
protocol: TCP
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].volumeMounts
value:
- name: prometheus-multiproc
mountPath: /tmp/litellm_prometheus_multiproc
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].readinessProbe.tcpSocket.port
value: metrics
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].livenessProbe.tcpSocket.port
value: metrics
template: gateway/deployment.yaml
- equal:
path: spec.template.spec.containers[1].resources.requests.cpu
value: 50m
template: gateway/deployment.yaml
- contains:
path: spec.template.spec.volumes
content:
name: prometheus-multiproc
emptyDir: {}
template: gateway/deployment.yaml
- lengthEqual:
path: spec.ports
count: 1
template: gateway/service.yaml
- equal:
path: spec.type
value: LoadBalancer
template: gateway/service.yaml
- equal:
path: metadata.name
value: RELEASE-NAME-litellm-gateway-metrics
template: gateway/service-metrics.yaml
- equal:
path: spec.type
value: ClusterIP
template: gateway/service-metrics.yaml
- equal:
path: spec.ports
value:
- port: 4101
targetPort: metrics
protocol: TCP
name: metrics
template: gateway/service-metrics.yaml
- equal:
path: spec.selector
value:
app.kubernetes.io/name: litellm
app.kubernetes.io/instance: RELEASE-NAME
app.kubernetes.io/component: gateway
template: gateway/service-metrics.yaml
- it: rejects a metrics port equal to the gateway port
template: gateway/deployment.yaml
set:
gateway.metricsServer.enabled: true
gateway.metricsServer.port: 4000
asserts:
- failedTemplate:
errorMessage: gateway.metricsServer.port must differ from the gateway port 4000

View file

@ -268,6 +268,22 @@ gateway:
config:
create: true
proxy_config: {}
# Serve Prometheus /metrics from a `metrics` sidecar container (same image,
# `python -m litellm.proxy.prometheus_metrics_server`) that aggregates the
# workers' PROMETHEUS_MULTIPROC_DIR samples over a shared emptyDir, so a
# scrape never runs on an inference worker. Adds a `metrics` port to the pod
# and a dedicated ClusterIP `<gateway>-metrics` Service; point your scrape
# config at it. The port has no virtual-key auth: keep it off public ingress.
# Needs the gateway image v1.101.0 or newer.
metricsServer:
enabled: false
port: 4001
resources:
requests:
cpu: 50m
memory: 128Mi
limits:
memory: 512Mi
image:
repository: ghcr.io/berriai/litellm-gateway
tag: "" # defaults to .Chart.AppVersion

View file

@ -0,0 +1,4 @@
-- Add org_id column to LiteLLM_ManagedObjectTable
-- Snapshots the creating key's organization at submission time, like team_id,
-- so CheckBatchCost can bill organization spend hours later without re-resolving
ALTER TABLE "LiteLLM_ManagedObjectTable" ADD COLUMN IF NOT EXISTS "org_id" TEXT;

View file

@ -1036,6 +1036,7 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t
created_at DateTime @default(now())
created_by String?
team_id String?
org_id String? // creating key's organization at submission time; CheckBatchCost bills org spend against it
api_key String?
request_tags Json? @default("[]")
updated_at DateTime @updatedAt

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.94"
version = "0.4.95"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.94"
version = "0.4.95"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -47,6 +47,7 @@ from typing import (
)
from litellm.types.integrations.datadog import DatadogInitParams
from litellm.types.integrations.newrelic import NewRelicInitParams
from litellm.litellm_core_utils.core_helpers import drop_params_env_flag
from litellm._logging import (
set_verbose,
_turn_on_debug,
@ -238,7 +239,7 @@ token: Optional[str] = (
)
telemetry = True
max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults
drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False))
drop_params = drop_params_env_flag(os.environ, verbose_logger)
modify_params = bool(os.getenv("LITELLM_MODIFY_PARAMS", False))
use_chat_completions_url_for_anthropic_messages: bool = bool(
os.getenv("LITELLM_USE_CHAT_COMPLETIONS_URL_FOR_ANTHROPIC_MESSAGES", False)
@ -325,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
@ -542,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
@ -2401,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

@ -10,7 +10,7 @@ from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
from litellm.types.llms.openai import Batch
from litellm.types.utils import CallTypes, ModelInfo, Usage
from litellm.types.utils import ModelInfo, Usage
from litellm.utils import token_counter
@ -23,6 +23,8 @@ class BatchCostUsageResult:
models: list[str]
successful_requests: int
failed_requests: int
prompt_cost: float = 0.0
completion_cost: float = 0.0
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
@ -151,7 +153,8 @@ class _LineOutcome(Enum):
@dataclass(frozen=True, slots=True)
class _BatchOutputLineStats:
cost: float
prompt_cost: float
completion_cost: float
prompt_tokens: int
completion_tokens: int
total_tokens: int
@ -214,15 +217,16 @@ def _compute_output_line_stats(
raw_model: Final = response_body.get("model")
response_model: Final = raw_model if isinstance(raw_model, str) and raw_model else None
completion_details: Final = usage.completion_tokens_details
line_prompt_cost, line_completion_cost = _output_line_cost(
usage=usage,
custom_llm_provider=custom_llm_provider,
model_name=model_name,
response_model=response_model,
model_info=model_info,
)
return _BatchOutputLineStats(
cost=_output_line_cost(
response_body=response_body,
usage=usage,
custom_llm_provider=custom_llm_provider,
model_name=model_name,
response_model=response_model,
model_info=model_info,
),
prompt_cost=line_prompt_cost,
completion_cost=line_completion_cost,
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
total_tokens=usage.total_tokens,
@ -234,31 +238,24 @@ def _compute_output_line_stats(
def _output_line_cost(
response_body: Mapping[str, object],
usage: Usage,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: str | None,
response_model: str | None,
model_info: ModelInfo | None,
) -> float:
) -> tuple[float, float]:
"""(prompt_cost, completion_cost) for one output line, priced at batch rates."""
from litellm.cost_calculator import batch_cost_calculator
if model_info is None and custom_llm_provider not in ("anthropic", "bedrock"):
return litellm.completion_cost(
completion_response=response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
cost_model: Final = (
model_name if custom_llm_provider == "bedrock" and model_name else response_model or model_name or ""
)
prompt_cost, completion_cost = batch_cost_calculator(
return batch_cost_calculator(
usage=usage,
model=cost_model,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
return prompt_cost + completion_cost
def _aggregate_batch_cost_usage_models(
@ -291,7 +288,9 @@ def _aggregate_batch_cost_usage_models(
**cache_token_params,
)
batch_models: Final = [model_name] if model_name else [stats.model for stats in line_stats if stats.model]
total_cost: Final = sum((stats.cost for stats in line_stats), 0.0)
total_prompt_cost: Final = sum((stats.prompt_cost for stats in line_stats), 0.0)
total_completion_cost: Final = sum((stats.completion_cost for stats in line_stats), 0.0)
total_cost: Final = total_prompt_cost + total_completion_cost
verbose_logger.debug(
"batch output aggregate: cost=%s usage=%s models=%s successful=%d failed=%d",
total_cost,
@ -306,6 +305,8 @@ def _aggregate_batch_cost_usage_models(
models=batch_models,
successful_requests=successful_requests,
failed_requests=failed_requests,
prompt_cost=total_prompt_cost,
completion_cost=total_completion_cost,
)
@ -330,7 +331,8 @@ def calculate_vertex_ai_batch_cost_and_usage(
"""
from litellm.cost_calculator import batch_cost_calculator
total_cost = 0.0
total_prompt_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
total_completion_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
total_tokens = 0
prompt_tokens = 0
completion_tokens = 0
@ -362,7 +364,8 @@ def calculate_vertex_ai_batch_cost_and_usage(
model=actual_model_name,
custom_llm_provider="vertex_ai",
)
total_cost += p_cost + c_cost
total_prompt_cost += p_cost
total_completion_cost += c_cost
except Exception as e:
verbose_logger.debug("vertex_ai batch cost calculation error for line: %s", str(e))
@ -370,6 +373,7 @@ def calculate_vertex_ai_batch_cost_and_usage(
completion_tokens += _completion
total_tokens += _total
total_cost: Final = total_prompt_cost + total_completion_cost
verbose_logger.info(
"vertex_ai batch cost: cost=%s, prompt=%d, completion=%d, total=%d, successful=%d, failed=%d",
total_cost,
@ -390,6 +394,8 @@ def calculate_vertex_ai_batch_cost_and_usage(
models=[actual_model_name],
successful_requests=successful_requests,
failed_requests=failed_requests,
prompt_cost=total_prompt_cost,
completion_cost=total_completion_cost,
)

View file

@ -370,7 +370,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
and isinstance(tool_call.get("custom"), dict)
)
for msg in messages:
leading_system_count: Final = next(
(index for index, msg in enumerate(messages) if msg.get("role") != "system"),
len(messages),
)
for index, msg in enumerate(messages):
role = msg.get("role")
content = msg.get("content", "")
tool_calls = msg.get("tool_calls")
@ -378,7 +383,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
if role == "system":
# Extract system message as instructions
if isinstance(content, str):
if isinstance(content, str) and index < leading_system_count:
if instructions:
# Concatenate multiple system prompts with a space
instructions = f"{instructions} {content}"

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",
]
)
@ -1461,7 +1465,10 @@ LITELLM_METADATA_FIELD: Final = "litellm_metadata"
OLD_LITELLM_METADATA_FIELD: Final = "metadata"
RETURN_RAW_MODEL_NAME_METADATA_KEY: Final = "_complexity_router_return_raw_model_name"
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY: Final = "_session_deployment_affinity_ttl"
OUTPUT_TOKEN_CEILING_PARAMS: Final = frozenset({"max_tokens", "max_completion_tokens", "max_output_tokens"})
CLIENT_OUTPUT_CEILING_METADATA_KEY: Final = "_client_output_ceiling"
CONSUMED_REQUEST_TAGS_METADATA_KEY: Final = "_consumed_request_tags"
ROUTING_REQUEST_TAGS_METADATA_KEY: Final = "_routing_request_tags"
INTERNAL_CALL_ORIGIN_METADATA_KEY: Final = "internal_call_origin"
SESSION_ID_GENERATED_METADATA_KEY: Final = "litellm_session_id_generated"
SESSION_ID_OMITTED_METADATA_KEY: Final = "litellm_session_id_omitted"
@ -1763,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
@ -2414,6 +2423,46 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
)
_RESPONSES_WS_BILLABLE_EVENT_TYPES: Final = frozenset({"response.completed", "response.incomplete"})
class _ResponsesWsEventResponse(BaseModel):
usage: Mapping[str, object] | None = None
class _ResponsesWsEvent(BaseModel):
type: str = ""
response: _ResponsesWsEventResponse | None = None
class ResponsesWebSocketTokenUsageProcessor(BaseTokenUsageProcessor):
@staticmethod
def collect_usage_from_responses_ws_results(
results: Sequence[Mapping[str, object]],
) -> tuple[Usage, ...]:
events: Final = tuple(_ResponsesWsEvent.model_validate(result) for result in results)
return tuple(
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( # pyright: ignore[reportPrivateUsage] # same shared transform the realtime processor uses
event.response.usage
)
for event in events
if event.type in _RESPONSES_WS_BILLABLE_EVENT_TYPES
and event.response is not None
and event.response.usage is not None
)
@staticmethod
def collect_and_combine_usage_from_responses_ws_results(
results: Sequence[Mapping[str, object]],
) -> Usage:
collected_usage_objects: Final = ResponsesWebSocketTokenUsageProcessor.collect_usage_from_responses_ws_results(
results
)
return ResponsesWebSocketTokenUsageProcessor.combine_usage_objects(
list(collected_usage_objects) # mutable-ok: combine_usage_objects requires a list parameter
)
_TRANSCRIPTION_COMPLETED_EVENT_TYPE: Final = "conversation.item.input_audio_transcription.completed"

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

@ -356,11 +356,24 @@
"description": "OpenTelemetry collector endpoint URL",
"required": true
},
"otel_traces_endpoint": {
"type": "text",
"ui_name": "Traces Endpoint URL",
"description": "Complete trace export URL used verbatim when the collector does not serve /v1/traces (OTel v2 only)",
"required": false
},
"otel_headers": {
"type": "text",
"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,9 +69,17 @@ 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(
default=None,
description=(
"Complete OTLP/HTTP trace URL, used verbatim. Set this when the "
"collector serves traces on a path other than ``/v1/traces``; "
"``endpoint`` is a base URL the signal path is appended to."
),
)
headers: str | None = None
owner: ExporterOwner | None = Field(
default=None,
@ -127,6 +135,14 @@ class OpenTelemetryV2Config(BaseSettings):
default=None,
validation_alias=AliasChoices("OTEL_ENDPOINT", "OTEL_EXPORTER_OTLP_ENDPOINT"),
)
traces_endpoint: str | None = Field(
default=None,
validation_alias=AliasChoices("OTEL_TRACES_ENDPOINT", "OTEL_EXPORTER_OTLP_TRACES_ENDPOINT"),
description=(
"Complete OTLP/HTTP trace URL for the single-destination shorthand, "
"used verbatim instead of ``endpoint`` + ``/v1/traces``."
),
)
headers: str | None = Field(
default=None,
validation_alias=AliasChoices("OTEL_HEADERS", "OTEL_EXPORTER_OTLP_HEADERS"),
@ -250,17 +266,22 @@ class OpenTelemetryV2Config(BaseSettings):
@model_validator(mode="after")
def _normalize(self) -> "OpenTelemetryV2Config":
# An endpoint with the default exporter kind implies OTLP/HTTP.
if self.endpoint and self.exporter == "console":
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(
kind=self.exporter,
endpoint=self.endpoint,
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,13 +188,20 @@ 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,
)
return HTTPExporter(
endpoint=_otlp_traces_endpoint(spec.endpoint),
endpoint=spec.traces_endpoint or _otlp_traces_endpoint(spec.endpoint),
headers=parse_headers(spec.headers),
)
if kind in _OTLP_GRPC_KINDS:
@ -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.
@ -201,7 +781,14 @@ def build_span_exporter(config: OpenTelemetryV2Config) -> SpanExporter:
``exporter`` / ``endpoint`` / ``headers`` fields. To configure multiple
exporters, populate ``config.exporters`` directly.
"""
return _exporter_from_spec(ExporterSpec(kind=config.exporter, endpoint=config.endpoint, headers=config.headers))
return _exporter_from_spec(
ExporterSpec(
kind=config.exporter,
endpoint=config.endpoint,
traces_endpoint=config.traces_endpoint,
headers=config.headers,
)
)
def _otlp_metrics_endpoint(endpoint: str | None) -> str | None:
@ -437,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`.
@ -445,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:
@ -461,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

@ -1,10 +1,12 @@
# What is this?
## Helper utilities
import copy
import logging
from collections.abc import Iterable, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.types.llms.openai import AllMessageValues, OpenAIChatCompletionFinishReason
@ -37,6 +39,41 @@ def safe_divide_seconds(seconds: float, denominator: float, default: float | Non
return float(seconds / denominator)
_DROP_PARAMS_BOOL: Final = TypeAdapter(bool)
def normalize_drop_params(value: object) -> bool | None:
if value is None or isinstance(value, bool):
return value
try:
return _DROP_PARAMS_BOOL.validate_python(value.strip() if isinstance(value, str) else value)
except ValidationError:
return None
def drop_params_flag(value: object, source: str, logger: logging.Logger) -> bool:
normalized: Final = normalize_drop_params(value)
if normalized is None and value is not None:
logger.warning("%s=%r is not a flag value, treating it as off", source, value)
return bool(normalized)
DROP_PARAMS_ENV_VAR: Final = "LITELLM_DROP_PARAMS"
def drop_params_env_flag(environ: Mapping[str, str], logger: logging.Logger) -> bool:
configured: Final = environ.get(DROP_PARAMS_ENV_VAR, "").strip()
if configured == "":
return False
normalized: Final = normalize_drop_params(configured)
if normalized is None:
logger.warning(
"%s=%r is not a flag value, treating it as on. Set it to true or false", DROP_PARAMS_ENV_VAR, configured
)
return True
return normalized
def safe_divide(
numerator: float,
denominator: float,

View file

@ -2,6 +2,7 @@ from collections.abc import Mapping, MutableMapping
from types import MappingProxyType
from typing import Final
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.llms.openai.data_residency import infer_openai_data_residency
AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
@ -113,7 +114,7 @@ def get_litellm_params(
custom_prompt_dict: dict | None = None,
litellm_metadata: dict | None = None,
disable_add_transform_inline_image_block: bool | None = None,
drop_params: bool | None = None,
drop_params: bool | str | None = None,
prompt_id: str | None = None,
prompt_variables: dict | None = None,
async_call: bool | None = None,
@ -175,7 +176,7 @@ def get_litellm_params(
"custom_prompt_dict": custom_prompt_dict,
"litellm_metadata": litellm_metadata,
"disable_add_transform_inline_image_block": disable_add_transform_inline_image_block,
"drop_params": drop_params,
"drop_params": normalize_drop_params(drop_params),
"prompt_id": prompt_id,
"prompt_variables": prompt_variables,
"async_call": async_call,

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

@ -48,6 +48,7 @@ from litellm.constants import (
)
from litellm.cost_calculator import (
RealtimeAPITokenUsageProcessor,
ResponsesWebSocketTokenUsageProcessor,
_select_model_name_for_cost_calc,
)
from litellm.exceptions import (
@ -201,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 (
@ -447,6 +450,13 @@ def _provider_response_id(source: object) -> str | None:
return candidate if isinstance(candidate, str) and candidate else None
def mask_api_base_credentials(api_base: str) -> str:
if "key=" not in api_base:
return api_base
key_end: Final = api_base.find("key=") + 4
return api_base[:key_end] + "*" * 5 + api_base[-4:]
class Logging(LiteLLMLoggingBaseClass):
global \
supabaseClient, \
@ -581,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
@ -1189,14 +1200,7 @@ class Logging(LiteLLMLoggingBaseClass):
return data
def _get_masked_api_base(self, api_base: str) -> str:
if "key=" in api_base:
# Find the position of "key=" in the string
key_index: Final = api_base.find("key=") + 4
# Mask the last 5 characters after "key="
masked_api_base = api_base[:key_index] + "*" * 5 + api_base[-4:]
else:
masked_api_base = api_base
return str(masked_api_base)
return str(mask_api_base_credentials(api_base))
def _pre_call(self, input, api_key, model=None, additional_args={}):
"""
@ -1585,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.
@ -1604,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,
@ -2028,6 +2035,17 @@ class Logging(LiteLLMLoggingBaseClass):
results=result,
)
elif self.call_type == CallTypes.aresponses_websocket.value and isinstance(result, list): # pyright: ignore[reportUnknownMemberType] # Logging.call_type is untyped
combined_ws_usage: Final = (
ResponsesWebSocketTokenUsageProcessor.collect_and_combine_usage_from_responses_ws_results(
results=result # pyright: ignore[reportUnknownArgumentType] # raw event dicts from the WS stream
)
)
logging_result = LiteLLMRealtimeStreamLoggingObject(
usage=combined_ws_usage,
results=result, # pyright: ignore[reportUnknownArgumentType] # raw event dicts from the WS stream
)
elif (
self.call_type == CallTypes.llm_passthrough_route.value
or self.call_type == CallTypes.allm_passthrough_route.value
@ -2938,6 +2956,19 @@ class Logging(LiteLLMLoggingBaseClass):
result._hidden_params["batch_successful_requests"] = batch_successful_requests # pyright: ignore[reportPrivateUsage] # rebind-ok: same result._hidden_params pattern as response_cost/batch_models above
result._hidden_params["batch_failed_requests"] = batch_failed_requests # pyright: ignore[reportPrivateUsage] # rebind-ok: same pattern as above
result.usage = batch_usage
batch_prompt_cost: Final = kwargs.get("batch_prompt_cost", None)
batch_completion_cost: Final = kwargs.get("batch_completion_cost", None)
if (
isinstance(batch_prompt_cost, float)
and isinstance(batch_completion_cost, float)
and isinstance(batch_cost, float)
):
self.set_cost_breakdown(
input_cost=batch_prompt_cost,
output_cost=batch_completion_cost,
total_cost=batch_cost,
cost_for_built_in_tools_cost_usd_dollar=0.0,
)
elif should_compute_batch_data:
batch_result: Final = await _handle_completed_batch(
@ -2953,6 +2984,12 @@ class Logging(LiteLLMLoggingBaseClass):
result._hidden_params["batch_successful_requests"] = batch_result.successful_requests # pyright: ignore[reportPrivateUsage] # rebind-ok: same pattern as above
result._hidden_params["batch_failed_requests"] = batch_result.failed_requests # pyright: ignore[reportPrivateUsage] # rebind-ok: same pattern as above
result.usage = batch_result.usage
self.set_cost_breakdown(
input_cost=batch_result.prompt_cost,
output_cost=batch_result.completion_cost,
total_cost=batch_result.cost,
cost_for_built_in_tools_cost_usd_dollar=0.0,
)
self.truncated_messages_for_logging = await truncate_base64_in_messages_async(
StandardLoggingPayloadSetup.append_system_prompt_messages(
@ -4819,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

@ -13,9 +13,11 @@ Pattern Overview:
"""
import json
from collections.abc import Mapping, Sequence
from collections.abc import Mapping, MutableSequence, Sequence
from copy import deepcopy
from dataclasses import dataclass
from itertools import chain, repeat
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, overload, runtime_checkable
from typing_extensions import ReadOnly, TypedDict, assert_never
@ -29,6 +31,7 @@ from litellm.llms.anthropic.experimental_pass_through.adapters.transformation im
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
StreamingScanKey,
StreamTransformSink,
)
from litellm.llms.base_llm.guardrail_translation.utils import (
anthropic_tool_name,
@ -39,6 +42,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
merge_guardrailed_scoped_messages,
merge_returned_tools_into_request_tools,
scoped_structured_message_indices,
stream_item_field,
stream_item_fingerprint,
)
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
@ -151,6 +155,46 @@ class ExtractedInput:
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
@dataclass(frozen=True, slots=True)
class _ToolCallShape:
name: str | None
arguments: str
@dataclass(frozen=True, slots=True)
class _SSEFieldRewrite:
"""One field of one nested section of a buffered SSE event, rewritten."""
section: str
field: str
value: object
class _SSEEventRewriter(Protocol):
def __call__(self, event: Mapping[str, object]) -> _SSEFieldRewrite | None: ...
def _rewritten_event(event: Mapping[str, object], rewrite_event: _SSEEventRewriter) -> Mapping[str, object]:
rewrite: Final = rewrite_event(event)
section: Final = None if rewrite is None else event.get(rewrite.section)
if rewrite is None or not isinstance(section, Mapping):
return event
return {**event, rewrite.section: {**section, rewrite.field: rewrite.value}} # mutable-ok: json.dumps needs a dict
def _tool_call_shapes(tool_calls: Sequence[object]) -> tuple[_ToolCallShape, ...]:
"""The guardrail-visible shape of each tool call, whether the guardrail handed
back the ``ChatCompletionMessageToolCall`` objects it was given or plain dicts."""
functions: Final = tuple(stream_item_field(tool_call, "function") for tool_call in tool_calls)
return tuple(
_ToolCallShape(
name=name if isinstance(name := stream_item_field(function, "name"), str) else None,
arguments=arguments if isinstance(arguments := stream_item_field(function, "arguments"), str) else "",
)
for function in functions
)
class _AnthropicSSEDelta(TypedDict, total=False):
type: ReadOnly[str]
text: ReadOnly[str]
@ -168,6 +212,8 @@ class AnthropicMessagesHandler(BaseTranslation):
them through guardrail rewrites; downstream provider handling is out of scope.
"""
delivers_ended_stream_rewrites = True
def __init__(self):
super().__init__()
self.adapter = LiteLLMAnthropicMessagesAdapter()
@ -1014,11 +1060,17 @@ class AnthropicMessagesHandler(BaseTranslation):
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
stream_transform_sink: StreamTransformSink | None = None,
deliver_ended_stream_rewrites: bool = False,
) -> Sequence[object]:
"""
Process output streaming response by applying guardrails to text content.
Get the string so far, check the apply guardrail to the string so far, and return the list of responses so far.
With ``deliver_ended_stream_rewrites``, an ended stream whose guardrail rewrote the text gets the rewrite
written back across the buffered chunks (full rewritten text in the first ``text_delta``, the rest blanked);
a rewrite on a stream that never reported a ``stop_reason`` has no write-back and is reported as
undeliverable, so the pipeline executor discards it and releases the original chunks.
"""
from litellm.integrations.custom_guardrail import ModifyResponseException
@ -1040,6 +1092,7 @@ class AnthropicMessagesHandler(BaseTranslation):
first_choice.message.tool_calls,
)
string_so_far = first_choice.message.content
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_list or ())
guardrail_inputs: Final = GenericGuardrailAPIInputs()
if string_so_far:
guardrail_inputs["texts"] = [string_so_far]
@ -1065,6 +1118,28 @@ class AnthropicMessagesHandler(BaseTranslation):
responses_so_far, request_data
)
raise
guardrailed_texts: Final = _guardrailed_inputs.get("texts")
if (
deliver_ended_stream_rewrites
and isinstance(string_so_far, str)
and string_so_far
and guardrailed_texts
and guardrailed_texts[0] != string_so_far
):
self._write_ended_stream_text_rewrite(responses_so_far, guardrailed_texts[0])
if deliver_ended_stream_rewrites:
returned_tool_calls: Final = _guardrailed_inputs.get("tool_calls")
self._write_ended_stream_tool_call_rewrites(
responses_so_far,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
post_guardrail_tool_calls=_tool_call_shapes(
returned_tool_calls
if isinstance(returned_tool_calls, list)
and len(returned_tool_calls) == len(pre_guardrail_tool_calls)
else tool_calls_list or ()
),
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
)
else:
verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
return responses_so_far
@ -1087,6 +1162,11 @@ class AnthropicMessagesHandler(BaseTranslation):
if e.original_response is None:
e.original_response = self._build_streaming_usage_response(responses_so_far, request_data)
raise
unended_texts: Final = _guardrailed_inputs.get("texts")
if deliver_ended_stream_rewrites and unended_texts and tuple(unended_texts) != (string_so_far,):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_to_apply.guardrail_name or "unknown")
return responses_so_far
def _prepare_request_data(
@ -1180,6 +1260,139 @@ class AnthropicMessagesHandler(BaseTranslation):
inputs["model"] = response_model
return inputs
@staticmethod
def _write_ended_stream_text_rewrite(
responses_so_far: MutableSequence[object], # mutable-ok: rewrites the caller's buffered chunks in place
rewritten_text: str,
) -> None:
"""Deliver an ended-stream guardrail text rewrite by rewriting the
buffered chunks in place: the first ``text_delta`` carries the full
rewritten text and every later one is blanked, leaving the surrounding
message and content-block framing untouched."""
replacements: Final = chain((rewritten_text,), repeat(""))
def rewrite_text_delta(event: Mapping[str, object]) -> _SSEFieldRewrite | None:
delta: Final = event.get("delta")
if event.get("type") != "content_block_delta" or not isinstance(delta, Mapping):
return None
if delta.get("type") != "text_delta":
return None
return _SSEFieldRewrite("delta", "text", next(replacements))
AnthropicMessagesHandler._rewrite_ended_stream_events(responses_so_far, rewrite_text_delta)
@classmethod
def _write_ended_stream_tool_call_rewrites(
cls,
responses_so_far: MutableSequence[object], # mutable-ok: rewrites the caller's buffered chunks in place
*,
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
post_guardrail_tool_calls: tuple[_ToolCallShape, ...],
guardrail_name: str,
) -> None:
"""Deliver ended-stream guardrail tool-call rewrites by rewriting the
buffered chunks in place: the rebuilt response lists tool calls in the
order of the stream's ``tool_use`` blocks, so the nth rewritten call lands
on the nth block, its first ``input_json_delta`` carrying the full rewritten
arguments, every later one blanked, and ``content_block_start`` carrying the
rewritten name. Blocks that do not line up with the rebuilt tool calls make
the rewrite undeliverable, so the pipeline executor discards it and releases
the original chunks."""
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
block_indices: Final = tuple(
index
for item in responses_so_far
for event in cls._iter_sse_events(item)
if event.get("type") == "content_block_start"
and isinstance(block := event.get("content_block"), Mapping)
and block.get("type") == "tool_use"
and isinstance(index := event.get("index"), int)
)
if len(block_indices) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
rewrites_by_block: Final = MappingProxyType(
{
index: after
for index, before, after in zip(block_indices, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if after != before
}
)
argument_replacements: Final = MappingProxyType(
{index: chain((rewrite.arguments,), repeat("")) for index, rewrite in rewrites_by_block.items()}
)
def rewrite_tool_use(event: Mapping[str, object]) -> _SSEFieldRewrite | None:
index: Final = event.get("index")
if not isinstance(index, int) or index not in rewrites_by_block:
return None
match event.get("type"):
case "content_block_start":
name: Final = rewrites_by_block[index].name
if name is None:
return None
return _SSEFieldRewrite("content_block", "name", name)
case "content_block_delta":
delta: Final = event.get("delta")
if not isinstance(delta, Mapping) or delta.get("type") != "input_json_delta":
return None
return _SSEFieldRewrite("delta", "partial_json", next(argument_replacements[index]))
case _:
return None
cls._rewrite_ended_stream_events(responses_so_far, rewrite_tool_use)
@staticmethod
def _rewrite_ended_stream_events(
responses_so_far: MutableSequence[object], # mutable-ok: rewrites the caller's buffered chunks in place
rewrite_event: _SSEEventRewriter,
) -> None:
"""Replace every buffered event ``rewrite_event`` returns a rewrite for, in
both chunk formats this stream carries (parsed event dicts and raw SSE
bytes), leaving every other event and the framing untouched."""
rewritten_items: Final = tuple(
AnthropicMessagesHandler._rewrite_buffered_item(item, rewrite_event) for item in responses_so_far
)
responses_so_far[:] = rewritten_items # rebind-ok: delivers the rewrites into the caller's buffer
@staticmethod
def _rewrite_buffered_item(item: object, rewrite_event: _SSEEventRewriter) -> object:
if isinstance(item, dict):
return _rewritten_event(_as_str_mapping(item), rewrite_event)
if isinstance(item, (bytes, bytearray)):
return AnthropicMessagesHandler._rewrite_sse_events(bytes(item), rewrite_event)
return item
@staticmethod
def _rewrite_sse_events(sse_bytes: bytes, rewrite_event: _SSEEventRewriter) -> bytes:
"""Rewrite the data lines of one SSE chunk that ``rewrite_event`` rewrites,
leaving all other events and framing byte-identical."""
try:
decoded: Final = sse_bytes.decode("utf-8")
except UnicodeDecodeError:
return sse_bytes
return "\n\n".join(
"\n".join(AnthropicMessagesHandler._rewrite_sse_line(line, rewrite_event) for line in block.split("\n"))
for block in decoded.split("\n\n")
).encode("utf-8")
@staticmethod
def _rewrite_sse_line(line: str, rewrite_event: _SSEEventRewriter) -> str:
if not line.startswith("data:"):
return line
try:
data: Final[str | int | float | bool | None | Sequence[object] | Mapping[str, object]] = json.loads(
line[len("data:") :].strip()
)
except json.JSONDecodeError:
return line
if not isinstance(data, dict):
return line
rewritten: Final = _rewritten_event(_as_str_mapping(data), rewrite_event)
return line if rewritten is data else "data: " + json.dumps(rewritten)
def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None:
stream_ended: Final = self._check_streaming_has_ended(responses_so_far)
return StreamingScanKey(

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

@ -1,7 +1,7 @@
from abc import ABC, abstractmethod
from collections.abc import Sequence
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Final, Optional
from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional
if TYPE_CHECKING:
from fastapi import HTTPException
@ -52,6 +52,15 @@ class StreamingScanKey:
class BaseTranslation(ABC):
delivers_ended_stream_rewrites: ClassVar[bool] = False
"""Whether ``process_output_streaming_response`` accepts
``deliver_ended_stream_rewrites=True`` and, on an ended (fully buffered)
stream, writes guardrail text and tool-call rewrites back across
``responses_so_far`` so a buffered pipeline can release rewritten chunks,
raising ``UndeliverableStreamRewrite`` for a shape it cannot place. Rewrites
on every other translation are undeliverable: the pipeline executor
discards them and releases the original chunks."""
@staticmethod
def transform_user_api_key_dict_to_metadata(
user_api_key_dict: Any | None,
@ -157,6 +166,7 @@ class BaseTranslation(ABC):
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
request_data: dict | None = None,
stream_transform_sink: StreamTransformSink | None = None,
deliver_ended_stream_rewrites: bool = False,
) -> Any:
"""
Process output streaming response with guardrails.
@ -164,6 +174,11 @@ class BaseTranslation(ABC):
Optional to override in subclasses. ``stream_transform_sink`` is the
out-parameter used by handlers that support streaming text
transformations (see ``StreamTransformSink``); base handlers ignore it.
``deliver_ended_stream_rewrites`` is passed True only when the caller
holds the whole buffered stream and the subclass declares
``delivers_ended_stream_rewrites``: the handler then writes
guardrail text and tool-call rewrites back across ``responses_so_far``
instead of discarding them.
"""
return responses_so_far

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

@ -36,7 +36,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
@ -254,7 +254,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(
@ -500,12 +500,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

@ -580,7 +580,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(
@ -627,7 +627,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(
@ -9837,7 +9837,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,
)
@ -9884,6 +9884,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)
@ -9968,7 +9974,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,
)
@ -10013,7 +10019,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)
@ -10043,6 +10056,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
)
@ -10113,6 +10128,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

@ -0,0 +1,9 @@
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from .transformation import HostedVLLMImageEditConfig
__all__ = ("HostedVLLMImageEditConfig",)
def get_hosted_vllm_image_edit_config(model: str) -> BaseImageEditConfig:
return HostedVLLMImageEditConfig()

View file

@ -0,0 +1,43 @@
from typing import Final
from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig
from litellm.secret_managers.main import get_secret_str
PARAMS_VLLM_OMNI_DOES_NOT_ACCEPT: Final = frozenset({"mask", "quality", "input_fidelity"})
class HostedVLLMImageEditConfig(OpenAIImageEditConfig):
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: BaseImageEditConfig contract
return [ # mutable-ok: BaseImageEditConfig returns list
param
for param in super().get_supported_openai_params(model)
if param not in PARAMS_VLLM_OMNI_DOES_NOT_ACCEPT
]
def validate_environment(
self,
headers: dict, # mutable-ok: BaseImageEditConfig contract
model: str,
api_key: str | None = None,
litellm_params: dict | None = None, # mutable-ok: BaseImageEditConfig contract
api_base: str | None = None,
) -> dict: # mutable-ok: BaseImageEditConfig contract
resolved_key: Final = api_key or get_secret_str("HOSTED_VLLM_API_KEY") or "fake-api-key"
return {**headers, "Authorization": f"Bearer {resolved_key}"} # mutable-ok: httpx headers are a dict
def get_complete_url(
self,
model: str,
api_base: str | None,
litellm_params: dict, # mutable-ok: BaseImageEditConfig contract
) -> str:
resolved_api_base: Final = api_base or get_secret_str("HOSTED_VLLM_API_BASE")
if resolved_api_base is None:
raise ValueError(
"api_base not set for Hosted VLLM images edits API. "
"Set via api_base parameter or HOSTED_VLLM_API_BASE environment variable"
)
trimmed: Final = resolved_api_base.rstrip("/")
if trimmed.endswith("/v1"):
return f"{trimmed}/images/edits"
return f"{trimmed}/v1/images/edits"

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

@ -49,6 +49,8 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
coerce_stream_holdback_value,
)
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
Choices,
GenericGuardrailAPIInputs,
ModelResponse,
@ -78,6 +80,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
Methods can be overridden to customize behavior for different message formats.
"""
delivers_ended_stream_rewrites = True
def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None:
"""
Convert chat completions request data to OpenAI-spec structured messages.
@ -453,6 +457,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
stream_transform_sink: StreamTransformSink | None = None,
deliver_ended_stream_rewrites: bool = False,
) -> list["ModelResponseStream"]:
"""
Process output streaming responses by applying guardrails to text content.
@ -467,6 +472,10 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
accumulated text (``responses_so_far`` is left untouched so it stays
a correct raw accumulator across rounds) and the guardrailed text
plus requested holdback are reported per choice on the sink.
deliver_ended_stream_rewrites: When True and the buffered stream has
ended, guardrail text rewrites are written back across
``responses_so_far`` (full rewritten text in each choice's first
content-carrying chunk, the rest blanked) instead of discarded.
Returns:
The (unmodified) list of responses.
@ -492,6 +501,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
litellm_logging_obj=litellm_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
deliver_ended_stream_rewrites=deliver_ended_stream_rewrites,
)
async def _process_streaming_block_only(
@ -502,27 +512,23 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
litellm_logging_obj: "LiteLLMLoggingObj | None",
user_api_key_dict: "UserAPIKeyAuth | None",
request_data: dict | None,
deliver_ended_stream_rewrites: bool = False,
) -> list["ModelResponseStream"]:
"""Block-only streaming path: run the guardrail so an in-flight BLOCK can
terminate the stream. Text rewrites are not propagated to the client here
(see ``_process_streaming_transform`` for the incremental_diff path)."""
(see ``_process_streaming_transform`` for the incremental_diff path) unless
``deliver_ended_stream_rewrites`` opts the ended-stream branch in."""
has_stream_ended: Final = self._first_choice_has_finished(responses_so_far)
if has_stream_ended:
# convert to model response
model_response: Final = cast(
ModelResponse,
stream_chunk_builder(chunks=responses_so_far, logging_obj=litellm_logging_obj),
)
# run process_output_response
await self.process_output_response(
response=model_response,
await self._process_ended_stream(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
deliver_ended_stream_rewrites=deliver_ended_stream_rewrites,
)
return responses_so_far
# Step 0: Check if any response has text content to process
@ -595,6 +601,48 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
return responses_so_far
async def _process_ended_stream(
self,
*,
responses_so_far: list["ModelResponseStream"], # mutable-ok: rewrites the caller's buffered chunks in place
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: "LiteLLMLoggingObj | None",
user_api_key_dict: "UserAPIKeyAuth | None",
request_data: dict[str, object] | None, # mutable-ok: same request-payload shape the hooks take
deliver_ended_stream_rewrites: bool,
) -> None:
"""Ended-stream path: rebuild the full response, run the non-streaming
output guardrail against it, and (when opted in) write any text or
tool-call rewrite back across the buffered chunks."""
model_response: Final = cast(
ModelResponse,
stream_chunk_builder(chunks=responses_so_far, logging_obj=litellm_logging_obj),
)
pre_guardrail_texts: Final = self._string_choice_contents(model_response)
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
await self.process_output_response(
response=model_response,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
if not deliver_ended_stream_rewrites:
return
guardrail_name: Final = guardrail_to_apply.guardrail_name or "unknown"
await self._write_ended_stream_text_rewrites(
responses_so_far=responses_so_far,
guardrailed_response=model_response,
pre_guardrail_texts=pre_guardrail_texts,
guardrail_name=guardrail_name,
)
self._write_ended_stream_tool_call_rewrites(
responses_so_far=responses_so_far,
guardrailed_response=model_response,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
guardrail_name=guardrail_name,
)
def build_stream_error_items(
self,
exc: "HTTPException",
@ -745,8 +793,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
"""
combined_texts: Final[dict[tuple[int, int | None], str]] = {}
for response_idx, response in enumerate(responses_so_far):
for choice_idx, choice in enumerate(response.choices):
for response in responses_so_far:
for choice in response.choices:
if isinstance(choice, litellm.StreamingChoices):
content = choice.delta.content
elif isinstance(choice, litellm.Choices):
@ -759,7 +807,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
if isinstance(content, str):
# String content - accumulate for this choice
str_key: tuple[int, int | None] = (choice_idx, None)
str_key: tuple[int, int | None] = (choice.index, None)
if str_key not in combined_texts:
combined_texts[str_key] = ""
combined_texts[str_key] += content
@ -770,7 +818,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
text_str = content_item.get("text")
if text_str:
list_key: tuple[int, int | None] = (
choice_idx,
choice.index,
content_idx,
)
if list_key not in combined_texts:
@ -960,6 +1008,117 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
if "name" in func_dict:
existing_tool_call.function.name = func_dict["name"]
@staticmethod
def _string_choice_contents(response: "ModelResponse") -> tuple[str | None, ...]:
return tuple(
choice.message.content if isinstance(choice.message.content, str) else None for choice in response.choices
)
async def _write_ended_stream_text_rewrites(
self,
responses_so_far: list["ModelResponseStream"], # mutable-ok: rewrites the caller's buffered chunks in place
guardrailed_response: "ModelResponse",
pre_guardrail_texts: tuple[str | None, ...],
guardrail_name: str,
) -> None:
"""Write ended-stream guardrail text rewrites back across the buffered
chunks: the full rewritten text lands in the choice's first
content-carrying chunk and the rest are blanked, the same shape the
in-flight write-back uses. Chunks carrying only finish_reason or usage
stay untouched. A rewrite on a stream carrying more than one distinct
choice index is reported as undeliverable, so the pipeline executor
discards it and releases the original chunks."""
post_guardrail_texts: Final = self._string_choice_contents(guardrailed_response)
changed: Final = tuple(
after
for before, after in zip(pre_guardrail_texts, post_guardrail_texts)
if before is not None and after is not None and after != before
)
if not changed:
return
stream_choice_indices: Final = frozenset(
choice.index for response in responses_so_far for choice in response.choices
)
if len(stream_choice_indices) != 1:
# stream_chunk_builder collapses every choice into one index-0
# choice, so a rewrite of the rebuilt response cannot be attributed
# back to a single choice on an n>1 stream: report it undeliverable
# rather than deliver the rewrite on the wrong choice
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
target_choice_index: Final = next(iter(stream_choice_indices))
await self._apply_guardrail_responses_to_output_streaming(
responses=responses_so_far,
guardrailed_texts=list(changed), # mutable-ok: callee takes lists
task_mappings=[(target_choice_index, None) for _ in changed], # mutable-ok: callee takes lists
)
@staticmethod
def _function_tool_call_shapes(response: "ModelResponse") -> tuple[tuple[str | None, str], ...]:
return tuple(
(tool_call.function.name, tool_call.function.arguments)
for choice in response.choices
for tool_call in choice.message.tool_calls or ()
if isinstance(tool_call, ChatCompletionMessageToolCall)
)
@staticmethod
def _function_tool_call_fragments(
responses_so_far: Sequence["ModelResponseStream"],
) -> tuple[tuple[ChatCompletionDeltaToolCall, ...], ...]:
"""Group the stream's function tool-call fragments by their tool-call index, in
the index order ``stream_chunk_builder`` lists the rebuilt tool calls, keeping
only the indices the builder keeps (an id and a name somewhere in the stream)."""
fragments: Final = tuple(
tool_call
for response in responses_so_far
for choice in response.choices
for tool_call in choice.delta.tool_calls or ()
if isinstance(tool_call, ChatCompletionDeltaToolCall)
)
identified: Final = frozenset(fragment.index for fragment in fragments if fragment.id)
named: Final = frozenset(fragment.index for fragment in fragments if fragment.function.name)
return tuple(
tuple(fragment for fragment in fragments if fragment.index == index) for index in sorted(identified & named)
)
def _write_ended_stream_tool_call_rewrites(
self,
responses_so_far: list["ModelResponseStream"], # mutable-ok: rewrites the caller's buffered chunks in place
guardrailed_response: "ModelResponse",
pre_guardrail_tool_calls: tuple[tuple[str | None, str], ...],
guardrail_name: str,
) -> None:
"""Write ended-stream guardrail tool-call rewrites back across the buffered
chunks: the rewritten name and full arguments land in the tool call's first
fragment and the arguments of its later fragments are blanked, mirroring the
text write-back. A rewrite on a stream carrying more than one distinct choice
index, or whose fragments do not line up with the rebuilt tool calls, is
reported as undeliverable, so the pipeline executor discards it and releases
the original chunks."""
post_guardrail_tool_calls: Final = self._function_tool_call_shapes(guardrailed_response)
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
stream_choice_indices: Final = frozenset(
choice.index for response in responses_so_far for choice in response.choices
)
fragments_by_tool_call: Final = self._function_tool_call_fragments(responses_so_far)
if len(stream_choice_indices) != 1 or len(fragments_by_tool_call) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
for before, (name, arguments), fragments in zip(
pre_guardrail_tool_calls, post_guardrail_tool_calls, fragments_by_tool_call
):
if (name, arguments) == before:
continue
head, *tail = fragments
head.function.name = name
head.function.arguments = arguments
for fragment in tail:
fragment.function.arguments = ""
async def _apply_guardrail_responses_to_output_streaming(
self,
responses: list["ModelResponseStream"],
@ -975,7 +1134,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
Args:
responses: List of ModelResponseStream objects to modify
guardrailed_texts: List of guardrailed text responses (combined from all chunks)
task_mappings: List of tuples (choice_idx, content_idx)
task_mappings: List of tuples (choice_idx, content_idx), where choice_idx
is the choice's ``index`` field, not its position in a chunk's list
Override this method to customize how responses are applied to streaming responses.
"""
@ -991,9 +1151,11 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
# Key: (choice_idx, content_idx), Value: boolean (True if already set)
already_set: Final[dict[tuple[int, int | None], bool]] = {}
# Iterate through all responses and update content
for response_idx, response in enumerate(responses):
for choice_idx_in_response, choice in enumerate(response.choices):
# Iterate through all responses and update content, matching each chunk's
# choice by its index field: on n>1 streams a chunk usually carries one
# choice at list position 0 whose index names the logical choice.
for response in responses:
for choice in response.choices:
if isinstance(choice, litellm.StreamingChoices):
content = choice.delta.content
elif isinstance(choice, litellm.Choices):
@ -1006,7 +1168,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
if isinstance(content, str):
# String content
str_key: tuple[int, int | None] = (choice_idx_in_response, None)
str_key: tuple[int, int | None] = (choice.index, None)
if str_key in guardrail_map:
if str_key not in already_set:
# First chunk - set the complete guardrailed text
@ -1027,7 +1189,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
for content_idx, content_item in enumerate(content):
if "text" in content_item:
list_key: tuple[int, int | None] = (
choice_idx_in_response,
choice.index,
content_idx,
)
if list_key in guardrail_map:

View file

@ -33,7 +33,7 @@ import time
import uuid
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from itertools import accumulate
from itertools import accumulate, chain, repeat
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Union, cast
@ -49,6 +49,7 @@ from litellm.completion_extras.litellm_responses_transformation.transformation i
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
StreamingScanKey,
StreamTransformSink,
)
from litellm.llms.base_llm.guardrail_translation.utils import (
blocked_responses_stream_usage,
@ -100,6 +101,18 @@ if TYPE_CHECKING:
from litellm.types.llms.openai import ResponseInputParam
class _ToolCallShape(NamedTuple):
name: str | None
arguments: str
def _tool_call_shapes(tool_calls: Sequence[ChatCompletionToolCallChunk]) -> tuple[_ToolCallShape, ...]:
return tuple(
_ToolCallShape(name=tool_call["function"].get("name"), arguments=tool_call["function"].get("arguments", ""))
for tool_call in tool_calls
)
class ResponseOutputEnvelope(TypedDict, total=False):
"""Dict form of a Responses API response, as far as guardrail write-back reads it."""
@ -118,6 +131,19 @@ class ResponsesStreamChunk(TypedDict, total=False):
content_index: ReadOnly[int]
_TERMINAL_ENVELOPE_EVENT_TYPES: Final = frozenset(
{
ResponsesAPIStreamEvents.RESPONSE_COMPLETED.value,
ResponsesAPIStreamEvents.RESPONSE_FAILED.value,
ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value,
}
)
_FUNCTION_CALL_ARGUMENT_EVENT_TYPES: Final = frozenset(
{"response.function_call_arguments.delta", "response.function_call_arguments.done"}
)
_OUTPUT_ITEM_EVENT_TYPES: Final = frozenset({"response.output_item.added", "response.output_item.done"})
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"function_call_output": "output", "message": "content"}
)
@ -330,6 +356,8 @@ class OpenAIResponsesHandler(BaseTranslation):
Methods can be overridden to customize behavior for different message formats.
"""
delivers_ended_stream_rewrites = True
def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None:
"""
Convert Responses API request data to OpenAI-spec structured messages.
@ -667,6 +695,8 @@ class OpenAIResponsesHandler(BaseTranslation):
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
user_api_key_dict: "UserAPIKeyAuth | None" = None,
request_data: dict | None = None,
stream_transform_sink: StreamTransformSink | None = None,
deliver_ended_stream_rewrites: bool = False,
) -> list[Any]:
"""
Process output streaming response by applying guardrails to text content.
@ -675,10 +705,18 @@ class OpenAIResponsesHandler(BaseTranslation):
chunk, apply the guardrail, then write the result back in-place so the
caller sees the modified content (e.g. PII tokens replaced).
For ``response.completed`` events (the normal end-of-stream signal) we
use the same per-item extraction + task-mapping approach as
``process_output_response`` so that unmasking / blocking works correctly
for every output item.
For terminal envelope events (``response.completed``, and equally
``response.incomplete`` / ``response.failed``, whose envelopes carry the
partial output) we use the same per-item extraction + task-mapping
approach as ``process_output_response`` so that unmasking / blocking
works correctly for every output item. With
``deliver_ended_stream_rewrites`` the earlier text-carrying events
(``response.output_text.delta`` / ``.done``,
``response.content_part.done``, ``response.output_item.done``) are synced
to the rewritten envelope too, so a client reading deltas sees the
rewrite instead of the raw model output; a rewrite observed where no
write-back is possible is reported as undeliverable, so the pipeline
executor discards it and releases the original events.
"""
if not responses_so_far:
return responses_so_far
@ -690,14 +728,16 @@ class OpenAIResponsesHandler(BaseTranslation):
return responses_so_far
# ------------------------------------------------------------------ #
# Case 1: response.completed — full response is available in the #
# final chunk; iterate output items, apply guardrail, write back. #
# Case 1: terminal envelope events (completed/incomplete/failed). #
# the accumulated response is available in the final chunk; iterate #
# output items, apply guardrail, write back. Falls through to the #
# string fallback when the envelope yields nothing to check. #
# ------------------------------------------------------------------ #
if final_chunk.get("type") == "response.completed":
if final_chunk.get("type") in _TERMINAL_ENVELOPE_EVENT_TYPES:
response_obj: Final[ResponseOutputEnvelope] = final_chunk.get("response") or {}
if not hasattr(response_obj, "get"):
return responses_so_far
outputs: Final[Sequence[object]] = response_obj.get("output") or []
outputs: Final[Sequence[object]] = (
(response_obj.get("output") or []) if hasattr(response_obj, "get") else []
)
texts_to_check: Final[list[str]] = []
tool_calls_to_check: Final[list[ChatCompletionToolCallChunk]] = []
@ -730,6 +770,7 @@ class OpenAIResponsesHandler(BaseTranslation):
if response_model:
inputs["model"] = response_model
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data,
@ -738,6 +779,12 @@ class OpenAIResponsesHandler(BaseTranslation):
)
guardrailed_texts: Final = guardrailed_inputs.get("texts", [])
returned_tool_calls: Final = guardrailed_inputs.get("tool_calls")
post_guardrail_tool_calls: Final = _tool_call_shapes(
returned_tool_calls
if isinstance(returned_tool_calls, list) and len(returned_tool_calls) == len(tool_calls_to_check)
else tool_calls_to_check
)
# Write guardrailed texts back into the output items in-place.
# final_chunk is a reference into responses_so_far so this
@ -747,11 +794,32 @@ class OpenAIResponsesHandler(BaseTranslation):
responses=guardrailed_texts,
task_mappings=task_mappings,
)
return responses_so_far
if deliver_ended_stream_rewrites:
rewrites_by_position: Final = MappingProxyType(
{
task_mappings[task_idx]: rewritten
for task_idx, rewritten in enumerate(guardrailed_texts)
if task_idx < len(texts_to_check) and rewritten != texts_to_check[task_idx]
}
)
if rewrites_by_position:
self._sync_stream_events_with_rewrites(
stream_events=responses_so_far[:-1],
rewrites_by_position=rewrites_by_position,
)
self._deliver_ended_stream_tool_call_rewrites(
responses_so_far=responses_so_far,
outputs=outputs,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
post_guardrail_tool_calls=post_guardrail_tool_calls,
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
)
return responses_so_far
# ------------------------------------------------------------------ #
# Case 2: response.output_item.done — extract tool calls only. #
# Case 2: response.output_item.done — extract tool calls only, then #
# fall through to the text fallback when a caller expects rewrites #
# delivered, so a truncated buffer still reports text undeliverable. #
# ------------------------------------------------------------------ #
if final_chunk.get("type") == "response.output_item.done":
model_response_stream: Final = (
@ -769,12 +837,14 @@ class OpenAIResponsesHandler(BaseTranslation):
input_type="response",
logging_obj=litellm_logging_obj,
)
return responses_so_far
if not deliver_ended_stream_rewrites:
return responses_so_far
# ------------------------------------------------------------------ #
# Fallback: apply guardrail to the accumulated text string. #
# No structured write-back is possible here; guardrails that only #
# need to block/flag (not rewrite) still work correctly. #
# need to block/flag (not rewrite) still work correctly, and a #
# rewrite a caller expects delivered is reported undeliverable. #
# ------------------------------------------------------------------ #
string_so_far: Final = self.get_streaming_string_so_far(responses_so_far)
if string_so_far:
@ -784,28 +854,206 @@ class OpenAIResponsesHandler(BaseTranslation):
)
if response_model:
fallback_inputs["model"] = response_model
await guardrail_to_apply.apply_guardrail(
fallback_outputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=fallback_inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
fallback_texts: Final = fallback_outputs.get("texts")
if deliver_ended_stream_rewrites and fallback_texts and tuple(fallback_texts) != (string_so_far,):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_to_apply.guardrail_name or "unknown")
return responses_so_far
@staticmethod
def _write_event_field(event: object, field: str, value: str) -> None:
if isinstance(event, dict):
event[field] = value # rebind-ok: delivering the rewrite means editing the buffered event in place
else:
setattr(event, field, value)
def _sync_stream_events_with_rewrites(
self,
stream_events: Sequence[Any],
rewrites_by_position: Mapping[tuple[int, int], str],
) -> None:
"""Sync pre-completion stream events with the rewritten completed
response, keyed by ``(output_index, content_index)``: the first
``output_text.delta`` for a rewritten item carries the full rewritten
text and the rest are blanked, while ``output_text.done``,
``content_part.done``, and ``output_item.done`` events carry the full
rewritten text, so every event a client may read agrees with the
rewritten ``response.completed`` payload."""
delta_replacements: Final = MappingProxyType(
{position: chain((rewritten,), repeat("")) for position, rewritten in rewrites_by_position.items()}
)
for event in stream_events:
if not (isinstance(event, dict) or hasattr(event, "get")):
continue
event_type = event.get("type")
output_index = event.get("output_index")
content_index = event.get("content_index")
if event_type == "response.output_item.done" and isinstance(output_index, int):
self._sync_output_item_done_event(event.get("item"), output_index, rewrites_by_position)
continue
if not isinstance(output_index, int) or not isinstance(content_index, int):
continue
position = (output_index, content_index)
if event_type == "response.output_text.delta" and position in delta_replacements:
self._write_event_field(event, "delta", next(delta_replacements[position]))
elif event_type == "response.output_text.done" and position in rewrites_by_position:
self._write_event_field(event, "text", rewrites_by_position[position])
elif event_type == "response.content_part.done" and position in rewrites_by_position:
part = event.get("part")
if isinstance(part, dict) or hasattr(part, "text"):
self._write_event_field(part, "text", rewrites_by_position[position])
@staticmethod
def _sync_output_item_done_event(
item: object,
output_index: int,
rewrites_by_position: Mapping[tuple[int, int], str],
) -> None:
content: Final = item.get("content") if isinstance(item, dict) else getattr(item, "content", None)
if not isinstance(content, list):
return
for (item_idx, content_idx), rewritten in rewrites_by_position.items():
if item_idx != output_index or content_idx >= len(content):
continue
OpenAIResponsesHandler._write_event_field(content[content_idx], "text", rewritten)
def _deliver_ended_stream_tool_call_rewrites(
self,
responses_so_far: Sequence[object],
outputs: Sequence[object],
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
post_guardrail_tool_calls: tuple[_ToolCallShape, ...],
guardrail_name: str,
) -> None:
"""Write ended-stream guardrail tool-call rewrites into the completed
envelope's ``function_call`` items and sync the earlier stream events,
keyed by ``call_id``. The guardrail sees the envelope's function calls
in output order, which is how a rewritten call finds its ``call_id``;
the stream events find their call through the ``call_id`` on
``output_item`` events and the ``item_id`` on argument events, since an
event's ``output_index`` need not match the envelope's (the chat bridge
numbers tool calls from 1 while the envelope lists them after the
message). A rewrite whose calls do not line up with the envelope, or
whose events cannot be found, is reported as undeliverable, so the
pipeline executor discards it and releases the original events."""
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
function_call_items: Final = tuple(
output_item for output_item in outputs if stream_item_field(output_item, "type") == "function_call"
)
call_ids: Final = tuple(
call_id
for output_item in function_call_items
if isinstance(call_id := stream_item_field(output_item, "call_id"), str) and call_id
)
stream_events: Final = responses_so_far[:-1]
call_id_by_item_id: Final = self._function_call_ids_by_item_id(stream_events)
event_call_ids: Final = tuple(
self._function_call_event_call_id(event, call_id_by_item_id) for event in stream_events
)
rewrites_by_call_id: Final = MappingProxyType(
{
call_id: after
for call_id, before, after in zip(call_ids, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if after != before
}
)
unresolved_argument_event: Final = any(
call_id is None and stream_item_field(event, "type") in _FUNCTION_CALL_ARGUMENT_EVENT_TYPES
for event, call_id in zip(stream_events, event_call_ids)
)
if (
len(call_ids) != len(function_call_items)
or len(frozenset(call_ids)) != len(call_ids)
or len(call_ids) != len(post_guardrail_tool_calls)
or unresolved_argument_event
or not rewrites_by_call_id.keys() <= frozenset(event_call_ids)
):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
for output_item, rewrite in (
(output_item, rewrites_by_call_id[call_id])
for output_item, call_id in zip(function_call_items, call_ids)
if call_id in rewrites_by_call_id
):
self._write_function_call_item(output_item, rewrite.name, rewrite.arguments)
delta_replacements: Final = MappingProxyType(
{call_id: chain((rewrite.arguments,), repeat("")) for call_id, rewrite in rewrites_by_call_id.items()}
)
for event, call_id in zip(stream_events, event_call_ids):
if call_id not in rewrites_by_call_id:
continue
match stream_item_field(event, "type"):
case "response.function_call_arguments.delta":
self._write_event_field(event, "delta", next(delta_replacements[call_id]))
case "response.function_call_arguments.done":
self._write_event_field(event, "arguments", rewrites_by_call_id[call_id].arguments)
case "response.output_item.added":
self._write_function_call_item(
stream_item_field(event, "item"), rewrites_by_call_id[call_id].name, None
)
case "response.output_item.done":
self._write_function_call_item(
stream_item_field(event, "item"),
rewrites_by_call_id[call_id].name,
rewrites_by_call_id[call_id].arguments,
)
case _:
pass
@staticmethod
def _function_call_ids_by_item_id(stream_events: Sequence[object]) -> Mapping[str, str]:
items: Final = tuple(
stream_item_field(event, "item")
for event in stream_events
if stream_item_field(event, "type") in _OUTPUT_ITEM_EVENT_TYPES
)
return MappingProxyType(
{
item_id: call_id
for item in items
if stream_item_field(item, "type") == "function_call"
and isinstance(item_id := stream_item_field(item, "id"), str)
and isinstance(call_id := stream_item_field(item, "call_id"), str)
}
)
@staticmethod
def _function_call_event_call_id(event: object, call_id_by_item_id: Mapping[str, str]) -> str | None:
event_type: Final = stream_item_field(event, "type")
if event_type in _FUNCTION_CALL_ARGUMENT_EVENT_TYPES:
item_id: Final = stream_item_field(event, "item_id")
return call_id_by_item_id.get(item_id) if isinstance(item_id, str) else None
if event_type not in _OUTPUT_ITEM_EVENT_TYPES:
return None
item: Final = stream_item_field(event, "item")
call_id: Final = stream_item_field(item, "call_id")
return call_id if stream_item_field(item, "type") == "function_call" and isinstance(call_id, str) else None
@staticmethod
def _write_function_call_item(item: object, name: str | None, arguments: str | None) -> None:
if item is None:
return
if name is not None:
OpenAIResponsesHandler._write_event_field(item, "name", name)
if arguments is not None:
OpenAIResponsesHandler._write_event_field(item, "arguments", arguments)
def _check_streaming_has_ended(self, responses_so_far: Sequence[object]) -> bool:
"""
Check if the streaming has ended.
"""
if not responses_so_far:
return False
terminal_types: Final = frozenset(
(
ResponsesAPIStreamEvents.RESPONSE_COMPLETED.value,
ResponsesAPIStreamEvents.RESPONSE_FAILED.value,
ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value,
)
)
return stream_item_field(responses_so_far[-1], "type") in terminal_types
return stream_item_field(responses_so_far[-1], "type") in _TERMINAL_ENVELOPE_EVENT_TYPES
def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None:
if not responses_so_far or not hasattr(responses_so_far[-1], "get"):

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

@ -32,6 +32,7 @@ class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase):
file_object: LiteLLMBatch | LiteLLMFineTuningJob | ResponsesAPIResponse
created_by: str | None = None
team_id: str | None = None
org_id: str | None = None
class LiteLLM_ManagedVectorStoreTable(LiteLLMPydanticObjectBase):

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:
@ -2819,6 +2836,18 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"UI username/password login. Default is False."
),
)
disable_env_credential_login: bool | None = Field(
None,
description=(
"If True, disables signing in to the Admin UI with the environment credentials: "
"UI_USERNAME/UI_PASSWORD, or the master key when UI_PASSWORD is unset (that fallback "
"means env-credential login is always live by default). Database users with passwords "
"are unaffected. LOCKOUT RISK: create at least one proxy admin user with a password "
"before enabling, or nobody can sign in to the UI. A locked-out admin can still "
"administer the proxy over the API with the master key, and can unset this setting "
"and restart the proxy to restore env-credential login. Default is False."
),
)
disable_budget_reservation: bool | None = Field(
None,
description=(
@ -3584,7 +3613,9 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
ui_callback_name="OpenTelemetry",
litellm_callback_params=[
"OTEL_EXPORTER",
"OTEL_EXPORTER_OTLP_PROTOCOL",
"OTEL_ENDPOINT",
"OTEL_TRACES_ENDPOINT",
"OTEL_HEADERS",
],
)
@ -5137,9 +5168,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."""
@ -5147,6 +5195,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
@ -5154,17 +5205,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,

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