mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
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:
commit
d644e4970f
416 changed files with 33970 additions and 6287 deletions
2
.github/workflows/_test-unit-base.yml
vendored
2
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
4
.github/workflows/test-litellm-ui-unit.yml
vendored
4
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -105,7 +105,7 @@
|
|||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38271
|
||||
"limit": 38269
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19584
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
17
helm/litellm-helm/templates/service-metrics.yaml
Normal file
17
helm/litellm-helm/templates/service-metrics.yaml
Normal 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 }}
|
||||
|
|
@ -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 }}
|
||||
|
|
|
|||
106
helm/litellm-helm/tests/metrics_server_tests.yaml
Normal file
106
helm/litellm-helm/tests/metrics_server_tests.yaml
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -441,3 +441,5 @@ ImplementationSpecific
|
|||
{{- .pathType -}}
|
||||
{{- end -}}
|
||||
{{- end -}}
|
||||
|
||||
{{- define "litellm.gateway.prometheusMultiprocDir" -}}/tmp/litellm_prometheus_multiproc{{- end -}}
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
18
helm/litellm/templates/gateway/service-metrics.yaml
Normal file
18
helm/litellm/templates/gateway/service-metrics.yaml
Normal 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 }}
|
||||
148
helm/litellm/tests/metrics_server_tests.yaml
Normal file
148
helm/litellm/tests/metrics_server_tests.yaml
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
49
litellm/integrations/otel/model/destination.py
Normal file
49
litellm/integrations/otel/model/destination.py
Normal 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, ...]] = ()
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
70
litellm/integrations/otel/plumbing/otlp_json.py
Normal file
70
litellm/integrations/otel/plumbing/otlp_json.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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 {}),
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
|
|
|||
152
litellm/integrations/otel/presets/destinations.py
Normal file
152
litellm/integrations/otel/presets/destinations.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
9
litellm/llms/hosted_vllm/image_edit/__init__.py
Normal file
9
litellm/llms/hosted_vllm/image_edit/__init__.py
Normal 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()
|
||||
43
litellm/llms/hosted_vllm/image_edit/transformation.py
Normal file
43
litellm/llms/hosted_vllm/image_edit/transformation.py
Normal 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"
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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={}
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue