mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
chore: merge litellm_internal_staging into litellm_project_write_routes_team_admin
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
7c0e586d78
215 changed files with 12866 additions and 3565 deletions
|
|
@ -1440,6 +1440,7 @@ jobs:
|
|||
TEST_FILES=$(printf "%s\n" \
|
||||
tests/local_testing/test_dual_cache.py \
|
||||
tests/local_testing/test_redis_batch_optimizations.py \
|
||||
tests/local_testing/test_redis_increment_with_floor.py \
|
||||
tests/local_testing/test_router_utils.py)
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
|
|
|
|||
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
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -8,14 +8,10 @@ import tempfile
|
|||
import time
|
||||
from dataclasses import dataclass, replace
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING, Final, Optional
|
||||
|
||||
from litellm_proxy_extras import prisma_toolchain
|
||||
from litellm_proxy_extras._logging import logger
|
||||
from litellm_proxy_extras.replica_identity import (
|
||||
REPLICA_IDENTITY_FULL_ENV_VAR,
|
||||
apply_replica_identity_full,
|
||||
)
|
||||
from litellm_proxy_extras.prisma_toolchain import (
|
||||
PRISMA_COMMAND_TIMEOUT_ENV_VAR,
|
||||
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR,
|
||||
|
|
@ -23,6 +19,14 @@ from litellm_proxy_extras.prisma_toolchain import (
|
|||
prisma_command_timeout,
|
||||
prisma_migrate_deploy_timeout,
|
||||
)
|
||||
from litellm_proxy_extras.replica_identity import (
|
||||
REPLICA_IDENTITY_FULL_ENV_VAR,
|
||||
apply_replica_identity_full,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import psycopg
|
||||
import psycopg.sql
|
||||
|
||||
|
||||
def str_to_bool(value: Optional[str]) -> bool:
|
||||
|
|
@ -46,6 +50,28 @@ def _get_prisma_env() -> dict:
|
|||
_MIGRATION_TS_RE = re.compile(r"^(\d{14})_")
|
||||
|
||||
_MIGRATION_DEADLOCK_MARKER = "deadlock detected"
|
||||
INDEX_REPAIR_ADVISORY_LOCK_KEY: Final = int.from_bytes(b"litellm", "big")
|
||||
_TRANSIENT_INDEX_SUFFIX_RE: Final = re.compile(r"_cc(?:new|old)\d*$")
|
||||
_INVALID_LITELLM_INDEXES_SQL: Final = (
|
||||
"SELECT n.nspname, c.relname, pg_size_pretty(pg_table_size(t.oid)) "
|
||||
"FROM pg_index i "
|
||||
"JOIN pg_class c ON c.oid = i.indexrelid "
|
||||
"JOIN pg_class t ON t.oid = i.indrelid "
|
||||
"JOIN pg_namespace n ON n.oid = t.relnamespace "
|
||||
"WHERE NOT i.indisvalid "
|
||||
" AND c.relkind = 'i' "
|
||||
" AND n.nspname = %s "
|
||||
" AND t.relname LIKE %s "
|
||||
" AND NOT EXISTS (SELECT 1 FROM pg_constraint k WHERE k.conindid = i.indexrelid) "
|
||||
"ORDER BY c.relname"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _InvalidIndex:
|
||||
schema: str
|
||||
name: str
|
||||
table_size: str
|
||||
|
||||
MAX_MIGRATE_DEPLOY_ATTEMPTS = 4
|
||||
|
||||
|
|
@ -624,7 +650,7 @@ class ProxyExtrasDBManager:
|
|||
def _strip_prisma_query_params(url: str) -> str:
|
||||
"""Remove Prisma-specific query params (connection_limit, pool_timeout,
|
||||
schema, etc.) from DATABASE_URL so psycopg can parse it."""
|
||||
from urllib.parse import urlparse, urlunparse, parse_qsl, urlencode
|
||||
from urllib.parse import parse_qsl, quote, urlencode, urlparse, urlunparse
|
||||
|
||||
parsed = urlparse(url)
|
||||
if not parsed.query:
|
||||
|
|
@ -645,7 +671,7 @@ class ProxyExtrasDBManager:
|
|||
"target_session_attrs",
|
||||
}
|
||||
kept = [(k, v) for k, v in parse_qsl(parsed.query) if k in libpq_params]
|
||||
return urlunparse(parsed._replace(query=urlencode(kept)))
|
||||
return urlunparse(parsed._replace(query=urlencode(kept, quote_via=quote)))
|
||||
|
||||
@staticmethod
|
||||
def _warn_if_db_ahead_of_head(migrations_dir: str) -> None:
|
||||
|
|
@ -719,6 +745,95 @@ class ProxyExtrasDBManager:
|
|||
", ".join(sorted_hostile[:5]) + (" ..." if len(sorted_hostile) > 5 else ""),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _invalid_litellm_indexes(
|
||||
conn: "psycopg.Connection[tuple[str, str, str]]", schema: str
|
||||
) -> tuple[_InvalidIndex, ...]:
|
||||
rows: Final = conn.execute(_INVALID_LITELLM_INDEXES_SQL, (schema, "LiteLLM\\_%")).fetchall()
|
||||
return tuple(_InvalidIndex(*row) for row in rows)
|
||||
|
||||
@staticmethod
|
||||
def _index_repair(index: _InvalidIndex) -> tuple["psycopg.sql.Composed", str]:
|
||||
from psycopg import sql
|
||||
|
||||
target: Final = sql.Identifier(index.schema, index.name)
|
||||
if _TRANSIENT_INDEX_SUFFIX_RE.search(index.name):
|
||||
return sql.SQL("DROP INDEX CONCURRENTLY IF EXISTS {}").format(target), "Dropped leftover"
|
||||
return sql.SQL("REINDEX INDEX CONCURRENTLY {}").format(target), "Rebuilt"
|
||||
|
||||
@staticmethod
|
||||
def _repair_index(conn: "psycopg.Connection[tuple[str, str, str]]", index: _InvalidIndex) -> None:
|
||||
import psycopg
|
||||
|
||||
statement, action = ProxyExtrasDBManager._index_repair(index)
|
||||
try:
|
||||
conn.execute(statement)
|
||||
except psycopg.Error as e:
|
||||
logger.warning(
|
||||
"Could not repair invalid index %s.%s, will retry on the next startup. "
|
||||
"If this keeps happening, run `%s` by hand as the index owner. Error: %s",
|
||||
index.schema,
|
||||
index.name,
|
||||
statement.as_string(conn),
|
||||
e,
|
||||
)
|
||||
return
|
||||
logger.info("%s invalid index %s.%s", action, index.schema, index.name)
|
||||
|
||||
@staticmethod
|
||||
def repair_invalid_indexes(lock_timeout: str = "30s") -> bool:
|
||||
"""Rebuild LiteLLM indexes an interrupted CREATE INDEX CONCURRENTLY left
|
||||
INVALID (a migration deadlock between replicas is the usual cause; the
|
||||
retried migration skips them because of IF NOT EXISTS). Never raises:
|
||||
returns True when no invalid index remains, False when the repair was
|
||||
skipped or failed and will be retried on the next startup. Looks in the
|
||||
schema DATABASE_URL names, the only URL Prisma migrates through, but
|
||||
connects over DIRECT_URL when set: the session settings, the advisory
|
||||
lock and REINDEX CONCURRENTLY all need one server session, which a
|
||||
transaction pooler does not give."""
|
||||
prisma_url: Final = os.getenv("DATABASE_URL")
|
||||
if not prisma_url:
|
||||
return False
|
||||
|
||||
try:
|
||||
import psycopg
|
||||
from psycopg import sql
|
||||
except ImportError:
|
||||
logger.warning(
|
||||
"psycopg is not installed; skipping the invalid index check. "
|
||||
"Install the litellm[extra_proxy] extra, which includes psycopg."
|
||||
)
|
||||
return False
|
||||
|
||||
schema: Final = ProxyExtrasDBManager._prisma_schema_param(prisma_url) or "public"
|
||||
cleaned_url: Final = ProxyExtrasDBManager._strip_prisma_query_params(os.getenv("DIRECT_URL") or prisma_url)
|
||||
try:
|
||||
with psycopg.connect(cleaned_url, connect_timeout=10, autocommit=True) as conn:
|
||||
conn.execute("SET statement_timeout = 0")
|
||||
conn.execute(sql.SQL("SET lock_timeout = {}").format(sql.Literal(lock_timeout)))
|
||||
found: Final = ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema)
|
||||
if not found:
|
||||
return True
|
||||
logger.warning(
|
||||
"Found %d invalid index(es) left by an interrupted CREATE INDEX "
|
||||
"CONCURRENTLY, rebuilding: %s",
|
||||
len(found),
|
||||
", ".join(f"{index.name} (table size {index.table_size})" for index in found),
|
||||
)
|
||||
lock_row: Final = conn.execute(
|
||||
"SELECT pg_try_advisory_lock(%s)", (INDEX_REPAIR_ADVISORY_LOCK_KEY,)
|
||||
).fetchone()
|
||||
if lock_row is None or not lock_row[0]:
|
||||
logger.info("Another replica is already rebuilding the invalid indexes, skipping")
|
||||
return False
|
||||
for index in ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema):
|
||||
ProxyExtrasDBManager._repair_index(conn, index)
|
||||
remaining: Final = ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema)
|
||||
except psycopg.Error as e:
|
||||
logger.warning("Could not check for invalid indexes, will retry on the next startup. Error: %s", e)
|
||||
return False
|
||||
return not remaining
|
||||
|
||||
@staticmethod
|
||||
def _setup_database_v2(use_migrate: bool) -> bool:
|
||||
"""
|
||||
|
|
@ -994,6 +1109,7 @@ class ProxyExtrasDBManager:
|
|||
use_migrate=use_migrate, use_v2_resolver=use_v2_resolver
|
||||
)
|
||||
if migrated:
|
||||
ProxyExtrasDBManager.repair_invalid_indexes()
|
||||
ProxyExtrasDBManager.apply_replica_identity_full_if_requested()
|
||||
return migrated
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ from contextvars import ContextVar
|
|||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -80,11 +82,29 @@ class _AsyncRedisCommands(Protocol):
|
|||
|
||||
def pipeline(self, transaction: bool = True) -> "Pipeline[bytes]": ...
|
||||
|
||||
def eval(self, script: str, numkeys: int, *keys_and_args: str | bytes | float) -> Awaitable[object]: ...
|
||||
|
||||
|
||||
_BREAKER_GUARD_FRAME_NAMES: Final = frozenset(
|
||||
{"<lambda>", "wrapper", "_run_under_circuit_breaker", "_run_under_circuit_breaker_sync"}
|
||||
)
|
||||
|
||||
_INCREMENT_WITH_FLOOR_LUA: Final = (
|
||||
"local count = redis.call('INCRBY', KEYS[1], ARGV[1]) "
|
||||
"if count < 0 then count = redis.call('INCRBY', KEYS[1], -count) end "
|
||||
"if redis.call('TTL', KEYS[1]) < 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]) end "
|
||||
"return count"
|
||||
)
|
||||
|
||||
_LUA_COUNT: Final = TypeAdapter(int)
|
||||
_OPTIONAL_COUNTS: Final = TypeAdapter(tuple[int | None, ...])
|
||||
|
||||
|
||||
def _decoded_counts(values: Sequence[bytes | str | None]) -> tuple[int | None, ...]:
|
||||
return _OPTIONAL_COUNTS.validate_python(
|
||||
tuple(value.decode("utf-8") if isinstance(value, bytes) else value for value in values)
|
||||
)
|
||||
|
||||
|
||||
def _get_call_stack_info(num_frames: int = 2) -> str:
|
||||
"""
|
||||
|
|
@ -736,6 +756,43 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
raise e
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
"""Add ``value`` to ``key``, clamp the result at zero, and give a new key ``ttl``, in one Lua call.
|
||||
|
||||
A counter whose key expired while a request was still in flight would otherwise be
|
||||
recreated negative by that request's decrement. Clamping inside the same call is what
|
||||
keeps it safe: a separate corrective write could land after another pod's increment and
|
||||
erase it.
|
||||
|
||||
The TTL is set only on a key that has none, so a counter expires ``ttl`` after it was
|
||||
created rather than ``ttl`` after it was last touched. Refreshing it on every touch
|
||||
would keep a count a dead worker never decremented alive for as long as the group
|
||||
takes traffic. Returns the resulting count.
|
||||
"""
|
||||
namespaced_key: Final = self.check_and_fix_namespace(key=key)
|
||||
count: Final[object] = self.redis_client.eval( # pyright: ignore[reportAttributeAccessIssue] # stubs omit eval
|
||||
_INCREMENT_WITH_FLOOR_LUA, 1, namespaced_key, value, ttl
|
||||
)
|
||||
return _LUA_COUNT.validate_python(count)
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
"""Read integer counters for ``key_list``, in order, raising when Redis cannot answer.
|
||||
|
||||
``batch_get_cache`` swallows every failure and returns an empty dict, which the caller
|
||||
cannot tell apart from "every counter is unset". A caller that has to fall back to its
|
||||
own numbers when Redis is unreachable needs the failure, not a dict of zeros.
|
||||
"""
|
||||
namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list]
|
||||
return _decoded_counts(self._run_redis_mget_operation(keys=namespaced_keys))
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
"""Async twin of ``batch_get_counts``, raising on failure the same way."""
|
||||
namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list]
|
||||
return _decoded_counts(await self._async_run_redis_mget_operation(keys=namespaced_keys))
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_scan_iter(self, pattern: str, count: int = 100) -> list:
|
||||
start_time: Final = time.time()
|
||||
|
|
@ -1241,6 +1298,14 @@ class RedisCache(BaseCache):
|
|||
result = result.decode()
|
||||
return float(result)
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
"""Async twin of ``increment_with_floor``, sharing its Lua script and its guarantees."""
|
||||
_redis_client: Final = self._async_commands()
|
||||
namespaced_key: Final = self.check_and_fix_namespace(key=key)
|
||||
count: Final = await _redis_client.eval(_INCREMENT_WITH_FLOOR_LUA, 1, namespaced_key, value, ttl)
|
||||
return _LUA_COUNT.validate_python(count)
|
||||
|
||||
async def flush_cache_buffer(self):
|
||||
print_verbose(f"flushing to redis....reached size of buffer {len(self.redis_batch_writing_buffer)}")
|
||||
await self.async_set_cache_pipeline(self.redis_batch_writing_buffer)
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -73,6 +73,9 @@ DEFAULT_MAX_TOKENS: Final = int(os.getenv("DEFAULT_MAX_TOKENS", 4096))
|
|||
DEFAULT_ALLOWED_FAILS: Final = int(os.getenv("DEFAULT_ALLOWED_FAILS", 3))
|
||||
DEFAULT_REDIS_SYNC_INTERVAL: Final = int(os.getenv("DEFAULT_REDIS_SYNC_INTERVAL", 1))
|
||||
DEFAULT_COOLDOWN_TIME_SECONDS: Final = int(os.getenv("DEFAULT_COOLDOWN_TIME_SECONDS", 5))
|
||||
DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS: Final = float(
|
||||
os.getenv("DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS", "1")
|
||||
)
|
||||
DEFAULT_REPLICATE_POLLING_RETRIES: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_RETRIES", 5))
|
||||
DEFAULT_REPLICATE_POLLING_DELAY_SECONDS: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1))
|
||||
DEFAULT_IMAGE_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250))
|
||||
|
|
@ -1458,7 +1461,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"
|
||||
|
|
|
|||
|
|
@ -2414,6 +2414,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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -601,6 +601,12 @@ class CustomGuardrail(CustomLogger):
|
|||
event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None,
|
||||
supported_event_hooks: list[GuardrailEventHooks],
|
||||
) -> None:
|
||||
allowed_hooks: Final = frozenset(supported_event_hooks) | (
|
||||
frozenset((GuardrailEventHooks.logging_only,))
|
||||
if self.uses_apply_guardrail_interface() and not self.use_native_lifecycle_hooks
|
||||
else frozenset()
|
||||
)
|
||||
|
||||
def _validate_event_hook_list_is_in_supported_event_hooks(
|
||||
event_hook: list[GuardrailEventHooks] | list[str],
|
||||
supported_event_hooks: list[GuardrailEventHooks],
|
||||
|
|
@ -608,7 +614,7 @@ class CustomGuardrail(CustomLogger):
|
|||
for hook in event_hook:
|
||||
if isinstance(hook, str):
|
||||
hook = GuardrailEventHooks(hook)
|
||||
if hook not in supported_event_hooks:
|
||||
if hook not in allowed_hooks:
|
||||
raise ValueError(f"Event hook {hook} is not in the supported event hooks {supported_event_hooks}")
|
||||
|
||||
if event_hook is None:
|
||||
|
|
@ -629,7 +635,7 @@ class CustomGuardrail(CustomLogger):
|
|||
default_list = event_hook.default if isinstance(event_hook.default, list) else [event_hook.default]
|
||||
_validate_event_hook_list_is_in_supported_event_hooks(default_list, supported_event_hooks)
|
||||
elif isinstance(event_hook, GuardrailEventHooks):
|
||||
if event_hook not in supported_event_hooks:
|
||||
if event_hook not in allowed_hooks:
|
||||
raise ValueError(f"Event hook {event_hook} is not in the supported event hooks {supported_event_hooks}")
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -773,7 +779,7 @@ class CustomGuardrail(CustomLogger):
|
|||
def uses_apply_guardrail_interface(self) -> bool:
|
||||
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
|
||||
|
||||
def _deployment_pre_call_target(self) -> "CustomLogger":
|
||||
def _deployment_hook_target(self) -> "CustomLogger":
|
||||
if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks:
|
||||
return self
|
||||
try:
|
||||
|
|
@ -802,7 +808,7 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
# CHECK IF GUARDRAIL REJECTS THE REQUEST
|
||||
if call_type == CallTypes.completion or call_type == CallTypes.acompletion:
|
||||
target: Final = self._deployment_pre_call_target()
|
||||
target: Final = self._deployment_hook_target()
|
||||
if target is not self:
|
||||
kwargs["guardrail_to_apply"] = self
|
||||
result: Final = await target.async_pre_call_hook(
|
||||
|
|
@ -845,7 +851,9 @@ class CustomGuardrail(CustomLogger):
|
|||
return None
|
||||
|
||||
# CHECK IF GUARDRAIL REJECTS THE REQUEST
|
||||
result: Final = await self.async_post_call_success_hook(
|
||||
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"),
|
||||
|
|
@ -853,7 +861,7 @@ class CustomGuardrail(CustomLogger):
|
|||
api_key=request_data.get("user_api_key_hash"),
|
||||
request_route=request_data.get("user_api_key_request_route"),
|
||||
),
|
||||
data=request_data,
|
||||
data=hook_request_data,
|
||||
response=response,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,7 +266,7 @@ 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.
|
||||
|
|
@ -259,6 +275,7 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
ExporterSpec(
|
||||
kind=self.exporter,
|
||||
endpoint=self.endpoint,
|
||||
traces_endpoint=self.traces_endpoint,
|
||||
headers=self.headers,
|
||||
)
|
||||
]
|
||||
|
|
|
|||
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)
|
||||
|
|
@ -136,7 +136,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 +165,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:
|
||||
|
|
@ -201,7 +209,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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -7,16 +7,19 @@ from litellm.types.utils import CredentialItem
|
|||
|
||||
|
||||
class CredentialAccessor:
|
||||
@staticmethod
|
||||
def find_credential(credential_name: str) -> CredentialItem | None:
|
||||
return next(
|
||||
(credential for credential in litellm.credential_list if credential.credential_name == credential_name),
|
||||
None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_credential_values(credential_name: str) -> dict:
|
||||
"""Safe accessor for credentials."""
|
||||
|
||||
if not litellm.credential_list:
|
||||
return {}
|
||||
for credential in litellm.credential_list:
|
||||
if credential.credential_name == credential_name:
|
||||
return credential.credential_values.copy()
|
||||
return {}
|
||||
credential: Final = CredentialAccessor.find_credential(credential_name)
|
||||
return {} if credential is None else credential.credential_values.copy()
|
||||
|
||||
@staticmethod
|
||||
def upsert_credentials(credentials: list[CredentialItem]):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
@ -447,6 +448,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, \
|
||||
|
|
@ -1189,14 +1197,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={}):
|
||||
"""
|
||||
|
|
@ -2028,6 +2029,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 +2950,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 +2978,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(
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -13,9 +13,10 @@ Pattern Overview:
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain, repeat
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, overload, runtime_checkable
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict, assert_never
|
||||
|
|
@ -29,6 +30,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,
|
||||
|
|
@ -168,6 +170,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
them through guardrail rewrites; downstream provider handling is out of scope.
|
||||
"""
|
||||
|
||||
delivers_ended_stream_text_rewrites = True
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
|
@ -1014,11 +1018,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
|
||||
|
||||
|
|
@ -1065,6 +1075,15 @@ 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])
|
||||
else:
|
||||
verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
|
||||
return responses_so_far
|
||||
|
|
@ -1087,6 +1106,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 +1204,63 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
inputs["model"] = response_model
|
||||
return inputs
|
||||
|
||||
@staticmethod
|
||||
def _write_ended_stream_text_rewrite(
|
||||
responses_so_far: list[Any], # 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. Handles both chunk formats
|
||||
this stream carries (parsed event dicts and raw SSE bytes)."""
|
||||
replacements: Final = chain((rewritten_text,), repeat(""))
|
||||
for idx, item in enumerate(responses_so_far):
|
||||
if isinstance(item, dict):
|
||||
delta = item.get("delta")
|
||||
if item.get("type") == "content_block_delta" and isinstance(delta, dict):
|
||||
if delta.get("type") == "text_delta":
|
||||
delta["text"] = next(replacements)
|
||||
elif isinstance(item, (bytes, bytearray)):
|
||||
responses_so_far[idx] = ( # rebind-ok: delivers the rewrite into the caller's buffer
|
||||
AnthropicMessagesHandler._rewrite_sse_text_deltas(bytes(item), replacements)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _rewrite_sse_text_deltas(sse_bytes: bytes, replacements: "Iterator[str]") -> bytes:
|
||||
"""Rewrite every ``text_delta`` data line in one SSE chunk with the next
|
||||
replacement text, 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(
|
||||
AnthropicMessagesHandler._rewrite_sse_block(block, replacements) for block in decoded.split("\n\n")
|
||||
).encode("utf-8")
|
||||
|
||||
@staticmethod
|
||||
def _rewrite_sse_block(block: str, replacements: "Iterator[str]") -> str:
|
||||
return "\n".join(AnthropicMessagesHandler._rewrite_sse_line(line, replacements) for line in block.split("\n"))
|
||||
|
||||
@staticmethod
|
||||
def _rewrite_sse_line(line: str, replacements: "Iterator[str]") -> 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) or data.get("type") != "content_block_delta":
|
||||
return line
|
||||
delta: Final = data.get("delta")
|
||||
if not isinstance(delta, dict) or delta.get("type") != "text_delta":
|
||||
return line
|
||||
return "data: " + json.dumps(
|
||||
{**data, "delta": {**delta, "text": next(replacements)}} # mutable-ok: json.dumps needs plain dicts
|
||||
)
|
||||
|
||||
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,
|
||||
|
|
|
|||
|
|
@ -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,14 @@ class StreamingScanKey:
|
|||
|
||||
|
||||
class BaseTranslation(ABC):
|
||||
delivers_ended_stream_text_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 rewrites back across ``responses_so_far`` so
|
||||
a buffered pipeline can release rewritten chunks. Tool-call rewrites, and
|
||||
text 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 +165,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 +173,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_text_rewrites``: the handler then writes
|
||||
guardrail text 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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -1741,6 +1741,12 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
logging_obj.post_call(
|
||||
api_key=api_key,
|
||||
original_response=response.text,
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
|
||||
return self._transform_ocr_response(
|
||||
provider_config=provider_config,
|
||||
model=model,
|
||||
|
|
@ -1804,6 +1810,12 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
logging_obj.post_call(
|
||||
api_key=api_key,
|
||||
original_response=response.text,
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
|
||||
# Use async response transform for async operations
|
||||
return await provider_config.async_transform_ocr_response(
|
||||
model=model,
|
||||
|
|
@ -9814,7 +9826,7 @@ class BaseLLMHTTPHandler:
|
|||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
litellm_params=MappingProxyType(dict(litellm_params, timeout=timeout)),
|
||||
extra_body=extra_body,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
|
|
@ -9859,6 +9871,12 @@ class BaseLLMHTTPHandler:
|
|||
data=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise vector_store_provider_config.get_error_class(
|
||||
error_message="Vector store search exceeded the caller timeout.",
|
||||
status_code=408,
|
||||
headers=httpx.Headers(),
|
||||
) from None
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
|
||||
|
||||
|
|
@ -9943,7 +9961,7 @@ class BaseLLMHTTPHandler:
|
|||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
litellm_params=MappingProxyType(dict(litellm_params, timeout=timeout)),
|
||||
extra_body=extra_body,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
|
|
@ -9988,7 +10006,14 @@ class BaseLLMHTTPHandler:
|
|||
url=url,
|
||||
headers=headers,
|
||||
data=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise vector_store_provider_config.get_error_class(
|
||||
error_message="Vector store search exceeded the caller timeout.",
|
||||
status_code=408,
|
||||
headers=httpx.Headers(),
|
||||
) from None
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
|
||||
|
||||
|
|
@ -10018,6 +10043,8 @@ class BaseLLMHTTPHandler:
|
|||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
vector_store_provider_config.validate_create_vector_store()
|
||||
|
||||
headers: Final = vector_store_provider_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=litellm_params
|
||||
)
|
||||
|
|
@ -10088,6 +10115,8 @@ class BaseLLMHTTPHandler:
|
|||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
vector_store_provider_config.validate_create_vector_store()
|
||||
|
||||
headers: Final = vector_store_provider_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=litellm_params
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -78,6 +78,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
Methods can be overridden to customize behavior for different message formats.
|
||||
"""
|
||||
|
||||
delivers_ended_stream_text_rewrites = True
|
||||
|
||||
def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None:
|
||||
"""
|
||||
Convert chat completions request data to OpenAI-spec structured messages.
|
||||
|
|
@ -453,6 +455,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 +470,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 +499,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 +510,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 +599,39 @@ 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 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)
|
||||
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 deliver_ended_stream_rewrites:
|
||||
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_to_apply.guardrail_name or "unknown",
|
||||
)
|
||||
|
||||
def build_stream_error_items(
|
||||
self,
|
||||
exc: "HTTPException",
|
||||
|
|
@ -745,8 +782,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 +796,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 +807,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 +997,52 @@ 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
|
||||
)
|
||||
|
||||
async def _apply_guardrail_responses_to_output_streaming(
|
||||
self,
|
||||
responses: list["ModelResponseStream"],
|
||||
|
|
@ -975,7 +1058,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 +1075,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 +1092,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 +1113,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,
|
||||
|
|
@ -118,6 +119,15 @@ 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,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{"function_call_output": "output", "message": "content"}
|
||||
)
|
||||
|
|
@ -330,6 +340,8 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
Methods can be overridden to customize behavior for different message formats.
|
||||
"""
|
||||
|
||||
delivers_ended_stream_text_rewrites = True
|
||||
|
||||
def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None:
|
||||
"""
|
||||
Convert Responses API request data to OpenAI-spec structured messages.
|
||||
|
|
@ -667,6 +679,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 +689,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 +712,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]] = []
|
||||
|
|
@ -747,11 +771,25 @@ 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,
|
||||
)
|
||||
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 +807,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 +824,83 @@ 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 _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={}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -835,6 +835,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",
|
||||
|
|
@ -3587,7 +3595,9 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
|
|||
ui_callback_name="OpenTelemetry",
|
||||
litellm_callback_params=[
|
||||
"OTEL_EXPORTER",
|
||||
"OTEL_EXPORTER_OTLP_PROTOCOL",
|
||||
"OTEL_ENDPOINT",
|
||||
"OTEL_TRACES_ENDPOINT",
|
||||
"OTEL_HEADERS",
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -497,10 +497,22 @@ class RouteChecks:
|
|||
|
||||
def _placeholder_to_regex(match: re.Match) -> str:
|
||||
placeholder: Final = match.group(0).strip("{}")
|
||||
if placeholder.endswith(":path"):
|
||||
# allow "/" in the placeholder value, but don't eat the route suffix after ":"
|
||||
return r"[^:]+"
|
||||
return r"[^/]+"
|
||||
if not placeholder.endswith(":path"):
|
||||
return r"[^/]+"
|
||||
# A ":path" placeholder takes whatever the router's own path
|
||||
# converter takes, slashes and colons alike, so an id spelled with
|
||||
# either (or both) still matches the template it was mounted under.
|
||||
#
|
||||
# Unless the template puts a ":" literal of its own after the
|
||||
# placeholder: the Google routes end in ":generateContent" and
|
||||
# friends, and there the value has to stop before that suffix
|
||||
# rather than swallow it and match a different verb.
|
||||
#
|
||||
# "[\s\S]" rather than ".", because "." stops at a newline and the
|
||||
# path converter does not: a %0A anywhere in the value would leave
|
||||
# the route unmatched here while still reaching the handler, which
|
||||
# turns this gate into a bypass for the lists built on it.
|
||||
return r"[^:]+" if ":" in match.string[match.end() :] else r"[\s\S]+"
|
||||
|
||||
pattern = re.sub(r"\{[^}]+\}", _placeholder_to_regex, pattern)
|
||||
# Anchor the pattern to match the entire string
|
||||
|
|
|
|||
|
|
@ -2038,7 +2038,7 @@ async def _user_api_key_auth_builder(
|
|||
fallback_spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
if team_member_spend > team_member_budget:
|
||||
if team_member_spend >= team_member_budget:
|
||||
_entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}"
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=team_member_spend,
|
||||
|
|
|
|||
|
|
@ -44,6 +44,91 @@ def _langfuse_environment_error(callback_vars: Mapping[str, str]) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
# Which credential family a dynamic variable belongs to. The families are the
|
||||
# integrations that share one account: every langfuse_* variable configures the
|
||||
# same Langfuse project whether it rides the classic callback or the OTel one,
|
||||
# and every dd_* variable configures the same Datadog account.
|
||||
_VAR_FAMILIES: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"arize_": "Arize",
|
||||
"dd_": "Datadog",
|
||||
"gcs_": "GCS",
|
||||
"humanloop_": "Humanloop",
|
||||
"langfuse_": "Langfuse",
|
||||
"langsmith_": "LangSmith",
|
||||
"newrelic_": "New Relic",
|
||||
"posthog_": "PostHog",
|
||||
"wandb_": "Weights & Biases",
|
||||
"weave_": "Weights & Biases",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _family_of(var: str) -> str | None:
|
||||
"""The credential family ``var`` configures, or ``None`` if it configures none.
|
||||
|
||||
``turn_off_message_logging`` and friends belong to no backend, so they carry
|
||||
no credentials anyone could redirect.
|
||||
"""
|
||||
return next((family for prefix, family in _VAR_FAMILIES.items() if var.startswith(prefix)), None)
|
||||
|
||||
|
||||
def cross_entry_family_error(
|
||||
callback_vars: Mapping[str, str] | None,
|
||||
stored_vars_by_entry: Sequence[Mapping[str, str]],
|
||||
) -> str | None:
|
||||
"""Reject an entry that changes what a family another entry holds resolves to.
|
||||
|
||||
Every stored entry's variables are flattened into one dict before a request
|
||||
reads them, and the flattened dict is what the exporter authenticates and
|
||||
addresses with. So an entry naming only a destination is enough to redirect
|
||||
credentials that were written somewhere else: a host on a second entry pairs
|
||||
with the key from the first, and the request carries that key to the new
|
||||
host.
|
||||
|
||||
Two rules together keep the flattened dict out of the caller's hands. A
|
||||
variable the family already configures has to keep the value it has, so
|
||||
nothing already in use can be moved. A variable the family does not yet
|
||||
configure may only carry a value the family already holds, which is what lets
|
||||
the same credential go in under its other spelling (``langfuse_secret`` and
|
||||
``langfuse_secret_key`` are one key) without anything here having to list the
|
||||
spellings. Between them, no value the caller chose can enter the family, and
|
||||
repeating the family as it stands is still allowed -- that is how one
|
||||
integration gets registered for both the success and the failure event.
|
||||
|
||||
A team admin who does want to move a family deletes the entry holding it
|
||||
first, which reveals nothing.
|
||||
|
||||
Only the writers this endpoint newly admits are held to this, because a proxy
|
||||
admin already holds every credential the proxy has.
|
||||
|
||||
``stored_vars_by_entry`` has to arrive decrypted; the credential values are
|
||||
encrypted at rest and ciphertext never equals the plaintext coming in.
|
||||
"""
|
||||
if not callback_vars:
|
||||
return None
|
||||
stored_by_var: Final = {
|
||||
var: value for entry in stored_vars_by_entry for var, value in entry.items() if _family_of(var) is not None
|
||||
}
|
||||
family_values: Final = frozenset(
|
||||
(family, value)
|
||||
for entry in stored_vars_by_entry
|
||||
for var, value in entry.items()
|
||||
if (family := _family_of(var)) is not None
|
||||
)
|
||||
held_families: Final = frozenset(family for family, _ in family_values)
|
||||
return next(
|
||||
(
|
||||
f"{family} is already configured by another callback entry on this team. "
|
||||
f"Remove that entry before setting {var} here."
|
||||
for var, value, family in ((v, callback_vars[v], _family_of(v)) for v in callback_vars)
|
||||
if family in held_families
|
||||
and (stored_by_var[var] != value if var in stored_by_var else (family, value) not in family_values)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def logging_metadata_config_error(metadata: Mapping[str, object] | None) -> str | None:
|
||||
"""Validate every ``logging`` entry of a team/key metadata payload."""
|
||||
if not metadata:
|
||||
|
|
|
|||
|
|
@ -10,8 +10,10 @@ import litellm
|
|||
from litellm import get_secret
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
CLIENT_OUTPUT_CEILING_METADATA_KEY,
|
||||
CONSUMED_REQUEST_TAGS_METADATA_KEY,
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -507,6 +509,8 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset(
|
|||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
CONSUMED_REQUEST_TAGS_METADATA_KEY,
|
||||
CLIENT_OUTPUT_CEILING_METADATA_KEY,
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY,
|
||||
"disable_global_guardrails",
|
||||
"disable_global_guardrail",
|
||||
"opted_out_global_guardrails",
|
||||
|
|
|
|||
|
|
@ -125,6 +125,7 @@ async def _resync_model_deployments(model_name: str) -> bool:
|
|||
)
|
||||
return proxy_server.llm_router is not None
|
||||
async with proxy_server.MODEL_RECONCILE_LOCK:
|
||||
await proxy_server.proxy_config.get_credentials(prisma_client=prisma_client)
|
||||
proxy_server.proxy_config._add_deployment(db_models=rows)
|
||||
proxy_server.llm_model_list = router.get_model_list()
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -46,17 +46,32 @@ def extract_sql_commands(diff_output: str) -> list[str]:
|
|||
|
||||
def check_prisma_schema_diff_helper(db_url: str) -> tuple[bool, list[str]]:
|
||||
"""Checks for differences between current database and Prisma schema.
|
||||
|
||||
Never raises: a diff that cannot be produced, because the runner is missing,
|
||||
because the command failed, or because it outlived its budget, is reported as
|
||||
"no diff" so boot continues.
|
||||
|
||||
Returns:
|
||||
A tuple containing:
|
||||
- A boolean indicating if differences were found (True) or not (False).
|
||||
- A string with the diff output or error message.
|
||||
Raises:
|
||||
subprocess.CalledProcessError: If the Prisma command fails.
|
||||
Exception: For any other errors during execution.
|
||||
- The SQL commands that would close the diff, empty when there is none.
|
||||
"""
|
||||
verbose_logger.debug("Checking for Prisma schema diff...")
|
||||
try:
|
||||
result: Final = subprocess.run(
|
||||
from litellm_proxy_extras.prisma_toolchain import (
|
||||
PRISMA_COMMAND_TIMEOUT_ENV_VAR,
|
||||
prisma_command_timeout,
|
||||
run_prisma,
|
||||
)
|
||||
except ImportError as e:
|
||||
print( # noqa: T201 # boot-time operator output, same channel as this helper's other messages
|
||||
f"Skipping the migration diff: litellm-proxy-extras has no Prisma runner. Error: {e}"
|
||||
)
|
||||
return False, []
|
||||
|
||||
verbose_logger.debug("Checking for Prisma schema diff...")
|
||||
timeout: Final = prisma_command_timeout()
|
||||
try:
|
||||
result: Final = run_prisma(
|
||||
[
|
||||
"prisma",
|
||||
"migrate",
|
||||
|
|
@ -67,12 +82,10 @@ def check_prisma_schema_diff_helper(db_url: str) -> tuple[bool, list[str]]:
|
|||
"./schema.prisma",
|
||||
"--script",
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=True,
|
||||
timeout=timeout,
|
||||
env=os.environ.copy(),
|
||||
)
|
||||
|
||||
# return True, "Migration diff generated successfully."
|
||||
sql_commands: Final = extract_sql_commands(result.stdout)
|
||||
|
||||
if sql_commands:
|
||||
|
|
@ -83,6 +96,12 @@ def check_prisma_schema_diff_helper(db_url: str) -> tuple[bool, list[str]]:
|
|||
return True, sql_commands
|
||||
else:
|
||||
return False, []
|
||||
except subprocess.TimeoutExpired:
|
||||
print( # noqa: T201 # boot-time operator output, same channel as this helper's other messages
|
||||
f"Timed out after {timeout}s generating the migration diff. "
|
||||
f"Raise {PRISMA_COMMAND_TIMEOUT_ENV_VAR} if this database needs longer."
|
||||
)
|
||||
return False, []
|
||||
except subprocess.CalledProcessError as e:
|
||||
error_message: Final = f"Failed to generate migration diff. Error: {e.stderr}"
|
||||
print(error_message) # noqa: T201
|
||||
|
|
|
|||
|
|
@ -1063,7 +1063,7 @@ class DBSpendUpdateWriter:
|
|||
|
||||
await enqueue_spend_logs(prisma_client, (payload,))
|
||||
if payload.get("call_type") in RESPONSES_SESSION_CALL_TYPES:
|
||||
request_spend_log_flush()
|
||||
request_spend_log_flush(prisma_client)
|
||||
else:
|
||||
verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.")
|
||||
|
||||
|
|
|
|||
|
|
@ -64,6 +64,7 @@ DISABLE_PREPARED_STATEMENTS_ENV_VAR: Final = "DATABASE_DISABLE_PREPARED_STATEMEN
|
|||
DisablePreparedStatementsFlag = Annotated[
|
||||
bool, BeforeValidator(partial(token_auth_flag_enabled, env_var=DISABLE_PREPARED_STATEMENTS_ENV_VAR))
|
||||
]
|
||||
MAX_IDLE_CONNECTION_LIFETIME_ENV_VAR: Final = "DATABASE_MAX_IDLE_CONNECTION_LIFETIME"
|
||||
|
||||
# schema.prisma pins `provider = "postgresql"`, so these are the only schemes
|
||||
# Prisma can actually connect with.
|
||||
|
|
@ -217,6 +218,9 @@ class DatabaseURLSettings(BaseSettings):
|
|||
disable_prepared_statements: DisablePreparedStatementsFlag = Field(
|
||||
default=False, validation_alias=DISABLE_PREPARED_STATEMENTS_ENV_VAR
|
||||
)
|
||||
max_idle_connection_lifetime: int | None = Field(
|
||||
default=None, validation_alias=MAX_IDLE_CONNECTION_LIFETIME_ENV_VAR
|
||||
)
|
||||
|
||||
# Writer
|
||||
database_url: str | None = Field(default=None, validation_alias="DATABASE_URL")
|
||||
|
|
@ -453,6 +457,12 @@ class DatabaseURLSettings(BaseSettings):
|
|||
if url:
|
||||
os.environ[env_var] = add_missing_query_params(url, MappingProxyType({"pgbouncer": "true"}))
|
||||
|
||||
lifetime_params: Final = idle_lifetime_params(self.max_idle_connection_lifetime)
|
||||
for env_var in ("DATABASE_URL", "DIRECT_URL"):
|
||||
url = os.environ.get(env_var)
|
||||
if url:
|
||||
os.environ[env_var] = add_missing_query_params(url, lifetime_params)
|
||||
|
||||
# The reader inherits the writer's connection params (pool size, timeouts,
|
||||
# pgbouncer mode). Without this the reader pool ignores the configured cap
|
||||
# and falls back to Prisma's `num_physical_cpus * 2 + 1` default.
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from datetime import datetime, timedelta
|
|||
from typing import Any, Final, Protocol
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.db.db_url_settings import add_missing_query_params, connection_params_from_url
|
||||
from litellm.proxy.db.token_auth import (
|
||||
DEFAULT_POSTGRES_PORT,
|
||||
DatabaseTokenAuth,
|
||||
|
|
@ -438,7 +439,10 @@ class PrismaWrapper:
|
|||
return None
|
||||
|
||||
endpoint: Final = self._iam_endpoint if self._iam_endpoint is not None else self._endpoint_from_env()
|
||||
db_url: Final = endpoint.build_url(mint_database_token(auth, endpoint))
|
||||
db_url: Final = add_missing_query_params(
|
||||
endpoint.build_url(mint_database_token(auth, endpoint)),
|
||||
connection_params_from_url(os.environ.get(self._db_url_env_var, "")),
|
||||
)
|
||||
os.environ[self._db_url_env_var] = db_url
|
||||
return db_url
|
||||
|
||||
|
|
@ -937,9 +941,17 @@ class PrismaManager:
|
|||
use_v2_resolver=use_v2_resolver,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
from litellm_proxy_extras.prisma_toolchain import (
|
||||
prisma_command_timeout,
|
||||
run_prisma,
|
||||
)
|
||||
except ImportError as e:
|
||||
verbose_proxy_logger.error("\x1b[1;31mLiteLLM: Failed to import proxy extras. Got %s\x1b[0m", e)
|
||||
return False
|
||||
|
||||
PrismaManager._raise_if_partitioned_spend_logs()
|
||||
# Use prisma db push with increased timeout
|
||||
subprocess.run(
|
||||
run_prisma(
|
||||
[
|
||||
"prisma",
|
||||
"db",
|
||||
|
|
@ -947,13 +959,15 @@ class PrismaManager:
|
|||
"--accept-data-loss",
|
||||
"--skip-generate",
|
||||
],
|
||||
timeout=60,
|
||||
check=True,
|
||||
timeout=prisma_command_timeout(),
|
||||
env=os.environ.copy(),
|
||||
stdout=None,
|
||||
stderr=None,
|
||||
)
|
||||
PrismaManager._apply_replica_identity_full_if_requested()
|
||||
return True
|
||||
except subprocess.TimeoutExpired:
|
||||
verbose_proxy_logger.warning("Attempt %s timed out", attempt + 1)
|
||||
except subprocess.TimeoutExpired as e:
|
||||
verbose_proxy_logger.warning("Attempt %s timed out after %.0fs", attempt + 1, e.timeout)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
except subprocess.CalledProcessError as e:
|
||||
attempts_left = 3 - attempt
|
||||
|
|
|
|||
|
|
@ -547,7 +547,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
|
||||
for _tool_call, is_allowed, _rule_id, message in checked:
|
||||
if not is_allowed and message is not None:
|
||||
verbose_proxy_logger.warning("Tool Permission Guardrail: %s", message)
|
||||
verbose_proxy_logger.info("Tool Permission Guardrail: %s", message)
|
||||
if self.on_disallowed_action == "block":
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name, message=message, blocked_content=True
|
||||
|
|
@ -809,7 +809,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
|
||||
new_tools: Final = self._collect_request_tools(data)
|
||||
if not new_tools:
|
||||
verbose_proxy_logger.warning(
|
||||
verbose_proxy_logger.debug(
|
||||
"Tool Permission Guardrail: not running guardrail. No tools or functions in data"
|
||||
)
|
||||
return data
|
||||
|
|
@ -820,7 +820,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
is_allowed, _, message = self._check_tool_permission(tool_name, tool_type)
|
||||
|
||||
if not is_allowed and message is not None:
|
||||
verbose_proxy_logger.warning("Tool Permission Guardrail: %s", message)
|
||||
verbose_proxy_logger.info("Tool Permission Guardrail: %s", message)
|
||||
if self.on_disallowed_action == "block":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ from litellm.cost_calculator import _infer_call_type
|
|||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -69,6 +69,36 @@ def _as_endpoint_translation(translation: _EndpointTranslation) -> _EndpointTran
|
|||
return translation
|
||||
|
||||
|
||||
def resolve_endpoint_translation(
|
||||
user_api_key_dict: UserAPIKeyAuth, first_response_item: object | None
|
||||
) -> "tuple[str, BaseTranslation] | None":
|
||||
"""
|
||||
Resolve the endpoint guardrail translation for a streamed response: the
|
||||
request route wins, falling back to inferring the call type from the first
|
||||
response chunk (the same resolution order the streaming iterator hook uses).
|
||||
Returns None when the call type is unresolvable or has no translation.
|
||||
"""
|
||||
route_call_types: Final = (
|
||||
get_call_types_for_route(user_api_key_dict.request_route) if user_api_key_dict.request_route else None
|
||||
)
|
||||
call_type: Final = (
|
||||
route_call_types[0].value
|
||||
if route_call_types
|
||||
else (
|
||||
_infer_call_type(call_type=None, completion_response=first_response_item)
|
||||
if first_response_item is not None
|
||||
else None
|
||||
)
|
||||
)
|
||||
if call_type is None:
|
||||
return None
|
||||
try:
|
||||
handler_cls: Final = get_guardrail_translation_mapping(CallTypes(call_type))
|
||||
except ValueError:
|
||||
return None
|
||||
return call_type, handler_cls()
|
||||
|
||||
|
||||
def _chunk_choices(item: object) -> Sequence[object]:
|
||||
choices: Final[Sequence[object]] = getattr(item, "choices", None) or []
|
||||
return choices
|
||||
|
|
@ -343,7 +373,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
|
||||
return response
|
||||
|
||||
async def _handle_streaming_block(
|
||||
async def handle_streaming_block(
|
||||
self,
|
||||
exc: "ModifyResponseException",
|
||||
endpoint_translation: _EndpointTranslation,
|
||||
|
|
@ -399,7 +429,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
return None
|
||||
return call_type
|
||||
|
||||
async def _emit_streaming_http_error(
|
||||
async def emit_streaming_http_error(
|
||||
self,
|
||||
exc: HTTPException,
|
||||
call_type: str | None,
|
||||
|
|
@ -592,7 +622,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
except ModifyResponseException as e:
|
||||
if e.original_response is None:
|
||||
e.original_response = responses_so_far
|
||||
async for block_chunk in self._handle_streaming_block(
|
||||
async for block_chunk in self.handle_streaming_block(
|
||||
e,
|
||||
endpoint_translation,
|
||||
stream_started=bool(responses_yielded),
|
||||
|
|
@ -601,7 +631,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
yield block_chunk
|
||||
raise _StreamTerminated()
|
||||
except HTTPException as e:
|
||||
async for error_item in self._emit_streaming_http_error(
|
||||
async for error_item in self.emit_streaming_http_error(
|
||||
e,
|
||||
call_type,
|
||||
responses_so_far,
|
||||
|
|
@ -781,7 +811,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
except ModifyResponseException as e:
|
||||
if e.original_response is None:
|
||||
e.original_response = responses_so_far
|
||||
async for block_chunk in self._handle_streaming_block(
|
||||
async for block_chunk in self.handle_streaming_block(
|
||||
e,
|
||||
endpoint_translation,
|
||||
stream_started=bool(responses_yielded),
|
||||
|
|
@ -869,6 +899,14 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
choices: Final = _chunk_choices(item)
|
||||
return any(getattr(choice, "finish_reason", None) is not None for choice in choices)
|
||||
|
||||
def resolve_streaming_flag(self, guardrail_to_apply: CustomGuardrail | None, name: str, default: object) -> object:
|
||||
"""Streaming flag resolution order (later wins): default < guardrail
|
||||
attribute < guardrail_config dict < this callback's optional_params."""
|
||||
attribute_value: Final = default if guardrail_to_apply is None else getattr(guardrail_to_apply, name, default)
|
||||
config: Final = None if guardrail_to_apply is None else getattr(guardrail_to_apply, "guardrail_config", None)
|
||||
config_value: Final = config.get(name, attribute_value) if isinstance(config, dict) else attribute_value
|
||||
return self.optional_params.get(name, config_value)
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -897,17 +935,8 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
if guardrail_to_apply is None:
|
||||
guardrail_to_apply = request_data.pop("guardrail_to_apply", None)
|
||||
|
||||
# Get streaming configuration. Resolution order (later wins): default
|
||||
# < guardrail attribute < guardrail_config dict < this callback's
|
||||
# optional_params.
|
||||
def _streaming_flag(name: str, default: object) -> Any:
|
||||
value = default
|
||||
if guardrail_to_apply is not None:
|
||||
value = getattr(guardrail_to_apply, name, value)
|
||||
config: Final[Mapping[str, object]] = getattr(guardrail_to_apply, "guardrail_config", {})
|
||||
if isinstance(config, dict):
|
||||
value = config.get(name, value)
|
||||
return self.optional_params.get(name, value)
|
||||
return self.resolve_streaming_flag(guardrail_to_apply, name, default)
|
||||
|
||||
sampling_rate: Final[int] = _streaming_flag("streaming_sampling_rate", 5)
|
||||
# Only apply the guardrail at end of stream (not per chunk).
|
||||
|
|
@ -1091,7 +1120,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
# The current chunk was appended to responses_so_far but not
|
||||
# yet yielded, so exclude it: the continuation must reflect
|
||||
# only what the client has actually received.
|
||||
async for block_chunk in self._handle_streaming_block(
|
||||
async for block_chunk in self.handle_streaming_block(
|
||||
e,
|
||||
endpoint_translation,
|
||||
stream_started=chunks_yielded,
|
||||
|
|
@ -1101,7 +1130,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
return
|
||||
except HTTPException as e:
|
||||
# Response already started (we already yielded chunks); cannot send 400.
|
||||
async for error_item in self._emit_streaming_http_error(
|
||||
async for error_item in self.emit_streaming_http_error(
|
||||
e,
|
||||
call_type,
|
||||
responses_so_far,
|
||||
|
|
@ -1175,7 +1204,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
# terminating SSE sequence with the block message rather than
|
||||
# propagating into a bare error blob that truncates the stream.
|
||||
# The withheld original chunks are never released.
|
||||
async for block_chunk in self._handle_streaming_block(
|
||||
async for block_chunk in self.handle_streaming_block(
|
||||
e,
|
||||
endpoint_translation,
|
||||
stream_started=bool(responses_yielded),
|
||||
|
|
@ -1184,7 +1213,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
yield block_chunk
|
||||
return
|
||||
except HTTPException as e:
|
||||
async for error_item in self._emit_streaming_http_error(
|
||||
async for error_item in self.emit_streaming_http_error(
|
||||
e,
|
||||
call_type,
|
||||
responses_so_far,
|
||||
|
|
|
|||
|
|
@ -18,11 +18,13 @@ from litellm._logging import verbose_logger, verbose_proxy_logger
|
|||
from litellm._service_logger import ServiceLogging
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import (
|
||||
CLIENT_OUTPUT_CEILING_METADATA_KEY,
|
||||
CONSUMED_REQUEST_TAGS_METADATA_KEY,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
LITELLM_PROXY_MASTER_KEY_ALIAS,
|
||||
OTEL_SERVICE_NAME_METADATA_KEYS,
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
SESSION_ID_GENERATED_METADATA_KEY,
|
||||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
|
|
@ -289,6 +291,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS: Final = (
|
|||
GATEWAY_INJECTED_CACHE_METADATA_KEY,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
CONSUMED_REQUEST_TAGS_METADATA_KEY,
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
"standard_logging_object",
|
||||
"proxy_server_request",
|
||||
|
|
@ -325,7 +328,9 @@ _CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logg
|
|||
# ``attempted_fallbacks`` and ``original_model_group`` are written by the router
|
||||
# and read by spend logs as fact; a client value has no legitimate meaning and no
|
||||
# key or team setting keeps it, so the strip is never gated.
|
||||
_ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset({"attempted_fallbacks", "original_model_group"})
|
||||
_ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset(
|
||||
{"attempted_fallbacks", "original_model_group", CLIENT_OUTPUT_CEILING_METADATA_KEY}
|
||||
)
|
||||
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override"
|
||||
|
||||
# Request fields whose value, when URL-valued, becomes the outbound destination
|
||||
|
|
@ -3075,10 +3080,9 @@ def _apply_resolved_guardrails_to_metadata(
|
|||
if metadata_variable_name not in data:
|
||||
data[metadata_variable_name] = {}
|
||||
|
||||
# Track pipeline-managed guardrails to exclude from independent execution
|
||||
pipeline_managed_guardrails: set = set()
|
||||
# Record the pipelines and the guardrails they step; the hook loops skip those per pipeline mode
|
||||
if pipelines:
|
||||
pipeline_managed_guardrails = PolicyResolver.get_pipeline_managed_guardrails(pipelines)
|
||||
pipeline_managed_guardrails: Final = PolicyResolver.get_pipeline_managed_guardrails(pipelines)
|
||||
data[metadata_variable_name]["_guardrail_pipelines"] = pipelines
|
||||
data[metadata_variable_name]["_pipeline_managed_guardrails"] = pipeline_managed_guardrails
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -3095,10 +3099,8 @@ def _apply_resolved_guardrails_to_metadata(
|
|||
existing_guardrails = []
|
||||
|
||||
# Combine existing guardrails with policy-resolved guardrails (no duplicates)
|
||||
# Exclude pipeline-managed guardrails from the flat list
|
||||
combined = set(existing_guardrails)
|
||||
combined.update(resolved_guardrails)
|
||||
combined -= pipeline_managed_guardrails
|
||||
data[metadata_variable_name]["guardrails"] = list(combined)
|
||||
|
||||
verbose_proxy_logger.debug("Policy engine: added guardrails to request metadata: %s", list(combined))
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm.proxy._types import (
|
|||
LiteLLM_AuditLogs,
|
||||
LiteLLM_TeamTable,
|
||||
LitellmTableNames,
|
||||
LitellmUserRoles,
|
||||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
TeamCallbackDeleteResponse,
|
||||
|
|
@ -28,7 +29,10 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.callback_config_validation import callback_config_error
|
||||
from litellm.proxy.common_utils.callback_config_validation import (
|
||||
callback_config_error,
|
||||
cross_entry_family_error,
|
||||
)
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
_CALLBACK_VAR_ENCRYPTED_PREFIX,
|
||||
decrypt_callback_vars,
|
||||
|
|
@ -230,6 +234,22 @@ def _callback_error(status_code: int, message: str) -> HTTPException:
|
|||
)
|
||||
|
||||
|
||||
def _unknown_team_error(team_id: str, user_api_key_dict: UserAPIKeyAuth, status_code: int) -> HTTPException:
|
||||
"""Report an unknown team without telling an unauthorized caller that it is unknown.
|
||||
|
||||
These routes are reachable by any authenticated caller so that a team admin can
|
||||
get as far as _verify_team_access. A distinct "does not exist" would therefore let
|
||||
any valid key probe which team ids exist, so a caller who could not have managed
|
||||
the team either way gets the same 403 body _verify_team_access raises.
|
||||
"""
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
return _callback_error(status_code, f"Team id = {team_id} does not exist.")
|
||||
return HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="You do not have access to this team",
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/team/{team_id:path}/callback",
|
||||
tags=["team management"],
|
||||
|
|
@ -304,10 +324,7 @@ async def add_team_callbacks(
|
|||
# Check if team_id exists already
|
||||
_existing_team = await prisma_client.get_data(team_id=team_id, table_name="team", query_type="find_unique")
|
||||
if _existing_team is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"Team id = {team_id} does not exist. Please use a different team id."},
|
||||
)
|
||||
raise _unknown_team_error(team_id, user_api_key_dict, status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
# IDOR guard: only proxy admins / org admins / team admins of THIS
|
||||
# team may write callback credentials. Without this, any
|
||||
|
|
@ -326,6 +343,28 @@ async def add_team_callbacks(
|
|||
if team_callback_settings is None or not isinstance(team_callback_settings, list):
|
||||
team_callback_settings = []
|
||||
|
||||
# One entry has to own a credential family end to end. The entries are
|
||||
# flattened into one dict before a request reads them, so an entry
|
||||
# naming only a destination would pair with a key written on another
|
||||
# entry and carry it to that destination -- a key a team admin can read
|
||||
# back nowhere. Repeating a value the owning entry already stores is
|
||||
# fine, which is how one integration covers both events. Proxy admins
|
||||
# are exempt: they already hold every credential the proxy has.
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
# Decrypted, because the check compares the incoming values against
|
||||
# the stored ones and the credentials are encrypted at rest.
|
||||
decrypted_logging: Final = decrypt_callback_vars(team_metadata).get("logging")
|
||||
stored_entries: Final = decrypted_logging if isinstance(decrypted_logging, list) else ()
|
||||
stored_entry_vars: Final = [ # mutable-ok: read-only input to the check, never stored
|
||||
entry.get("callback_vars") or {} for entry in stored_entries
|
||||
]
|
||||
family_error: Final = cross_entry_family_error(data.callback_vars, stored_entry_vars)
|
||||
if family_error is not None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=family_error,
|
||||
)
|
||||
|
||||
## check if it already exists, for the same callback event
|
||||
for callback in team_callback_settings:
|
||||
if (
|
||||
|
|
@ -452,7 +491,7 @@ async def delete_team_callback(
|
|||
team_id=team_id, table_name="team", query_type="find_unique"
|
||||
)
|
||||
if _existing_team is None:
|
||||
raise _callback_error(404, f"Team id = {team_id} does not exist.")
|
||||
raise _unknown_team_error(team_id, user_api_key_dict, status.HTTP_404_NOT_FOUND)
|
||||
|
||||
# IDOR guard: only proxy admins / org admins / team admins of THIS team may
|
||||
# deregister its callbacks, otherwise any authenticated key holder could
|
||||
|
|
@ -726,10 +765,7 @@ async def get_team_callbacks(
|
|||
# Check if team_id exists
|
||||
_existing_team = await prisma_client.get_data(team_id=team_id, table_name="team", query_type="find_unique")
|
||||
if _existing_team is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Team id = {team_id} does not exist."},
|
||||
)
|
||||
raise _unknown_team_error(team_id, user_api_key_dict, status.HTTP_404_NOT_FOUND)
|
||||
|
||||
# IDOR guard: callback metadata holds third-party API credentials
|
||||
# (Langfuse / Langsmith / GCS). Only proxy admins / org admins /
|
||||
|
|
|
|||
|
|
@ -5,18 +5,27 @@ Runs guardrails sequentially per pipeline step definitions, handling
|
|||
pass/fail actions (allow, block, next, modify_response) and data forwarding.
|
||||
"""
|
||||
|
||||
import copy
|
||||
import time
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, Final, Literal
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import LOGS_GUARDRAIL_INFORMATION_MARKER
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import independent_snapshot
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_or_create_metadata_bucket,
|
||||
independent_snapshot,
|
||||
)
|
||||
from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
|
@ -25,6 +34,14 @@ from litellm.types.proxy.policy_engine.pipeline_types import (
|
|||
PipelineStep,
|
||||
PipelineStepResult,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, StandardLoggingGuardrailInformation
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
||||
BaseTranslation,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
try:
|
||||
from fastapi.exceptions import HTTPException
|
||||
|
|
@ -32,6 +49,121 @@ except ImportError:
|
|||
HTTPException = None
|
||||
|
||||
|
||||
class UndeliverableStreamRewrite(Exception):
|
||||
def __init__(self, guardrail_name: str) -> None:
|
||||
super().__init__(
|
||||
f"Guardrail '{guardrail_name}' rewrote the streamed response in a way this endpoint's "
|
||||
"streaming pipeline cannot deliver"
|
||||
)
|
||||
self.guardrail_name: Final = guardrail_name
|
||||
|
||||
|
||||
def _tool_call_shape(tool_call: object) -> tuple[object, object]:
|
||||
plain: Final = tool_call.model_dump() if isinstance(tool_call, BaseModel) else tool_call
|
||||
function: Final = plain.get("function") if isinstance(plain, Mapping) else None
|
||||
if not isinstance(function, Mapping):
|
||||
return (None, None)
|
||||
return (function.get("name"), function.get("arguments"))
|
||||
|
||||
|
||||
def _text_snapshot(texts: Sequence[str] | None) -> tuple[str, ...] | None:
|
||||
return None if texts is None else tuple(texts)
|
||||
|
||||
|
||||
def _tool_call_shapes(tool_calls: Sequence[object] | None) -> tuple[tuple[object, object], ...] | None:
|
||||
return None if tool_calls is None else tuple(_tool_call_shape(tool_call) for tool_call in tool_calls)
|
||||
|
||||
|
||||
def _rewrote(sent: tuple[object, ...] | None, returned: tuple[object, ...] | None) -> bool:
|
||||
return sent is not None and returned is not None and returned != sent
|
||||
|
||||
|
||||
_GuardrailMethodT = TypeVar("_GuardrailMethodT", bound=Callable[..., object])
|
||||
|
||||
|
||||
def _logged_by_inner_guardrail(method: _GuardrailMethodT) -> _GuardrailMethodT:
|
||||
vars(method)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True # rebind-ok: stamps the method the class body just defined
|
||||
return method
|
||||
|
||||
|
||||
class _StreamRewriteObserver(CustomGuardrail):
|
||||
"""Stand-in handed to the endpoint translation in place of a streaming pipeline step's
|
||||
guardrail. It records whether the guardrail returned different output than it was given,
|
||||
which for guardrails like Bedrock's ANONYMIZED action is only known at runtime. Text
|
||||
rewrites are deliverable on translations that write them back across the buffered chunks
|
||||
(``delivers_ended_stream_text_rewrites``); tool-call rewrites and text rewrites on any
|
||||
other translation are discarded by the executor, which releases the original chunks.
|
||||
The inner guardrail's ``apply_guardrail`` already records the guardrail information
|
||||
and span, so the observer's stays out of ``log_guardrail_information``."""
|
||||
|
||||
def __init__(self, inner: CustomGuardrail) -> None:
|
||||
super().__init__(guardrail_name=inner.guardrail_name)
|
||||
self.inner: Final = inner
|
||||
self.rewrote_texts = False
|
||||
self.rewrote_tool_calls = False
|
||||
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
return self.inner.structured_messages_cover_full_request()
|
||||
|
||||
@_logged_by_inner_guardrail
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict, # mutable-ok: matches CustomGuardrail.apply_guardrail
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
sent_texts: Final = _text_snapshot(inputs.get("texts"))
|
||||
sent_tool_shapes: Final = _tool_call_shapes(inputs.get("tool_calls"))
|
||||
outputs: Final = await self.inner.apply_guardrail(
|
||||
inputs=inputs, request_data=request_data, input_type=input_type, logging_obj=logging_obj
|
||||
)
|
||||
self.rewrote_texts = self.rewrote_texts or _rewrote(sent_texts, _text_snapshot(outputs.get("texts")))
|
||||
self.rewrote_tool_calls = self.rewrote_tool_calls or _rewrote(
|
||||
sent_tool_shapes, _tool_call_shapes(outputs.get("tool_calls"))
|
||||
)
|
||||
return outputs
|
||||
|
||||
|
||||
def _prepare_hook_input(
|
||||
step: PipelineStep,
|
||||
callback: CustomGuardrail,
|
||||
data: dict, # mutable-ok: same request-payload shape the hooks mutate
|
||||
raw_request_snapshot: dict | None, # mutable-ok: same request-payload shape as data
|
||||
) -> tuple[dict, bool]: # mutable-ok: returns that same request-payload dict
|
||||
"""Inject the step's guardrail name into metadata so should_run_guardrail() allows it,
|
||||
and pick the payload the step scans: a scan_raw_request step evaluates the pristine
|
||||
pre-pipeline snapshot instead of `data` (which earlier pass_data steps in this same
|
||||
pipeline may have already rewritten), same reason the normal sequential/parallel
|
||||
guardrail loops do this."""
|
||||
if "metadata" not in data:
|
||||
data["metadata"] = {} # mutable-ok: request metadata bucket, hooks mutate it
|
||||
data["metadata"]["guardrails"] = [
|
||||
step.guardrail
|
||||
] # mutable-ok: guardrails list is part of the request-payload shape
|
||||
|
||||
scans_raw_request: Final = callback.scan_raw_request
|
||||
hook_input: Final[dict] = ( # mutable-ok: same request-payload shape as data
|
||||
independent_snapshot(raw_request_snapshot) if scans_raw_request and raw_request_snapshot is not None else data
|
||||
)
|
||||
if hook_input is not data:
|
||||
hook_input.setdefault("metadata", {})["guardrails"] = [step.guardrail] # mutable-ok: request metadata shape
|
||||
return hook_input, scans_raw_request
|
||||
|
||||
|
||||
def _release_original_chunks(
|
||||
guardrail_name: str,
|
||||
streaming_chunks: list[object], # mutable-ok: shared buffered-stream chunks, restored in place
|
||||
originals: Sequence[object],
|
||||
) -> None:
|
||||
streaming_chunks[:] = originals # rebind-ok: the caller's buffer is the stream the client receives
|
||||
verbose_proxy_logger.warning(
|
||||
"Pipeline: guardrail '%s' rewrote the streamed response in a way this endpoint's streaming "
|
||||
"pipeline cannot deliver yet; the rewrite was discarded and the original stream released",
|
||||
guardrail_name,
|
||||
)
|
||||
|
||||
|
||||
class PipelineExecutor:
|
||||
"""Executes guardrail pipelines with ordered, conditional step logic."""
|
||||
|
||||
|
|
@ -44,6 +176,8 @@ class PipelineExecutor:
|
|||
call_type: str,
|
||||
policy_name: str,
|
||||
raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data
|
||||
streaming_chunks: list[Any] | None = None, # mutable-ok: shared buffered-stream chunks, read per step
|
||||
endpoint_translation: "BaseTranslation | None" = None,
|
||||
) -> PipelineExecutionResult:
|
||||
"""
|
||||
Execute pipeline steps sequentially with conditional actions.
|
||||
|
|
@ -60,6 +194,12 @@ class PipelineExecutor:
|
|||
step whose guardrail opted into ``scan_raw_request`` evaluates
|
||||
the original request instead of whatever an earlier
|
||||
``pass_data`` step in this same pipeline already rewrote.
|
||||
streaming_chunks: buffered chunks of a completed stream. When set
|
||||
(with ``endpoint_translation``), post_call steps scan the
|
||||
assembled streamed output through the endpoint translation
|
||||
instead of calling ``async_post_call_success_hook``.
|
||||
endpoint_translation: the guardrail translation for the streamed
|
||||
endpoint, resolved by the caller.
|
||||
|
||||
Returns:
|
||||
PipelineExecutionResult with terminal action and step results
|
||||
|
|
@ -84,6 +224,8 @@ class PipelineExecutor:
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
raw_request_snapshot=raw_request_snapshot,
|
||||
streaming_chunks=streaming_chunks,
|
||||
endpoint_translation=endpoint_translation,
|
||||
)
|
||||
|
||||
duration = time.perf_counter() - start_time
|
||||
|
|
@ -109,8 +251,10 @@ class PipelineExecutor:
|
|||
action,
|
||||
)
|
||||
|
||||
# Forward modified data to next step if pass_data is True
|
||||
if step.pass_data and modified_data is not None:
|
||||
# Forward modified data to the next step if pass_data is True;
|
||||
# post_call response replacements always chain, matching the flat
|
||||
# callback loop where each hook sees the previous hook's response
|
||||
if modified_data is not None and (step.pass_data or mode == "post_call"):
|
||||
working_data = {**working_data, **modified_data}
|
||||
|
||||
# Handle terminal actions
|
||||
|
|
@ -118,18 +262,22 @@ class PipelineExecutor:
|
|||
return _allow_result(step_results=step_results, working_data=working_data, request_data=data)
|
||||
|
||||
if action == "block":
|
||||
_carry_working_guardrail_information(working_data=working_data, request_data=data)
|
||||
return PipelineExecutionResult(
|
||||
terminal_action="block",
|
||||
step_results=step_results,
|
||||
error_message=error_detail,
|
||||
original_exception=original_exception,
|
||||
modified_data=working_data if working_data != data else None,
|
||||
)
|
||||
|
||||
if action == "modify_response":
|
||||
_carry_working_guardrail_information(working_data=working_data, request_data=data)
|
||||
return PipelineExecutionResult(
|
||||
terminal_action="modify_response",
|
||||
step_results=step_results,
|
||||
modify_response_message=step.modify_response_message or error_detail,
|
||||
modified_data=working_data if working_data != data else None,
|
||||
)
|
||||
|
||||
# action == "next" → continue to next step
|
||||
|
|
@ -137,6 +285,51 @@ class PipelineExecutor:
|
|||
# Ran out of steps without a terminal action → default allow
|
||||
return _allow_result(step_results=step_results, working_data=working_data, request_data=data)
|
||||
|
||||
@staticmethod
|
||||
async def _run_streaming_step(
|
||||
step: PipelineStep,
|
||||
callback: CustomGuardrail,
|
||||
endpoint_translation: "BaseTranslation",
|
||||
streaming_chunks: list[object], # mutable-ok: shared buffered-stream chunks the translation rewrites in place
|
||||
hook_input: dict[str, object], # mutable-ok: same request-payload shape as data
|
||||
user_api_key_dict: "UserAPIKeyAuth | None",
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None",
|
||||
) -> None:
|
||||
"""Run one streaming post_call step through the endpoint translation, delivering
|
||||
text rewrites on translations that support ended-stream write-back. A rewrite that
|
||||
cannot reach the client yet (a tool-call rewrite, a text rewrite on a translation
|
||||
without write-back, or one the translation refused with
|
||||
``UndeliverableStreamRewrite``) is discarded: the buffered chunks go back to the
|
||||
originals and the step passes, so the client gets the stream the merge base sent."""
|
||||
observer: Final = _StreamRewriteObserver(callback)
|
||||
deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_text_rewrites
|
||||
originals: Final = copy.deepcopy(streaming_chunks)
|
||||
try:
|
||||
if deliver_rewrites:
|
||||
await endpoint_translation.process_output_streaming_response(
|
||||
responses_so_far=streaming_chunks,
|
||||
guardrail_to_apply=observer,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=hook_input,
|
||||
deliver_ended_stream_rewrites=True,
|
||||
)
|
||||
else:
|
||||
await endpoint_translation.process_output_streaming_response(
|
||||
responses_so_far=streaming_chunks,
|
||||
guardrail_to_apply=observer,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=hook_input,
|
||||
)
|
||||
except UndeliverableStreamRewrite:
|
||||
_release_original_chunks(step.guardrail, streaming_chunks, originals)
|
||||
else:
|
||||
if observer.rewrote_tool_calls or (observer.rewrote_texts and not deliver_rewrites):
|
||||
_release_original_chunks(step.guardrail, streaming_chunks, originals)
|
||||
if not callback.records_own_guardrail_information:
|
||||
add_guardrail_to_applied_guardrails_header(request_data=hook_input, guardrail_name=step.guardrail)
|
||||
|
||||
@staticmethod
|
||||
async def _run_step(
|
||||
step: PipelineStep,
|
||||
|
|
@ -145,6 +338,8 @@ class PipelineExecutor:
|
|||
user_api_key_dict: Any,
|
||||
call_type: str,
|
||||
raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data
|
||||
streaming_chunks: list[Any] | None = None, # mutable-ok: shared buffered-stream chunks, read per step
|
||||
endpoint_translation: "BaseTranslation | None" = None,
|
||||
) -> tuple[
|
||||
Literal["pass", "fail", "error"],
|
||||
dict | None,
|
||||
|
|
@ -168,34 +363,17 @@ class PipelineExecutor:
|
|||
verbose_proxy_logger.warning("Pipeline: guardrail '%s' not found in callbacks", step.guardrail)
|
||||
return ("error", None, f"Guardrail '{step.guardrail}' not found", None)
|
||||
|
||||
hook_input, scans_raw_request = _prepare_hook_input(step, callback, data, raw_request_snapshot)
|
||||
snapshot_entries_before: Final = len(_recorded_guardrail_information(hook_input))
|
||||
|
||||
# Use unified_guardrail path if callback implements apply_guardrail
|
||||
target: CustomLogger = callback
|
||||
use_unified: Final = PipelineExecutor.supports_unified_execution(callback)
|
||||
if use_unified and streaming_chunks is None:
|
||||
hook_input["guardrail_to_apply"] = callback
|
||||
target = UnifiedLLMGuardrails()
|
||||
|
||||
try:
|
||||
# Inject guardrail name into metadata so should_run_guardrail() allows it
|
||||
if "metadata" not in data:
|
||||
data["metadata"] = {}
|
||||
data["metadata"]["guardrails"] = [step.guardrail]
|
||||
|
||||
# A scan_raw_request step evaluates the pristine pre-pipeline
|
||||
# snapshot instead of `data` (which earlier pass_data steps in
|
||||
# this same pipeline may have already rewritten), same reason
|
||||
# the normal sequential/parallel guardrail loops do this.
|
||||
scans_raw_request: Final = callback.scan_raw_request
|
||||
hook_input: Final[dict] = ( # mutable-ok: same request-payload shape as data
|
||||
independent_snapshot(raw_request_snapshot)
|
||||
if scans_raw_request and raw_request_snapshot is not None
|
||||
else data
|
||||
)
|
||||
if hook_input is not data:
|
||||
hook_input.setdefault("metadata", {})["guardrails"] = [step.guardrail]
|
||||
|
||||
# Use unified_guardrail path if callback implements apply_guardrail
|
||||
target: CustomLogger = callback
|
||||
use_unified: Final = (
|
||||
"apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks
|
||||
)
|
||||
if use_unified:
|
||||
hook_input["guardrail_to_apply"] = callback
|
||||
target = UnifiedLLMGuardrails()
|
||||
|
||||
if mode == "pre_call":
|
||||
response = await target.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -207,6 +385,24 @@ class PipelineExecutor:
|
|||
callback.mark_pre_call_hook_ran(data)
|
||||
if isinstance(response, dict):
|
||||
callback.mark_pre_call_hook_ran(response)
|
||||
elif mode == "post_call" and streaming_chunks is not None:
|
||||
if not use_unified or endpoint_translation is None:
|
||||
return (
|
||||
"error",
|
||||
None,
|
||||
f"Guardrail '{step.guardrail}' does not support streaming pipeline execution",
|
||||
None,
|
||||
)
|
||||
await PipelineExecutor._run_streaming_step(
|
||||
step=step,
|
||||
callback=callback,
|
||||
endpoint_translation=endpoint_translation,
|
||||
streaming_chunks=streaming_chunks,
|
||||
hook_input=hook_input,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_logging_obj=data.get("litellm_logging_obj"),
|
||||
)
|
||||
response = None
|
||||
elif mode == "post_call":
|
||||
response = await target.async_post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -220,11 +416,19 @@ class PipelineExecutor:
|
|||
# same contract as run_in_parallel/scan_raw_request elsewhere: any
|
||||
# data it returned is discarded, since applying it on top of the
|
||||
# raw snapshot would silently undo whatever an earlier step in
|
||||
# this pipeline already did.
|
||||
modified_data = None
|
||||
if response is not None and isinstance(response, dict) and not scans_raw_request:
|
||||
modified_data = response
|
||||
return ("pass", modified_data, None, None)
|
||||
# this pipeline already did. A post_call hook's non-None return is
|
||||
# a replacement response (the flat callback-loop contract), carried
|
||||
# under the same "response" key the step input uses.
|
||||
if response is None or scans_raw_request:
|
||||
return ("pass", None, None, None)
|
||||
if mode == "post_call":
|
||||
return (
|
||||
"pass",
|
||||
{"response": response},
|
||||
None,
|
||||
None,
|
||||
) # mutable-ok: modified-data contract is a plain dict
|
||||
return ("pass", response if isinstance(response, dict) else None, None, None)
|
||||
|
||||
except Exception as e:
|
||||
if CustomGuardrail._is_guardrail_intervention(e):
|
||||
|
|
@ -233,6 +437,18 @@ class PipelineExecutor:
|
|||
else:
|
||||
verbose_proxy_logger.error("Pipeline: unexpected error from guardrail '%s': %s", step.guardrail, e)
|
||||
return ("error", None, str(e), e)
|
||||
finally:
|
||||
if hook_input is not data:
|
||||
_append_guardrail_information(
|
||||
request_data=data,
|
||||
entries=_recorded_guardrail_information(hook_input)[snapshot_entries_before:],
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def supports_unified_execution(callback: CustomGuardrail) -> bool:
|
||||
"""Whether this guardrail runs through the unified apply_guardrail path,
|
||||
the interface streaming pipeline execution requires."""
|
||||
return "apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks
|
||||
|
||||
@staticmethod
|
||||
def find_guardrail_callback(guardrail_name: str) -> CustomGuardrail | None:
|
||||
|
|
@ -283,6 +499,40 @@ def _restore_request_guardrails(
|
|||
return {**working_data, "metadata": stripped} # mutable-ok: request dict
|
||||
|
||||
|
||||
_GUARDRAIL_INFORMATION_KEY: Final = "standard_logging_guardrail_information"
|
||||
|
||||
|
||||
def _recorded_guardrail_information(source: Mapping[str, object]) -> list[StandardLoggingGuardrailInformation]:
|
||||
bucket: Final = source.get(get_metadata_variable_name_from_kwargs(source))
|
||||
recorded: Final = bucket.get(_GUARDRAIL_INFORMATION_KEY) if isinstance(bucket, dict) else None
|
||||
return recorded if isinstance(recorded, list) else []
|
||||
|
||||
|
||||
def _append_guardrail_information(
|
||||
request_data: dict[str, object], # mutable-ok: same request-payload shape as execute_steps' data
|
||||
entries: Sequence[StandardLoggingGuardrailInformation],
|
||||
) -> None:
|
||||
if not entries:
|
||||
return
|
||||
_, request_bucket = get_or_create_metadata_bucket(request_data)
|
||||
existing: Final = request_bucket.get(_GUARDRAIL_INFORMATION_KEY)
|
||||
if isinstance(existing, list):
|
||||
existing.extend(entries)
|
||||
return
|
||||
request_bucket[_GUARDRAIL_INFORMATION_KEY] = list(entries)
|
||||
|
||||
|
||||
def _carry_working_guardrail_information(
|
||||
working_data: Mapping[str, object],
|
||||
request_data: dict[str, object], # mutable-ok: same request-payload shape as execute_steps' data
|
||||
) -> None:
|
||||
recorded: Final = _recorded_guardrail_information(working_data)
|
||||
existing: Final = _recorded_guardrail_information(request_data)
|
||||
if recorded is existing:
|
||||
return
|
||||
_append_guardrail_information(request_data=request_data, entries=[e for e in recorded if e not in existing])
|
||||
|
||||
|
||||
def _pipeline_action_for_outcome(step: PipelineStep, outcome: str) -> str:
|
||||
"""
|
||||
Map pipeline step outcome to the configured action.
|
||||
|
|
|
|||
|
|
@ -8,10 +8,13 @@ from __future__ import annotations
|
|||
|
||||
import glob
|
||||
import os
|
||||
import re
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
_LIVE_GAUGE_PID: Final = re.compile(r"gauge_live[a-z]*_(\d+)\.db$")
|
||||
|
||||
|
||||
def wipe_directory(directory: str) -> None:
|
||||
"""Delete all .db files in the directory. Called once before workers fork."""
|
||||
|
|
@ -38,3 +41,35 @@ def mark_worker_exit(worker_pid: int) -> None:
|
|||
verbose_proxy_logger.info("Prometheus cleanup: marked worker %s as dead", worker_pid)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("Failed to mark prometheus worker %s as dead: %s", worker_pid, e)
|
||||
|
||||
|
||||
def _is_running(pid: int) -> bool:
|
||||
try:
|
||||
os.kill(pid, 0)
|
||||
except ProcessLookupError:
|
||||
return False
|
||||
except PermissionError:
|
||||
return True
|
||||
return True
|
||||
|
||||
|
||||
def mark_dead_workers(directory: str) -> tuple[int, ...]:
|
||||
"""Drop the live-gauge files of workers that no longer exist and return their pids.
|
||||
|
||||
Uvicorn's multi-worker supervisor has no exit hook, so a replacement worker calls this at startup; without it
|
||||
a crashed worker's in-flight gauges stay in the aggregate forever.
|
||||
"""
|
||||
owners: Final = frozenset(
|
||||
int(match.group(1))
|
||||
for match in map(_LIVE_GAUGE_PID.search, glob.glob(os.path.join(directory, "gauge_live*_*.db")))
|
||||
if match is not None
|
||||
)
|
||||
dead: Final = tuple(sorted(pid for pid in owners if pid != os.getpid() and not _is_running(pid)))
|
||||
if not dead:
|
||||
return dead
|
||||
from prometheus_client import multiprocess
|
||||
|
||||
for pid in dead:
|
||||
multiprocess.mark_process_dead(pid, path=directory)
|
||||
verbose_proxy_logger.info("Prometheus cleanup: marked dead workers %s in %s", dead, directory)
|
||||
return dead
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from litellm.integrations.prometheus_metrics_endpoint import make_metrics_asgi_a
|
|||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
METRICS_PATH: Final = "/metrics"
|
||||
HEALTH_PATH: Final = "/health"
|
||||
PID_HEADER: Final = "x-litellm-metrics-pid"
|
||||
_PARENT_POLL_INTERVAL_SECONDS: Final = 1.0
|
||||
_STARTUP_TIMEOUT_SECONDS: Final = 30.0
|
||||
|
|
@ -77,6 +78,10 @@ def build_metrics_app(multiproc_dir: str) -> FastAPI:
|
|||
app: Final = FastAPI(title="LiteLLM Prometheus metrics", docs_url=None, redoc_url=None, openapi_url=None)
|
||||
app.mount(METRICS_PATH, _add_pid_header(make_metrics_asgi_app(registry)))
|
||||
|
||||
@app.get(HEALTH_PATH)
|
||||
def health() -> dict[str, str]:
|
||||
return {"status": "healthy", "multiproc_dir": multiproc_dir}
|
||||
|
||||
return app
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -278,6 +278,7 @@ from litellm.litellm_core_utils.asyncify import asyncify
|
|||
from litellm.litellm_core_utils.audio_utils.utils import resolve_speech_media_type
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
drop_params_flag,
|
||||
get_litellm_metadata_from_kwargs,
|
||||
)
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
|
|
@ -631,6 +632,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
|||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
router as pass_through_router,
|
||||
)
|
||||
from litellm.proxy.prometheus_cleanup import mark_dead_workers, mark_worker_exit
|
||||
from litellm.proxy.public_endpoints import router as public_endpoints_router
|
||||
from litellm.proxy.public_endpoints.public_v1 import router as public_v1_router
|
||||
from litellm.proxy.rag_endpoints.endpoints import router as rag_router
|
||||
|
|
@ -1059,6 +1061,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
|
|||
|
||||
init_verbose_loggers()
|
||||
|
||||
prometheus_multiproc_dir: Final = os.environ.get("PROMETHEUS_MULTIPROC_DIR")
|
||||
if prometheus_multiproc_dir:
|
||||
mark_dead_workers(prometheus_multiproc_dir)
|
||||
|
||||
## RUN WORKER STARTUP HOOKS (e.g., gflags initialization) ##
|
||||
_startup_hooks_env: Final = os.environ.get("LITELLM_WORKER_STARTUP_HOOKS", "")
|
||||
if _startup_hooks_env:
|
||||
|
|
@ -1365,6 +1371,9 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
|
|||
|
||||
await proxy_shutdown_event(worker_heartbeat=worker_heartbeat)
|
||||
|
||||
if prometheus_multiproc_dir:
|
||||
mark_worker_exit(os.getpid())
|
||||
|
||||
|
||||
def _generate_stable_operation_id(route: "APIRoute") -> str:
|
||||
operation_id = re.sub(r"\W", "_", f"{route.name}{route.path_format}")
|
||||
|
|
@ -2574,10 +2583,11 @@ async def _repair_stale_spend_counter(counter_key: str, db_spend: float) -> None
|
|||
)
|
||||
|
||||
|
||||
async def reseed_spend_counter_from_db(counter_key: str) -> None:
|
||||
async def reseed_spend_counter_from_db(counter_key: str) -> bool:
|
||||
"""Recover a counter that the reservation reconcile found in an inconsistent
|
||||
state (missing, or where applying the reconcile delta would drive it
|
||||
negative) by reseeding it from the DB instead of deleting it.
|
||||
negative) by reseeding it from the DB instead of deleting it. Returns
|
||||
whether a DB row was found and the counter was reseeded.
|
||||
|
||||
The DB row is a LAGGING authoritative floor, not post-request truth: the
|
||||
entity .spend column is flushed in batches (every PROXY_BATCH_WRITE_AT), so
|
||||
|
|
@ -2592,8 +2602,9 @@ async def reseed_spend_counter_from_db(counter_key: str) -> None:
|
|||
"""
|
||||
db_spend: Final = await SpendCounterReseed.from_db(prisma_client=prisma_client, counter_key=counter_key)
|
||||
if db_spend is None:
|
||||
return
|
||||
return False
|
||||
await _repair_stale_spend_counter(counter_key=counter_key, db_spend=db_spend)
|
||||
return True
|
||||
|
||||
|
||||
async def _floor_spend_from_db(
|
||||
|
|
@ -2655,21 +2666,14 @@ async def _authoritative_floor_spend(
|
|||
return db_spend
|
||||
|
||||
|
||||
async def _read_spend_counter_estimate(counter_key: str, fallback_spend: float) -> tuple[float, bool]:
|
||||
"""Return (spend, authoritative). ``authoritative`` is True when the value
|
||||
came from Redis or a fresh DB read (cross-pod truth), False when it came
|
||||
from the per-pod in-memory copy or the caller's fallback. Only the
|
||||
fail-closed path reads the flag; normal callers ignore it."""
|
||||
# 1. Redis first (cross-pod authoritative). On clean miss, skip
|
||||
# in-memory: per-pod in-memory only has this pod's writes, so it
|
||||
# would mask cross-pod increments.
|
||||
redis_clean_miss = False
|
||||
async def read_spend_counter_cache_value(counter_key: str) -> tuple[float | None, bool]:
|
||||
"""Return (value, authoritative) for the live counter, None when absent. A clean
|
||||
Redis miss is final: the per-pod in-memory copy outlives the Redis TTL and only
|
||||
holds this pod's writes, so it is consulted only when Redis is unreachable."""
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
try:
|
||||
val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key)
|
||||
if val is not None:
|
||||
return float(val), True
|
||||
redis_clean_miss = True
|
||||
redis_val: Final = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key)
|
||||
return (float(redis_val) if redis_val is not None else None), True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"get_current_spend: Redis read failed for %s, falling back to in-memory: %s",
|
||||
|
|
@ -2677,13 +2681,20 @@ async def _read_spend_counter_estimate(counter_key: str, fallback_spend: float)
|
|||
e,
|
||||
)
|
||||
|
||||
# 2. In-memory only when Redis is unreachable.
|
||||
if not redis_clean_miss:
|
||||
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
|
||||
if val is not None:
|
||||
return float(val), False
|
||||
in_memory_val: Final = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
|
||||
return (float(in_memory_val) if in_memory_val is not None else None), False
|
||||
|
||||
# 3. Reseed from DB - fallback_spend lags cross-pod, would allow bypass.
|
||||
|
||||
async def _read_spend_counter_estimate(counter_key: str, fallback_spend: float) -> tuple[float, bool]:
|
||||
"""Return (spend, authoritative). ``authoritative`` is True when the value
|
||||
came from Redis or a fresh DB read (cross-pod truth), False when it came
|
||||
from the per-pod in-memory copy or the caller's fallback. Only the
|
||||
fail-closed path reads the flag; normal callers ignore it."""
|
||||
cached_val, cached_authoritative = await read_spend_counter_cache_value(counter_key=counter_key)
|
||||
if cached_val is not None:
|
||||
return cached_val, cached_authoritative
|
||||
|
||||
# Reseed from DB - fallback_spend lags cross-pod, would allow bypass.
|
||||
db_spend: Final = await SpendCounterReseed.coalesced(
|
||||
prisma_client=prisma_client,
|
||||
spend_counter_cache=spend_counter_cache,
|
||||
|
|
@ -5509,6 +5520,8 @@ class ProxyConfig:
|
|||
|
||||
parse_budget_reset_time(value)
|
||||
setattr(litellm, key, value)
|
||||
elif key == "drop_params":
|
||||
litellm.drop_params = drop_params_flag(value, "litellm_settings.drop_params", verbose_proxy_logger)
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"%s setting litellm.%s=%s%s",
|
||||
|
|
@ -7109,11 +7122,10 @@ class ProxyConfig:
|
|||
],
|
||||
)
|
||||
|
||||
# Only load models from DB if "models" is in supported_db_objects (or if supported_db_objects is not set)
|
||||
if self._should_load_db_object(object_type="models"):
|
||||
new_models: Final = await self._get_models_from_db(prisma_client=prisma_client)
|
||||
|
||||
# update llm router
|
||||
load_models: Final = self._should_load_db_object(object_type="models")
|
||||
new_models: Final = await self._get_models_from_db(prisma_client=prisma_client) if load_models else None
|
||||
await self.get_credentials(prisma_client=prisma_client)
|
||||
if load_models:
|
||||
still_desired_ids = await self._update_llm_router(
|
||||
new_models=new_models, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
|
|
@ -7153,12 +7165,9 @@ class ProxyConfig:
|
|||
async def _resync_config_from_db() -> None:
|
||||
await self.add_deployment(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj)
|
||||
|
||||
async def _resync_credentials_from_db() -> None:
|
||||
await self.get_credentials(prisma_client=prisma_client)
|
||||
|
||||
subscriber: Final = ConfigSyncSubscriber(
|
||||
redis_cache=redis_cache,
|
||||
resync_callbacks=(_resync_config_from_db, _resync_credentials_from_db),
|
||||
resync_callbacks=(_resync_config_from_db,),
|
||||
)
|
||||
self.config_sync_subscriber = subscriber
|
||||
subscriber.start()
|
||||
|
|
@ -8013,7 +8022,7 @@ class ProxyConfig:
|
|||
|
||||
async def get_credentials(self, prisma_client: PrismaClient):
|
||||
try:
|
||||
credentials = await CredentialsRepository(prisma_client).find_all()
|
||||
credentials = await CredentialsRepository(WriterPinnedClient(prisma_client.db)).find_all()
|
||||
credentials = [self.decrypt_credentials(cred) for cred in credentials]
|
||||
await self.delete_credentials(credentials) # delete credentials that are not in the all-up list
|
||||
CredentialAccessor.upsert_credentials(credentials) # upsert credentials that are in the all-up list
|
||||
|
|
@ -9597,19 +9606,6 @@ class ProxyStartupEvent:
|
|||
)
|
||||
|
||||
if store_model_in_db is True:
|
||||
### GET STORED CREDENTIALS ###
|
||||
scheduler.add_job(
|
||||
proxy_config.get_credentials,
|
||||
"interval",
|
||||
seconds=config_reload_interval_seconds,
|
||||
# REMOVED jitter parameter - major cause of memory leak
|
||||
args=[prisma_client],
|
||||
id="get_credentials_job",
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
await proxy_config.get_credentials(prisma_client=prisma_client)
|
||||
|
||||
# MEMORY LEAK FIX: Increase interval from 10s to 30s minimum
|
||||
# Frequent polling was causing excessive memory allocations
|
||||
scheduler.add_job(
|
||||
|
|
@ -9623,7 +9619,7 @@ class ProxyStartupEvent:
|
|||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
|
||||
# this will load all existing models on proxy startup
|
||||
# this will load all existing credentials and models on proxy startup
|
||||
await proxy_config.add_deployment(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj)
|
||||
|
||||
proxy_config.start_config_sync_subscriber(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -905,13 +905,13 @@ async def _set_reserved_entry_actual_cost(
|
|||
increment=adjustment,
|
||||
)
|
||||
elif reseed_on_inconsistent:
|
||||
# Post-call reconcile / release: the counter was flushed or reseeded
|
||||
# between reservation and reconcile (Redis restart / cross-pod reset),
|
||||
# so the optimistic delta no longer applies. Recover by reseeding from
|
||||
# the DB's lagging authoritative floor rather than deleting the counter
|
||||
# and failing open — deleting it is what left budgets unenforced after a
|
||||
# Redis reload.
|
||||
await reseed_spend_counter_from_db(counter_key=counter_key)
|
||||
# Post-call reconcile / release: the counter was flushed, expired or reseeded
|
||||
# between reservation and reconcile, so the optimistic delta no longer applies.
|
||||
# Reseed from the DB floor (which cannot include this request's cost yet) and
|
||||
# add the settled cost, since increment_spend_counters skips reserved keys.
|
||||
reseeded: Final = await reseed_spend_counter_from_db(counter_key=counter_key)
|
||||
if reseeded and actual_cost > 0:
|
||||
await _increment_spend_counter_cache(counter_key=counter_key, increment=actual_cost)
|
||||
else:
|
||||
# Pre-call admission resize: the in-flight reservation cost is not yet
|
||||
# persisted, so the DB floor would discard it. Keep the original
|
||||
|
|
@ -925,18 +925,16 @@ async def _counter_can_apply_adjustment(
|
|||
counter_key: str,
|
||||
adjustment: float,
|
||||
) -> bool:
|
||||
from litellm.proxy.proxy_server import spend_counter_cache
|
||||
from litellm.proxy.proxy_server import read_spend_counter_cache_value
|
||||
|
||||
current_value: Final = await spend_counter_cache.async_get_cache(key=counter_key)
|
||||
try:
|
||||
current_value, _ = await read_spend_counter_cache_value(counter_key=counter_key)
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
if current_value is None:
|
||||
return False
|
||||
|
||||
try:
|
||||
current_float: Final = float(current_value)
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
|
||||
return not (adjustment < 0 and current_float + adjustment < -1e-12)
|
||||
return not (adjustment < 0 and current_value + adjustment < -1e-12)
|
||||
|
||||
|
||||
async def _release_applied_entries_best_effort(
|
||||
|
|
|
|||
|
|
@ -590,6 +590,7 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
|
|||
metadata=metadata,
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
omit_when_missing=_omits_session_id_when_missing(metadata),
|
||||
batch_trace_session_id=_get_batch_trace_session_id(call_type=call_type, request_id=id),
|
||||
),
|
||||
request_duration_ms=_get_request_duration_ms(start_time, end_time),
|
||||
status=_get_status_for_spend_log(
|
||||
|
|
@ -628,20 +629,44 @@ def _omits_session_id_when_missing(metadata: Mapping[str, object] | None) -> boo
|
|||
return general_settings.get("missing_session_id") == "omit"
|
||||
|
||||
|
||||
_BATCH_TRACE_CALL_TYPES: Final = frozenset(
|
||||
{
|
||||
CallTypes.create_batch.value,
|
||||
CallTypes.acreate_batch.value,
|
||||
CallTypes.retrieve_batch.value,
|
||||
CallTypes.aretrieve_batch.value,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _get_batch_trace_session_id(call_type: str | None, request_id: str | None) -> str | None:
|
||||
"""A batch's create row and its poller-written cost row both derive their request id
|
||||
from the same batch id (the cost row appends BATCH_COST_REQUEST_ID_SUFFIX), so using
|
||||
that id as the session groups the batch lifecycle into one trace on the logs UI. The
|
||||
poller builds its own logging context, so per-request trace ids can never link them."""
|
||||
if call_type not in _BATCH_TRACE_CALL_TYPES or not request_id:
|
||||
return None
|
||||
return request_id.removesuffix(BATCH_COST_REQUEST_ID_SUFFIX)
|
||||
|
||||
|
||||
def _get_session_id_for_spend_log(
|
||||
kwargs: Mapping[str, object],
|
||||
metadata: Mapping[str, object] | None,
|
||||
standard_logging_payload: StandardLoggingPayload | None,
|
||||
omit_when_missing: bool,
|
||||
batch_trace_session_id: str | None = None,
|
||||
) -> str | None:
|
||||
"""Under `omit` only `metadata.session_id`, the key Langfuse reads, counts as a session; `litellm_session_id` may
|
||||
be a copied trace id."""
|
||||
be a copied trace id. Batch call types carry a deterministic session derived from the batch id, which outranks
|
||||
the per-request trace ids because those differ between the create call and the cost poller's row."""
|
||||
if omit_when_missing:
|
||||
session_id: Final = metadata.get("session_id") if metadata else None
|
||||
return str(session_id) if session_id else None
|
||||
|
||||
from litellm._uuid import uuid
|
||||
|
||||
if batch_trace_session_id is not None:
|
||||
return batch_trace_session_id
|
||||
if standard_logging_payload is not None and standard_logging_payload.get("trace_id") is not None:
|
||||
return str(standard_logging_payload.get("trace_id"))
|
||||
if kwargs.get("litellm_trace_id") is not None:
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import sys
|
|||
import threading
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Collection, Coroutine, Mapping, Sequence
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
|
|
@ -139,6 +139,7 @@ from litellm.proxy.db.token_auth import (
|
|||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
resolve_endpoint_translation,
|
||||
)
|
||||
from litellm.proxy.hooks import PROXY_HOOKS, get_proxy_hook
|
||||
from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck
|
||||
|
|
@ -177,6 +178,9 @@ from litellm.types.mcp import (
|
|||
)
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult
|
||||
from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams
|
||||
from litellm.utils import (
|
||||
_add_custom_logger_callback_to_specific_event, # pyright: ignore[reportPrivateUsage] # only string-to-logger helper
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.types import CallToolResult
|
||||
|
|
@ -446,12 +450,161 @@ def _policy_pipelines(data: Mapping[str, object]) -> tuple[tuple[str, "Guardrail
|
|||
)
|
||||
|
||||
|
||||
def _pipeline_managed_guardrail_names(data: Mapping[str, object]) -> frozenset[str]:
|
||||
managed: Final = _policy_state_metadata(data).get("_pipeline_managed_guardrails")
|
||||
return (
|
||||
frozenset(cast("Collection[str]", managed)) # cast-ok: the policy engine wrote these guardrail names
|
||||
if managed
|
||||
else frozenset()
|
||||
def _pipeline_step_guardrail_names(pipelines: Sequence[tuple[str, "GuardrailPipeline"]]) -> frozenset[str]:
|
||||
return frozenset(step.guardrail for _policy_name, pipeline in pipelines for step in pipeline.steps)
|
||||
|
||||
|
||||
def _pipeline_managed_guardrail_names(
|
||||
data: Mapping[str, object], mode: Literal["pre_call", "post_call"]
|
||||
) -> frozenset[str]:
|
||||
return _pipeline_step_guardrail_names(
|
||||
tuple((policy_name, pipeline) for policy_name, pipeline in _policy_pipelines(data) if pipeline.mode == mode)
|
||||
)
|
||||
|
||||
|
||||
def _partition_post_call_callbacks() -> tuple[tuple[CustomGuardrail, ...], tuple[CustomLogger, ...]]:
|
||||
resolved: Final = tuple(
|
||||
litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
|
||||
cast( # cast-ok: the resolver returns None for unknown names, filtered below
|
||||
_custom_logger_compatible_callbacks_literal, callback
|
||||
)
|
||||
)
|
||||
if isinstance(callback, str)
|
||||
else callback
|
||||
for callback in litellm.callbacks
|
||||
)
|
||||
present: Final = tuple(callback for callback in resolved if callback is not None)
|
||||
guardrails: Final = tuple(callback for callback in present if isinstance(callback, CustomGuardrail))
|
||||
others: Final = cast( # cast-ok: mirrors the legacy loop, which treated every non-guardrail entry as a CustomLogger
|
||||
"tuple[CustomLogger, ...]",
|
||||
tuple(callback for callback in present if not isinstance(callback, CustomGuardrail)),
|
||||
)
|
||||
return (guardrails, others)
|
||||
|
||||
|
||||
def _merge_pipeline_metadata_bucket(
|
||||
data: dict, bucket_key: str, modified_bucket_value: object
|
||||
) -> None: # mutable-ok: request payload dict, written in place
|
||||
if not isinstance(modified_bucket_value, dict):
|
||||
return
|
||||
modified_bucket: Final = cast("dict[str, object]", modified_bucket_value) # cast-ok: metadata buckets are str-keyed
|
||||
surviving_writes: Final = {
|
||||
key: value for key, value in modified_bucket.items() if key != "guardrails"
|
||||
} # mutable-ok: merged into the live request metadata bucket in place
|
||||
existing_bucket: Final = data.get(bucket_key)
|
||||
if isinstance(existing_bucket, dict):
|
||||
cast("dict[str, object]", existing_bucket).update(surviving_writes) # cast-ok: metadata buckets are str-keyed
|
||||
else:
|
||||
data[bucket_key] = surviving_writes
|
||||
|
||||
|
||||
def _merge_pipeline_metadata_writes(
|
||||
data: dict, modified_data: Mapping[str, object]
|
||||
) -> None: # mutable-ok: request payload dict, written in place
|
||||
"""
|
||||
Copy metadata-bucket writes from a pipeline's working copy back onto the request.
|
||||
|
||||
Post_call pipelines run step hooks against a copied request dict so the payload
|
||||
already sent upstream stays untouched, but hooks record proxy-internal logging
|
||||
state in the metadata buckets (``applied_guardrails`` for response headers,
|
||||
``standard_logging_guardrail_information`` for spend logs), and those writes
|
||||
must reach the request dict the proxy keeps reading after the pipeline returns.
|
||||
|
||||
The ``guardrails`` key is the executor's per-step activation flag for
|
||||
``should_run_guardrail``, not a hook write, so it stays in the working copy.
|
||||
"""
|
||||
for bucket_key in ("metadata", "litellm_metadata"):
|
||||
_merge_pipeline_metadata_bucket(data, bucket_key, modified_data.get(bucket_key))
|
||||
|
||||
|
||||
def _pipeline_step_supports_unified_streaming(guardrail_name: str) -> bool:
|
||||
callback: Final = PipelineExecutor.find_guardrail_callback(guardrail_name)
|
||||
return callback is not None and PipelineExecutor.supports_unified_execution(callback)
|
||||
|
||||
|
||||
def _post_call_pipelines(data: Mapping[str, object]) -> tuple[tuple[str, "GuardrailPipeline"], ...]:
|
||||
return tuple(
|
||||
(policy_name, pipeline) for policy_name, pipeline in _policy_pipelines(data) if pipeline.mode == "post_call"
|
||||
)
|
||||
|
||||
|
||||
def _warn_background_skips_post_call_pipelines(data: Mapping[str, object]) -> None:
|
||||
if data.get("background") is not True:
|
||||
return
|
||||
policy_names: Final = tuple(policy_name for policy_name, _pipeline in _post_call_pipelines(data))
|
||||
if not policy_names:
|
||||
return
|
||||
verbose_proxy_logger.warning(
|
||||
"Policies with post_call guardrail pipelines do not run on background responses yet; "
|
||||
"the response is released ungoverned by them: %s",
|
||||
", ".join(policy_names),
|
||||
)
|
||||
|
||||
|
||||
def _pipeline_is_streamable(policy_name: str, pipeline: "GuardrailPipeline") -> bool:
|
||||
unsupported: Final = tuple(
|
||||
dict.fromkeys(
|
||||
step.guardrail for step in pipeline.steps if not _pipeline_step_supports_unified_streaming(step.guardrail)
|
||||
)
|
||||
)
|
||||
if not unsupported:
|
||||
return True
|
||||
verbose_proxy_logger.warning(
|
||||
"Policy '%s' has post_call pipeline guardrails without the unified apply_guardrail interface, "
|
||||
"which streaming pipelines need; the stream skips the pipeline and its guardrails run on their own: %s",
|
||||
policy_name,
|
||||
", ".join(unsupported),
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def _route_supports_streaming_pipelines(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
return not user_api_key_dict.request_route or resolve_endpoint_translation(user_api_key_dict, None) is not None
|
||||
|
||||
|
||||
def _stream_gated_guardrail_names(
|
||||
request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth
|
||||
) -> frozenset[str]:
|
||||
if not _route_supports_streaming_pipelines(user_api_key_dict):
|
||||
return frozenset()
|
||||
return _pipeline_step_guardrail_names(
|
||||
tuple(
|
||||
(policy_name, pipeline)
|
||||
for policy_name, pipeline in _post_call_pipelines(request_data)
|
||||
if all(_pipeline_step_supports_unified_streaming(step.guardrail) for step in pipeline.steps)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _streamable_post_call_pipelines(
|
||||
request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth
|
||||
) -> tuple[tuple[str, "GuardrailPipeline"], ...]:
|
||||
"""
|
||||
The post_call pipelines a streaming response can be gated through.
|
||||
|
||||
Streaming pipelines scan the buffered stream through the endpoint guardrail
|
||||
translation of the request route, so every step's guardrail needs the
|
||||
unified apply_guardrail interface and the route needs a translation. A
|
||||
pipeline that cannot be run that way yet is left out and its guardrails
|
||||
run on the stream on their own, the way they did before pipelines ran on
|
||||
streams at all, with a warning naming the pipeline.
|
||||
"""
|
||||
post_call_pipelines: Final = _post_call_pipelines(request_data)
|
||||
if not post_call_pipelines:
|
||||
return ()
|
||||
if not _route_supports_streaming_pipelines(user_api_key_dict):
|
||||
verbose_proxy_logger.warning(
|
||||
"Policies with post_call guardrail pipelines cannot scan streaming responses on route %s yet "
|
||||
"(no endpoint guardrail translation); the stream skips the pipelines and their guardrails run "
|
||||
"on their own: %s",
|
||||
user_api_key_dict.request_route,
|
||||
", ".join(policy_name for policy_name, _pipeline in post_call_pipelines),
|
||||
)
|
||||
return ()
|
||||
return tuple(
|
||||
(policy_name, pipeline)
|
||||
for policy_name, pipeline in post_call_pipelines
|
||||
if _pipeline_is_streamable(policy_name, pipeline)
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -857,6 +1010,14 @@ class ProxyLogging:
|
|||
litellm.logging_callback_manager.add_litellm_async_success_callback(callback)
|
||||
litellm.logging_callback_manager.add_litellm_async_failure_callback(callback)
|
||||
|
||||
# Runs after load_config applied every litellm_settings key: logger __init__s read e.g. s3_callback_params
|
||||
success_callbacks: Final = tuple(cb for cb in litellm.success_callback if isinstance(cb, str))
|
||||
failure_callbacks: Final = tuple(cb for cb in litellm.failure_callback if isinstance(cb, str))
|
||||
for callback in success_callbacks:
|
||||
_add_custom_logger_callback_to_specific_event(callback, "success")
|
||||
for callback in failure_callbacks:
|
||||
_add_custom_logger_callback_to_specific_event(callback, "failure")
|
||||
|
||||
async def update_request_status(self, litellm_call_id: str, status: Literal["success", "fail"]):
|
||||
# only use this if slack alerting is being used
|
||||
if self.alerting is None:
|
||||
|
|
@ -1585,7 +1746,8 @@ class ProxyLogging:
|
|||
call_type: str,
|
||||
event_hook: str,
|
||||
raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data
|
||||
) -> dict:
|
||||
response: LLMResponseTypes | None = None,
|
||||
) -> tuple[dict, LLMResponseTypes | None]: # mutable-ok: returns the request-payload dict onward
|
||||
"""
|
||||
Execute guardrail pipelines if any are configured for this request.
|
||||
|
||||
|
|
@ -1597,20 +1759,27 @@ class ProxyLogging:
|
|||
``scan_raw_request`` evaluates the pristine request, not whatever an
|
||||
earlier ``pass_data`` step in the same pipeline already rewrote.
|
||||
|
||||
Returns the (possibly modified) data dict.
|
||||
Returns the (possibly modified) data dict, plus the replacement
|
||||
response when a post_call pipeline step returned one (None when the
|
||||
response is unchanged), matching the flat callback-loop contract.
|
||||
"""
|
||||
pipelines: Final = _policy_pipelines(data)
|
||||
if not pipelines:
|
||||
return data
|
||||
return data, None
|
||||
|
||||
current_response = response # rebind-ok: chains each pipeline's replacement response into the next
|
||||
for policy_name, pipeline in pipelines:
|
||||
if pipeline.mode != event_hook:
|
||||
continue
|
||||
|
||||
step_input: dict = (
|
||||
{**data, "response": current_response} if current_response is not None else data
|
||||
) # mutable-ok: same request-payload shape as data
|
||||
|
||||
result: PipelineExecutionResult = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data=data,
|
||||
data=step_input,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
policy_name=policy_name,
|
||||
|
|
@ -1621,26 +1790,46 @@ class ProxyLogging:
|
|||
result=result,
|
||||
data=data,
|
||||
policy_name=policy_name,
|
||||
original_response=current_response,
|
||||
)
|
||||
|
||||
return data
|
||||
if current_response is not None and result.modified_data is not None:
|
||||
current_response = result.modified_data.get("response", current_response)
|
||||
|
||||
return data, current_response if current_response is not response else None
|
||||
|
||||
@staticmethod
|
||||
def _handle_pipeline_result(
|
||||
result: PipelineExecutionResult,
|
||||
data: dict,
|
||||
policy_name: str,
|
||||
original_response: "LLMResponseTypes | Sequence[object] | None" = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Handle a PipelineExecutionResult — allow, block, or modify_response.
|
||||
|
||||
Returns data dict if allowed, raises on block/modify_response.
|
||||
``original_response`` is set on the post_call path, where the request
|
||||
payload (already sent upstream) must stay untouched; a replacement
|
||||
response carried in ``modified_data`` is adopted by the caller, and
|
||||
metadata-bucket writes (applied guardrails, guardrail logging info)
|
||||
are merged back so headers and spend logs still see them, on block
|
||||
and modify_response too, so failure spend records keep guardrail
|
||||
cost and status. On the
|
||||
streaming path it is the buffered chunk list, carried into
|
||||
``ModifyResponseException.original_response`` for usage reporting.
|
||||
"""
|
||||
if result.terminal_action == "allow":
|
||||
if result.modified_data is not None:
|
||||
data.update(result.modified_data)
|
||||
if original_response is None:
|
||||
data.update(result.modified_data)
|
||||
else:
|
||||
_merge_pipeline_metadata_writes(data, result.modified_data)
|
||||
return data
|
||||
|
||||
if result.modified_data is not None:
|
||||
_merge_pipeline_metadata_writes(data, result.modified_data)
|
||||
|
||||
if result.terminal_action == "block":
|
||||
original_exception: Final = result.original_exception
|
||||
if original_exception is not None and not _exception_changes_request_flow(original_exception):
|
||||
|
|
@ -1678,6 +1867,7 @@ class ProxyLogging:
|
|||
request_data=data,
|
||||
guardrail_name=f"pipeline:{policy_name}",
|
||||
detection_info=None,
|
||||
original_response=original_response,
|
||||
)
|
||||
|
||||
return data
|
||||
|
|
@ -1794,8 +1984,10 @@ class ProxyLogging:
|
|||
)
|
||||
|
||||
try:
|
||||
_warn_background_skips_post_call_pipelines(data)
|
||||
|
||||
# Execute guardrail pipelines before the normal callback loop
|
||||
data = await self._maybe_execute_pipelines(
|
||||
data, _ = await self._maybe_execute_pipelines( # rebind-ok: pipeline edits feed the callback loop below
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
|
|
@ -1804,7 +1996,7 @@ class ProxyLogging:
|
|||
)
|
||||
|
||||
# Get pipeline-managed guardrails to skip in normal loop
|
||||
pipeline_managed: Final = _pipeline_managed_guardrail_names(data)
|
||||
pipeline_managed: Final = _pipeline_managed_guardrail_names(data, "pre_call")
|
||||
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
# Skip the per-request callback walk entirely when nothing in
|
||||
|
|
@ -2782,36 +2974,35 @@ class ProxyLogging:
|
|||
from litellm.proxy.proxy_server import llm_router
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
guardrail_callbacks: Final[list[CustomGuardrail]] = []
|
||||
other_callbacks: Final[list[CustomLogger]] = []
|
||||
_, pipeline_response = await self._maybe_execute_pipelines(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=getattr(data.get("litellm_logging_obj"), "call_type", None) or "acompletion",
|
||||
event_hook="post_call",
|
||||
response=response,
|
||||
)
|
||||
if pipeline_response is not None:
|
||||
response = pipeline_response # rebind-ok: adopt the pipeline's replacement response, same contract as the callback loops below
|
||||
|
||||
pipeline_managed: Final = _pipeline_managed_guardrail_names(data, "post_call")
|
||||
guardrail_callbacks, other_callbacks = _partition_post_call_callbacks()
|
||||
try:
|
||||
for callback in litellm.callbacks:
|
||||
_callback: CustomLogger | None = None
|
||||
if isinstance(callback, str):
|
||||
_callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
|
||||
cast(_custom_logger_compatible_callbacks_literal, callback)
|
||||
)
|
||||
else:
|
||||
_callback = callback
|
||||
|
||||
if _callback is not None:
|
||||
if isinstance(_callback, CustomGuardrail):
|
||||
guardrail_callbacks.append(_callback)
|
||||
else:
|
||||
other_callbacks.append(_callback)
|
||||
############## Handle Guardrails ########################################
|
||||
#############################################################################
|
||||
|
||||
# Merge model-level guardrails before checking which guardrails to run
|
||||
guardrail_data: Final = _check_and_merge_model_level_guardrails(data=data, llm_router=llm_router)
|
||||
|
||||
parallel_guardrails: Final[tuple[CustomGuardrail, ...]] = tuple(
|
||||
callback for callback in guardrail_callbacks if getattr(callback, "run_in_parallel", False)
|
||||
callback
|
||||
for callback in guardrail_callbacks
|
||||
if getattr(callback, "run_in_parallel", False)
|
||||
and not (callback.guardrail_name and callback.guardrail_name in pipeline_managed)
|
||||
)
|
||||
|
||||
for callback in guardrail_callbacks:
|
||||
# Main - V2 Guardrails implementation
|
||||
|
||||
if callback.guardrail_name and callback.guardrail_name in pipeline_managed:
|
||||
continue
|
||||
|
||||
if getattr(callback, "run_in_parallel", False):
|
||||
continue
|
||||
|
||||
|
|
@ -3108,11 +3299,16 @@ class ProxyLogging:
|
|||
# dict lookups + llm_router.get_deployment() per callback per chunk.
|
||||
_cached_guardrail_data: dict | None = None
|
||||
_guardrail_data_computed = False
|
||||
pipeline_gated: Final = (
|
||||
_stream_gated_guardrail_names(data, user_api_key_dict) if caps.has_guardrail else frozenset()
|
||||
)
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
try:
|
||||
_callback: CustomLogger | None = None
|
||||
if isinstance(callback, CustomGuardrail):
|
||||
if callback.guardrail_name in pipeline_gated:
|
||||
continue
|
||||
# Main - V2 Guardrails implementation
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
|
|
@ -3169,12 +3365,13 @@ class ProxyLogging:
|
|||
1. /chat/completions
|
||||
"""
|
||||
caps: Final = ProxyLogging._callback_capabilities()
|
||||
post_call_pipelines: Final = _streamable_post_call_pipelines(request_data, user_api_key_dict)
|
||||
# Fast path: no real overrides. Internal proxy CustomLogger callbacks
|
||||
# (e.g. _PROXY_MaxBudgetLimiter, ManagedFiles) inherit the default
|
||||
# ``async for chunk: yield chunk`` body, so wrapping the iterator
|
||||
# through each of them adds N pass-through trampolines per chunk for
|
||||
# zero behavior change. Skip the chain entirely and stream through.
|
||||
if not caps.iterator_overrides:
|
||||
if not caps.iterator_overrides and not post_call_pipelines:
|
||||
try:
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
|
|
@ -3194,8 +3391,11 @@ class ProxyLogging:
|
|||
current_response = response
|
||||
stream_needs_translation: Final = ProxyLogging._stream_requires_guardrail_translation(user_api_key_dict)
|
||||
|
||||
pipeline_gated_names: Final = _pipeline_step_guardrail_names(post_call_pipelines)
|
||||
for resolved_callback, kind in caps.iterator_overrides:
|
||||
if isinstance(resolved_callback, CustomGuardrail):
|
||||
if resolved_callback.guardrail_name in pipeline_gated_names:
|
||||
continue
|
||||
if (
|
||||
resolved_callback.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call)
|
||||
is not True
|
||||
|
|
@ -3235,6 +3435,14 @@ class ProxyLogging:
|
|||
),
|
||||
)
|
||||
|
||||
if post_call_pipelines:
|
||||
current_response = self._pipeline_gated_stream(
|
||||
response=current_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
pipelines=post_call_pipelines,
|
||||
)
|
||||
|
||||
try:
|
||||
async for chunk in current_response:
|
||||
yield chunk
|
||||
|
|
@ -3250,6 +3458,81 @@ class ProxyLogging:
|
|||
# we reach this point the metadata is fully populated.
|
||||
ProxyLogging._fire_deferred_stream_logging(request_data)
|
||||
|
||||
async def _pipeline_gated_stream(
|
||||
self,
|
||||
response: "AsyncGenerator[object, None]",
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_data: dict, # mutable-ok: same request-payload shape the hooks mutate
|
||||
pipelines: "tuple[tuple[str, GuardrailPipeline], ...]",
|
||||
) -> "AsyncGenerator[Any, None]":
|
||||
"""
|
||||
Execute post_call policy pipelines against a streamed response.
|
||||
|
||||
Buffers the whole stream (nothing reaches the client until every
|
||||
pipeline allows it), then runs each pipeline's steps against the
|
||||
assembled output through the endpoint guardrail translation, the same
|
||||
machinery flat post_call guardrails use at end of stream. An allow
|
||||
releases the buffered chunks: verbatim when no guardrail rewrote the
|
||||
output, rewritten in place when one rewrote text and the translation
|
||||
delivers ended-stream rewrites (later steps then re-scan the rewritten
|
||||
chunks, so rewrites chain). A rewrite the translation cannot deliver
|
||||
yet (a tool-call rewrite, or a text rewrite on a route without
|
||||
write-back) is discarded by the executor and the original chunks are
|
||||
released, as is a buffered shape no translation resolves; a block or
|
||||
modify_response terminates with the translation's block chunks or the
|
||||
raised error.
|
||||
"""
|
||||
buffered: Final[list[object]] = [] # mutable-ok: accumulates the stream before the pipeline verdict
|
||||
async for item in response:
|
||||
buffered.append(item)
|
||||
if not buffered:
|
||||
return
|
||||
|
||||
resolved: Final = resolve_endpoint_translation(user_api_key_dict, buffered[0])
|
||||
if resolved is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Policies with post_call guardrail pipelines cannot scan this streaming response shape yet; "
|
||||
"the stream is released ungoverned by them: %s",
|
||||
", ".join(policy_name for policy_name, _pipeline in pipelines),
|
||||
)
|
||||
for buffered_item in buffered:
|
||||
yield buffered_item
|
||||
return
|
||||
call_type, endpoint_translation = resolved
|
||||
|
||||
for policy_name, pipeline in pipelines:
|
||||
result: PipelineExecutionResult = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode="post_call",
|
||||
data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
policy_name=policy_name,
|
||||
streaming_chunks=buffered,
|
||||
endpoint_translation=endpoint_translation,
|
||||
)
|
||||
try:
|
||||
ProxyLogging._handle_pipeline_result(
|
||||
result, data=request_data, policy_name=policy_name, original_response=buffered
|
||||
)
|
||||
except ModifyResponseException as e:
|
||||
if e.original_response is None:
|
||||
e.original_response = buffered
|
||||
async for block_chunk in unified_guardrail.handle_streaming_block(
|
||||
e, endpoint_translation, stream_started=False, responses_so_far=()
|
||||
):
|
||||
yield block_chunk
|
||||
return
|
||||
except HTTPException as e:
|
||||
async for error_chunk in unified_guardrail.emit_streaming_http_error(
|
||||
e, call_type, buffered, request_data
|
||||
):
|
||||
yield error_chunk
|
||||
return
|
||||
|
||||
for buffered_item in buffered:
|
||||
yield buffered_item
|
||||
|
||||
@staticmethod
|
||||
def _fire_deferred_stream_logging(request_data: dict) -> None:
|
||||
"""
|
||||
|
|
@ -3528,7 +3811,7 @@ class _StaleReadEngine:
|
|||
class PrismaClient:
|
||||
spend_log_transactions: list = []
|
||||
_spend_log_transactions_lock = asyncio.Lock()
|
||||
spend_log_flush_requested: ClassVar[asyncio.Event] = asyncio.Event()
|
||||
spend_log_flush_requested: "asyncio.Event | None" = None
|
||||
spend_log_queue_bytes: ClassVar[int] = 0
|
||||
spend_logs_queue_monitor_task: "asyncio.Task[None] | None" = None
|
||||
tool_usage_transactions: list["ToolUsageTransaction"] = []
|
||||
|
|
@ -6254,23 +6537,27 @@ async def enqueue_spend_logs(
|
|||
)
|
||||
|
||||
|
||||
def request_spend_log_flush() -> None:
|
||||
"""Wake the queue monitor now rather than leaving the rows for its next poll.
|
||||
def request_spend_log_flush(prisma_client: PrismaClient) -> None:
|
||||
"""Wake this client's queue monitor now rather than leaving the rows for its next poll.
|
||||
|
||||
The Responses API hands the client an id it can chain from straight away, and that
|
||||
lookup reads the DB, so the row cannot sit in this worker's queue for a poll interval.
|
||||
Repeated requests coalesce into the monitor's next pass, so the batching holds.
|
||||
A request made before the monitor is running is dropped, and loses nothing: the
|
||||
monitor reads the queue on its first pass, before it ever waits on a request.
|
||||
"""
|
||||
PrismaClient.spend_log_flush_requested.set()
|
||||
flush_requested: Final = prisma_client.spend_log_flush_requested
|
||||
if flush_requested is not None:
|
||||
flush_requested.set()
|
||||
|
||||
|
||||
async def _wait_for_spend_log_flush_request(interval: float) -> bool:
|
||||
async def _wait_for_spend_log_flush_request(flush_requested: asyncio.Event, interval: float) -> bool:
|
||||
"""Wait out ``interval``, returning early and True when a flush was requested."""
|
||||
try:
|
||||
await asyncio.wait_for(PrismaClient.spend_log_flush_requested.wait(), timeout=interval)
|
||||
await asyncio.wait_for(flush_requested.wait(), timeout=interval)
|
||||
except asyncio.TimeoutError:
|
||||
return False
|
||||
PrismaClient.spend_log_flush_requested.clear()
|
||||
flush_requested.clear()
|
||||
return True
|
||||
|
||||
|
||||
|
|
@ -6697,6 +6984,8 @@ async def _monitor_spend_logs_queue(
|
|||
max_backoff: Final = 30.0 # Maximum backoff interval in seconds
|
||||
backoff_multiplier: Final = 1.5 # Exponential backoff multiplier
|
||||
current_interval = base_interval
|
||||
flush_requested: Final = asyncio.Event()
|
||||
prisma_client.spend_log_flush_requested = flush_requested # rebind-ok: the client owns its monitor's flush signal
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"Starting spend logs queue monitor (threshold: %s, poll_interval: %ss)", threshold, base_interval
|
||||
|
|
@ -6735,7 +7024,7 @@ async def _monitor_spend_logs_queue(
|
|||
# Exponential backoff when no logs to process
|
||||
current_interval = min(current_interval * backoff_multiplier, max_backoff)
|
||||
|
||||
if await _wait_for_spend_log_flush_request(current_interval):
|
||||
if await _wait_for_spend_log_flush_request(flush_requested, current_interval):
|
||||
current_interval = base_interval
|
||||
except Exception as e:
|
||||
spend_log_error("Error in spend logs queue monitor: %s", str(e), exc=e)
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from litellm.completion_extras.litellm_responses_transformation.transformation i
|
|||
from litellm.constants import request_timeout
|
||||
from litellm.integrations.anthropic_cache_control_hook import CARRY_UNMATCHED_MESSAGE_POINTS
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
update_responses_input_with_model_file_ids,
|
||||
|
|
@ -1257,7 +1258,7 @@ def responses(
|
|||
responses_api_provider_config=responses_api_provider_config,
|
||||
response_api_optional_params=response_api_optional_params,
|
||||
allowed_openai_params=allowed_openai_params,
|
||||
drop_params=request_drop_params if isinstance(request_drop_params, bool) else None,
|
||||
drop_params=normalize_drop_params(request_drop_params),
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
|
|
@ -2085,7 +2086,7 @@ def compact_responses(
|
|||
responses_api_provider_config=responses_api_provider_config,
|
||||
response_api_optional_params=response_api_optional_params,
|
||||
allowed_openai_params=None,
|
||||
drop_params=request_drop_params if isinstance(request_drop_params, bool) else None,
|
||||
drop_params=normalize_drop_params(request_drop_params),
|
||||
)
|
||||
|
||||
# Pre Call logging
|
||||
|
|
|
|||
|
|
@ -21,7 +21,16 @@ import time
|
|||
import traceback
|
||||
import weakref
|
||||
from collections import defaultdict
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Iterator, Mapping, Sequence
|
||||
from collections.abc import (
|
||||
AsyncGenerator,
|
||||
AsyncIterator,
|
||||
Callable,
|
||||
Generator,
|
||||
Iterator,
|
||||
Mapping,
|
||||
MutableMapping,
|
||||
Sequence,
|
||||
)
|
||||
from functools import lru_cache, partial
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
|
||||
|
|
@ -45,12 +54,15 @@ from litellm.caching.caching import (
|
|||
RedisClusterCache,
|
||||
)
|
||||
from litellm.constants import (
|
||||
CLIENT_OUTPUT_CEILING_METADATA_KEY,
|
||||
CONSUMED_REQUEST_TAGS_METADATA_KEY,
|
||||
DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS,
|
||||
DEFAULT_HEALTH_CHECK_INTERVAL,
|
||||
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
|
||||
DEFAULT_MAX_LRU_CACHE_SIZE,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
OUTPUT_TOKEN_CEILING_PARAMS,
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY,
|
||||
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
)
|
||||
|
|
@ -648,6 +660,18 @@ class FallbackAwareStreamWrapper(CustomStreamWrapper):
|
|||
self.fallback_headers_adopted = True
|
||||
|
||||
|
||||
def as_output_cap(value: object) -> int | None:
|
||||
"""A client-sent output cap coerced to an int: ints, floats and numeric strings, never bools
|
||||
or negatives."""
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float, str)):
|
||||
return None
|
||||
try:
|
||||
cap: Final = int(float(value))
|
||||
except (ValueError, OverflowError):
|
||||
return None
|
||||
return cap if cap >= 0 else None
|
||||
|
||||
|
||||
class Router:
|
||||
model_names: set = set()
|
||||
cache_responses: bool | None = False
|
||||
|
|
@ -955,7 +979,6 @@ class Router:
|
|||
DEFAULT_HEALTH_CHECK_INTERVAL * DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER
|
||||
)
|
||||
self.health_state_cache = DeploymentHealthCache(cache=self.cache, staleness_threshold=float(_staleness))
|
||||
self.failed_calls = InMemoryCache() # cache to track failed call per deployment, if num failed calls within 1 minute > allowed fails, then add it to cooldown
|
||||
|
||||
if num_retries is not None:
|
||||
self.num_retries = num_retries
|
||||
|
|
@ -1248,7 +1271,7 @@ class Router:
|
|||
selector = LeastBusyLoggingHandler(router_cache=self.cache)
|
||||
if register_callbacks:
|
||||
if isinstance(litellm.input_callback, list):
|
||||
litellm.input_callback.append(selector)
|
||||
litellm.logging_callback_manager.add_litellm_input_callback(selector)
|
||||
else:
|
||||
litellm.input_callback = [selector]
|
||||
case RoutingStrategy.USAGE_BASED_ROUTING.value:
|
||||
|
|
@ -1523,9 +1546,33 @@ class Router:
|
|||
self._override_selectors[strategy] = self._build_strategy_selector(
|
||||
strategy=strategy,
|
||||
routing_strategy_args={},
|
||||
register_callbacks=False,
|
||||
)
|
||||
return self._override_selectors[strategy]
|
||||
|
||||
def _override_selector_pre_call_check(
|
||||
self, strategy: str | None, selector: RouterStrategySelector | None, deployment: dict
|
||||
) -> None:
|
||||
"""
|
||||
Override selectors are not in `litellm.callbacks`, so the pre-call check that
|
||||
`routing_strategy_pre_call_checks` runs for the router's own selectors (rpm
|
||||
accounting for `usage-based-routing-v2`) runs here, for the overriding request only.
|
||||
"""
|
||||
if selector is None or strategy is None or selector is not self._override_selectors.get(strategy):
|
||||
return
|
||||
selector.pre_call_check(deployment)
|
||||
|
||||
async def _async_override_selector_pre_call_check(
|
||||
self,
|
||||
strategy: str | None,
|
||||
selector: RouterStrategySelector | None,
|
||||
deployment: dict,
|
||||
parent_otel_span: Span | None,
|
||||
) -> None:
|
||||
if selector is None or strategy is None or selector is not self._override_selectors.get(strategy):
|
||||
return
|
||||
await selector.async_pre_call_check(deployment, parent_otel_span)
|
||||
|
||||
def _get_routing_context(
|
||||
self, model: str, request_kwargs: dict | None = None
|
||||
) -> tuple[str | None, RouterStrategySelector | None]:
|
||||
|
|
@ -3609,6 +3656,13 @@ class Router:
|
|||
effective_model_info: Final = kwargs.get("model_info") or deployment.get("model_info") or MappingProxyType({})
|
||||
self._set_failed_deployment_id_on_exception(exception, MappingProxyType({"model_info": effective_model_info}))
|
||||
|
||||
@staticmethod
|
||||
def _stamp_retry_skip_deployment_id(exception: Exception, kwargs: Mapping[str, object]) -> None:
|
||||
effective_model_info: Final = kwargs.get("model_info")
|
||||
deployment_id: Final = effective_model_info.get("id") if isinstance(effective_model_info, Mapping) else None
|
||||
if isinstance(deployment_id, str) and deployment_id:
|
||||
exception.retry_skip_deployment_id = deployment_id # pyright: ignore[reportAttributeAccessIssue] # dynamic stamp, read by _deployment_ids_to_skip_on_retry
|
||||
|
||||
def _update_kwargs_with_default_litellm_params(
|
||||
self, kwargs: dict, metadata_variable_name: str | None = "metadata"
|
||||
) -> None:
|
||||
|
|
@ -3726,6 +3780,11 @@ class Router:
|
|||
refund_stale_reservation_before_retry(self.cache, kwargs)
|
||||
set_io_token_rate_limit_request_kwargs(kwargs, store_in_context=deployment_has_io_token_limits(deployment))
|
||||
|
||||
kwargs[metadata_variable_name].setdefault(
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY,
|
||||
tuple(_get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name)),
|
||||
)
|
||||
|
||||
## DEPLOYMENT-LEVEL TAGS
|
||||
deployment_tags: Final = deployment.get("litellm_params", {}).get("tags")
|
||||
if deployment_tags:
|
||||
|
|
@ -4214,10 +4273,12 @@ class Router:
|
|||
}
|
||||
)
|
||||
litellm_logging_object = cast(LiteLLMLogging, litellm_logging_object)
|
||||
prompt_management_deployment: Final = self.get_available_deployment(
|
||||
specific_deployment: Final = kwargs.pop("specific_deployment", None)
|
||||
prompt_management_deployment: Final = await self.async_get_available_deployment(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "prompt"}],
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
messages=cast(list[dict[str, str]], messages), # cast-ok: selection reads messages structurally
|
||||
specific_deployment=specific_deployment,
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
|
||||
self._update_kwargs_with_deployment(deployment=prompt_management_deployment, kwargs=kwargs)
|
||||
|
|
@ -4304,6 +4365,7 @@ class Router:
|
|||
model=model,
|
||||
messages=[{"role": "user", "content": "prompt"}],
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
|
||||
data: Final = deployment["litellm_params"].copy()
|
||||
|
|
@ -4334,6 +4396,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.image_generation(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aimage_generation(self, prompt: str, model: str, **kwargs):
|
||||
|
|
@ -4418,6 +4481,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.aimage_generation(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def atranscription(self, file: FileTypes, model: str, **kwargs):
|
||||
|
|
@ -4522,6 +4586,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.atranscription(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aspeech(self, model: str, input: str, voice: str | None = None, **kwargs):
|
||||
|
|
@ -4636,6 +4701,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.aspeech(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def arerank(self, model: str, **kwargs):
|
||||
|
|
@ -4694,6 +4760,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.arerank(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
def text_completion(
|
||||
|
|
@ -4828,6 +4895,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.atext_completion(model=%s)\x1b[31m Exception %s\x1b[0m", model, e)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aadapter_completion(
|
||||
|
|
@ -4918,6 +4986,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.aadapter_completion(model=%s)\x1b[31m Exception %s\x1b[0m", model, e)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def _asearch_with_fallbacks(self, original_function: Callable, **kwargs):
|
||||
|
|
@ -5700,6 +5769,7 @@ class Router:
|
|||
model=model,
|
||||
input=input,
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
|
||||
data: Final = deployment["litellm_params"].copy()
|
||||
|
|
@ -5738,6 +5808,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.embedding(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aembedding(
|
||||
|
|
@ -5825,6 +5896,7 @@ class Router:
|
|||
verbose_router_logger.info("litellm.aembedding(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
#### FILES API ####
|
||||
|
|
@ -6198,6 +6270,7 @@ class Router:
|
|||
)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def aretrieve_batch(
|
||||
|
|
@ -6418,6 +6491,7 @@ class Router:
|
|||
)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
self._stamp_retry_skip_deployment_id(e, kwargs)
|
||||
raise e
|
||||
|
||||
async def alist_batches(
|
||||
|
|
@ -7535,7 +7609,9 @@ class Router:
|
|||
|
||||
@staticmethod
|
||||
def _deployment_ids_to_skip_on_retry(exception: Exception, already_skipped: object) -> tuple[str, ...]:
|
||||
failed_deployment_id: Final[str | None] = getattr(exception, "failed_deployment_id", None)
|
||||
failed_deployment_id: Final[str | None] = getattr(exception, "retry_skip_deployment_id", None) or getattr(
|
||||
exception, "failed_deployment_id", None
|
||||
)
|
||||
status_code: Final = getattr(exception, "status_code", None)
|
||||
if not failed_deployment_id or not isinstance(status_code, int):
|
||||
return ()
|
||||
|
|
@ -9366,6 +9442,12 @@ class Router:
|
|||
#### VALIDATE MODEL ########
|
||||
# Check if this is a prompt management model before validating as LLM provider
|
||||
litellm_model: Final = deployment.litellm_params.model
|
||||
if isinstance(deployment.litellm_params.drop_params, str):
|
||||
verbose_router_logger.warning(
|
||||
"model=%s drop_params=%r is not a flag value, treating it as unset",
|
||||
deployment.model_name,
|
||||
deployment.litellm_params.drop_params,
|
||||
)
|
||||
is_prompt_management_model = False
|
||||
|
||||
if "/" in litellm_model:
|
||||
|
|
@ -12604,7 +12686,83 @@ class Router:
|
|||
request_kwargs.pop(carrier, None)
|
||||
|
||||
@staticmethod
|
||||
def _drop_client_effort_carriers_a_tier_pin_supersedes(
|
||||
def _tier_ceiling_under_the_surface_name(
|
||||
tier_litellm_params: Mapping[str, object], responses_call: bool
|
||||
) -> Mapping[str, object]:
|
||||
"""``max_tokens``, ``max_completion_tokens`` and ``max_output_tokens`` are one
|
||||
ceiling under three names, and each surface reads exactly one of them: the
|
||||
Responses bridge builds its internal ``max_tokens`` from ``max_output_tokens``
|
||||
and would overwrite the tier's, chat and /v1/messages never read
|
||||
``max_output_tokens``, and litellm already renames ``max_tokens`` to
|
||||
``max_completion_tokens`` for the OpenAI models that require it. Collapse
|
||||
whatever the tier carries onto the surface's own name, preferring a value the
|
||||
operator already wrote under that name."""
|
||||
surface_key: Final = "max_output_tokens" if responses_call else "max_tokens"
|
||||
carried: Final = tuple(
|
||||
key
|
||||
for key in (surface_key, "max_tokens", "max_completion_tokens", "max_output_tokens")
|
||||
if key in tier_litellm_params
|
||||
)
|
||||
if not carried:
|
||||
return tier_litellm_params
|
||||
return MappingProxyType(
|
||||
{
|
||||
**{k: v for k, v in tier_litellm_params.items() if k not in OUTPUT_TOKEN_CEILING_PARAMS},
|
||||
surface_key: tier_litellm_params[carried[0]],
|
||||
}
|
||||
)
|
||||
|
||||
def _pin_tier_params_onto_request(
|
||||
self,
|
||||
model: str,
|
||||
tier_litellm_params: Mapping[str, object] | None,
|
||||
request_kwargs: dict,
|
||||
responses_call: bool,
|
||||
) -> bool:
|
||||
"""Apply a routing strategy's per-tier litellm_params on top of the request and report
|
||||
whether they pinned an output ceiling, so the caller can hand the request its own ceiling
|
||||
back on a routing pass that pins none."""
|
||||
if not tier_litellm_params:
|
||||
return False
|
||||
accepted_tier_params: Final = self._tier_params_the_target_accepts(model, tier_litellm_params, request_kwargs)
|
||||
surface_tier_params: Final = self._tier_ceiling_under_the_surface_name(
|
||||
accepted_tier_params, responses_call=responses_call
|
||||
)
|
||||
self._drop_client_carriers_a_tier_pin_supersedes(request_kwargs, surface_tier_params)
|
||||
request_kwargs.update(surface_tier_params)
|
||||
return not OUTPUT_TOKEN_CEILING_PARAMS.isdisjoint(surface_tier_params)
|
||||
|
||||
@staticmethod
|
||||
def _restore_client_ceiling_no_tier_pins(request_kwargs: MutableMapping[str, object]) -> None:
|
||||
"""A model-group fallback re-enters routing with the kwargs an earlier auto-router pass
|
||||
already rewrote, so a ceiling sized for that pass's tier would ride onto a group no tier
|
||||
chose. When this pass pins none, hand the request back exactly the carriers the caller
|
||||
sent, which the first pinning pass stamped. The stamp lives in a metadata bucket a
|
||||
caller can also write, so the proxy strips the key at ingestion and this read takes
|
||||
nothing but the three ceiling carriers as integers: no other key ever reaches kwargs."""
|
||||
stamped: Final = next(
|
||||
(
|
||||
bucket.get(CLIENT_OUTPUT_CEILING_METADATA_KEY)
|
||||
for bucket in (request_kwargs.get("metadata"), request_kwargs.get("litellm_metadata"))
|
||||
if isinstance(bucket, dict) and CLIENT_OUTPUT_CEILING_METADATA_KEY in bucket
|
||||
),
|
||||
None,
|
||||
)
|
||||
if not isinstance(stamped, dict):
|
||||
return
|
||||
callers_ceiling: Final = MappingProxyType(
|
||||
{
|
||||
carrier: cap
|
||||
for carrier, value in stamped.items()
|
||||
if carrier in OUTPUT_TOKEN_CEILING_PARAMS and (cap := as_output_cap(value)) is not None
|
||||
}
|
||||
)
|
||||
for carrier in OUTPUT_TOKEN_CEILING_PARAMS:
|
||||
request_kwargs.pop(carrier, None)
|
||||
request_kwargs.update(callers_ceiling)
|
||||
|
||||
@staticmethod
|
||||
def _drop_client_carriers_a_tier_pin_supersedes(
|
||||
request_kwargs: dict[str, object],
|
||||
tier_litellm_params: Mapping[str, object],
|
||||
) -> None:
|
||||
|
|
@ -12614,7 +12772,22 @@ class Router:
|
|||
the ``reasoning_effort`` alias, so a pinned effort only reaches the wire
|
||||
if the client's other encodings are removed before the merge. Non-effort
|
||||
fields a carrier also holds (``output_config.format``,
|
||||
``reasoning.summary``) are kept."""
|
||||
``reasoning.summary``) are kept. An output ceiling has the same shape:
|
||||
``max_tokens``, ``max_completion_tokens`` and ``max_output_tokens`` are
|
||||
one setting under three names, and a provider handed two of them either
|
||||
rejects the request or picks one by iteration order."""
|
||||
if not OUTPUT_TOKEN_CEILING_PARAMS.isdisjoint(tier_litellm_params):
|
||||
_, metadata_bucket = get_or_create_metadata_bucket(request_kwargs)
|
||||
metadata_bucket.setdefault(
|
||||
CLIENT_OUTPUT_CEILING_METADATA_KEY,
|
||||
{
|
||||
carrier: request_kwargs[carrier]
|
||||
for carrier in OUTPUT_TOKEN_CEILING_PARAMS
|
||||
if carrier in request_kwargs
|
||||
},
|
||||
)
|
||||
for carrier in OUTPUT_TOKEN_CEILING_PARAMS:
|
||||
request_kwargs.pop(carrier, None)
|
||||
if "reasoning_effort" not in tier_litellm_params:
|
||||
return
|
||||
request_kwargs.pop("thinking", None)
|
||||
|
|
@ -12655,6 +12828,7 @@ class Router:
|
|||
# Execute Pre-Routing Hooks
|
||||
# this hook can modify the model, messages before the routing decision is made
|
||||
#########################################################
|
||||
responses_call: Final = input is not None and messages is None
|
||||
pre_routing_hook_response: Final = await self.async_pre_routing_hook(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
|
|
@ -12666,12 +12840,14 @@ class Router:
|
|||
model = pre_routing_hook_response.model
|
||||
messages = pre_routing_hook_response.messages
|
||||
record_pre_routing_selection(request_kwargs, model)
|
||||
if pre_routing_hook_response.litellm_params:
|
||||
accepted_tier_params: Final = self._tier_params_the_target_accepts(
|
||||
model, pre_routing_hook_response.litellm_params, request_kwargs
|
||||
)
|
||||
self._drop_client_effort_carriers_a_tier_pin_supersedes(request_kwargs, accepted_tier_params)
|
||||
request_kwargs.update(accepted_tier_params)
|
||||
tier_pins_ceiling: Final = self._pin_tier_params_onto_request(
|
||||
model=model,
|
||||
tier_litellm_params=pre_routing_hook_response.litellm_params if pre_routing_hook_response else None,
|
||||
request_kwargs=request_kwargs,
|
||||
responses_call=responses_call,
|
||||
)
|
||||
if not tier_pins_ceiling:
|
||||
self._restore_client_ceiling_no_tier_pins(request_kwargs)
|
||||
#########################################################
|
||||
|
||||
# Resolve the strategy and logger AFTER the pre-routing hook, since
|
||||
|
|
@ -12688,10 +12864,16 @@ class Router:
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
if isinstance(healthy_deployments, dict):
|
||||
await self._async_override_selector_pre_call_check(
|
||||
strategy, strategy_selector, healthy_deployments, parent_otel_span
|
||||
)
|
||||
return healthy_deployments
|
||||
|
||||
# When encrypted content affinity pins to a specific deployment,
|
||||
if request_kwargs.get("_encrypted_content_affinity_pinned") and len(healthy_deployments) == 1:
|
||||
await self._async_override_selector_pre_call_check(
|
||||
strategy, strategy_selector, healthy_deployments[0], parent_otel_span
|
||||
)
|
||||
return healthy_deployments[0]
|
||||
|
||||
start_time: Final = time.time()
|
||||
|
|
@ -12717,6 +12899,9 @@ class Router:
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
raise exception
|
||||
await self._async_override_selector_pre_call_check(
|
||||
strategy, strategy_selector, deployment, parent_otel_span
|
||||
)
|
||||
verbose_router_logger.info(
|
||||
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
|
||||
model,
|
||||
|
|
@ -12771,6 +12956,7 @@ class Router:
|
|||
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs)
|
||||
|
||||
# 1. Execute pre-routing hook
|
||||
responses_call: Final = input is not None and messages is None
|
||||
pre_routing_hook_response: Final = await self.async_pre_routing_hook(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
|
|
@ -12782,12 +12968,14 @@ class Router:
|
|||
model = pre_routing_hook_response.model
|
||||
messages = pre_routing_hook_response.messages
|
||||
record_pre_routing_selection(request_kwargs, model)
|
||||
if pre_routing_hook_response.litellm_params:
|
||||
accepted_tier_params: Final = self._tier_params_the_target_accepts(
|
||||
model, pre_routing_hook_response.litellm_params, request_kwargs
|
||||
)
|
||||
self._drop_client_effort_carriers_a_tier_pin_supersedes(request_kwargs, accepted_tier_params)
|
||||
request_kwargs.update(accepted_tier_params)
|
||||
tier_pins_ceiling: Final = self._pin_tier_params_onto_request(
|
||||
model=model,
|
||||
tier_litellm_params=pre_routing_hook_response.litellm_params if pre_routing_hook_response else None,
|
||||
request_kwargs=request_kwargs,
|
||||
responses_call=responses_call,
|
||||
)
|
||||
if not tier_pins_ceiling:
|
||||
self._restore_client_ceiling_no_tier_pins(request_kwargs)
|
||||
|
||||
# 2. Get healthy deployments
|
||||
healthy_deployments: Final = await self.async_get_healthy_deployments(
|
||||
|
|
@ -12799,6 +12987,8 @@ class Router:
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
|
||||
|
||||
# 3. If specific deployment returned, verify if it supports pass-through
|
||||
if isinstance(healthy_deployments, dict):
|
||||
if (healthy_deployments.get("model_info") or {}).get("blocked") is True:
|
||||
|
|
@ -12809,6 +12999,9 @@ class Router:
|
|||
)
|
||||
litellm_params: Final = healthy_deployments.get("litellm_params", {})
|
||||
if litellm_params.get("use_in_pass_through"):
|
||||
await self._async_override_selector_pre_call_check(
|
||||
strategy, strategy_selector, healthy_deployments, parent_otel_span
|
||||
)
|
||||
return healthy_deployments
|
||||
else:
|
||||
raise litellm.BadRequestError(
|
||||
|
|
@ -12829,7 +13022,6 @@ class Router:
|
|||
|
||||
# 5. Apply load balancing strategy
|
||||
start_time: Final = time.perf_counter()
|
||||
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
|
||||
if strategy == "simple-shuffle":
|
||||
return simple_shuffle(
|
||||
llm_router_instance=self,
|
||||
|
|
@ -12853,6 +13045,9 @@ class Router:
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
raise exception
|
||||
await self._async_override_selector_pre_call_check(
|
||||
strategy, strategy_selector, deployment, parent_otel_span
|
||||
)
|
||||
|
||||
verbose_router_logger.info(
|
||||
"async_get_available_deployment_for_pass_through model: %s, selected deployment: %s",
|
||||
|
|
@ -13428,6 +13623,7 @@ class Router:
|
|||
specific_deployment=specific_deployment,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
|
||||
|
||||
if isinstance(healthy_deployments, dict):
|
||||
if (healthy_deployments.get("model_info") or {}).get("blocked") is True:
|
||||
|
|
@ -13436,6 +13632,7 @@ class Router:
|
|||
model=model,
|
||||
llm_provider="",
|
||||
)
|
||||
self._override_selector_pre_call_check(strategy, strategy_selector, healthy_deployments)
|
||||
return healthy_deployments
|
||||
|
||||
parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs)
|
||||
|
|
@ -13511,7 +13708,6 @@ class Router:
|
|||
cooldown_list=_cooldown_list,
|
||||
)
|
||||
|
||||
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
|
||||
if strategy == "simple-shuffle":
|
||||
# if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm
|
||||
############## Check 'weight' param set for weighted pick #################
|
||||
|
|
@ -13543,6 +13739,7 @@ class Router:
|
|||
enable_pre_call_checks=self.enable_pre_call_checks,
|
||||
cooldown_list=_cooldown_list,
|
||||
)
|
||||
self._override_selector_pre_call_check(strategy, strategy_selector, deployment)
|
||||
verbose_router_logger.info(
|
||||
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
|
||||
model,
|
||||
|
|
@ -13586,6 +13783,8 @@ class Router:
|
|||
specific_deployment=specific_deployment,
|
||||
)
|
||||
|
||||
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
|
||||
|
||||
# 2. If the returned is a specific deployment (Dict), verify and return directly
|
||||
if isinstance(healthy_deployments, dict):
|
||||
if (healthy_deployments.get("model_info") or {}).get("blocked") is True:
|
||||
|
|
@ -13596,6 +13795,7 @@ class Router:
|
|||
)
|
||||
litellm_params: Final = healthy_deployments.get("litellm_params", {})
|
||||
if litellm_params.get("use_in_pass_through"):
|
||||
self._override_selector_pre_call_check(strategy, strategy_selector, healthy_deployments)
|
||||
return healthy_deployments
|
||||
else:
|
||||
# Specific deployment does not support pass-through
|
||||
|
|
@ -13655,7 +13855,6 @@ class Router:
|
|||
)
|
||||
|
||||
# 6. Apply load balancing strategy
|
||||
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
|
||||
if strategy == "simple-shuffle":
|
||||
return simple_shuffle(
|
||||
llm_router_instance=self,
|
||||
|
|
@ -13687,6 +13886,7 @@ class Router:
|
|||
enable_pre_call_checks=self.enable_pre_call_checks,
|
||||
cooldown_list=_cooldown_list,
|
||||
)
|
||||
self._override_selector_pre_call_check(strategy, strategy_selector, deployment)
|
||||
|
||||
verbose_router_logger.info(
|
||||
"get_available_deployment_for_pass_through model: %s, selected deployment: %s",
|
||||
|
|
|
|||
|
|
@ -207,17 +207,27 @@ custom_dimensions:
|
|||
- name: sqlMigration
|
||||
weight: 0.7
|
||||
patterns: ['\b(create|alter|drop)\s{1,4}table\b']
|
||||
- name: dataPipeline
|
||||
weight: 0.4
|
||||
scoring_mode: match_count
|
||||
keywords: [airflow, dbt, snowflake]
|
||||
```
|
||||
|
||||
Each dimension contributes its weight once when any matcher hits the current ask. Repeated matches do not increase it. The built-in score and tier boundaries are unchanged, and the total score is not renormalized. Keywords use the existing case-insensitive word-boundary and CJK rules. Regexes search the first 2048 characters case-insensitively and compile during configuration validation and router initialization, never per request
|
||||
|
||||
`scoring_mode` is optional and defaults to `binary`, the behavior above. `match_count` grades the dimension by how many distinct matchers hit: none scores 0 and emits no signal, one scores half the weight, two or more score the full weight. Repeated occurrences of one matcher never raise the count, keywords are distinct case-insensitively, patterns are distinct by source, and a keyword and a pattern are always distinct from each other. Matching stops as soon as the selected mode's maximum is reached, so a binary dimension still stops at its first hit. Existing configurations without the field keep binary scoring and the same tuning fingerprint, so the field only counts as a tuning change when set to `match_count`
|
||||
|
||||
### Weights through the API versus the dashboard
|
||||
|
||||
The API and YAML store exactly the weights written. A `dimension_weights` map and inline custom weights are read literally, missing recognized built-in names score zero, and nothing renormalizes the vector, so a total other than 1 is legal and scores accordingly. The dashboard's heuristic scoring editor is the one place that rebalances: editing one weight there holds it and redistributes the remainder across the other active dimensions in the draft, then Save sends the resulting explicit values, which the backend stores and scores as written. Opening a router, applying a preset, editing matchers, changing `scoring_mode`, or saving unrelated fields never normalizes existing weights
|
||||
|
||||
Only `heuristic`, `heuristic_first` and `hybrid` accept custom dimensions. Each name must be a unique ASCII identifier starting with a letter, at most 64 characters, and cannot reuse a built-in dimension name or a key in `dimension_weights`. Set its weight inline, greater than zero and at most one
|
||||
|
||||
Patterns are checked at configuration time against a grammar whose worst case stays a few milliseconds on 2048 characters. Every quantifier needs an explicit upper bound of at most 64 and must repeat a single character or character class, so `\s{1,4}` is accepted while `\s+`, `(a|aa){0,12}` and `(?:ab){0,64}` are refused. Backreferences, lookarounds, atomic groups and possessive quantifiers are refused as well. Each pattern is then costed: alternation branches and repeat lengths multiply the ways the engine can retry, and every later piece of the pattern is charged once per path that can reach it, so `a?a?a?a?a?a?a?a?` followed by a long fixed tail is refused even though each quantifier is small. The budget is 2048 work units per pattern and 8192 across the router. An invalid or over-budget pattern fails the write with a message naming the pattern and the rule it broke
|
||||
|
||||
Limits are 16 dimensions, 32 combined keywords/patterns per dimension, 256 characters per matcher and 4096 matcher characters per dimension. Matching runs inline on the request path with no timeout and no worker thread, because the grammar is what bounds the cost. These are routing hints, not security enforcement rules
|
||||
|
||||
The existing heuristic-v1 tuning quota covers custom dimensions and their weights: one changed router without an auto-router license, unlimited with the entitlement. Omitting `custom_dimensions` preserves existing scoring. Routing decisions and spend logs include signals such as `custom (sqlMigration)` without recording the configured pattern or matched text. The field is configured through YAML or the model API; this change adds no dashboard editor
|
||||
The existing heuristic-v1 tuning quota covers custom dimensions, their weights and their scoring mode: one changed router without an auto-router license, unlimited with the entitlement. Omitting `custom_dimensions` preserves existing scoring. Routing decisions and spend logs include signals such as `custom (sqlMigration)` without recording the configured pattern or matched text. The field is configured through YAML, the model API, or the dashboard's heuristic scoring editor
|
||||
|
||||
## Usage
|
||||
|
||||
|
|
|
|||
|
|
@ -20,7 +20,7 @@ import random
|
|||
import re
|
||||
import time
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from itertools import accumulate, islice, takewhile
|
||||
from itertools import accumulate, chain, islice, takewhile
|
||||
from threading import Lock
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
|
||||
|
|
@ -31,6 +31,7 @@ from litellm._logging import verbose_router_logger
|
|||
from litellm.constants import (
|
||||
EMPTY_MAPPING,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
OUTPUT_TOKEN_CEILING_PARAMS,
|
||||
RETURN_RAW_MODEL_NAME_METADATA_KEY,
|
||||
SESSION_ID_GENERATED_METADATA_KEY,
|
||||
)
|
||||
|
|
@ -82,6 +83,7 @@ from .config import (
|
|||
ClassificationRubric,
|
||||
ComplexityRouterConfig,
|
||||
ComplexityTier,
|
||||
CustomDimension,
|
||||
TierDefinition,
|
||||
)
|
||||
from .stall_detector import detect_stalled_task
|
||||
|
|
@ -879,6 +881,15 @@ class DimensionScore:
|
|||
self.signal = signal
|
||||
|
||||
|
||||
class _CustomDimensionMatchers(NamedTuple):
|
||||
"""One custom dimension's distinct matchers and the number of hits that saturates its score."""
|
||||
|
||||
dimension: CustomDimension
|
||||
keywords: tuple[str, ...]
|
||||
patterns: tuple[re.Pattern[str], ...]
|
||||
saturation: int
|
||||
|
||||
|
||||
class KeywordOverride(NamedTuple):
|
||||
"""A keyword_tier_rules match: the winning tier and, on the lexical path, the keyword that fired."""
|
||||
|
||||
|
|
@ -1121,7 +1132,12 @@ class ComplexityRouter(CustomLogger):
|
|||
)
|
||||
self.simple_keywords = self.config.simple_keywords or DEFAULT_SIMPLE_KEYWORDS
|
||||
self._custom_dimensions = tuple(
|
||||
(dimension, tuple(re.compile(pattern, re.IGNORECASE) for pattern in dimension.patterns))
|
||||
_CustomDimensionMatchers(
|
||||
dimension,
|
||||
tuple(dict.fromkeys(keyword.lower() for keyword in dimension.keywords)),
|
||||
tuple(re.compile(pattern, re.IGNORECASE) for pattern in dict.fromkeys(dimension.patterns)),
|
||||
2 if dimension.scoring_mode == "match_count" else 1,
|
||||
)
|
||||
for dimension in self.config.custom_dimensions
|
||||
)
|
||||
if self.config.has_custom_tiers:
|
||||
|
|
@ -1325,15 +1341,26 @@ class ComplexityRouter(CustomLogger):
|
|||
score: Final = score_high if match_count >= high_threshold else score_low
|
||||
return DimensionScore(name, score, f"{signal_label} ({detail})"), match_count
|
||||
|
||||
def _count_custom_hits(self, matchers: _CustomDimensionMatchers, user_text: str, scanned: str) -> int:
|
||||
hits: Final = chain(
|
||||
(self._keyword_matches(user_text, keyword) for keyword in matchers.keywords),
|
||||
(pattern.search(scanned) is not None for pattern in matchers.patterns),
|
||||
)
|
||||
return sum(islice((1 for hit in hits if hit), matchers.saturation))
|
||||
|
||||
def _score_custom_dimensions(self, prompt: str, user_text: str) -> tuple[tuple[DimensionScore, float], ...]:
|
||||
if not self._custom_dimensions:
|
||||
return ()
|
||||
scanned: Final = prompt[:CUSTOM_PATTERN_SCAN_CHARS]
|
||||
return tuple(
|
||||
(DimensionScore(dimension.name, 1.0, f"custom ({dimension.name})"), dimension.weight)
|
||||
for dimension, patterns in self._custom_dimensions
|
||||
if any(self._keyword_matches(user_text, keyword) for keyword in dimension.keywords)
|
||||
or any(pattern.search(scanned) is not None for pattern in patterns)
|
||||
(
|
||||
DimensionScore(
|
||||
matchers.dimension.name, hits / matchers.saturation, f"custom ({matchers.dimension.name})"
|
||||
),
|
||||
matchers.dimension.weight,
|
||||
)
|
||||
for matchers in self._custom_dimensions
|
||||
if (hits := self._count_custom_hits(matchers, user_text, scanned))
|
||||
)
|
||||
|
||||
def _score_multi_step(self, text: str) -> DimensionScore:
|
||||
|
|
@ -2084,11 +2111,15 @@ class ComplexityRouter(CustomLogger):
|
|||
raise ValueError(f"No model configured for tier {tier_key} and no default_model set")
|
||||
|
||||
def _litellm_params_for_model(self, tier: ComplexityTier | str | None, model: str) -> Mapping[str, object]:
|
||||
if tier is None:
|
||||
return MappingProxyType({})
|
||||
entries: Final = self.config.tier_model_configs.get(_tier_name(tier), ())
|
||||
entries: Final = self.config.tier_model_configs.get(_tier_name(tier), ()) if tier is not None else ()
|
||||
entry: Final = next((candidate for candidate in entries if candidate.model_name == model), None)
|
||||
return entry.litellm_params if entry is not None else MappingProxyType({})
|
||||
explicit: Final = entry.litellm_params if entry is not None else MappingProxyType({})
|
||||
if not self.config.max_tokens_from_tier_model or not OUTPUT_TOKEN_CEILING_PARAMS.isdisjoint(explicit):
|
||||
return explicit
|
||||
ceiling: Final = self._group_output_ceiling(model)
|
||||
if ceiling is None:
|
||||
return explicit
|
||||
return MappingProxyType({**explicit, "max_tokens": ceiling})
|
||||
|
||||
@staticmethod
|
||||
def _pick_from_tier_value(model: str | Sequence[str], tier_key: str) -> str:
|
||||
|
|
@ -2423,12 +2454,15 @@ class ComplexityRouter(CustomLogger):
|
|||
return name if self.config.has_custom_tiers else ComplexityTier(name)
|
||||
|
||||
def _deployment_window(self, group: str, deployment: Mapping[str, object]) -> int | None:
|
||||
return self._deployment_limit(group, deployment, "max_input_tokens")
|
||||
|
||||
def _deployment_limit(
|
||||
self, group: str, deployment: Mapping[str, object], key: Literal["max_input_tokens", "max_output_tokens"]
|
||||
) -> int | None:
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
|
||||
|
||||
deployment_model_info: Final = deployment.get("model_info")
|
||||
declared: Final = (
|
||||
deployment_model_info.get("max_input_tokens") if isinstance(deployment_model_info, Mapping) else None
|
||||
)
|
||||
declared: Final = deployment_model_info.get(key) if isinstance(deployment_model_info, Mapping) else None
|
||||
if isinstance(declared, int):
|
||||
return declared
|
||||
litellm_params: Final = deployment.get("litellm_params")
|
||||
|
|
@ -2445,18 +2479,34 @@ class ComplexityRouter(CustomLogger):
|
|||
deployment=cast(dict, deployment), # cast-ok: router deployments are plain dicts
|
||||
received_model_name=group,
|
||||
)
|
||||
window: Final = model_info.get("max_input_tokens")
|
||||
limit: Final = model_info.get(key)
|
||||
except Exception: # noqa: BLE001 # best-effort: an unmappable deployment must not hide the others
|
||||
return None
|
||||
return window if isinstance(window, int) else None
|
||||
return limit if isinstance(limit, int) else None
|
||||
|
||||
def _group_deployments(self, group: str) -> Sequence[Mapping[str, object]]:
|
||||
list_models: Final = getattr(self.litellm_router_instance, "get_model_list", None)
|
||||
deployments: Final = list_models(model_name=group) if callable(list_models) else None
|
||||
return tuple(deployments) if isinstance(deployments, list) else ()
|
||||
|
||||
def _group_output_ceiling(self, group: str) -> int | None:
|
||||
"""Smallest max_output_tokens across the group's deployments, or None when any deployment
|
||||
declares none: the core router picks within the group without a fit check, and a ceiling
|
||||
above an unmapped member's real limit is a provider 400 on that member."""
|
||||
deployments: Final = self._group_deployments(group)
|
||||
ceilings: Final = tuple(
|
||||
ceiling
|
||||
for deployment in deployments
|
||||
if (ceiling := self._deployment_limit(group, deployment, "max_output_tokens")) is not None
|
||||
)
|
||||
return min(ceilings) if ceilings and len(ceilings) == len(deployments) else None
|
||||
|
||||
def _group_window_facts(self, group: str) -> tuple[int | None, bool]:
|
||||
"""(smallest declared context window across the group's deployments, whether any deployment
|
||||
declares none). The core router picks a deployment within the group without a fit check, so
|
||||
the group is only as safe as its smallest member."""
|
||||
list_models: Final = getattr(self.litellm_router_instance, "get_model_list", None)
|
||||
deployments: Final = list_models(model_name=group) if callable(list_models) else None
|
||||
if not isinstance(deployments, list) or not deployments:
|
||||
deployments: Final = self._group_deployments(group)
|
||||
if not deployments:
|
||||
return (None, True)
|
||||
windows: Final = tuple(
|
||||
window for deployment in deployments if (window := self._deployment_window(group, deployment)) is not None
|
||||
|
|
@ -3505,6 +3555,7 @@ class ComplexityRouter(CustomLogger):
|
|||
ComplexityTier.MEDIUM, messages, resolved_messages, request_kwargs
|
||||
)
|
||||
fallback_tier: Final = None if default_model_first else ComplexityTier.MEDIUM
|
||||
default_tier_params: Final = self._litellm_params_for_model(fallback_tier, routed_model)
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
|
|
@ -3513,7 +3564,9 @@ class ComplexityRouter(CustomLogger):
|
|||
cause="default_fallback",
|
||||
tier=fallback_tier,
|
||||
conversation_continuing=conversation_continuing,
|
||||
tier_litellm_params=default_tier_params,
|
||||
),
|
||||
litellm_params=default_tier_params,
|
||||
)
|
||||
|
||||
ask: Final = user_message or ""
|
||||
|
|
@ -3540,6 +3593,7 @@ class ComplexityRouter(CustomLogger):
|
|||
_tier_name(plan_floor),
|
||||
routed_model,
|
||||
)
|
||||
plan_tier_params: Final = self._litellm_params_for_model(plan_floor, routed_model)
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
|
|
@ -3551,7 +3605,9 @@ class ComplexityRouter(CustomLogger):
|
|||
matched_keyword=plan_mode_sentinel,
|
||||
escalation_keyword=escalation_keyword,
|
||||
escalated=False,
|
||||
tier_litellm_params=plan_tier_params,
|
||||
),
|
||||
litellm_params=plan_tier_params,
|
||||
)
|
||||
|
||||
override: Final = await self._resolve_keyword_tier_override(ask, request_kwargs)
|
||||
|
|
@ -3644,6 +3700,7 @@ class ComplexityRouter(CustomLogger):
|
|||
outcome.signals,
|
||||
fallback_model,
|
||||
)
|
||||
fallback_tier_params: Final = self._litellm_params_for_model(None, fallback_model)
|
||||
return PreRoutingHookResponse(
|
||||
model=fallback_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
|
|
@ -3654,7 +3711,9 @@ class ComplexityRouter(CustomLogger):
|
|||
signals=outcome.signals,
|
||||
escalation_keyword=escalation_keyword,
|
||||
escalated=False,
|
||||
tier_litellm_params=fallback_tier_params,
|
||||
),
|
||||
litellm_params=fallback_tier_params,
|
||||
)
|
||||
if self.config.adaptive:
|
||||
# hard_floor rather than a hard pick, and passed whenever the sentinel is present
|
||||
|
|
|
|||
|
|
@ -667,6 +667,14 @@ class CustomDimension(BaseModel):
|
|||
weight: float = Field(gt=0, le=1, allow_inf_nan=False)
|
||||
keywords: tuple[Annotated[str, Field(min_length=1, max_length=256)], ...] = Field(default=(), max_length=32)
|
||||
patterns: tuple[Annotated[str, Field(min_length=1, max_length=256)], ...] = Field(default=(), max_length=32)
|
||||
scoring_mode: Literal["binary", "match_count"] = Field(
|
||||
default="binary",
|
||||
description=(
|
||||
"'binary' scores 1 when any matcher hits. 'match_count' scores 0.5 when one distinct matcher hits and 1 "
|
||||
"when two or more do; repeated occurrences of one matcher never raise it. Keywords are distinct "
|
||||
"case-insensitively, patterns by source, and a keyword and a pattern are always distinct from each other."
|
||||
),
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_matchers(self) -> "CustomDimension":
|
||||
|
|
@ -794,8 +802,9 @@ class ComplexityRouterConfig(BaseModel):
|
|||
default=(),
|
||||
max_length=16,
|
||||
description=(
|
||||
"Named binary dimensions added to the heuristic-v1 score. Each contributes its inline weight once "
|
||||
"when any keyword matches the current ask or a case-insensitive regex matches its first 2048 characters. "
|
||||
"Named dimensions added to the heuristic-v1 score. Each contributes its inline weight once "
|
||||
"when any keyword matches the current ask or a case-insensitive regex matches its first 2048 characters; "
|
||||
"scoring_mode 'match_count' instead grades half weight for one distinct matcher and full for two or more. "
|
||||
"Regex quantifiers repeat one character or class at most 64 times. Unbounded quantifiers, repeated groups, "
|
||||
"backreferences and lookarounds are rejected. Conservative work limits include alternation paths, "
|
||||
"repeat lengths and subsequent matching: 2048 units per pattern, 8192 across the router. "
|
||||
|
|
@ -1080,6 +1089,20 @@ class ComplexityRouterConfig(BaseModel):
|
|||
"wording the built-ins don't cover, or after a client release changes its strings."
|
||||
),
|
||||
)
|
||||
max_tokens_from_tier_model: bool = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"Set max_tokens on every routed request to the output ceiling of the tier model it "
|
||||
"lands on, replacing whatever the caller sent. A caller behind an auto-router cannot "
|
||||
"pick one value that fits every tier: the smallest tier's ceiling starves a bigger "
|
||||
"tier's thinking budget, and a bigger tier's ceiling is rejected by the smallest. The "
|
||||
"ceiling is the smallest max_output_tokens across the tier model's deployments, read "
|
||||
"from each deployment's model_info and then the model cost map; a tier model with a "
|
||||
"deployment whose ceiling is unknown keeps the caller's value. A max_tokens, "
|
||||
"max_completion_tokens or max_output_tokens in the tier's own litellm_params still "
|
||||
"wins. Set false to forward the caller's value unchanged."
|
||||
),
|
||||
)
|
||||
route_housekeeping_to_cheapest_tier: bool = Field(
|
||||
default=True,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -1,17 +1,103 @@
|
|||
#### What this does ####
|
||||
# identifies least busy deployment
|
||||
# How is this achieved?
|
||||
# - Before each call, have the router print the state of requests {"deployment": "requests_in_flight"}
|
||||
# - use litellm.input_callbacks to log when a request is just about to be made to a model - {"deployment-id": traffic}
|
||||
# - use litellm.success + failure callbacks to log when a request completed
|
||||
# - in get_available_deployment, for a given model group name -> pick based on traffic
|
||||
|
||||
import random
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
IN_FLIGHT_COUNT_TTL_SECONDS: Final = 60 * 60
|
||||
|
||||
|
||||
class _ModelInfo(TypedDict, total=False):
|
||||
id: ReadOnly[str | int | None]
|
||||
|
||||
|
||||
class _Metadata(TypedDict, total=False):
|
||||
model_group: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _LitellmParams(TypedDict, total=False):
|
||||
metadata: ReadOnly[_Metadata | None]
|
||||
model_info: ReadOnly[_ModelInfo | None]
|
||||
|
||||
|
||||
class _CallKwargs(TypedDict, total=False):
|
||||
litellm_params: ReadOnly[_LitellmParams | None]
|
||||
|
||||
|
||||
class _DeploymentModelInfo(TypedDict):
|
||||
id: ReadOnly[str | int]
|
||||
|
||||
|
||||
class _Deployment(TypedDict):
|
||||
model_info: ReadOnly[_DeploymentModelInfo]
|
||||
|
||||
|
||||
_CALL_KWARGS: Final = TypeAdapter(_CallKwargs)
|
||||
_DEPLOYMENTS: Final = TypeAdapter(list[_Deployment])
|
||||
_MEMORY_COUNTS: Final = TypeAdapter(tuple[float | None, ...] | None)
|
||||
|
||||
|
||||
def _request_count_key(model_group: str, deployment_id: str) -> str:
|
||||
return f"{model_group}_request_count:{deployment_id}"
|
||||
|
||||
|
||||
def _deployment_ref(kwargs: Mapping[str, object]) -> tuple[str, str] | None:
|
||||
try:
|
||||
call: Final = _CALL_KWARGS.validate_python(kwargs)
|
||||
except ValidationError:
|
||||
return None
|
||||
litellm_params: Final = call.get("litellm_params")
|
||||
metadata: Final = litellm_params.get("metadata") if litellm_params else None
|
||||
model_info: Final = litellm_params.get("model_info") if litellm_params else None
|
||||
model_group: Final = metadata.get("model_group") if metadata else None
|
||||
deployment_id: Final = model_info.get("id") if model_info else None
|
||||
if model_group is None or deployment_id is None:
|
||||
return None
|
||||
return model_group, str(deployment_id)
|
||||
|
||||
|
||||
def _request_count_keys(model_group: str, healthy_deployments: Sequence[Mapping[str, object]]) -> tuple[str, ...]:
|
||||
return tuple(
|
||||
_request_count_key(model_group, str(deployment["model_info"]["id"]))
|
||||
for deployment in _DEPLOYMENTS.validate_python(healthy_deployments)
|
||||
)
|
||||
|
||||
|
||||
def _as_counts(values: Sequence[float | None]) -> tuple[int, ...]:
|
||||
return tuple(0 if value is None else int(value) for value in values)
|
||||
|
||||
|
||||
def _local_counts(raw: object, keys: tuple[str, ...]) -> tuple[int, ...]:
|
||||
values: Final = _MEMORY_COUNTS.validate_python(raw)
|
||||
if values is None or len(values) != len(keys):
|
||||
return (0,) * len(keys)
|
||||
return _as_counts(values)
|
||||
|
||||
|
||||
def _least_busy(
|
||||
healthy_deployments: Sequence[Mapping[str, object]], counts: tuple[int, ...]
|
||||
) -> Mapping[str, object] | None:
|
||||
if not healthy_deployments:
|
||||
return None
|
||||
return healthy_deployments[min(range(len(healthy_deployments)), key=lambda index: counts[index])]
|
||||
|
||||
|
||||
def _warn_unreadable(model_group: str, error: Exception) -> None:
|
||||
verbose_router_logger.warning(
|
||||
"least-busy routing could not read the shared in-flight counts for %s, "
|
||||
"falling back to this worker's own counts: %s",
|
||||
model_group,
|
||||
error,
|
||||
)
|
||||
|
||||
|
||||
def _warn_unwritable(key: str, error: Exception) -> None:
|
||||
verbose_router_logger.warning("least-busy routing could not update the in-flight count under %s: %s", key, error)
|
||||
|
||||
|
||||
class LeastBusyLoggingHandler(CustomLogger):
|
||||
test_flag: bool = False
|
||||
|
|
@ -20,195 +106,101 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
|
||||
def __init__(self, router_cache: DualCache):
|
||||
self.router_cache = router_cache
|
||||
self.router_cache_id = str(id(router_cache))
|
||||
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
"""
|
||||
Log when a model is being used.
|
||||
def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None:
|
||||
self._increment(kwargs, 1)
|
||||
|
||||
Caching based on model group.
|
||||
"""
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
def log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
self._increment(kwargs, -1)
|
||||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# update cache
|
||||
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
|
||||
request_count_dict[id] = request_count_dict.get(id, 0) + 1
|
||||
def log_failure_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
self._increment(kwargs, -1)
|
||||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
|
||||
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
except Exception:
|
||||
pass
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
await self._async_increment(kwargs, -1)
|
||||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# decrement count in cache
|
||||
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
|
||||
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
|
||||
if request_count_value is None:
|
||||
return
|
||||
request_count_dict[id] = request_count_value - 1
|
||||
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
|
||||
### TESTING ###
|
||||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# decrement count in cache
|
||||
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
|
||||
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
|
||||
if request_count_value is None:
|
||||
return
|
||||
request_count_dict[id] = request_count_value - 1
|
||||
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
|
||||
### TESTING ###
|
||||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# decrement count in cache
|
||||
request_count_dict: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
|
||||
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
|
||||
if request_count_value is None:
|
||||
return
|
||||
request_count_dict[id] = request_count_value - 1
|
||||
await self.router_cache.async_set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
|
||||
### TESTING ###
|
||||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
# decrement count in cache
|
||||
request_count_dict: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
|
||||
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
|
||||
if request_count_value is None:
|
||||
return
|
||||
request_count_dict[id] = request_count_value - 1
|
||||
await self.router_cache.async_set_cache(key=request_count_api_key, value=request_count_dict)
|
||||
|
||||
### TESTING ###
|
||||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _get_available_deployments(
|
||||
self,
|
||||
healthy_deployments: list,
|
||||
all_deployments: dict,
|
||||
):
|
||||
"""
|
||||
Helper to get deployments using least busy strategy
|
||||
"""
|
||||
for d in healthy_deployments:
|
||||
## if healthy deployment not yet used
|
||||
if d["model_info"]["id"] not in all_deployments:
|
||||
all_deployments[d["model_info"]["id"]] = 0
|
||||
# map deployment to id
|
||||
# pick least busy deployment
|
||||
min_traffic = float("inf")
|
||||
min_deployment = None
|
||||
for k, v in all_deployments.items():
|
||||
if v < min_traffic:
|
||||
min_traffic = v
|
||||
min_deployment = k
|
||||
if min_deployment is not None:
|
||||
## check if min deployment is a string, if so, cast it to int
|
||||
for m in healthy_deployments:
|
||||
if m["model_info"]["id"] == min_deployment:
|
||||
return m
|
||||
min_deployment = random.choice(healthy_deployments)
|
||||
else:
|
||||
min_deployment = random.choice(healthy_deployments)
|
||||
return min_deployment
|
||||
async def async_log_failure_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
await self._async_increment(kwargs, -1)
|
||||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
|
||||
def get_available_deployments(
|
||||
self,
|
||||
model_group: str,
|
||||
healthy_deployments: list,
|
||||
):
|
||||
"""
|
||||
Sync helper to get deployments using least busy strategy
|
||||
"""
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
all_deployments: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
|
||||
return self._get_available_deployments(
|
||||
healthy_deployments=healthy_deployments,
|
||||
all_deployments=all_deployments,
|
||||
)
|
||||
self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]]
|
||||
) -> Mapping[str, object] | None:
|
||||
keys: Final = _request_count_keys(model_group, healthy_deployments)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
shared: Final = _as_counts(redis_cache.batch_get_counts(list(keys)))
|
||||
except Exception as e:
|
||||
_warn_unreadable(model_group, e)
|
||||
else:
|
||||
return _least_busy(healthy_deployments, shared)
|
||||
local: Final = _local_counts(self.router_cache.batch_get_cache(list(keys), local_only=True), keys)
|
||||
return _least_busy(healthy_deployments, local)
|
||||
|
||||
async def async_get_available_deployments(self, model_group: str, healthy_deployments: list):
|
||||
"""
|
||||
Async helper to get deployments using least busy strategy
|
||||
"""
|
||||
request_count_api_key: Final = f"{model_group}_request_count"
|
||||
all_deployments: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
|
||||
return self._get_available_deployments(
|
||||
healthy_deployments=healthy_deployments,
|
||||
all_deployments=all_deployments,
|
||||
)
|
||||
async def async_get_available_deployments(
|
||||
self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]]
|
||||
) -> Mapping[str, object] | None:
|
||||
keys: Final = _request_count_keys(model_group, healthy_deployments)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
shared: Final = _as_counts(await redis_cache.async_batch_get_counts(list(keys)))
|
||||
except Exception as e:
|
||||
_warn_unreadable(model_group, e)
|
||||
else:
|
||||
return _least_busy(healthy_deployments, shared)
|
||||
local: Final = _local_counts(await self.router_cache.async_batch_get_cache(list(keys), local_only=True), keys)
|
||||
return _least_busy(healthy_deployments, local)
|
||||
|
||||
def _increment(self, kwargs: Mapping[str, object], delta: int) -> None:
|
||||
ref: Final = _deployment_ref(kwargs)
|
||||
if ref is None:
|
||||
return
|
||||
key: Final = _request_count_key(*ref)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
try:
|
||||
local: Final = self.router_cache.increment_cache(
|
||||
key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS
|
||||
)
|
||||
if local < 0:
|
||||
self.router_cache.set_cache(key, 0, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
if redis_cache is None:
|
||||
return
|
||||
redis_cache.increment_with_floor(key, delta, IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
except Exception as e:
|
||||
_warn_unwritable(key, e)
|
||||
|
||||
async def _async_increment(self, kwargs: Mapping[str, object], delta: int) -> None:
|
||||
ref: Final = _deployment_ref(kwargs)
|
||||
if ref is None:
|
||||
return
|
||||
key: Final = _request_count_key(*ref)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
try:
|
||||
local: Final = await self.router_cache.async_increment_cache(
|
||||
key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS
|
||||
)
|
||||
if local is not None and local < 0:
|
||||
await self.router_cache.async_set_cache(key, 0, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
if redis_cache is None:
|
||||
return
|
||||
await redis_cache.async_increment_with_floor(key, delta, IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
except Exception as e:
|
||||
_warn_unwritable(key, e)
|
||||
|
|
|
|||
|
|
@ -39,7 +39,7 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
# ------------
|
||||
"""
|
||||
{
|
||||
{model_group}_map: {
|
||||
cost_map:{model_group}: {
|
||||
id: {
|
||||
f"{date:hour:minute}" : {"tpm": 34, "rpm": 3}
|
||||
}
|
||||
|
|
@ -50,7 +50,7 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
current_hour: Final = datetime.now().strftime("%H")
|
||||
current_minute: Final = datetime.now().strftime("%M")
|
||||
precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
|
||||
cost_key: Final = f"{model_group}_map"
|
||||
cost_key: Final = f"cost_map:{model_group}"
|
||||
|
||||
total_tokens = 0
|
||||
|
||||
|
|
@ -112,15 +112,14 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
# ------------
|
||||
"""
|
||||
{
|
||||
{model_group}_map: {
|
||||
cost_map:{model_group}: {
|
||||
id: {
|
||||
"cost": [..]
|
||||
f"{date:hour:minute}" : {"tpm": 34, "rpm": 3}
|
||||
}
|
||||
}
|
||||
}
|
||||
"""
|
||||
cost_key: Final = f"{model_group}_map"
|
||||
cost_key: Final = f"cost_map:{model_group}"
|
||||
|
||||
current_date: Final = datetime.now().strftime("%Y-%m-%d")
|
||||
current_hour: Final = datetime.now().strftime("%H")
|
||||
|
|
@ -176,7 +175,7 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
"""
|
||||
Returns a deployment with the lowest cost
|
||||
"""
|
||||
cost_key: Final = f"{model_group}_map"
|
||||
cost_key: Final = f"cost_map:{model_group}"
|
||||
|
||||
request_count_dict: Final = await self.router_cache.async_get_cache(key=cost_key) or {}
|
||||
|
||||
|
|
|
|||
|
|
@ -32,6 +32,12 @@ def _average_latency(samples: Sequence[float]) -> float:
|
|||
return sum(samples) / len(samples)
|
||||
|
||||
|
||||
def _ttft_seconds(elapsed: timedelta | float) -> float:
|
||||
if isinstance(elapsed, timedelta):
|
||||
return elapsed.total_seconds()
|
||||
return float(elapsed)
|
||||
|
||||
|
||||
class LowestLatencyLoggingHandler(CustomLogger):
|
||||
test_flag: bool = False
|
||||
logged_success: int = 0
|
||||
|
|
@ -86,14 +92,13 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
# breaks JSON serialization when the router cache syncs to
|
||||
# Redis (issue #33169)
|
||||
response_ms = response_ms.total_seconds()
|
||||
time_to_first_token_response_time = None
|
||||
time_to_first_token: float | None = None
|
||||
|
||||
if kwargs.get("stream", None) is not None and kwargs["stream"] is True:
|
||||
# only log ttft for streaming request
|
||||
time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time
|
||||
time_to_first_token = _ttft_seconds(kwargs.get("completion_start_time", end_time) - start_time)
|
||||
|
||||
final_value: float = response_ms
|
||||
time_to_first_token: float | None = None
|
||||
total_tokens = 0
|
||||
|
||||
if isinstance(response_obj, ModelResponse):
|
||||
|
|
@ -111,13 +116,6 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
else:
|
||||
final_value = response_seconds
|
||||
|
||||
if time_to_first_token_response_time is not None:
|
||||
if isinstance(time_to_first_token_response_time, timedelta):
|
||||
ttft_seconds = time_to_first_token_response_time.total_seconds()
|
||||
else:
|
||||
ttft_seconds = time_to_first_token_response_time
|
||||
time_to_first_token = safe_divide_seconds(ttft_seconds, completion_tokens)
|
||||
|
||||
# ------------
|
||||
# Update usage
|
||||
# ------------
|
||||
|
|
@ -138,14 +136,14 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
## Time to first token
|
||||
if time_to_first_token is not None:
|
||||
if (
|
||||
len(request_count_dict[id].get("time_to_first_token", []))
|
||||
len(request_count_dict[id].get("time_to_first_token_seconds", []))
|
||||
< self.routing_args.max_latency_list_size
|
||||
):
|
||||
request_count_dict[id].setdefault("time_to_first_token", []).append(time_to_first_token)
|
||||
request_count_dict[id].setdefault("time_to_first_token_seconds", []).append(time_to_first_token)
|
||||
else:
|
||||
request_count_dict[id]["time_to_first_token"] = request_count_dict[id]["time_to_first_token"][
|
||||
1:
|
||||
] + [time_to_first_token]
|
||||
request_count_dict[id]["time_to_first_token_seconds"] = request_count_dict[id][
|
||||
"time_to_first_token_seconds"
|
||||
][1:] + [time_to_first_token]
|
||||
|
||||
if precise_minute not in request_count_dict[id]:
|
||||
request_count_dict[id][precise_minute] = {}
|
||||
|
|
@ -252,7 +250,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
{model_group}_map: {
|
||||
id: {
|
||||
"latency": [..]
|
||||
"time_to_first_token": [..]
|
||||
"time_to_first_token_seconds": [..]
|
||||
f"{date:hour:minute}" : {"tpm": 34, "rpm": 3}
|
||||
}
|
||||
}
|
||||
|
|
@ -273,14 +271,13 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
# breaks JSON serialization when the router cache syncs to
|
||||
# Redis (issue #33169)
|
||||
response_ms = response_ms.total_seconds()
|
||||
time_to_first_token_response_time = None
|
||||
time_to_first_token: float | None = None
|
||||
if kwargs.get("stream", None) is not None and kwargs["stream"] is True:
|
||||
# only log ttft for streaming request
|
||||
time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time
|
||||
time_to_first_token = _ttft_seconds(kwargs.get("completion_start_time", end_time) - start_time)
|
||||
|
||||
final_value: float = response_ms
|
||||
total_tokens = 0
|
||||
time_to_first_token: float | None = None
|
||||
|
||||
if isinstance(response_obj, ModelResponse):
|
||||
_usage: Final = getattr(response_obj, "usage", None)
|
||||
|
|
@ -296,13 +293,6 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
final_value = float(normalized_value)
|
||||
else:
|
||||
final_value = response_seconds
|
||||
|
||||
if time_to_first_token_response_time is not None:
|
||||
if isinstance(time_to_first_token_response_time, timedelta):
|
||||
ttft_seconds = time_to_first_token_response_time.total_seconds()
|
||||
else:
|
||||
ttft_seconds = time_to_first_token_response_time
|
||||
time_to_first_token = safe_divide_seconds(ttft_seconds, completion_tokens)
|
||||
# ------------
|
||||
# Update usage
|
||||
# ------------
|
||||
|
|
@ -328,14 +318,14 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
## Time to first token
|
||||
if time_to_first_token is not None:
|
||||
if (
|
||||
len(request_count_dict[id].get("time_to_first_token", []))
|
||||
len(request_count_dict[id].get("time_to_first_token_seconds", []))
|
||||
< self.routing_args.max_latency_list_size
|
||||
):
|
||||
request_count_dict[id].setdefault("time_to_first_token", []).append(time_to_first_token)
|
||||
request_count_dict[id].setdefault("time_to_first_token_seconds", []).append(time_to_first_token)
|
||||
else:
|
||||
request_count_dict[id]["time_to_first_token"] = request_count_dict[id]["time_to_first_token"][
|
||||
1:
|
||||
] + [time_to_first_token]
|
||||
request_count_dict[id]["time_to_first_token_seconds"] = request_count_dict[id][
|
||||
"time_to_first_token_seconds"
|
||||
][1:] + [time_to_first_token]
|
||||
|
||||
if precise_minute not in request_count_dict[id]:
|
||||
request_count_dict[id][precise_minute] = {}
|
||||
|
|
@ -433,7 +423,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
or float("inf")
|
||||
)
|
||||
item_latency = item_map.get("latency", [])
|
||||
item_ttft_latency = item_map.get("time_to_first_token", [])
|
||||
item_ttft_latency = item_map.get("time_to_first_token_seconds", [])
|
||||
item_rpm = item_map.get(precise_minute, {}).get("rpm", 0)
|
||||
item_tpm = item_map.get(precise_minute, {}).get("tpm", 0)
|
||||
|
||||
|
|
|
|||
|
|
@ -41,8 +41,7 @@ def simple_shuffle(
|
|||
|
||||
############## Check if 'weight' or 'rpm' or 'tpm' param set for a weighted pick #################
|
||||
for weight_by in ["weight", "rpm", "tpm"]:
|
||||
weight = healthy_deployments[0].get("litellm_params").get(weight_by, None)
|
||||
if weight is not None:
|
||||
if any(m["litellm_params"].get(weight_by) is not None for m in healthy_deployments):
|
||||
weights = [m["litellm_params"].get(weight_by, 0) for m in healthy_deployments]
|
||||
verbose_router_logger.debug("\nweight %s", weights)
|
||||
total_weight = sum(weights)
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import CONSUMED_REQUEST_TAGS_METADATA_KEY
|
||||
from litellm.constants import CONSUMED_REQUEST_TAGS_METADATA_KEY, ROUTING_REQUEST_TAGS_METADATA_KEY
|
||||
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
|
||||
from litellm.types.router import ConsumedRequestTagsStamp, DeploymentTypedDict, RouterErrors
|
||||
|
||||
|
|
@ -461,7 +461,10 @@ def _request_tags_after_router_consumption(metadata: object, model: str) -> Sequ
|
|||
if not isinstance(metadata, Mapping):
|
||||
return None
|
||||
typed_metadata: Final[Mapping[str, object]] = metadata
|
||||
request_tags: Final = _tags_in_metadata(typed_metadata)
|
||||
request_tags: Final = _tags_in_metadata(
|
||||
typed_metadata,
|
||||
key=ROUTING_REQUEST_TAGS_METADATA_KEY if ROUTING_REQUEST_TAGS_METADATA_KEY in typed_metadata else "tags",
|
||||
)
|
||||
stamp: Final = typed_metadata.get(CONSUMED_REQUEST_TAGS_METADATA_KEY)
|
||||
if not isinstance(stamp, ConsumedRequestTagsStamp) or stamp.model_group != model:
|
||||
return request_tags
|
||||
|
|
@ -646,7 +649,7 @@ async def get_deployments_for_tag(
|
|||
return healthy_deployments
|
||||
|
||||
|
||||
def _tags_in_metadata(metadata: object) -> list[str]:
|
||||
def _tags_in_metadata(metadata: object, key: str = "tags") -> list[str]:
|
||||
"""
|
||||
Tags out of a metadata bucket the caller controls the shape of.
|
||||
|
||||
|
|
@ -657,7 +660,7 @@ def _tags_in_metadata(metadata: object) -> list[str]:
|
|||
if not isinstance(metadata, Mapping):
|
||||
return []
|
||||
typed_metadata: Final[Mapping[str, object]] = metadata
|
||||
tags: Final = typed_metadata.get("tags")
|
||||
tags: Final = typed_metadata.get(key)
|
||||
if isinstance(tags, str) or not isinstance(tags, Sequence):
|
||||
return []
|
||||
typed_tags: Final[Sequence[object]] = tags
|
||||
|
|
|
|||
|
|
@ -52,7 +52,17 @@ def tuning_fingerprint(complexity_router_config: object) -> str | None:
|
|||
supplied: Final = ((_TUNING_FIELD_SET - frozenset(("tier_model_configs",))) & frozenset(raw)) | (
|
||||
frozenset(("tier_model_configs",)) if validated.tier_model_configs else frozenset()
|
||||
)
|
||||
payload: Final = validated.model_dump(mode="json", include=supplied)
|
||||
payload: Final = validated.model_dump(
|
||||
mode="json",
|
||||
include=supplied,
|
||||
exclude={
|
||||
"custom_dimensions": {
|
||||
index: {"scoring_mode"}
|
||||
for index, dimension in enumerate(validated.custom_dimensions)
|
||||
if dimension.scoring_mode == "binary"
|
||||
}
|
||||
},
|
||||
)
|
||||
return hashlib.sha256(json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing_extensions import TypedDict
|
|||
from litellm import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -36,10 +37,19 @@ _MAX_CORRECTED_IN_MEMORY_TTL_SECONDS: Final = 60.0
|
|||
|
||||
|
||||
class CooldownCache:
|
||||
def __init__(self, cache: DualCache, default_cooldown_time: float):
|
||||
def __init__(
|
||||
self,
|
||||
cache: DualCache,
|
||||
default_cooldown_time: float,
|
||||
redis_read_interval_seconds: float = DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS,
|
||||
):
|
||||
self.cache = cache
|
||||
self.default_cooldown_time = default_cooldown_time
|
||||
self.in_memory_cache = InMemoryCache()
|
||||
self._cooldown_store = DualCache(
|
||||
in_memory_cache=self.in_memory_cache,
|
||||
default_redis_batch_cache_expiry=redis_read_interval_seconds,
|
||||
)
|
||||
# Initialize the masker with custom settings for exception strings
|
||||
self.exception_masker = SensitiveDataMasker(
|
||||
visible_prefix=50, # Show first 50 characters
|
||||
|
|
@ -48,6 +58,21 @@ class CooldownCache:
|
|||
mask_short_values=False, # Truncate long messages only; keep short ones readable
|
||||
)
|
||||
|
||||
@property
|
||||
def cooldown_store(self) -> DualCache:
|
||||
"""
|
||||
The cache cooldown entries live in, with the router's Redis attached on first use.
|
||||
|
||||
It is kept separate from the router-wide cache so that a key missing from memory is
|
||||
re-read from Redis every `redis_read_interval_seconds` rather than on the router
|
||||
cache's much longer batch interval, which is what lets a sibling replica see a
|
||||
cooldown another replica wrote, and so that unrelated router keys cannot evict a
|
||||
cooldown from the in-memory tier before it expires. Redis is attached lazily because
|
||||
the router builds its cooldown cache before it wires up the shared Redis client.
|
||||
"""
|
||||
self._cooldown_store.attach_redis_cache(self.cache.redis_cache)
|
||||
return self._cooldown_store
|
||||
|
||||
def _common_add_cooldown_logic(
|
||||
self, model_id: str, original_exception, exception_status, cooldown_time: float
|
||||
) -> tuple[str, CooldownCacheValue]:
|
||||
|
|
@ -93,7 +118,7 @@ class CooldownCache:
|
|||
)
|
||||
|
||||
# Set the cache with a TTL equal to the cooldown time
|
||||
self.cache.set_cache(
|
||||
self.cooldown_store.set_cache(
|
||||
value=cooldown_data,
|
||||
key=cooldown_key,
|
||||
ttl=_cooldown_time,
|
||||
|
|
@ -122,13 +147,13 @@ class CooldownCache:
|
|||
cooldown_cache_value: Final = CooldownCacheValue(**result) # pyright: ignore[reportUnknownArgumentType] - result comes from an untyped cache read, not from our own code
|
||||
remaining: Final = (cooldown_cache_value["timestamp"] + cooldown_cache_value["cooldown_time"]) - current_time
|
||||
if remaining <= 0:
|
||||
self.cache.in_memory_cache.delete_cache(key)
|
||||
self.in_memory_cache.delete_cache(key)
|
||||
return None
|
||||
current_expiry: Final = self.cache.in_memory_cache.ttl_dict.get(key)
|
||||
current_expiry: Final = self.in_memory_cache.ttl_dict.get(key)
|
||||
if current_expiry is not None and current_expiry > current_time + remaining + 5:
|
||||
corrected_ttl: Final = min(remaining, _MAX_CORRECTED_IN_MEMORY_TTL_SECONDS)
|
||||
self.cache.in_memory_cache.delete_cache(key)
|
||||
self.cache.in_memory_cache.set_cache(key, result, ttl=corrected_ttl)
|
||||
self.in_memory_cache.delete_cache(key)
|
||||
self.in_memory_cache.set_cache(key, result, ttl=corrected_ttl)
|
||||
return cooldown_cache_value
|
||||
|
||||
async def async_get_active_cooldowns(
|
||||
|
|
@ -137,12 +162,7 @@ class CooldownCache:
|
|||
# Generate the keys for the deployments
|
||||
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
|
||||
|
||||
# Retrieve the values for the keys using mget
|
||||
## more likely to be none if no models ratelimited. So just check redis every 1s
|
||||
## each redis call adds ~100ms latency.
|
||||
|
||||
## check in memory cache first
|
||||
results: Final = await self.cache.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
|
||||
results: Final = await self.cooldown_store.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
|
||||
active_cooldowns: Final[list[tuple[str, CooldownCacheValue]]] = []
|
||||
|
||||
if results is None or all(v is None for v in results):
|
||||
|
|
@ -164,7 +184,7 @@ class CooldownCache:
|
|||
# Generate the keys for the deployments
|
||||
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
|
||||
# Retrieve the values for the keys using mget
|
||||
results: Final = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
|
||||
active_cooldowns: Final = []
|
||||
current_time: Final = time.time()
|
||||
|
|
@ -184,7 +204,7 @@ class CooldownCache:
|
|||
keys: Final = [f"deployment:{model_id}:cooldown" for model_id in model_ids]
|
||||
|
||||
# Retrieve the values for the keys using mget
|
||||
results: Final = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
|
||||
min_cooldown_time: float | None = None
|
||||
# Process the results
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import (
|
||||
DEFAULT_COOLDOWN_TIME_SECONDS,
|
||||
DEFAULT_FAILURE_THRESHOLD_MINIMUM_REQUESTS,
|
||||
|
|
@ -558,9 +559,12 @@ def should_cooldown_based_on_allowed_fails_policy(
|
|||
When *allowed_fails_override* / *cooldown_time_override* are supplied they
|
||||
take precedence over the router-level values (used by deployment-level overrides).
|
||||
|
||||
The counter lives in the router's shared ``DualCache`` (Redis when configured), so
|
||||
every worker process increments the same key and the threshold applies fleet-wide.
|
||||
|
||||
When *cache_key_suffix* is supplied the fail counter is keyed as
|
||||
``{deployment}:{cache_key_suffix}`` so that different exception types are
|
||||
tracked independently per deployment.
|
||||
``deployment:{deployment}:allowed_fails:{cache_key_suffix}`` so that different
|
||||
exception types are tracked independently per deployment.
|
||||
|
||||
Returns:
|
||||
- True if fails exceed the allowed limit (should cooldown)
|
||||
|
|
@ -584,16 +588,25 @@ def should_cooldown_based_on_allowed_fails_policy(
|
|||
else (litellm_router_instance.cooldown_time or DEFAULT_COOLDOWN_TIME_SECONDS)
|
||||
)
|
||||
|
||||
cache_key: Final = f"{deployment}:{cache_key_suffix}" if cache_key_suffix else deployment
|
||||
current_fails: Final = litellm_router_instance.failed_calls.get_cache(key=cache_key) or 0
|
||||
updated_fails: Final = current_fails + 1
|
||||
base_key: Final = f"deployment:{deployment}:allowed_fails"
|
||||
cache_key: Final = f"{base_key}:{cache_key_suffix}" if cache_key_suffix else base_key
|
||||
updated_fails: Final = _increment_allowed_fails(
|
||||
cache=litellm_router_instance.cache, cache_key=cache_key, ttl=cooldown_time
|
||||
)
|
||||
return updated_fails > allowed_fails
|
||||
|
||||
if updated_fails > allowed_fails:
|
||||
return True
|
||||
else:
|
||||
litellm_router_instance.failed_calls.set_cache(key=cache_key, value=updated_fails, ttl=cooldown_time)
|
||||
|
||||
return False
|
||||
def _increment_allowed_fails(cache: DualCache, cache_key: str, ttl: float) -> int:
|
||||
"""
|
||||
Return the fleet-wide fail count. ``DualCache.increment_cache`` bumps the in-memory tier
|
||||
before Redis and re-raises a Redis error, so a Redis outage degrades to this worker's own count.
|
||||
"""
|
||||
try:
|
||||
return cache.increment_cache(key=cache_key, value=1, ttl=ttl)
|
||||
except Exception as e: # noqa: BLE001 # a Redis outage must not stop failing deployments from cooling down
|
||||
verbose_router_logger.warning("allowed_fails counter fell back to this worker's in-memory count: %s", e)
|
||||
local_fails: Final = cache.get_cache(key=cache_key, local_only=True)
|
||||
return local_fails if isinstance(local_fails, int) else 0
|
||||
|
||||
|
||||
def _is_allowed_fails_set_on_router(
|
||||
|
|
|
|||
|
|
@ -12,7 +12,9 @@ import httpx
|
|||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_checkable
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
|
@ -314,6 +316,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
|||
timeout: float | str | httpx.Timeout | None = None # if str, pass in as os.environ/
|
||||
stream_timeout: float | str | None = None # timeout when making stream=True calls, if str, pass in as os.environ/
|
||||
max_retries: int | None = None
|
||||
drop_params: bool | str | None = None
|
||||
organization: str | None = None # for openai orgs
|
||||
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None
|
||||
litellm_credential_name: str | None = None
|
||||
|
|
@ -404,6 +407,18 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
|||
return filtered
|
||||
return data
|
||||
|
||||
@field_validator("drop_params", mode="before")
|
||||
@classmethod
|
||||
def coerce_drop_params(cls, value: object) -> bool | str | None:
|
||||
normalized: Final = normalize_drop_params(value)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if value is not None:
|
||||
verbose_logger.warning("drop_params=%r is not a flag value, treating it as unset", value)
|
||||
return None
|
||||
|
||||
def __contains__(self, key) -> bool:
|
||||
# Define custom behavior for the 'in' operator
|
||||
return hasattr(self, key)
|
||||
|
|
|
|||
|
|
@ -580,6 +580,7 @@ CallTypesLiteral = Literal[
|
|||
"search",
|
||||
"asearch",
|
||||
"_arealtime",
|
||||
"_aresponses_websocket",
|
||||
"create_batch",
|
||||
"acreate_batch",
|
||||
"create_file",
|
||||
|
|
|
|||
140
litellm/utils.py
140
litellm/utils.py
|
|
@ -80,6 +80,7 @@ from litellm.constants import (
|
|||
PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO,
|
||||
TOOL_CHOICE_OBJECT_TOKEN_COUNT,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
|
||||
from litellm.litellm_core_utils.fallback_generalizations import (
|
||||
match_capability_generalizations,
|
||||
)
|
||||
|
|
@ -672,11 +673,19 @@ def load_credentials_from_list(kwargs: dict):
|
|||
CredentialAccessor: Final = getattr(sys.modules[__name__], "CredentialAccessor")
|
||||
|
||||
credential_name: Final = kwargs.get("litellm_credential_name")
|
||||
if credential_name and litellm.credential_list:
|
||||
credential_accessor: Final[Mapping[str, object]] = CredentialAccessor.get_credential_values(credential_name)
|
||||
for key, value in credential_accessor.items():
|
||||
if key not in kwargs:
|
||||
kwargs[key] = value
|
||||
if not credential_name:
|
||||
return
|
||||
credential: Final = CredentialAccessor.find_credential(credential_name)
|
||||
if credential is None:
|
||||
verbose_logger.warning(
|
||||
"litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it",
|
||||
credential_name,
|
||||
len(litellm.credential_list),
|
||||
)
|
||||
return
|
||||
for key, value in credential.credential_values.items():
|
||||
if key not in kwargs:
|
||||
kwargs[key] = value
|
||||
|
||||
|
||||
def get_dynamic_callbacks(
|
||||
|
|
@ -3231,7 +3240,7 @@ def get_optional_params_transcription(
|
|||
|
||||
passed_params.pop("OPENAI_TRANSCRIPTION_PARAMS")
|
||||
custom_llm_provider = passed_params.pop("custom_llm_provider")
|
||||
drop_params = passed_params.pop("drop_params")
|
||||
drop_params = normalize_drop_params(passed_params.pop("drop_params"))
|
||||
special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs")
|
||||
for k, v in special_params.items():
|
||||
passed_params[k] = v
|
||||
|
|
@ -3339,7 +3348,7 @@ def get_optional_params_image_gen(
|
|||
model = passed_params.pop("model", None)
|
||||
custom_llm_provider = passed_params.pop("custom_llm_provider")
|
||||
provider_config = passed_params.pop("provider_config", None)
|
||||
drop_params = passed_params.pop("drop_params", None)
|
||||
drop_params = normalize_drop_params(passed_params.pop("drop_params", None))
|
||||
additional_drop_params = passed_params.pop("additional_drop_params", None)
|
||||
special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs")
|
||||
for k, v in special_params.items():
|
||||
|
|
@ -3467,7 +3476,7 @@ def get_optional_params_embeddings(
|
|||
custom_llm_provider = passed_params.pop("custom_llm_provider", None)
|
||||
special_params: Final = passed_params.pop("kwargs")
|
||||
|
||||
drop_params = passed_params.pop("drop_params", None)
|
||||
drop_params = normalize_drop_params(passed_params.pop("drop_params", None))
|
||||
additional_drop_params = passed_params.pop("additional_drop_params", None)
|
||||
allowed_openai_params = passed_params.pop("allowed_openai_params", None) or []
|
||||
# Remove function objects from passed_params to avoid JSON serialization errors
|
||||
|
|
@ -4194,6 +4203,7 @@ def get_optional_params(
|
|||
base_model: str | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
drop_params = normalize_drop_params(drop_params) # rebind-ok: config and DB deployments pass "true" as a string
|
||||
passed_params: Final = locals().copy()
|
||||
special_params: Final = passed_params.pop("kwargs")
|
||||
# Remove base_model from passed_params so it doesn't interfere with
|
||||
|
|
@ -4271,20 +4281,20 @@ def get_optional_params(
|
|||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "anthropic_text":
|
||||
optional_params = litellm.AnthropicTextConfig().map_openai_params(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
optional_params = litellm.AnthropicTextConfig().map_openai_params(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "cohere_chat" or custom_llm_provider == "cohere":
|
||||
|
|
@ -4293,14 +4303,14 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "triton":
|
||||
optional_params = litellm.TritonConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=drop_params if drop_params is not None else False,
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "maritalk":
|
||||
|
|
@ -4308,35 +4318,35 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "replicate":
|
||||
optional_params = litellm.ReplicateConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "predibase":
|
||||
optional_params = litellm.PredibaseConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "huggingface":
|
||||
optional_params = litellm.HuggingFaceChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "together_ai":
|
||||
optional_params = litellm.TogetherAIChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai" and (
|
||||
model in litellm.vertex_chat_models
|
||||
|
|
@ -4350,7 +4360,7 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "gemini":
|
||||
|
|
@ -4358,21 +4368,21 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai_beta" or (custom_llm_provider == "vertex_ai" and "gemini" in model):
|
||||
optional_params = litellm.VertexGeminiConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif litellm.VertexAIAnthropicConfig.is_supported_model(model=model, custom_llm_provider=custom_llm_provider):
|
||||
optional_params = litellm.VertexAIAnthropicConfig().map_openai_params(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
if model in litellm.vertex_mistral_models:
|
||||
|
|
@ -4381,35 +4391,35 @@ def get_optional_params(
|
|||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
else:
|
||||
optional_params = litellm.MistralConfig().map_openai_params(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif model in litellm.vertex_ai_ai21_models:
|
||||
optional_params = litellm.VertexAIAi21Config().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif provider_config is not None:
|
||||
optional_params = provider_config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
else: # use generic openai-like param mapping
|
||||
optional_params = litellm.VertexAILlama3Config().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "sagemaker":
|
||||
|
|
@ -4418,7 +4428,7 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "bedrock":
|
||||
BedrockModelInfo: Final = getattr(sys.modules[__name__], "BedrockModelInfo")
|
||||
|
|
@ -4429,14 +4439,14 @@ def get_optional_params(
|
|||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif bedrock_route == "openai":
|
||||
optional_params = litellm.AmazonBedrockOpenAIConfig().map_openai_params(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif "anthropic" in bedrock_base_model and bedrock_route == "invoke":
|
||||
if bedrock_base_model in litellm.AmazonAnthropicConfig.get_legacy_anthropic_model_names():
|
||||
|
|
@ -4444,21 +4454,21 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
else:
|
||||
optional_params = litellm.AmazonAnthropicClaudeConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif provider_config is not None:
|
||||
optional_params = provider_config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
if bedrock_route == "claude_platform":
|
||||
optional_params = BedrockModelInfo.map_claude_platform_auth_params(
|
||||
|
|
@ -4469,28 +4479,28 @@ def get_optional_params(
|
|||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "ollama":
|
||||
optional_params = litellm.OllamaConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "ollama_chat":
|
||||
optional_params = litellm.OllamaChatConfig().map_openai_params(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "nlp_cloud":
|
||||
optional_params = litellm.NLPCloudConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "petals":
|
||||
|
|
@ -4498,35 +4508,35 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "deepinfra":
|
||||
optional_params = litellm.DeepInfraConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "perplexity" and provider_config is not None:
|
||||
optional_params = provider_config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "mistral" or custom_llm_provider == "codestral":
|
||||
optional_params = litellm.MistralConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "text-completion-codestral":
|
||||
optional_params = litellm.CodestralTextCompletionConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "text-completion-inception":
|
||||
|
|
@ -4534,7 +4544,7 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "databricks":
|
||||
|
|
@ -4542,21 +4552,21 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "nvidia_nim":
|
||||
optional_params = litellm.NvidiaNimConfig().map_openai_params(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "cerebras":
|
||||
optional_params = litellm.CerebrasConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "xai":
|
||||
optional_params = litellm.XAIChatConfig().map_openai_params(
|
||||
|
|
@ -4569,77 +4579,77 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "fireworks_ai":
|
||||
optional_params = litellm.FireworksAIConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "volcengine":
|
||||
optional_params = litellm.VolcEngineConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "hosted_vllm":
|
||||
optional_params = litellm.HostedVLLMChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "vllm":
|
||||
optional_params = litellm.VLLMConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "groq":
|
||||
optional_params = litellm.GroqChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "bedrock_mantle":
|
||||
optional_params = litellm.BedrockMantleChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "deepseek":
|
||||
optional_params = litellm.DeepSeekChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "tencent":
|
||||
optional_params = litellm.TencentChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "openrouter":
|
||||
optional_params = litellm.OpenrouterConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "watsonx":
|
||||
optional_params = litellm.IBMWatsonXChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
# WatsonX-text param check
|
||||
for param in passed_params:
|
||||
|
|
@ -4652,21 +4662,21 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "openai":
|
||||
optional_params = litellm.OpenAIConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "nebius":
|
||||
optional_params = litellm.NebiusConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
_azure_detection_model: Final = base_model or model
|
||||
|
|
@ -4675,14 +4685,14 @@ def get_optional_params(
|
|||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=_azure_detection_model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=_azure_detection_model):
|
||||
optional_params = litellm.AzureOpenAIGPT5Config().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=_azure_detection_model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
|
|
@ -4701,21 +4711,21 @@ def get_optional_params(
|
|||
optional_params=optional_params,
|
||||
model=_azure_detection_model,
|
||||
api_version=api_version,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
elif provider_config is not None:
|
||||
optional_params = provider_config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
else: # assume passing in params for openai-like api
|
||||
optional_params = litellm.OpenAILikeChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
drop_params=bool(drop_params),
|
||||
)
|
||||
# if user passed in non-default kwargs for specific providers/models, pass them along
|
||||
optional_params = add_provider_specific_params_to_optional_params(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm"
|
||||
version = "1.101.0"
|
||||
version = "1.102.0"
|
||||
description = "Library to easily interface with LLM API providers"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10, <3.15"
|
||||
|
|
@ -67,7 +67,7 @@ proxy = [
|
|||
"azure-identity>=1.25.2,<2.0",
|
||||
"azure-storage-blob>=12.28.0,<13.0",
|
||||
"mcp>=1.28.1,<2.0",
|
||||
"litellm-proxy-extras==0.4.94",
|
||||
"litellm-proxy-extras==0.4.95",
|
||||
"litellm-enterprise==0.1.65",
|
||||
"RestrictedPython>=8.5,<9.0",
|
||||
"rich>=13.9.4,<14.0",
|
||||
|
|
@ -114,7 +114,6 @@ caching = ["diskcache>=5.6.3,<6.0"]
|
|||
mcp = ["mcp>=1.28.1,<2.0"]
|
||||
# Driver for the MongoDB Atlas vector store; Atlas Vector Search has no HTTP query API.
|
||||
# The floor is 4.9 because that is the release AsyncMongoClient landed in.
|
||||
mongodb = ["pymongo>=4.9,<5.0"]
|
||||
# SAML SSO for the admin UI. python3-saml pulls in xmlsec/lxml, whose wheels
|
||||
# bundle the native libxmlsec1/libxml2 libraries, so no system packages are
|
||||
# required. Kept out of the base `proxy` extra so it stays optional.
|
||||
|
|
@ -328,7 +327,7 @@ members = ["enterprise", "litellm-proxy-extras"]
|
|||
profile = "black"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "1.101.0"
|
||||
version = "1.102.0"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 2956
|
||||
"limit": 2918
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
|
|
@ -9,10 +9,10 @@
|
|||
"limit": 806
|
||||
},
|
||||
"ANN201": {
|
||||
"limit": 1979
|
||||
"limit": 1965
|
||||
},
|
||||
"ANN202": {
|
||||
"limit": 831
|
||||
"limit": 829
|
||||
},
|
||||
"ANN204": {
|
||||
"limit": 683
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 3
|
||||
},
|
||||
"BLE001": {
|
||||
"limit": 2916
|
||||
"limit": 2914
|
||||
},
|
||||
"C401": {
|
||||
"limit": 8
|
||||
|
|
@ -189,7 +189,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"S110": {
|
||||
"limit": 217
|
||||
"limit": 207
|
||||
},
|
||||
"S112": {
|
||||
"limit": 22
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -242,6 +242,22 @@ this with `litellm_license`. To tune the export cadence, set
|
|||
`LITELLM_BILLING_METRICS_EXPORT_INTERVAL_MS` through `gateway_extra_env` /
|
||||
`backend_extra_env`
|
||||
|
||||
### Prometheus metrics sidecar
|
||||
|
||||
`gateway_metrics_port` adds a `metrics` sidecar
|
||||
(`python -m litellm.proxy.prometheus_metrics_server`) to the gateway task that
|
||||
aggregates the workers' samples over a shared task volume, so a scrape never
|
||||
runs on an inference worker. The ALB never routes to that port and the tasks
|
||||
security group only opens it to `gateway_metrics_scrape_cidrs`. Needs
|
||||
`gateway_image` v1.101.0 or newer. See
|
||||
[Prometheus metrics](https://docs.litellm.ai/docs/proxy/prometheus) for the
|
||||
metrics themselves.
|
||||
|
||||
```hcl
|
||||
gateway_metrics_port = 4001
|
||||
gateway_metrics_scrape_cidrs = ["10.0.0.0/16"]
|
||||
```
|
||||
|
||||
## Tenant deployment
|
||||
|
||||
Every resource the stack creates is named `${tenant}-litellm-${env}` (or
|
||||
|
|
|
|||
|
|
@ -212,6 +212,45 @@ locals {
|
|||
# pull the config from S3 first, so the command goes through `sh -c`;
|
||||
# otherwise we keep the image's ENTRYPOINT and only override `command`.
|
||||
gateway_uvicorn_args = "--host 0.0.0.0 --port 4000 --workers ${var.gateway_num_workers}"
|
||||
|
||||
metrics_enabled = var.gateway_metrics_port != null
|
||||
metrics_multiproc_dir = "/tmp/litellm_prometheus_multiproc"
|
||||
metrics_volume = "prometheus-multiproc"
|
||||
metrics_env = local.metrics_enabled ? [{ name = "PROMETHEUS_MULTIPROC_DIR", value = local.metrics_multiproc_dir }] : []
|
||||
metrics_mount_points = local.metrics_enabled ? [{ sourceVolume = local.metrics_volume, containerPath = local.metrics_multiproc_dir }] : []
|
||||
metrics_health_cmd = "import socket; socket.create_connection(('127.0.0.1', ${coalesce(var.gateway_metrics_port, 0)}), timeout=2).close()"
|
||||
|
||||
gateway_metrics_container = local.metrics_enabled ? [
|
||||
{
|
||||
name = "metrics"
|
||||
image = var.gateway_image
|
||||
essential = false
|
||||
entryPoint = ["python", "-m", "litellm.proxy.prometheus_metrics_server"]
|
||||
command = ["--port", tostring(var.gateway_metrics_port)]
|
||||
|
||||
portMappings = [{ containerPort = var.gateway_metrics_port, protocol = "tcp" }]
|
||||
environment = local.metrics_env
|
||||
mountPoints = local.metrics_mount_points
|
||||
|
||||
healthCheck = {
|
||||
command = ["CMD", "python", "-c", local.metrics_health_cmd]
|
||||
interval = 30
|
||||
timeout = 5
|
||||
retries = 3
|
||||
startPeriod = 30
|
||||
}
|
||||
|
||||
logConfiguration = {
|
||||
logDriver = "awslogs"
|
||||
options = {
|
||||
awslogs-group = aws_cloudwatch_log_group.gateway.name
|
||||
awslogs-region = var.region
|
||||
awslogs-stream-prefix = "metrics"
|
||||
}
|
||||
}
|
||||
}
|
||||
] : []
|
||||
|
||||
backend_uvicorn_args = "--host 0.0.0.0 --port 4001"
|
||||
|
||||
gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args};; *) exec uvicorn gateway.main:app ${local.gateway_uvicorn_args};; esac"
|
||||
|
|
@ -269,7 +308,7 @@ resource "aws_ecs_task_definition" "gateway" {
|
|||
execution_role_arn = aws_iam_role.task_execution.arn
|
||||
task_role_arn = aws_iam_role.task.arn
|
||||
|
||||
container_definitions = jsonencode([
|
||||
container_definitions = jsonencode(concat([
|
||||
merge(
|
||||
{
|
||||
name = "gateway"
|
||||
|
|
@ -283,8 +322,10 @@ resource "aws_ecs_task_definition" "gateway" {
|
|||
local.billing_metrics_env,
|
||||
local.gateway_extra_env_list,
|
||||
local.proxy_config_env,
|
||||
local.metrics_env,
|
||||
)
|
||||
secrets = concat(local.shared_secrets, local.gateway_extra_secrets_list)
|
||||
secrets = concat(local.shared_secrets, local.gateway_extra_secrets_list)
|
||||
mountPoints = local.metrics_mount_points
|
||||
|
||||
# Container-level healthCheck intentionally omitted — the wolfi
|
||||
# runtime image doesn't ship curl/wget. The ALB target group polls
|
||||
|
|
@ -301,7 +342,14 @@ resource "aws_ecs_task_definition" "gateway" {
|
|||
},
|
||||
local.gateway_proxy_overrides,
|
||||
)
|
||||
])
|
||||
], local.gateway_metrics_container))
|
||||
|
||||
dynamic "volume" {
|
||||
for_each = local.metrics_enabled ? [1] : []
|
||||
content {
|
||||
name = local.metrics_volume
|
||||
}
|
||||
}
|
||||
|
||||
tags = local.tags
|
||||
}
|
||||
|
|
|
|||
|
|
@ -48,4 +48,7 @@ module "litellm" {
|
|||
backend_extra_env = var.backend_extra_env
|
||||
gateway_extra_secrets = var.gateway_extra_secrets
|
||||
backend_extra_secrets = var.backend_extra_secrets
|
||||
|
||||
gateway_metrics_port = var.gateway_metrics_port
|
||||
gateway_metrics_scrape_cidrs = var.gateway_metrics_scrape_cidrs
|
||||
}
|
||||
|
|
|
|||
|
|
@ -102,6 +102,13 @@ env = "stage"
|
|||
# }
|
||||
# }
|
||||
|
||||
# ---------- Prometheus metrics sidecar ----------
|
||||
# Serve /metrics from a sidecar in the gateway task instead of the inference
|
||||
# workers. The port is not behind the ALB and has no auth: open it only to
|
||||
# your Prometheus subnets.
|
||||
# gateway_metrics_port = 4001
|
||||
# gateway_metrics_scrape_cidrs = ["10.0.0.0/16"]
|
||||
|
||||
# ---------- Extra env / secrets ----------
|
||||
# Plain-text env vars (non-sensitive). Land directly in the ECS task def.
|
||||
# gateway_extra_env = {
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue