mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge branch 'litellm_internal_staging' into litellm_model_update_null_clear
Staging landed a general replica read-back primitive (await_everywhere, replicas_for, read_body_back_everywhere) that subsumes the model-specific one this branch added, so the duplicate polling helpers and their unit tests are dropped and the model lifecycle test reads back through the shared helper instead. Also keeps both blocks of new coverage registry rows and takes staging's DataDog reader, which already contains this branch's credential-hiding change.
This commit is contained in:
commit
1a1ca0eb5a
284 changed files with 21093 additions and 4807 deletions
2
.github/workflows/_test-unit-base.yml
vendored
2
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -113,7 +113,7 @@ jobs:
|
|||
if: steps.changes.outputs.decision != 'skip'
|
||||
timeout-minutes: 8
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra mongodb
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
|
||||
uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]'
|
||||
|
||||
- name: Cache Prisma binaries
|
||||
|
|
|
|||
4
.github/workflows/test-litellm-ui-unit.yml
vendored
4
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -66,7 +66,7 @@ jobs:
|
|||
BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
run: |
|
||||
full_suite() { npm run test -- --run --pool forks --poolOptions.forks.maxForks=14; }
|
||||
full_suite() { npm run test -- --run --pool forks --maxWorkers=14; }
|
||||
|
||||
if [ -z "$BASE_SHA" ]; then
|
||||
echo "Push to $GITHUB_REF_NAME: running the full suite"
|
||||
|
|
@ -95,4 +95,4 @@ jobs:
|
|||
|
||||
echo "Pull request: running tests related to ${#changed_files[@]} changed UI files"
|
||||
npm run test -- related "${changed_files[@]}" --run --passWithNoTests \
|
||||
--pool forks --poolOptions.forks.maxForks=14
|
||||
--pool forks --maxWorkers=14
|
||||
|
|
|
|||
|
|
@ -67,7 +67,6 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
# Copy full source tree
|
||||
|
|
@ -90,7 +89,6 @@ RUN uv sync --frozen --no-default-groups --no-editable \
|
|||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.65"
|
||||
version = "0.1.66"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.65"
|
||||
version = "0.1.66"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -47,7 +47,6 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
# Stage 2 — copy source and install the project + workspace members.
|
||||
|
|
@ -60,7 +59,6 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra bedrock-realtime \
|
||||
--extra mongodb \
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -47,6 +47,7 @@ from typing import (
|
|||
)
|
||||
from litellm.types.integrations.datadog import DatadogInitParams
|
||||
from litellm.types.integrations.newrelic import NewRelicInitParams
|
||||
from litellm.litellm_core_utils.core_helpers import drop_params_env_flag
|
||||
from litellm._logging import (
|
||||
set_verbose,
|
||||
_turn_on_debug,
|
||||
|
|
@ -238,7 +239,7 @@ token: Optional[str] = (
|
|||
)
|
||||
telemetry = True
|
||||
max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults
|
||||
drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False))
|
||||
drop_params = drop_params_env_flag(os.environ, verbose_logger)
|
||||
modify_params = bool(os.getenv("LITELLM_MODIFY_PARAMS", False))
|
||||
use_chat_completions_url_for_anthropic_messages: bool = bool(
|
||||
os.getenv("LITELLM_USE_CHAT_COMPLETIONS_URL_FOR_ANTHROPIC_MESSAGES", False)
|
||||
|
|
@ -325,6 +326,9 @@ ssl_certificate: Optional[str] = None
|
|||
user_url_validation: bool = True
|
||||
user_url_allowed_hosts: List[str] = []
|
||||
provider_url_destination_allowed_hosts: List[str] = []
|
||||
#: "override" (default) or "additive": whether a key or team destination replaces
|
||||
#: the operator's exporter for that backend or exports alongside it.
|
||||
otel_tenant_destination_mode: str | None = None
|
||||
ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance
|
||||
disable_streaming_logging: bool = False
|
||||
disable_token_counter: bool = False
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1374,6 +1374,7 @@ bedrock_embedding_models: Final[set] = set(
|
|||
"cohere.embed-multilingual-v3",
|
||||
"cohere.embed-v4:0",
|
||||
"twelvelabs.marengo-embed-2-7-v1:0",
|
||||
"twelvelabs.marengo-embed-3-0-v1:0",
|
||||
]
|
||||
)
|
||||
|
||||
|
|
@ -1464,6 +1465,7 @@ SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY: Final = "_session_deployment_affin
|
|||
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"
|
||||
|
|
@ -1765,6 +1767,10 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [
|
|||
SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"]
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
|
||||
DEFAULT_ACCESS_GROUP_CACHE_TTL: Final = int(os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600))
|
||||
SPEND_LOG_KEY_METADATA_CACHE_TTL: Final = 600
|
||||
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL: Final = 30
|
||||
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS: Final = 10000
|
||||
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS: Final = 5000
|
||||
# Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated
|
||||
# callers from forcing a DB query per request for unknown names, while bounding
|
||||
# staleness so a transient DB error (which surfaces as an empty list) cannot
|
||||
|
|
|
|||
|
|
@ -45,6 +45,9 @@ from litellm.llms.azure.cost_calculation import (
|
|||
from litellm.llms.azure_ai.cost_calculator import (
|
||||
cost_per_token as azure_ai_cost_per_token,
|
||||
)
|
||||
from litellm.llms.azure_ai.cost_calculator import (
|
||||
is_azure_model_router as azure_ai_is_model_router_name,
|
||||
)
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse
|
||||
from litellm.llms.bedrock.cost_calculation import (
|
||||
cost_per_token as bedrock_cost_per_token,
|
||||
|
|
@ -1659,11 +1662,10 @@ def completion_cost(
|
|||
data_residency=data_residency,
|
||||
vertex_location=vertex_location,
|
||||
response=completion_response,
|
||||
request_model=request_model_for_cost,
|
||||
)
|
||||
|
||||
# Get additional costs from provider (e.g., routing fees, infrastructure costs)
|
||||
if custom_llm_provider == "azure_ai":
|
||||
if custom_llm_provider == "azure_ai" and not azure_ai_is_model_router_name(model):
|
||||
model_for_additional_costs = request_model_for_cost
|
||||
if completion_response is not None:
|
||||
hidden_params = getattr(completion_response, "_hidden_params", None) or {}
|
||||
|
|
|
|||
|
|
@ -367,6 +367,13 @@
|
|||
"ui_name": "Headers",
|
||||
"description": "Headers for OTEL exporter (e.g., x-honeycomb-team=YOUR_API_KEY)",
|
||||
"required": false
|
||||
},
|
||||
"otel_exporter_otlp_protocol": {
|
||||
"type": "select",
|
||||
"ui_name": "Export Protocol",
|
||||
"description": "OTLP wire format for trace exports. Use http/json for collectors that cannot decode protobuf",
|
||||
"options": ["http/protobuf", "http/json"],
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"description": "OpenTelemetry Logging Integration"
|
||||
|
|
|
|||
|
|
@ -133,17 +133,17 @@ class MlflowLogger(CustomLogger):
|
|||
if final_response:
|
||||
end_time_ns: Final = int(end_time.timestamp() * 1e9)
|
||||
|
||||
self._extract_and_set_chat_attributes(span, kwargs, final_response)
|
||||
self._end_span_or_trace(
|
||||
span=span,
|
||||
outputs=final_response,
|
||||
status=SpanStatusCode.OK,
|
||||
end_time_ns=end_time_ns,
|
||||
)
|
||||
|
||||
# Remove the stream_id from the map
|
||||
with self._lock:
|
||||
self._stream_id_to_span.pop(litellm_call_id)
|
||||
try:
|
||||
self._extract_and_set_chat_attributes(span, kwargs, final_response)
|
||||
self._end_span_or_trace(
|
||||
span=span,
|
||||
outputs=final_response,
|
||||
status=SpanStatusCode.OK,
|
||||
end_time_ns=end_time_ns,
|
||||
)
|
||||
finally:
|
||||
with self._lock:
|
||||
self._stream_id_to_span.pop(litellm_call_id, None)
|
||||
|
||||
def _add_chunk_events(self, span, response_obj):
|
||||
from mlflow.entities import SpanEvent
|
||||
|
|
@ -282,15 +282,15 @@ class MlflowLogger(CustomLogger):
|
|||
"""End an MLflow span or a trace."""
|
||||
if span.parent_id is None:
|
||||
self._client.end_trace(
|
||||
trace_id=span.request_id,
|
||||
span.request_id,
|
||||
outputs=outputs,
|
||||
status=status,
|
||||
end_time_ns=end_time_ns,
|
||||
)
|
||||
else:
|
||||
self._client.end_span(
|
||||
trace_id=span.request_id,
|
||||
span_id=span.span_id,
|
||||
span.request_id,
|
||||
span.span_id,
|
||||
outputs=outputs,
|
||||
status=status,
|
||||
end_time_ns=end_time_ns,
|
||||
|
|
|
|||
|
|
@ -16,9 +16,11 @@ from opentelemetry.trace import (
|
|||
Span,
|
||||
Tracer,
|
||||
get_current_span,
|
||||
get_tracer_provider,
|
||||
set_span_in_context,
|
||||
use_span,
|
||||
)
|
||||
from opentelemetry.trace import TracerProvider as ApiTracerProvider
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -63,6 +65,7 @@ from litellm.integrations.otel.plumbing.metrics import (
|
|||
create_genai_metrics,
|
||||
)
|
||||
from litellm.integrations.otel.plumbing.providers import (
|
||||
attach_tenant_fan_out,
|
||||
build_tracer_provider,
|
||||
get_event_logger,
|
||||
get_meter,
|
||||
|
|
@ -85,6 +88,7 @@ if TYPE_CHECKING:
|
|||
)
|
||||
|
||||
LITELLM_TRACER_NAME: Final = "litellm"
|
||||
_published_v2_provider: ApiTracerProvider | None = None
|
||||
|
||||
|
||||
def _span_error_from_exception(
|
||||
|
|
@ -180,7 +184,9 @@ class OpenTelemetryV2(CustomLogger):
|
|||
self.config: OpenTelemetryV2Config = config or OpenTelemetryV2Config(**kwargs)
|
||||
self.callback_name = callback_name
|
||||
self._tracer_provider: TracerProvider = (
|
||||
tracer_provider if tracer_provider is not None else build_tracer_provider(self.config)
|
||||
tracer_provider
|
||||
if tracer_provider is not None
|
||||
else build_tracer_provider(self.config, tenant_overrides=True)
|
||||
)
|
||||
self.tracer: Tracer = get_tracer(self._tracer_provider, LITELLM_TRACER_NAME)
|
||||
self._metrics_recorder = self._init_metrics(meter_provider)
|
||||
|
|
@ -195,6 +201,11 @@ class OpenTelemetryV2(CustomLogger):
|
|||
self._open_llm_calls: OrderedDict[str, _LLMCallSpan] = OrderedDict()
|
||||
self._init_otel_logger_on_litellm_proxy()
|
||||
|
||||
@property
|
||||
def tracer_provider(self) -> TracerProvider:
|
||||
"""The provider this logger emits through, read-only to its callers."""
|
||||
return self._tracer_provider
|
||||
|
||||
def _init_metrics(self, meter_provider: "MeterProvider | None") -> "GenAIMetricRecorder | None":
|
||||
"""Create the six GenAI histograms when metrics are enabled, else ``None``.
|
||||
|
||||
|
|
@ -863,12 +874,33 @@ def publish_global_otel_v2_provider(
|
|||
``opentelemetry.trace.set_tracer_provider``) are injected so the publish step is
|
||||
unit-testable without reading or mutating real global OTel state. Returns the
|
||||
logger whose provider was published.
|
||||
|
||||
The published provider is also the one that fans spans out to key/team
|
||||
destinations, because it is the only provider the whole request tree passes
|
||||
through; see :func:`attach_tenant_fan_out`. It is remembered for
|
||||
:func:`fan_out_provider` because neither the OTel global (``set_tracer_provider``
|
||||
keeps the first provider it was ever handed) nor
|
||||
``proxy_server.open_telemetry_logger`` (a legacy v1 logger can hold that slot)
|
||||
reliably leads back to it.
|
||||
"""
|
||||
global _published_v2_provider
|
||||
logger: Final = select_global_otel_v2_logger(in_memory_loggers, registered=registered)
|
||||
set_global_provider(logger._tracer_provider)
|
||||
attach_tenant_fan_out(logger.tracer_provider, *_v2_configs(in_memory_loggers, logger))
|
||||
set_global_provider(logger.tracer_provider)
|
||||
_published_v2_provider = logger.tracer_provider # rebind-ok: startup records the one provider carrying the fan-out
|
||||
return logger
|
||||
|
||||
|
||||
def _v2_configs(in_memory_loggers: Sequence[object], logger: "OpenTelemetryV2") -> tuple[OpenTelemetryV2Config, ...]:
|
||||
"""Every v2 logger's config, the published logger's first.
|
||||
|
||||
Each preset keeps its own provider and exporters, so the accounts the operator
|
||||
writes to are spread over all of them, not held by the published logger alone.
|
||||
"""
|
||||
others: Final = tuple(cb.config for cb in in_memory_loggers if isinstance(cb, OpenTelemetryV2) and cb is not logger)
|
||||
return (logger.config, *others)
|
||||
|
||||
|
||||
def _registered_v2_logger() -> "OpenTelemetryV2 | None":
|
||||
try:
|
||||
from litellm.proxy import proxy_server
|
||||
|
|
@ -904,6 +936,25 @@ def seed_request_identity(user_api_key_dict: object, model: str | None = None) -
|
|||
logger.seed_request_identity(user_api_key_dict, model=model)
|
||||
|
||||
|
||||
def fan_out_provider() -> ApiTracerProvider:
|
||||
"""The provider :func:`publish_global_otel_v2_provider` gave the tenant fan-out.
|
||||
|
||||
Read off the publish itself, not the OTel global and not the registered logger:
|
||||
the global keeps whichever provider claimed it first (auto-instrumentation, a
|
||||
legacy logger), and the registered slot can hold a v1 logger while the publish
|
||||
picked a v2 one from ``_in_memory_loggers``. Either detour lands on a provider
|
||||
with no fan-out and drops every destination at auth.
|
||||
"""
|
||||
published: Final = _published_v2_provider
|
||||
if published is not None:
|
||||
return published
|
||||
logger: Final = _registered_v2_logger()
|
||||
if logger is not None:
|
||||
attach_tenant_fan_out(logger.tracer_provider, logger.config)
|
||||
return logger.tracer_provider
|
||||
return get_tracer_provider()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def phase_span(name: str) -> "Iterator[Span | None]":
|
||||
logger: Final = _registered_v2_logger()
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from litellm.integrations.otel.model.payloads import (
|
|||
ServiceSpanData,
|
||||
ToolDefinition,
|
||||
)
|
||||
from litellm.integrations.otel.model.semconv import Error
|
||||
|
||||
# Attribute keys in the semconv-ai / Traceloop vocabulary.
|
||||
_LEGACY_SYSTEM: Final = "gen_ai.system"
|
||||
|
|
@ -36,7 +37,7 @@ _LEGACY_PRESENCE_PENALTY: Final = "llm.presence_penalty"
|
|||
_LEGACY_STOP_SEQUENCES: Final = "llm.chat.stop_sequences"
|
||||
_LEGACY_SERVICE: Final = "service"
|
||||
_LEGACY_CALL_TYPE: Final = "call_type"
|
||||
_LEGACY_ERROR: Final = "error"
|
||||
_LEGACY_ERROR: Final = Error.MESSAGE_LEGACY
|
||||
|
||||
|
||||
class LegacyMapper:
|
||||
|
|
|
|||
|
|
@ -69,7 +69,7 @@ class ExporterSpec(BaseModel):
|
|||
|
||||
kind: str = Field(
|
||||
default="console",
|
||||
description="console | in_memory | otlp_http | otlp_grpc | <factory kind>",
|
||||
description="console | in_memory | otlp_http | http/json | otlp_grpc | <factory kind>",
|
||||
)
|
||||
endpoint: str | None = None
|
||||
traces_endpoint: str | None = Field(
|
||||
|
|
@ -269,7 +269,9 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
if (self.endpoint or self.traces_endpoint) and self.exporter == "console":
|
||||
self.exporter = "otlp_http"
|
||||
# When no explicit destinations are given, fold the single-destination
|
||||
# shorthand into one spec so the provider always has a destination.
|
||||
# shorthand into one spec so the provider always has a destination. A spec
|
||||
# with no fields set is how the presets tell "nothing configured" from an
|
||||
# operator who asked for the console by name.
|
||||
if not self.exporters:
|
||||
self.exporters = [
|
||||
ExporterSpec(
|
||||
|
|
@ -278,6 +280,8 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
traces_endpoint=self.traces_endpoint,
|
||||
headers=self.headers,
|
||||
)
|
||||
if not self.model_fields_set.isdisjoint(("exporter", "endpoint", "headers"))
|
||||
else ExporterSpec()
|
||||
]
|
||||
# Ensure ``genai`` is always present and first.
|
||||
names = list(self.mapper_names)
|
||||
|
|
|
|||
49
litellm/integrations/otel/model/destination.py
Normal file
49
litellm/integrations/otel/model/destination.py
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
"""The resolved OTLP destination a request's traces export to.
|
||||
|
||||
Backend-agnostic on purpose: every OTEL backend reduces to an endpoint plus auth
|
||||
headers. The per-backend field mapping lives in ``presets.destinations``.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
from urllib.parse import quote
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class OtelDestination(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
endpoint: str
|
||||
headers: Mapping[str, str] = Field(default_factory=dict)
|
||||
resource_attributes: Mapping[str, str] = Field(default_factory=dict)
|
||||
callback_name: str | None = None
|
||||
protocol: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"OTLP transport, defaulting to the backend's own. Not derivable from the "
|
||||
"scheme: Arize's ``https://otlp.arize.com/v1`` is gRPC."
|
||||
),
|
||||
)
|
||||
|
||||
def header_string(self) -> str:
|
||||
"""Render headers as the ``k=v,k2=v2`` form an ``ExporterSpec`` expects.
|
||||
|
||||
Values are percent-encoded because ``providers.parse_headers`` decodes them
|
||||
with the SDK's W3C-Baggage parser: a value carrying a ``,`` or ``=`` (a
|
||||
Langfuse project name, a base64 Authorization payload ending in ``==``)
|
||||
would otherwise be split into bogus pairs on the way back out.
|
||||
"""
|
||||
return ",".join(f"{key}={quote(value, safe='')}" for key, value in self.headers.items())
|
||||
|
||||
def cache_key(self) -> tuple[str, tuple[tuple[str, str], ...], tuple[tuple[str, str], ...], str | None]:
|
||||
"""Identity for processor reuse, so one destination means one exporter."""
|
||||
return (
|
||||
self.endpoint,
|
||||
tuple(sorted(self.headers.items())),
|
||||
tuple(sorted(self.resource_attributes.items())),
|
||||
self.protocol,
|
||||
)
|
||||
|
||||
|
||||
NO_DESTINATIONS: Final[tuple[OtelDestination, ...]] = ()
|
||||
|
|
@ -204,6 +204,9 @@ class Error:
|
|||
|
||||
TYPE: Final = "error.type"
|
||||
MESSAGE: Final = "error.message"
|
||||
# The same text under the bare key the semconv-ai / Traceloop vocabulary uses
|
||||
# (see ``LegacyMapper``), so anything reading or redacting error text covers both.
|
||||
MESSAGE_LEGACY: Final = "error"
|
||||
|
||||
|
||||
class LiteLLMError:
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
"""Trace-context + Baggage helpers."""
|
||||
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from contextvars import ContextVar, Token
|
||||
from typing import Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from opentelemetry import baggage
|
||||
from opentelemetry.context import Context, get_current
|
||||
|
|
@ -21,6 +22,9 @@ from opentelemetry.trace.propagation.tracecontext import (
|
|||
|
||||
from litellm.integrations.otel.model.semconv import HTTP
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.otel.model.destination import OtelDestination
|
||||
|
||||
_PROPAGATOR: Final = TraceContextTextMapPropagator()
|
||||
|
||||
# The request's root span — the FastAPI-owned SERVER span — captured ONCE when the
|
||||
|
|
@ -304,3 +308,65 @@ def extract_traceparent(headers: Mapping[str, str]) -> Context | None:
|
|||
return None
|
||||
carrier: Final = {str(key).lower(): value for key, value in headers.items()}
|
||||
return _PROPAGATOR.extract(carrier)
|
||||
|
||||
|
||||
# The OTLP destinations this request's key or team pointed its traces at, resolved
|
||||
# once during auth. A ``ContextVar`` for the same reason the root span above is one:
|
||||
# it rides the request task's context into the ``asyncio.create_task`` children that
|
||||
# close the LLM span, and it is visible to every ``SpanProcessor.on_end`` that fires
|
||||
# on the request task. Stateful MCP handlers set and reset it per message; the
|
||||
# request-task value otherwise dies with that task.
|
||||
_request_destinations: Final['ContextVar[tuple["OtelDestination", ...]]'] = ContextVar(
|
||||
"litellm_otel_request_destinations", default=()
|
||||
)
|
||||
|
||||
|
||||
def set_request_destinations(destinations: 'tuple["OtelDestination", ...]') -> "Token[tuple[OtelDestination, ...]]":
|
||||
"""Anchor the destinations this request exports to and return a reset token."""
|
||||
return _request_destinations.set(destinations)
|
||||
|
||||
|
||||
def reset_request_destinations(token: "Token[tuple[OtelDestination, ...]]") -> None:
|
||||
_request_destinations.reset(token)
|
||||
|
||||
|
||||
def request_destinations() -> 'tuple["OtelDestination", ...]':
|
||||
"""The destinations resolved for this request, empty outside a proxy request."""
|
||||
return _request_destinations.get()
|
||||
|
||||
|
||||
#: ``litellm_settings: otel_tenant_destination_mode`` and its env equivalent.
|
||||
ADDITIVE_DESTINATION_MODE: Final = "additive"
|
||||
OTEL_TENANT_DESTINATION_MODE_ENV: Final = "LITELLM_OTEL_TENANT_DESTINATION_MODE"
|
||||
|
||||
|
||||
def tenant_destinations_are_additive() -> bool:
|
||||
"""Whether a tenant destination exports alongside the operator's own exporter.
|
||||
|
||||
Override is the default: the tenant's traffic reaches the tenant's account and
|
||||
nowhere else. Operators running one org-wide backend across every team set this
|
||||
to ``additive`` so the same trace lands in both places.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
configured: Final = litellm.otel_tenant_destination_mode or os.environ.get(OTEL_TENANT_DESTINATION_MODE_ENV)
|
||||
return isinstance(configured, str) and configured.strip().lower() == ADDITIVE_DESTINATION_MODE
|
||||
|
||||
|
||||
def destination_backends() -> frozenset[str]:
|
||||
"""Backends this request resolved a tenant destination for.
|
||||
|
||||
The fan-out already carries the whole trace to those destinations, so the
|
||||
per-request tracer route must never send a second copy, in either mode.
|
||||
"""
|
||||
return frozenset(d.callback_name for d in _request_destinations.get() if d.callback_name)
|
||||
|
||||
|
||||
def suppressed_backends() -> frozenset[str]:
|
||||
"""Backends whose operator-level exporters this request must NOT reach.
|
||||
|
||||
Empty under ``additive``, where the operator keeps its copy of every span.
|
||||
"""
|
||||
if tenant_destinations_are_additive():
|
||||
return frozenset()
|
||||
return destination_backends()
|
||||
|
|
|
|||
70
litellm/integrations/otel/plumbing/otlp_json.py
Normal file
70
litellm/integrations/otel/plumbing/otlp_json.py
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
"""OTLP/HTTP span exporter that sends the OTLP/JSON encoding instead of protobuf.
|
||||
|
||||
The SDK only ships a protobuf OTLP/HTTP exporter; this reuses its transport and
|
||||
retry loop and swaps the payload for OTLP/JSON (enums as integers, ids as hex).
|
||||
"""
|
||||
|
||||
import base64
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
from google.protobuf.json_format import MessageToDict
|
||||
from opentelemetry.exporter.otlp.proto.common.trace_encoder import encode_spans
|
||||
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
|
||||
from opentelemetry.sdk.trace import ReadableSpan
|
||||
|
||||
JSON_CONTENT_TYPE: Final = "application/json"
|
||||
_HEX_ID_KEYS: Final = frozenset({"traceId", "spanId", "parentSpanId"})
|
||||
|
||||
_JsonValue: TypeAlias = "Mapping[str, _JsonValue] | Sequence[_JsonValue] | str | int | float | bool | None"
|
||||
_JsonObject: TypeAlias = Mapping[str, "_JsonValue"]
|
||||
|
||||
|
||||
def _objects(node: _JsonObject, key: str) -> tuple[_JsonObject, ...]:
|
||||
items: Final = node.get(key)
|
||||
if isinstance(items, str) or not isinstance(items, Sequence):
|
||||
return ()
|
||||
return tuple(item for item in items if isinstance(item, Mapping))
|
||||
|
||||
|
||||
def _hex_ids(node: _JsonObject) -> _JsonObject:
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: base64.b64decode(item).hex() if key in _HEX_ID_KEYS and isinstance(item, str) else item
|
||||
for key, item in node.items()
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _hex_span(span: _JsonObject) -> _JsonObject:
|
||||
links: Final = _objects(span, "links")
|
||||
if not links:
|
||||
return _hex_ids(span)
|
||||
return MappingProxyType({**_hex_ids(span), "links": tuple(_hex_ids(link) for link in links)})
|
||||
|
||||
|
||||
def _hex_scope_spans(scope: _JsonObject) -> _JsonObject:
|
||||
return MappingProxyType({**scope, "spans": tuple(_hex_span(span) for span in _objects(scope, "spans"))})
|
||||
|
||||
|
||||
def _hex_resource_spans(resource: _JsonObject) -> _JsonObject:
|
||||
scope_spans: Final = tuple(_hex_scope_spans(scope) for scope in _objects(resource, "scopeSpans"))
|
||||
return MappingProxyType({**resource, "scopeSpans": scope_spans})
|
||||
|
||||
|
||||
def encode_spans_json(spans: Sequence[ReadableSpan]) -> bytes:
|
||||
payload: Final[_JsonObject] = MessageToDict(encode_spans(spans), use_integers_for_enums=True)
|
||||
resource_spans: Final = tuple(_hex_resource_spans(resource) for resource in _objects(payload, "resourceSpans"))
|
||||
hexed: Final[_JsonObject] = MappingProxyType({**payload, "resourceSpans": resource_spans})
|
||||
return json.dumps(hexed, default=dict, separators=(",", ":")).encode()
|
||||
|
||||
|
||||
class OTLPJsonSpanExporter(OTLPSpanExporter):
|
||||
def __init__(self, endpoint: str | None, headers: dict[str, str]) -> None: # mutable-ok: SDK __init__ takes Dict
|
||||
super().__init__(endpoint=endpoint, headers=headers)
|
||||
self._session.headers["Content-Type"] = JSON_CONTENT_TYPE
|
||||
|
||||
def _serialize_spans(self, spans: Sequence[ReadableSpan]) -> bytes:
|
||||
return encode_spans_json(spans)
|
||||
|
|
@ -1,9 +1,14 @@
|
|||
"""Provider / exporter factory + the Baggage span processor."""
|
||||
|
||||
from collections.abc import Callable, Iterable
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
from opentelemetry import _logs, baggage, metrics
|
||||
from opentelemetry import _logs, baggage, metrics, trace
|
||||
from opentelemetry._events import EventLogger
|
||||
from opentelemetry._logs import LoggerProvider, NoOpLoggerProvider
|
||||
from opentelemetry.context import Context
|
||||
|
|
@ -19,7 +24,8 @@ from opentelemetry.sdk._logs.export import (
|
|||
)
|
||||
from opentelemetry.sdk.metrics import MeterProvider as SDKMeterProvider
|
||||
from opentelemetry.sdk.resources import Resource
|
||||
from opentelemetry.sdk.trace import ReadableSpan, SpanProcessor, TracerProvider
|
||||
from opentelemetry.sdk.trace import Event, ReadableSpan, SpanProcessor, TracerProvider
|
||||
from opentelemetry.sdk.trace import Span as SDKSpan
|
||||
from opentelemetry.sdk.trace.export import (
|
||||
BatchSpanProcessor,
|
||||
ConsoleSpanExporter,
|
||||
|
|
@ -29,18 +35,35 @@ from opentelemetry.sdk.trace.export import (
|
|||
from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
|
||||
InMemorySpanExporter,
|
||||
)
|
||||
from opentelemetry.trace import Span, SpanKind, Tracer
|
||||
from opentelemetry.trace import Span, SpanKind, Status, Tracer
|
||||
from opentelemetry.util.re import parse_env_headers
|
||||
from opentelemetry.util.types import Attributes, AttributeValue
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._version import version as litellm_version
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.semconv import LiteLLM
|
||||
from litellm.integrations.otel.model.semconv import (
|
||||
DB,
|
||||
MCP,
|
||||
Error,
|
||||
ExceptionEvent,
|
||||
GenAI,
|
||||
LiteLLM,
|
||||
LiteLLMError,
|
||||
Server,
|
||||
)
|
||||
from litellm.integrations.otel.model.spans import LiteLLMSpanKind
|
||||
from litellm.integrations.otel.plumbing.context import (
|
||||
request_destinations,
|
||||
suppressed_backends,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.metrics import Meter
|
||||
from opentelemetry.sdk.metrics.export import MetricReader
|
||||
|
||||
from litellm.integrations.otel.model.destination import OtelDestination
|
||||
|
||||
_SPAN_KIND_BY_ROLE_KIND: Final[dict[LiteLLMSpanKind, SpanKind]] = {
|
||||
LiteLLMSpanKind.SERVER: SpanKind.SERVER,
|
||||
LiteLLMSpanKind.CLIENT: SpanKind.CLIENT,
|
||||
|
|
@ -136,7 +159,8 @@ def parse_headers(raw: str | None) -> dict[str, str]:
|
|||
|
||||
|
||||
_IN_MEMORY_KINDS: Final = ("in_memory", "inmemory", "memory")
|
||||
_OTLP_HTTP_KINDS: Final = ("otlp_http", "http", "http/protobuf", "http/json")
|
||||
_OTLP_HTTP_JSON_KINDS: Final = ("http/json",)
|
||||
_OTLP_HTTP_KINDS: Final = ("otlp_http", "http", "http/protobuf", *_OTLP_HTTP_JSON_KINDS)
|
||||
_OTLP_GRPC_KINDS: Final = ("otlp_grpc", "grpc")
|
||||
|
||||
|
||||
|
|
@ -164,6 +188,13 @@ def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter:
|
|||
return factory(spec)
|
||||
if kind in _IN_MEMORY_KINDS:
|
||||
return InMemorySpanExporter()
|
||||
if kind in _OTLP_HTTP_JSON_KINDS:
|
||||
from litellm.integrations.otel.plumbing.otlp_json import OTLPJsonSpanExporter
|
||||
|
||||
return OTLPJsonSpanExporter(
|
||||
endpoint=spec.traces_endpoint or _otlp_traces_endpoint(spec.endpoint),
|
||||
headers=parse_headers(spec.headers),
|
||||
)
|
||||
if kind in _OTLP_HTTP_KINDS:
|
||||
from opentelemetry.exporter.otlp.proto.http.trace_exporter import (
|
||||
OTLPSpanExporter as HTTPExporter,
|
||||
|
|
@ -194,6 +225,555 @@ def _processor_for(exporter: SpanExporter, use_simple: bool | None) -> SpanProce
|
|||
return SimpleSpanProcessor(exporter) if use_simple else BatchSpanProcessor(exporter)
|
||||
|
||||
|
||||
#: Distinct tenant destinations whose exporters stay alive. Each holds a connection
|
||||
#: pool and a batch thread, so the cache is bounded and evicts least-recently-used.
|
||||
_MAX_CACHED_DESTINATION_PROCESSORS: Final = 32
|
||||
|
||||
#: Workers closing shed destination processors, bounding the threads a tenant can
|
||||
#: create by cycling its destination config.
|
||||
_DRAIN_WORKERS: Final = 2
|
||||
|
||||
#: Shed processors waiting to be closed before the fan-out stops building new ones.
|
||||
#: Each still owns a batch thread until its close returns, and a collector that never
|
||||
#: answers makes every close take the exporter's full timeout, so past this many the
|
||||
#: operator's exporter keeps the span instead (see ``deliverable``).
|
||||
_MAX_PENDING_DRAINS: Final = 64
|
||||
|
||||
#: How long ``shutdown`` waits for spans already being forwarded, so teardown closes
|
||||
#: no processor under one. Bounded: an exporter that never returns must not hold the
|
||||
#: proxy open.
|
||||
_SHUTDOWN_DRAIN_SECONDS: Final = 5.0
|
||||
|
||||
#: An exporter's account: its normalized endpoint and the credentials it presents.
|
||||
_SinkKey = tuple[str, tuple[tuple[str, str], ...]]
|
||||
|
||||
#: Header names that spell one credential two ways. Arize's operator exporter sends
|
||||
#: ``space_id`` where a tenant destination sends ``arize-space-id``.
|
||||
_CREDENTIAL_ALIASES: Final = MappingProxyType({"arize_space_id": "space_id"})
|
||||
|
||||
|
||||
class _DrainPool:
|
||||
"""Closes shed destination processors off the span-export path.
|
||||
|
||||
``shutdown`` flushes over the network and is reached from ``on_end``, so closing
|
||||
one inline would let a single unreachable tenant collector stall every other
|
||||
tenant's spans behind it. A fixed set of workers rather than a thread per
|
||||
processor means a tenant cycling its destination config cannot spawn threads as
|
||||
fast as it can send requests; slow shutdowns queue behind each other.
|
||||
|
||||
The workers are daemons and belong to the fan-out that sheds the processors, so
|
||||
neither an unreachable collector nor a lazily built process-wide singleton can
|
||||
hold the proxy open on the way down.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
workers: int = _DRAIN_WORKERS,
|
||||
pending: "queue.Queue[SpanProcessor | None] | None" = None,
|
||||
capacity: int = _MAX_PENDING_DRAINS,
|
||||
) -> None:
|
||||
self._workers: Final = workers
|
||||
self._capacity: Final = capacity
|
||||
self._lock: Final = threading.Lock()
|
||||
self._closed = False
|
||||
self._backlog = 0 # guarded by ``_lock``: submitted processors whose close has not returned
|
||||
self._pending: Final[queue.Queue[SpanProcessor | None]] = pending if pending is not None else queue.Queue()
|
||||
self._threads: Final = tuple(
|
||||
threading.Thread(target=self._drain_until_closed, daemon=True, name="litellm-otel-destination-drain")
|
||||
for _ in range(workers)
|
||||
)
|
||||
for worker in self._threads:
|
||||
worker.start()
|
||||
|
||||
def submit(self, processor: SpanProcessor) -> None:
|
||||
"""Queue ``processor`` for closing, or hand it off once the pool is retired.
|
||||
|
||||
The check and the put share one lock. Reading a closed flag on its own leaves
|
||||
room for :meth:`close` to run in between, and the processor would land behind
|
||||
the sentinels every worker has already exited on.
|
||||
|
||||
Past close there is no worker left to take it, and the caller is whichever
|
||||
thread just ended a span, so closing it inline would park that thread on a
|
||||
network flush the shutdown deadline has already stopped waiting for. The extra
|
||||
thread is bounded by the same close: the fan-out stops handing processors out
|
||||
at that point, so only the ones already exporting when it happened arrive here.
|
||||
"""
|
||||
with self._lock:
|
||||
if not self._closed:
|
||||
self._backlog += 1
|
||||
self._pending.put(processor)
|
||||
return
|
||||
threading.Thread(
|
||||
target=_shutdown_quietly,
|
||||
args=(processor,),
|
||||
daemon=True,
|
||||
name="litellm-otel-destination-drain-straggler",
|
||||
).start()
|
||||
|
||||
def saturated(self) -> bool:
|
||||
"""Whether enough closes are outstanding that building another processor must wait.
|
||||
|
||||
The workers close in order and each close blocks for as long as its exporter
|
||||
does, so a collector that stopped answering would otherwise turn every new
|
||||
destination into one more batch thread parked behind them, for as long as the
|
||||
tenants keep rotating. Holding the count here rather than reading the queue
|
||||
keeps the two processors a worker is mid-close on in the total.
|
||||
"""
|
||||
with self._lock:
|
||||
return self._backlog >= self._capacity
|
||||
|
||||
def close(self, timeout: float | None = None) -> None:
|
||||
"""Retire the workers once they have closed everything already queued.
|
||||
|
||||
A proxy that rebuilds its telemetry builds another fan-out, so workers that
|
||||
outlive the one that started them are two more threads per reload, forever.
|
||||
|
||||
``timeout`` bounds how long the caller waits for that draining to finish. The
|
||||
workers are daemons, so whatever is still flushing when it expires is dropped
|
||||
by the interpreter rather than holding it open.
|
||||
"""
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
for _ in range(self._workers):
|
||||
self._pending.put(None)
|
||||
if timeout is None:
|
||||
return
|
||||
deadline: Final = time.monotonic() + timeout
|
||||
for worker in self._threads:
|
||||
worker.join(timeout=max(0.0, deadline - time.monotonic()))
|
||||
|
||||
def _drain_until_closed(self) -> None:
|
||||
while True:
|
||||
processor: SpanProcessor | None = self._pending.get() # rebind-ok: loop variable
|
||||
if processor is None:
|
||||
return
|
||||
_shutdown_quietly(processor)
|
||||
with self._lock:
|
||||
self._backlog -= 1
|
||||
|
||||
|
||||
_NO_ATTRIBUTES: Final[Mapping[str, AttributeValue]] = MappingProxyType({})
|
||||
_DB_SYSTEM_KEYS: Final = frozenset({DB.SYSTEM_NAME, DB.SYSTEM_LEGACY})
|
||||
# Keys on a database span that describe the proxy's own datastore: its host, its
|
||||
# port, and its schema.
|
||||
_DATASTORE_ENDPOINT_KEYS: Final = frozenset({Server.ADDRESS, Server.PORT, DB.NAMESPACE})
|
||||
# A span carrying one of these describes the tenant's own call (the model call, the
|
||||
# MCP call, the guardrail), so its error text is theirs to see. Every other span is
|
||||
# the proxy's own work, whose error text names the operator's infrastructure.
|
||||
_TENANT_OWNED_KEYS: Final = frozenset({GenAI.OPERATION_NAME, MCP.METHOD_NAME, LiteLLM.GUARDRAIL_NAME})
|
||||
_PROXY_ERROR_TEXT_KEYS: Final = frozenset({Error.MESSAGE, Error.MESSAGE_LEGACY})
|
||||
# A guardrail that never answered carries the exception it raised as its response,
|
||||
# which names the operator's guardrail endpoint. The second spelling is the legacy
|
||||
# status the request-level logger still maps.
|
||||
_GUARDRAIL_UNREACHABLE_STATUSES: Final = frozenset({"guardrail_failed_to_respond", "failure"})
|
||||
# Attribute prefixes the FastAPI instrumentor uses for headers the operator opted to
|
||||
# capture (``OTEL_INSTRUMENTATION_HTTP_CAPTURE_HEADERS_SERVER_*``). The request
|
||||
# side carries the caller's bearer token verbatim.
|
||||
_CAPTURED_HEADER_PREFIXES: Final = ("http.request.header.", "http.response.header.")
|
||||
# The instrumentor stamps the request URL on the server span with its query string,
|
||||
# under the old convention and the new one, and litellm accepts a virtual key as a
|
||||
# ``?key=`` query parameter.
|
||||
_URL_KEYS: Final = frozenset({"http.url", "http.target", "url.full"})
|
||||
_URL_QUERY_KEY: Final = "url.query"
|
||||
|
||||
|
||||
class _TenantSpanView(ReadableSpan):
|
||||
"""A ``ReadableSpan`` view for one destination, leaving the operator's own span alone."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
inner: ReadableSpan,
|
||||
resource: Resource,
|
||||
attributes: Attributes,
|
||||
events: Sequence[Event],
|
||||
status: Status,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
name=inner.name,
|
||||
context=inner.context,
|
||||
parent=inner.parent,
|
||||
resource=resource,
|
||||
attributes=attributes,
|
||||
events=events,
|
||||
links=inner.links,
|
||||
kind=inner.kind,
|
||||
status=status,
|
||||
start_time=inner.start_time,
|
||||
end_time=inner.end_time,
|
||||
instrumentation_scope=inner.instrumentation_scope,
|
||||
)
|
||||
|
||||
|
||||
def _is_database_span(attributes: Mapping[str, AttributeValue]) -> bool:
|
||||
return any(key in attributes for key in _DB_SYSTEM_KEYS)
|
||||
|
||||
|
||||
def _is_tenant_owned_span(attributes: Mapping[str, AttributeValue]) -> bool:
|
||||
return any(key in attributes for key in _TENANT_OWNED_KEYS)
|
||||
|
||||
|
||||
def _guardrail_unreachable(attributes: Mapping[str, AttributeValue]) -> bool:
|
||||
return attributes.get(LiteLLM.GUARDRAIL_STATUS) in _GUARDRAIL_UNREACHABLE_STATUSES
|
||||
|
||||
|
||||
def _tenant_visible(key: str, database: bool, owned: bool, unreachable_guardrail: bool) -> bool:
|
||||
if key.startswith(_CAPTURED_HEADER_PREFIXES) or key in (LiteLLMError.STACK_TRACE, _URL_QUERY_KEY):
|
||||
return False
|
||||
if database and key in _DATASTORE_ENDPOINT_KEYS:
|
||||
return False
|
||||
if unreachable_guardrail and key == LiteLLM.GUARDRAIL_RESPONSE:
|
||||
return False
|
||||
return owned or key not in _PROXY_ERROR_TEXT_KEYS
|
||||
|
||||
|
||||
def _without_query(key: str, value: AttributeValue) -> AttributeValue:
|
||||
if key not in _URL_KEYS or not isinstance(value, str):
|
||||
return value
|
||||
return value.partition("?")[0]
|
||||
|
||||
|
||||
def _same_attributes(kept: Mapping[str, AttributeValue], attributes: Mapping[str, AttributeValue]) -> bool:
|
||||
return len(kept) == len(attributes) and all(kept[key] is value for key, value in attributes.items())
|
||||
|
||||
|
||||
def _without_stack_trace(event: Event) -> Event:
|
||||
attributes: Final = event.attributes or _NO_ATTRIBUTES
|
||||
if ExceptionEvent.STACKTRACE not in attributes:
|
||||
return event
|
||||
return Event(
|
||||
name=event.name,
|
||||
attributes=MappingProxyType(
|
||||
{key: value for key, value in attributes.items() if key != ExceptionEvent.STACKTRACE}
|
||||
),
|
||||
timestamp=event.timestamp,
|
||||
)
|
||||
|
||||
|
||||
def _for_destination(span: ReadableSpan, destination: "OtelDestination") -> ReadableSpan:
|
||||
"""The view of ``span`` a tenant destination receives.
|
||||
|
||||
A span the tenant's own call produced keeps its error text. Every other span is
|
||||
the proxy's own work (the request root, auth, the database), and its error text,
|
||||
its events and its status description come off, since a Prisma failure there
|
||||
spells out the operator's Postgres endpoint. A database span loses that endpoint
|
||||
too, and a guardrail that failed to respond loses its response text, which is the
|
||||
exception it raised and names the operator's guardrail endpoint. Stack traces walk
|
||||
the operator's install and come off every span, as do the headers the operator
|
||||
captures on the server span, whose request side holds the caller's bearer token,
|
||||
and the query string of the request URL, which can hold the same key. The span
|
||||
itself stays, so the tenant still gets the whole trace tree.
|
||||
"""
|
||||
extra: Final = destination.resource_attributes
|
||||
attributes: Final = span.attributes or _NO_ATTRIBUTES
|
||||
database: Final = _is_database_span(attributes)
|
||||
owned: Final = _is_tenant_owned_span(attributes)
|
||||
unreachable: Final = _guardrail_unreachable(attributes)
|
||||
kept: Final = MappingProxyType(
|
||||
{
|
||||
key: _without_query(key, value)
|
||||
for key, value in attributes.items()
|
||||
if _tenant_visible(key, database, owned, unreachable)
|
||||
}
|
||||
)
|
||||
recorded: Final = span.events
|
||||
events: Final = tuple(_without_stack_trace(event) for event in recorded) if owned else ()
|
||||
unchanged: Final = owned and _same_attributes(kept, attributes) and all(a is b for a, b in zip(events, recorded))
|
||||
if not extra and unchanged:
|
||||
return span
|
||||
resource: Final = span.resource.merge(Resource(extra)) if extra else span.resource
|
||||
status: Final = span.status if owned else Status(span.status.status_code)
|
||||
return _TenantSpanView(span, resource, kept, events, status)
|
||||
|
||||
|
||||
class TenantFanOutSpanProcessor(SpanProcessor):
|
||||
"""Export every finished span to each destination this request resolved.
|
||||
|
||||
Destinations ride a request-scoped ``ContextVar`` set during auth, so concurrent
|
||||
requests stay isolated. The forwarded view keeps the original trace and parent
|
||||
ids, so the tenant gets the same tree the operator would have received.
|
||||
|
||||
Exactly one provider carries this processor, the one published as the OTel global
|
||||
(see :func:`attach_tenant_fan_out`). That provider is the only one every span
|
||||
passes through: the FastAPI server span, the auth span and the post-call database
|
||||
spans are emitted on the global, while a second v2 logger's provider sees only
|
||||
that logger's own gen-AI span. Attaching the fan-out per logger would hand a
|
||||
tenant a one-span trace whenever its backend is not the global one, and two
|
||||
copies of the model call whenever it is.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
processor_factory: 'Callable[["OtelDestination"], SpanProcessor | None] | None' = None,
|
||||
shutdown_drain_seconds: float = _SHUTDOWN_DRAIN_SECONDS,
|
||||
operator_sinks: frozenset[_SinkKey] = frozenset(),
|
||||
pending_drains: int = _MAX_PENDING_DRAINS,
|
||||
drain_pool: _DrainPool | None = None,
|
||||
) -> None:
|
||||
self._operator_sinks: Final = operator_sinks
|
||||
self._drain_seconds: Final = shutdown_drain_seconds
|
||||
self._lock: Final = threading.Condition()
|
||||
self._closed = False # guarded by ``_lock``: an unlocked read races the teardown it gates
|
||||
self._build: Final = processor_factory if processor_factory is not None else _destination_processor
|
||||
self._processors: OrderedDict[object, SpanProcessor] = OrderedDict() # mutable-ok: bounded LRU
|
||||
self._retired: OrderedDict[int, SpanProcessor] = OrderedDict() # mutable-ok: drains as exports finish
|
||||
self._exporting: dict[int, int] = {} # mutable-ok: per-processor in-flight export count
|
||||
self._drain: Final = drain_pool if drain_pool is not None else _DrainPool(capacity=pending_drains)
|
||||
|
||||
def on_start(self, span: SDKSpan, parent_context: Context | None = None) -> None:
|
||||
return None
|
||||
|
||||
def on_end(self, span: ReadableSpan) -> None:
|
||||
suppressed: Final = suppressed_backends()
|
||||
for destination in request_destinations():
|
||||
if self._operator_already_writes(destination, suppressed):
|
||||
continue
|
||||
processor = self._acquire(destination) # rebind-ok: loop variable; pyright forbids Final in a loop
|
||||
if processor is None:
|
||||
continue
|
||||
try:
|
||||
processor.on_end(_for_destination(span, destination))
|
||||
except Exception as exc: # noqa: BLE001 # one destination's failure must not cost the others their span
|
||||
verbose_logger.debug("OTel V2 fan-out: forwarding to %s failed: %s", destination.endpoint, exc)
|
||||
finally:
|
||||
self._release(processor)
|
||||
|
||||
def _operator_already_writes(self, destination: "OtelDestination", suppressed: frozenset[str]) -> bool:
|
||||
"""Whether the operator's own exporter is sending this span to the same account.
|
||||
|
||||
Only reachable under ``additive``, where nothing is suppressed: a team that
|
||||
names the operator's own project would otherwise have every span written
|
||||
there twice, once by the operator's exporter and once by the fan-out.
|
||||
"""
|
||||
return (
|
||||
destination.callback_name not in suppressed
|
||||
and _sink_key(destination.endpoint, destination.headers) in self._operator_sinks
|
||||
)
|
||||
|
||||
def shutdown(self) -> None:
|
||||
"""Close every destination processor, once the spans in flight have landed.
|
||||
|
||||
``on_end`` runs on whichever thread ends a span and can reach this fan-out
|
||||
while the SDK is tearing the provider down, so closing blind would drop a
|
||||
trace mid-forward and would hand the next caller a fresh exporter nothing
|
||||
will ever close. Refusing new work and then waiting out the in-flight ones
|
||||
keeps both from happening. A straggler past the bound is retired instead of
|
||||
closed: the thread still exporting it closes it through the drain as soon as
|
||||
its export returns, so no span is dropped mid-forward.
|
||||
|
||||
Every close then goes to the drain rather than running here. Closing a
|
||||
destination processor flushes it over the network and the SDK joins its own
|
||||
worker with no timeout of its own, so one tenant collector that answers but
|
||||
never finishes a response would otherwise hold process teardown open for as
|
||||
long as it likes. The drain's workers are daemons, and the whole teardown
|
||||
shares one deadline.
|
||||
"""
|
||||
deadline: Final = time.monotonic() + self._drain_seconds
|
||||
with self._lock:
|
||||
self._closed = True
|
||||
self._lock.wait_for(lambda: not self._exporting, timeout=self._drain_seconds)
|
||||
live: Final = tuple((id(p), p) for p in (*self._processors.values(), *self._retired.values()))
|
||||
closing: Final = tuple(p for ident, p in live if ident not in self._exporting)
|
||||
self._processors.clear()
|
||||
self._retired = OrderedDict( # mutable-ok: the same bounded map, keeping only what is still exporting
|
||||
(ident, p) for ident, p in live if ident in self._exporting
|
||||
)
|
||||
for processor in closing:
|
||||
self._drain.submit(processor)
|
||||
self._drain.close(timeout=max(0.0, deadline - time.monotonic()))
|
||||
|
||||
def force_flush(self, timeout_millis: int = 30000) -> bool:
|
||||
results: Final = tuple(self._flush_one(processor, timeout_millis) for processor in self._snapshot())
|
||||
return all(results)
|
||||
|
||||
def _snapshot(self) -> tuple[SpanProcessor, ...]:
|
||||
with self._lock:
|
||||
return (*self._processors.values(), *self._retired.values())
|
||||
|
||||
@staticmethod
|
||||
def _flush_one(processor: SpanProcessor, timeout_millis: int) -> bool:
|
||||
try:
|
||||
return processor.force_flush(timeout_millis)
|
||||
except Exception: # noqa: BLE001 # one exporter's flush failure must not fail the whole flush
|
||||
return False
|
||||
|
||||
def deliverable(self, destinations: Iterable["OtelDestination"]) -> tuple["OtelDestination", ...]:
|
||||
"""The subset of ``destinations`` this fan-out can actually export to.
|
||||
|
||||
A destination whose exporter will not build (a protocol whose package is not
|
||||
installed, a malformed endpoint) has to be dropped before the request anchors
|
||||
it, not when its first span ends. By then the operator's own exporter has been
|
||||
told to hold that backend's spans back for this request, so dropping there
|
||||
loses the span outright instead of leaving it where it would have gone with no
|
||||
override at all.
|
||||
"""
|
||||
return tuple(destination for destination in destinations if self._buildable(destination))
|
||||
|
||||
def _buildable(self, destination: "OtelDestination") -> bool:
|
||||
"""Whether a processor for ``destination`` exists or can be built right now."""
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
return False
|
||||
built: Final = self._cached_or_built_locked(destination, anchored=False)
|
||||
drained: Final = self._drainable_locked()
|
||||
for shed in drained:
|
||||
self._drain.submit(shed)
|
||||
return built is not None
|
||||
|
||||
def _acquire(self, destination: "OtelDestination") -> SpanProcessor | None:
|
||||
"""The processor for ``destination``, marked busy until ``_release``.
|
||||
|
||||
The build happens under the same lock that reads the cache, so a cold cache
|
||||
met by a burst of concurrent requests yields one exporter rather than one per
|
||||
thread with all but the winner shed. Building an exporter opens no connection,
|
||||
so the cost of holding the lock is a constructor, once per destination.
|
||||
"""
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
return None
|
||||
processor: Final = self._cached_or_built_locked(destination, anchored=True)
|
||||
if processor is None:
|
||||
return None
|
||||
self._exporting[id(processor)] = self._exporting.get(id(processor), 0) + 1
|
||||
drained: Final = self._drainable_locked()
|
||||
for shed in drained:
|
||||
self._drain.submit(shed)
|
||||
return processor
|
||||
|
||||
def _cached_or_built_locked(self, destination: "OtelDestination", *, anchored: bool) -> SpanProcessor | None:
|
||||
"""The cached processor for ``destination``, or a new one if the drain can take it.
|
||||
|
||||
Every build past the cache cap sheds one processor into the drain, so while the
|
||||
shed ones are stuck closing against a collector that stopped answering, a
|
||||
destination that is not yet anchored is refused rather than parked behind them:
|
||||
``deliverable`` then leaves its spans with the operator's exporter until the
|
||||
drain catches up. One the request already anchored is rebuilt regardless. The
|
||||
operator's exporter has stood down for it, so refusing here would drop the span,
|
||||
and other tenants' auths can evict it in the meantime, with that eviction being
|
||||
what tips the drain over. Eviction holds while the drain is saturated, so such a
|
||||
rebuild costs the cache one entry rather than shedding another processor, and
|
||||
the total stays at one per destination in flight.
|
||||
"""
|
||||
key: Final = destination.cache_key()
|
||||
if (cached := self._processors.get(key)) is not None:
|
||||
self._processors.move_to_end(key)
|
||||
self._retire_overflow_locked()
|
||||
return cached
|
||||
if not anchored and self._drain.saturated():
|
||||
verbose_logger.debug("OTel V2 fan-out: drain saturated, not building for %s", destination.endpoint)
|
||||
return None
|
||||
return self._build_locked(destination, key)
|
||||
|
||||
def _build_locked(self, destination: "OtelDestination", key: object) -> SpanProcessor | None:
|
||||
built: Final = self._build(destination)
|
||||
if built is None:
|
||||
return None
|
||||
self._processors[key] = built
|
||||
self._retire_overflow_locked()
|
||||
return built
|
||||
|
||||
def _release(self, processor: SpanProcessor) -> None:
|
||||
with self._lock:
|
||||
remaining: Final = self._exporting.get(id(processor), 1) - 1
|
||||
if remaining > 0:
|
||||
self._exporting[id(processor)] = remaining
|
||||
else:
|
||||
self._exporting.pop(id(processor), None)
|
||||
if not self._exporting:
|
||||
self._lock.notify_all()
|
||||
drained: Final = self._drainable_locked()
|
||||
for retired in drained:
|
||||
self._drain.submit(retired)
|
||||
|
||||
def _retire_overflow_locked(self) -> None:
|
||||
"""Move the LRU processor out of the cache once it is past the cap, drain permitting.
|
||||
|
||||
Eviction is what feeds the drain, and a destination a request already anchored
|
||||
is rebuilt on its next span, which would shed another one. While the shed ones
|
||||
are stuck closing against a collector that stopped answering, evicting would
|
||||
churn the cache at one more processor, and one more batch thread, per span.
|
||||
Holding above the cap instead keeps the total at one processor per destination
|
||||
in flight, since ``deliverable`` anchors no new destination while the drain is
|
||||
saturated. Once it has room again, every hit and build trims one entry.
|
||||
"""
|
||||
if len(self._processors) <= _MAX_CACHED_DESTINATION_PROCESSORS or self._drain.saturated():
|
||||
return
|
||||
_, evicted = self._processors.popitem(last=False)
|
||||
self._retired[id(evicted)] = evicted
|
||||
|
||||
def _drainable_locked(self) -> tuple[SpanProcessor, ...]:
|
||||
"""Retired processors no thread is exporting through, removed from the list.
|
||||
|
||||
``on_end`` holds a processor across an export, so closing an evicted one there
|
||||
drops the span it is holding. A retiree is out of the cache and can never be
|
||||
handed out again, so once its export count reaches zero it stays there.
|
||||
"""
|
||||
idle: Final = tuple(key for key in self._retired if self._exporting.get(key, 0) == 0)
|
||||
return tuple(self._retired.pop(key) for key in idle)
|
||||
|
||||
|
||||
def _destination_processor(destination: "OtelDestination") -> SpanProcessor | None:
|
||||
"""A batching OTLP processor aimed at ``destination``, or ``None`` if unbuildable.
|
||||
|
||||
A protocol that resolves to a headerless exporter is unbuildable too: the
|
||||
console fallback would swallow the tenant's credentials and print its spans to
|
||||
the proxy's stdout while the operator's exporter stands down for them.
|
||||
"""
|
||||
kind: Final = destination.protocol or "otlp_http"
|
||||
if exporter_transport(kind) == "headerless":
|
||||
verbose_logger.debug("OTel V2 fan-out: no OTLP transport for protocol %r at %s", kind, destination.endpoint)
|
||||
return None
|
||||
try:
|
||||
spec: Final = ExporterSpec(
|
||||
kind=kind,
|
||||
endpoint=destination.endpoint,
|
||||
headers=destination.header_string(),
|
||||
owner=None,
|
||||
)
|
||||
return _processor_for(_exporter_from_spec(spec), use_simple=False)
|
||||
except Exception as exc: # noqa: BLE001 # a malformed destination must not break the request or the other destinations
|
||||
verbose_logger.debug("OTel V2 fan-out: no processor for %s: %s", destination.endpoint, exc)
|
||||
return None
|
||||
|
||||
|
||||
def _shutdown_quietly(processor: SpanProcessor) -> None:
|
||||
try:
|
||||
processor.shutdown()
|
||||
except Exception as exc: # noqa: BLE001 # defensive: shedding a spare processor must not raise
|
||||
verbose_logger.debug("OTel V2 fan-out: discarding processor failed: %s", exc)
|
||||
|
||||
|
||||
class _OverriddenBackendFilter(SpanProcessor):
|
||||
"""Hold a span back from ``owner``'s operator-level exporter when the request
|
||||
pointed ``owner`` at a tenant's own account.
|
||||
|
||||
Wrapping is the only place this works: ``SynchronousMultiSpanProcessor.on_end``
|
||||
ignores return values, so a sibling processor can never veto the export.
|
||||
|
||||
Under ``additive`` mode nothing is suppressed, so the wrapper passes every span
|
||||
straight through and the operator keeps its copy.
|
||||
"""
|
||||
|
||||
def __init__(self, inner: SpanProcessor, owner: str) -> None:
|
||||
self._inner: Final = inner
|
||||
self._owner: Final = owner
|
||||
|
||||
def on_start(self, span: SDKSpan, parent_context: Context | None = None) -> None:
|
||||
self._inner.on_start(span, parent_context)
|
||||
|
||||
def on_end(self, span: ReadableSpan) -> None:
|
||||
if self._owner in suppressed_backends():
|
||||
return
|
||||
self._inner.on_end(span)
|
||||
|
||||
def shutdown(self) -> None:
|
||||
self._inner.shutdown()
|
||||
|
||||
def force_flush(self, timeout_millis: int = 30000) -> bool:
|
||||
return self._inner.force_flush(timeout_millis)
|
||||
|
||||
|
||||
def build_span_exporter(config: OpenTelemetryV2Config) -> SpanExporter:
|
||||
"""Build a single exporter from the top-level config fields.
|
||||
|
||||
|
|
@ -444,6 +1024,7 @@ def build_tracer_provider(
|
|||
exporter: SpanExporter | None = None,
|
||||
baggage_processor: SpanProcessor | None = None,
|
||||
use_simple_processor: bool | None = None,
|
||||
tenant_overrides: bool = False,
|
||||
) -> TracerProvider:
|
||||
"""Build the shared :class:`TracerProvider`.
|
||||
|
||||
|
|
@ -452,6 +1033,13 @@ def build_tracer_provider(
|
|||
``config.exporters`` entry — this is what fans spans out to multiple
|
||||
backends. ``exporter`` and ``use_simple_processor`` are explicit overrides:
|
||||
pass a single exporter to attach exactly that one (used by tests).
|
||||
|
||||
``tenant_overrides`` wraps each owned exporter so a request that pointed that
|
||||
backend at a key's or team's own account skips it. Every v2 logger's provider
|
||||
wants it, since any of them may own the overridden backend; delivering to the
|
||||
tenant is a separate job, done once by :func:`attach_tenant_fan_out`. The
|
||||
per-tenant providers this same function builds must leave it off, or they would
|
||||
filter out the very spans they exist to carry.
|
||||
"""
|
||||
provider: Final = TracerProvider(resource=build_resource(config))
|
||||
if baggage_processor is None:
|
||||
|
|
@ -468,15 +1056,107 @@ def build_tracer_provider(
|
|||
if spec.requires_headers and not spec.headers:
|
||||
continue
|
||||
exp = _exporter_from_spec(spec)
|
||||
processor = _processor_for(
|
||||
exp,
|
||||
(spec.use_simple_processor if spec.use_simple_processor is not None else use_simple_processor),
|
||||
)
|
||||
owner = spec.owner.value if spec.owner is not None else None
|
||||
provider.add_span_processor(
|
||||
_processor_for(
|
||||
exp,
|
||||
(spec.use_simple_processor if spec.use_simple_processor is not None else use_simple_processor),
|
||||
)
|
||||
_OverriddenBackendFilter(processor, owner) if tenant_overrides and owner is not None else processor
|
||||
)
|
||||
return provider
|
||||
|
||||
|
||||
_FAN_OUT_ATTACH_LOCK: Final = threading.Lock()
|
||||
|
||||
|
||||
def attach_tenant_fan_out(provider: TracerProvider, *configs: OpenTelemetryV2Config) -> None:
|
||||
"""Give ``provider`` the fan-out that delivers spans to key/team destinations.
|
||||
|
||||
Called on the one provider published as the OTel global, and idempotent so a
|
||||
second publish (a test, a re-initialized proxy) cannot double-export. Concurrent
|
||||
first calls (requests racing to anchor before any publish) serialize on one lock
|
||||
so exactly one fan-out lands. ``configs`` name the operator's own exporters, one
|
||||
config per v2 logger since each keeps its own provider and still writes its
|
||||
account, so an additive destination pointing at any of them is delivered once
|
||||
rather than twice.
|
||||
"""
|
||||
with _FAN_OUT_ATTACH_LOCK:
|
||||
if any(isinstance(processor, TenantFanOutSpanProcessor) for processor in _attached_processors(provider)):
|
||||
return
|
||||
provider.add_span_processor(TenantFanOutSpanProcessor(operator_sinks=operator_sink_keys(*configs)))
|
||||
|
||||
|
||||
def deliverable_destinations(
|
||||
destinations: Iterable["OtelDestination"],
|
||||
provider: trace.TracerProvider | None = None,
|
||||
) -> tuple["OtelDestination", ...]:
|
||||
"""The destinations a request can anchor, given what is published to carry them.
|
||||
|
||||
Anchoring a destination is what tells the operator's own exporter to stand down
|
||||
for that backend, so one nothing can deliver has to be dropped here: with no
|
||||
fan-out attached, or with an exporter that will not build, the request keeps
|
||||
exactly the routing it would have had without any override.
|
||||
"""
|
||||
fan_out: Final = next(
|
||||
(
|
||||
processor
|
||||
for processor in _attached_processors(provider if provider is not None else trace.get_tracer_provider())
|
||||
if isinstance(processor, TenantFanOutSpanProcessor)
|
||||
),
|
||||
None,
|
||||
)
|
||||
return fan_out.deliverable(destinations) if fan_out is not None else ()
|
||||
|
||||
|
||||
def operator_sink_keys(*configs: OpenTelemetryV2Config) -> frozenset[_SinkKey]:
|
||||
"""The accounts the operator's own exporters write to, in destination terms.
|
||||
|
||||
Every v2 logger's config counts, since each logger exports through its own
|
||||
provider. An exporter with no endpoint of its own resolves one from the
|
||||
environment at export time, so it has no comparable identity and is left out,
|
||||
and so is one that never reaches the wire: a console kind ignores the endpoint,
|
||||
and a header-gated spec with no credentials is skipped when the provider is built.
|
||||
"""
|
||||
return frozenset(
|
||||
key
|
||||
for config in configs
|
||||
for spec in config.exporters
|
||||
if _exports_to_the_wire(spec) and (key := _sink_key(spec.endpoint, parse_headers(spec.headers))) is not None
|
||||
)
|
||||
|
||||
|
||||
def _exports_to_the_wire(spec: ExporterSpec) -> bool:
|
||||
"""Whether ``build_tracer_provider`` gives ``spec`` an exporter that sends OTLP."""
|
||||
return exporter_transport(spec.kind) != "headerless" and not (spec.requires_headers and not spec.headers)
|
||||
|
||||
|
||||
def _sink_key(endpoint: str | None, headers: Mapping[str, str]) -> "_SinkKey | None":
|
||||
"""The account an exporter writes to, or ``None`` when it has no fixed one.
|
||||
|
||||
Normalized on the three counts that make one account look like two: the operator's
|
||||
spec carries the signal path a tenant destination leaves for the exporter to
|
||||
append, header names survive one round trip lowercased and the other not, and one
|
||||
credential answers to more than one name (see :data:`_CREDENTIAL_ALIASES`).
|
||||
"""
|
||||
normalized: Final = _otlp_traces_endpoint(endpoint)
|
||||
if normalized is None:
|
||||
return None
|
||||
return (normalized, tuple(sorted((_credential_name(name), value) for name, value in headers.items())))
|
||||
|
||||
|
||||
def _credential_name(header: str) -> str:
|
||||
"""The credential a header carries, under whichever name the backend spells it."""
|
||||
normalized: Final = header.strip().lower().replace("-", "_")
|
||||
return _CREDENTIAL_ALIASES.get(normalized, normalized)
|
||||
|
||||
|
||||
def _attached_processors(provider: trace.TracerProvider) -> "tuple[SpanProcessor, ...]":
|
||||
"""The processors already on ``provider``, or empty when the SDK hides them."""
|
||||
multi: Final = getattr(provider, "_active_span_processor", None)
|
||||
return tuple(getattr(multi, "_span_processors", ()))
|
||||
|
||||
|
||||
def get_tracer(provider: TracerProvider, name: str = "litellm") -> Tracer:
|
||||
# Stamp the instrumentation scope with the LiteLLM package version so every
|
||||
# emitted span carries a deterministic ``scope.version`` (the standard OTel
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from opentelemetry.trace import Tracer
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import OTEL_SERVICE_NAME_METADATA_KEYS
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.plumbing.context import destination_backends
|
||||
from litellm.integrations.otel.plumbing.providers import (
|
||||
build_tracer_provider,
|
||||
exporter_transport,
|
||||
|
|
@ -231,10 +232,21 @@ class TenantTracerCache:
|
|||
concurrent overflow eviction can't shut it down between selection and
|
||||
the caller's span start. The caller must ``release`` it exactly once.
|
||||
"""
|
||||
# A backend with a destination is delivered by the fan-out processor, which
|
||||
# carries the whole trace and already carries this tenant's credentials and
|
||||
# service name. Routing here too would detach this span onto a second provider,
|
||||
# so the tenant would get the request tree plus a stray one-span trace.
|
||||
if self._callback_name is not None and self._callback_name in destination_backends():
|
||||
return TenantRoute(tracer=default, detached=False)
|
||||
credential_headers: Final = self._credential_headers(dynamic_params)
|
||||
project_headers: Final = self._project_headers(auth_metadata)
|
||||
service_name: Final = tenant_service_name(auth_metadata)
|
||||
if not credential_headers and not project_headers and service_name is None:
|
||||
tenant_account: Final = bool(credential_headers) or bool(project_headers)
|
||||
# A service name on its own only relabels the operator's own backend, so moving
|
||||
# the span to a second provider for it while some other backend has a
|
||||
# destination would drop the model call out of the trace the fan-out delivers.
|
||||
# The destination stamps the same service name itself.
|
||||
if not tenant_account and (service_name is None or destination_backends()):
|
||||
return TenantRoute(tracer=default, detached=False)
|
||||
# A fixed per-integration region endpoint (New Relic us/eu), never a
|
||||
# caller-supplied host; ``None`` keeps the preset's own endpoint.
|
||||
|
|
@ -255,7 +267,7 @@ class TenantTracerCache:
|
|||
_shutdown_provider(evicted)
|
||||
return TenantRoute(
|
||||
tracer=get_tracer(provider, self._tracer_name),
|
||||
detached=bool(project_headers) or bool(credential_headers),
|
||||
detached=tenant_account,
|
||||
provider=provider,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ class _AgentOpsSettings(BaseSettings):
|
|||
def agentops_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
"""Build the AgentOps config without any network I/O.
|
||||
|
||||
|
|
|
|||
|
|
@ -26,10 +26,12 @@ class _ArizeSettings(BaseSettings):
|
|||
def arize_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
base: Final = config_overrides or OpenTelemetryV2Config()
|
||||
mappers: Final = ensure_mappers(base.mapper_names, "openinference")
|
||||
arize_cfg: Final = _V1ArizeLogger.get_arize_config()
|
||||
headers: Final = _arize_headers(arize_cfg)
|
||||
base: Final = config_overrides or OpenTelemetryV2Config()
|
||||
return base.model_copy(
|
||||
update={
|
||||
"exporters": [
|
||||
|
|
@ -41,7 +43,7 @@ def arize_preset(
|
|||
owner=ExporterOwner.ARIZE_AX,
|
||||
),
|
||||
],
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "openinference"),
|
||||
"mapper_names": mappers,
|
||||
"resource_attributes": {
|
||||
**base.resource_attributes,
|
||||
**({"model_id": arize_cfg.project_name} if arize_cfg.project_name else {}),
|
||||
|
|
|
|||
|
|
@ -18,6 +18,18 @@ class Preset(Protocol):
|
|||
|
||||
``config_overrides`` lets one preset layer onto another's config (or onto
|
||||
test-supplied defaults); the factory calls presets with no arguments.
|
||||
|
||||
``allow_missing_credentials`` lets a credential-mandatory backend (langfuse and
|
||||
weave) degrade to an exporter-less, mapper-only config instead of raising when the
|
||||
operator set no env credentials of their own. That is a real
|
||||
deployment: every team brings its own account and the operator keeps none, and
|
||||
without it the whole V2 path silently falls back to the legacy integration, so
|
||||
no team destination is ever reached. Credential-optional backends ignore it.
|
||||
"""
|
||||
|
||||
def __call__(self, *, config_overrides: OpenTelemetryV2Config | None = None) -> OpenTelemetryV2Config: ...
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config: ...
|
||||
|
|
|
|||
152
litellm/integrations/otel/presets/destinations.py
Normal file
152
litellm/integrations/otel/presets/destinations.py
Normal file
|
|
@ -0,0 +1,152 @@
|
|||
"""Map a key's or team's callback vars to the OTLP destination its traces export to.
|
||||
|
||||
Header building is delegated to each preset's existing ``*_dynamic_headers`` builder,
|
||||
so a destination authenticates exactly the way the per-request tracer route already
|
||||
did; only the endpoint and transport need a per-backend rule.
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Callable, Mapping
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.otel.model.destination import OtelDestination
|
||||
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
#: An endpoint plus the OTLP transport to reach it with, or ``None`` when the backend
|
||||
#: names no destination. The transport is ``None`` where the backend has only one.
|
||||
_Destination = tuple[str, str | None]
|
||||
|
||||
|
||||
@lru_cache(maxsize=128)
|
||||
def _warn_host_not_allowlisted(host: str) -> None:
|
||||
"""Cached so one misconfigured team logs once rather than once per request."""
|
||||
verbose_logger.warning(
|
||||
"OTel V2: not exporting to key/team Langfuse host '%s'. Add it to "
|
||||
"litellm_settings.provider_url_destination_allowed_hosts to permit it",
|
||||
host,
|
||||
)
|
||||
|
||||
|
||||
def _langfuse_destination(params: StandardCallbackDynamicParams) -> "_Destination | None":
|
||||
"""The tenant's own Langfuse host, else the operator's, else Langfuse US cloud.
|
||||
|
||||
A host the tenant named has to be allowlisted by the operator, the same way a
|
||||
URL-valued ``model`` is: anyone who can mint a key can write it, and it becomes an
|
||||
endpoint the proxy posts the request's whole trace to, carrying the tenant's own
|
||||
credentials. The operator's own ``LANGFUSE_HOST`` is not checked, since an internal
|
||||
collector there is a deployment choice.
|
||||
"""
|
||||
from litellm.integrations.langfuse.langfuse_otel import (
|
||||
LANGFUSE_CLOUD_US_ENDPOINT,
|
||||
LangfuseOtelLogger,
|
||||
)
|
||||
|
||||
tenant_host: Final = params.get("langfuse_host") or None
|
||||
host: Final = tenant_host or LangfuseOtelLogger._get_langfuse_otel_host() # pyright: ignore[reportPrivateUsage] # reuse the backend's own env host resolver rather than duplicating it
|
||||
if not host:
|
||||
return (LANGFUSE_CLOUD_US_ENDPOINT, None)
|
||||
normalized: Final = host if host.startswith("http") else f"https://{host}"
|
||||
endpoint: Final = f"{normalized.rstrip('/')}/api/public/otel"
|
||||
if tenant_host is None:
|
||||
return (endpoint, None)
|
||||
if not is_url_destination_allowed_by_host(endpoint, litellm.provider_url_destination_allowed_hosts):
|
||||
_warn_host_not_allowlisted(host)
|
||||
return None
|
||||
return (endpoint, None)
|
||||
|
||||
|
||||
def _arize_destination(params: StandardCallbackDynamicParams) -> "_Destination | None":
|
||||
from litellm.integrations.arize.arize import ArizeLogger
|
||||
|
||||
config: Final = ArizeLogger.get_arize_config()
|
||||
return (config.endpoint, config.protocol)
|
||||
|
||||
|
||||
def _weave_destination(params: StandardCallbackDynamicParams) -> "_Destination | None":
|
||||
from litellm.integrations.weave.weave_otel import weave_otel_endpoint
|
||||
|
||||
return (weave_otel_endpoint(os.environ.get("WANDB_HOST")), None)
|
||||
|
||||
|
||||
def _newrelic_destination(params: StandardCallbackDynamicParams) -> "_Destination | None":
|
||||
from litellm.integrations.otel.presets.newrelic import newrelic_dynamic_endpoint
|
||||
|
||||
endpoint: Final = newrelic_dynamic_endpoint(params)
|
||||
return (endpoint, None) if endpoint else None
|
||||
|
||||
|
||||
#: Callback name -> destination resolver. A backend is destination-capable exactly
|
||||
#: when it appears here AND in ``DYNAMIC_HEADERS_BY_CALLBACK``: without a header
|
||||
#: builder the destination would carry no tenant credentials, and the exporter
|
||||
#: would post the tenant's traffic to the operator's account.
|
||||
_DESTINATION_BY_CALLBACK: Final[Mapping[str, Callable[[StandardCallbackDynamicParams], "_Destination | None"]]] = (
|
||||
MappingProxyType(
|
||||
{
|
||||
"langfuse_otel": _langfuse_destination,
|
||||
"arize": _arize_destination,
|
||||
"weave_otel": _weave_destination,
|
||||
"newrelic": _newrelic_destination,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
#: Headers a destination must carry to authenticate. Several dynamic-header builders
|
||||
#: gate each credential independently, so a half-configured backend yields a non-empty
|
||||
#: but unusable header set; accepting it would suppress the operator's own exporter and
|
||||
#: send the request's whole trace where it cannot be stored.
|
||||
_REQUIRED_HEADERS_BY_CALLBACK: Final[Mapping[str, frozenset[str]]] = MappingProxyType(
|
||||
{
|
||||
"langfuse_otel": frozenset({"Authorization"}),
|
||||
"arize": frozenset({"arize-space-id", "api_key"}),
|
||||
"weave_otel": frozenset({"Authorization", "project_id"}),
|
||||
"newrelic": frozenset({"api-key"}),
|
||||
}
|
||||
)
|
||||
|
||||
_NO_ATTRS: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
def destination_capable_backends() -> frozenset[str]:
|
||||
"""Backends a key or team can point at its own account."""
|
||||
from litellm.integrations.otel.presets import DYNAMIC_HEADERS_BY_CALLBACK
|
||||
|
||||
return frozenset(_DESTINATION_BY_CALLBACK) & frozenset(DYNAMIC_HEADERS_BY_CALLBACK)
|
||||
|
||||
|
||||
def destination_for(
|
||||
callback_name: str,
|
||||
params: StandardCallbackDynamicParams,
|
||||
service_name: str | None = None,
|
||||
) -> OtelDestination | None:
|
||||
"""The destination ``params`` names for ``callback_name``, or ``None``.
|
||||
|
||||
``None`` means the caller configured nothing usable for this backend, so the
|
||||
request keeps the operator's global exporters. ``service_name`` is the key's or
|
||||
team's ``otel_service_name``, which the per-request tracer route applies when the
|
||||
backend is not overridden and the destination has to apply once it is.
|
||||
"""
|
||||
from litellm.integrations.otel.presets import DYNAMIC_HEADERS_BY_CALLBACK
|
||||
|
||||
header_builder: Final = DYNAMIC_HEADERS_BY_CALLBACK.get(callback_name)
|
||||
destination_builder: Final = _DESTINATION_BY_CALLBACK.get(callback_name)
|
||||
if header_builder is None or destination_builder is None:
|
||||
return None
|
||||
headers: Final = header_builder(params)
|
||||
if not headers or not _REQUIRED_HEADERS_BY_CALLBACK[callback_name] <= frozenset(headers):
|
||||
return None
|
||||
resolved: Final = destination_builder(params)
|
||||
if resolved is None:
|
||||
return None
|
||||
endpoint, protocol = resolved
|
||||
return OtelDestination(
|
||||
endpoint=endpoint,
|
||||
headers=MappingProxyType(dict(headers)), # mutable-ok: MappingProxyType needs a concrete mapping to wrap
|
||||
resource_attributes=MappingProxyType({"service.name": service_name}) if service_name else _NO_ATTRS,
|
||||
callback_name=callback_name,
|
||||
protocol=protocol,
|
||||
)
|
||||
|
|
@ -10,17 +10,32 @@ from litellm.integrations.otel.model.config import (
|
|||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.integrations.otel.presets.utils import (
|
||||
credential_gated_exporters,
|
||||
ensure_mappers,
|
||||
)
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
|
||||
def langfuse_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
cfg: Final = _V1Langfuse.get_langfuse_otel_config()
|
||||
kind: Final = cfg.exporter if isinstance(cfg.exporter, str) else "otlp_http"
|
||||
base: Final = config_overrides or OpenTelemetryV2Config()
|
||||
mappers: Final = ensure_mappers(base.mapper_names, "langfuse")
|
||||
try:
|
||||
cfg: Final = _V1Langfuse.get_langfuse_otel_config()
|
||||
except Exception:
|
||||
if not allow_missing_credentials:
|
||||
raise
|
||||
return base.model_copy(
|
||||
update={ # mutable-ok: pydantic model_copy takes a plain update mapping
|
||||
"exporters": credential_gated_exporters(base.exporters, ExporterOwner.LANGFUSE_OTEL),
|
||||
"mapper_names": mappers,
|
||||
}
|
||||
)
|
||||
kind: Final = cfg.exporter if isinstance(cfg.exporter, str) else "otlp_http"
|
||||
return base.model_copy(
|
||||
update={
|
||||
"exporters": [
|
||||
|
|
@ -32,7 +47,7 @@ def langfuse_preset(
|
|||
owner=ExporterOwner.LANGFUSE_OTEL,
|
||||
),
|
||||
],
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "langfuse"),
|
||||
"mapper_names": mappers,
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm.integrations.otel.presets.utils import ensure_mappers
|
|||
def langtrace_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
"""Compose the Langtrace mapper on top of the customer's OTLP destination.
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.integrations.otel.model.config import (
|
|||
def levo_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
cfg: Final = _V1Levo.get_levo_config()
|
||||
base: Final = config_overrides or OpenTelemetryV2Config()
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ class _NewRelicSettings(BaseSettings):
|
|||
def newrelic_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
settings: Final = _NewRelicSettings()
|
||||
base: Final = config_overrides or OpenTelemetryV2Config()
|
||||
|
|
|
|||
|
|
@ -60,6 +60,7 @@ def phoenix_project_headers(auth_metadata: Mapping[str, str] | None) -> Mapping[
|
|||
def phoenix_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
cfg: Final = _V1Phoenix.get_arize_phoenix_config()
|
||||
headers: Final = cfg.otlp_auth_headers if hasattr(cfg, "otlp_auth_headers") else None
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@
|
|||
from collections.abc import Iterable
|
||||
from typing import Final
|
||||
|
||||
from litellm.integrations.otel.model.config import ExporterOwner, ExporterSpec
|
||||
|
||||
|
||||
def ensure_mappers(mapper_names: Iterable[str], *names: str) -> list[str]:
|
||||
"""Return ``mapper_names`` with each of ``names`` appended if not already present.
|
||||
|
|
@ -15,3 +17,32 @@ def ensure_mappers(mapper_names: Iterable[str], *names: str) -> list[str]:
|
|||
if name not in result:
|
||||
result.append(name)
|
||||
return result
|
||||
|
||||
|
||||
def credential_gated_exporters(
|
||||
exporters: "Iterable[ExporterSpec]", owner: "ExporterOwner"
|
||||
) -> "tuple[ExporterSpec, ...]":
|
||||
"""``exporters`` with the operator's destination replaced by a header-gated one.
|
||||
|
||||
Used when a credential-mandatory backend is asked to build without the operator's
|
||||
own credentials, so only key/team destinations receive spans. Two things have to
|
||||
happen for that to mean "export nowhere": the placeholder console spec that
|
||||
``OpenTelemetryV2Config`` folds in for an empty exporter list is dropped, or every
|
||||
span would be printed to stdout, and the gated spec keeps the owner so the
|
||||
override filter still recognises which backend this provider speaks for.
|
||||
"""
|
||||
return (
|
||||
*(spec for spec in exporters if not is_unconfigured_placeholder(spec)),
|
||||
ExporterSpec(owner=owner, requires_headers=True),
|
||||
)
|
||||
|
||||
|
||||
def is_unconfigured_placeholder(spec: "ExporterSpec") -> bool:
|
||||
"""Whether ``spec`` is the one ``_normalize`` folds in when nothing was configured.
|
||||
|
||||
No field set is what says the operator asked for nothing: an exporter they did
|
||||
configure survives, even ``OTEL_EXPORTER=console`` whose value matches the default,
|
||||
and so does the gated spec this module appends, which would otherwise eat itself
|
||||
when one preset layers onto another.
|
||||
"""
|
||||
return not spec.model_fields_set
|
||||
|
|
|
|||
|
|
@ -7,7 +7,10 @@ from litellm.integrations.otel.model.config import (
|
|||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.integrations.otel.presets.utils import (
|
||||
credential_gated_exporters,
|
||||
ensure_mappers,
|
||||
)
|
||||
from litellm.integrations.weave.weave_otel import (
|
||||
_get_weave_authorization_header,
|
||||
get_weave_otel_config,
|
||||
|
|
@ -18,9 +21,21 @@ from litellm.types.utils import StandardCallbackDynamicParams
|
|||
def weave_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
weave_cfg: Final = get_weave_otel_config()
|
||||
base: Final = config_overrides or OpenTelemetryV2Config()
|
||||
mappers: Final = ensure_mappers(base.mapper_names, "openinference", "weave")
|
||||
try:
|
||||
weave_cfg: Final = get_weave_otel_config()
|
||||
except Exception:
|
||||
if not allow_missing_credentials:
|
||||
raise
|
||||
return base.model_copy(
|
||||
update={ # mutable-ok: pydantic model_copy takes a plain update mapping
|
||||
"exporters": credential_gated_exporters(base.exporters, ExporterOwner.WEAVE_OTEL),
|
||||
"mapper_names": mappers,
|
||||
}
|
||||
)
|
||||
return base.model_copy(
|
||||
update={
|
||||
"exporters": [
|
||||
|
|
@ -33,7 +48,7 @@ def weave_preset(
|
|||
),
|
||||
],
|
||||
# Weave consumes OpenInference + a small Weave-specific overlay.
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "openinference", "weave"),
|
||||
"mapper_names": mappers,
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -117,6 +117,14 @@ def _get_weave_authorization_header(api_key: str) -> str:
|
|||
return f"Basic {auth_header}"
|
||||
|
||||
|
||||
def weave_otel_endpoint(host: str | None) -> str:
|
||||
"""The OTLP traces endpoint for a self-managed ``host``, else Weave cloud."""
|
||||
if not host:
|
||||
return WEAVE_BASE_URL + WEAVE_OTEL_ENDPOINT
|
||||
normalized: Final = host if host.startswith("http") else f"https://{host}"
|
||||
return normalized.rstrip("/") + WEAVE_OTEL_ENDPOINT
|
||||
|
||||
|
||||
def get_weave_otel_config() -> WeaveOtelConfig:
|
||||
"""
|
||||
Retrieves the Weave OpenTelemetry configuration based on environment variables.
|
||||
|
|
@ -134,7 +142,6 @@ def get_weave_otel_config() -> WeaveOtelConfig:
|
|||
"""
|
||||
api_key: Final = os.getenv("WANDB_API_KEY")
|
||||
project_id: Final = os.getenv("WANDB_PROJECT_ID")
|
||||
host = os.getenv("WANDB_HOST")
|
||||
|
||||
if not api_key:
|
||||
raise ValueError("WANDB_API_KEY must be set for Weave OpenTelemetry integration.")
|
||||
|
|
@ -144,15 +151,8 @@ def get_weave_otel_config() -> WeaveOtelConfig:
|
|||
"WANDB_PROJECT_ID must be set for Weave OpenTelemetry integration. Format: <entity>/<project_name>"
|
||||
)
|
||||
|
||||
if host:
|
||||
if not host.startswith("http"):
|
||||
host = "https://" + host
|
||||
# Self-managed instances use a different path
|
||||
endpoint = host.rstrip("/") + WEAVE_OTEL_ENDPOINT
|
||||
verbose_logger.debug("Using Weave OTEL endpoint from host: %s", endpoint)
|
||||
else:
|
||||
endpoint = WEAVE_BASE_URL + WEAVE_OTEL_ENDPOINT
|
||||
verbose_logger.debug("Using Weave cloud endpoint: %s", endpoint)
|
||||
endpoint: Final = weave_otel_endpoint(os.getenv("WANDB_HOST"))
|
||||
verbose_logger.debug("Using Weave OTEL endpoint: %s", endpoint)
|
||||
|
||||
# Weave uses Basic auth with format: api:<WANDB_API_KEY>
|
||||
auth_header: Final = _get_weave_authorization_header(api_key=api_key)
|
||||
|
|
|
|||
|
|
@ -1,10 +1,12 @@
|
|||
# What is this?
|
||||
## Helper utilities
|
||||
import copy
|
||||
import logging
|
||||
from collections.abc import Iterable, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.llms.openai import AllMessageValues, OpenAIChatCompletionFinishReason
|
||||
|
|
@ -37,6 +39,41 @@ def safe_divide_seconds(seconds: float, denominator: float, default: float | Non
|
|||
return float(seconds / denominator)
|
||||
|
||||
|
||||
_DROP_PARAMS_BOOL: Final = TypeAdapter(bool)
|
||||
|
||||
|
||||
def normalize_drop_params(value: object) -> bool | None:
|
||||
if value is None or isinstance(value, bool):
|
||||
return value
|
||||
try:
|
||||
return _DROP_PARAMS_BOOL.validate_python(value.strip() if isinstance(value, str) else value)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def drop_params_flag(value: object, source: str, logger: logging.Logger) -> bool:
|
||||
normalized: Final = normalize_drop_params(value)
|
||||
if normalized is None and value is not None:
|
||||
logger.warning("%s=%r is not a flag value, treating it as off", source, value)
|
||||
return bool(normalized)
|
||||
|
||||
|
||||
DROP_PARAMS_ENV_VAR: Final = "LITELLM_DROP_PARAMS"
|
||||
|
||||
|
||||
def drop_params_env_flag(environ: Mapping[str, str], logger: logging.Logger) -> bool:
|
||||
configured: Final = environ.get(DROP_PARAMS_ENV_VAR, "").strip()
|
||||
if configured == "":
|
||||
return False
|
||||
normalized: Final = normalize_drop_params(configured)
|
||||
if normalized is None:
|
||||
logger.warning(
|
||||
"%s=%r is not a flag value, treating it as on. Set it to true or false", DROP_PARAMS_ENV_VAR, configured
|
||||
)
|
||||
return True
|
||||
return normalized
|
||||
|
||||
|
||||
def safe_divide(
|
||||
numerator: float,
|
||||
denominator: float,
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ from collections.abc import Mapping, MutableMapping
|
|||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
|
||||
from litellm.llms.openai.data_residency import infer_openai_data_residency
|
||||
|
||||
AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
|
||||
|
|
@ -113,7 +114,7 @@ def get_litellm_params(
|
|||
custom_prompt_dict: dict | None = None,
|
||||
litellm_metadata: dict | None = None,
|
||||
disable_add_transform_inline_image_block: bool | None = None,
|
||||
drop_params: bool | None = None,
|
||||
drop_params: bool | str | None = None,
|
||||
prompt_id: str | None = None,
|
||||
prompt_variables: dict | None = None,
|
||||
async_call: bool | None = None,
|
||||
|
|
@ -175,7 +176,7 @@ def get_litellm_params(
|
|||
"custom_prompt_dict": custom_prompt_dict,
|
||||
"litellm_metadata": litellm_metadata,
|
||||
"disable_add_transform_inline_image_block": disable_add_transform_inline_image_block,
|
||||
"drop_params": drop_params,
|
||||
"drop_params": normalize_drop_params(drop_params),
|
||||
"prompt_id": prompt_id,
|
||||
"prompt_variables": prompt_variables,
|
||||
"async_call": async_call,
|
||||
|
|
|
|||
|
|
@ -9,17 +9,19 @@ export LITELLM_LOCAL_MODEL_COST_MAP=True
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import datetime, timezone
|
||||
from importlib.resources import files
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -42,6 +44,10 @@ def _count_model_entries(model_cost: dict) -> int:
|
|||
return sum(1 for key in model_cost if key not in RESERVED_TOP_LEVEL_KEYS)
|
||||
|
||||
|
||||
def git_blob_id(body: bytes) -> str:
|
||||
return hashlib.sha1(b"blob %d\0" % len(body) + body, usedforsecurity=False).hexdigest()
|
||||
|
||||
|
||||
class GetModelCostMap:
|
||||
"""
|
||||
Handles fetching, validating, and loading the model cost map.
|
||||
|
|
@ -53,15 +59,24 @@ class GetModelCostMap:
|
|||
|
||||
_backup_model_count: int = -1 # -1 = not yet loaded
|
||||
|
||||
@staticmethod
|
||||
def read_local_model_cost_map_bytes() -> bytes:
|
||||
return files("litellm").joinpath("model_prices_and_context_window_backup.json").read_bytes()
|
||||
|
||||
@staticmethod
|
||||
def read_local_model_cost_map_text() -> str:
|
||||
return files("litellm").joinpath("model_prices_and_context_window_backup.json").read_text(encoding="utf-8")
|
||||
return GetModelCostMap.read_local_model_cost_map_bytes().decode("utf-8")
|
||||
|
||||
@staticmethod
|
||||
def load_local_model_cost_map_with_revision() -> "ModelCostMapReloaded":
|
||||
body: Final = GetModelCostMap.read_local_model_cost_map_bytes()
|
||||
content: Final = json.loads(body)
|
||||
return ModelCostMapReloaded(model_cost_map=content, revision=git_blob_id(body))
|
||||
|
||||
@staticmethod
|
||||
def load_local_model_cost_map() -> dict:
|
||||
"""Load the local backup model cost map bundled with the package."""
|
||||
content: Final = json.loads(GetModelCostMap.read_local_model_cost_map_text())
|
||||
return content
|
||||
return GetModelCostMap.load_local_model_cost_map_with_revision().model_cost_map
|
||||
|
||||
@classmethod
|
||||
def _get_backup_model_count(cls) -> int:
|
||||
|
|
@ -166,6 +181,8 @@ MODEL_COST_MAP_FETCH_MAX_WAIT_SECONDS: Final = 30.0
|
|||
@dataclass(frozen=True, slots=True)
|
||||
class ModelCostMapReloaded:
|
||||
model_cost_map: dict # mutable-ok: adopted as litellm.model_cost, whose consumer contract is a plain mutable dict
|
||||
revision: str | None = None
|
||||
etag: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -254,7 +271,9 @@ def _classify_fetch_response(response: httpx.Response, url: str) -> _FetchAttemp
|
|||
return ModelCostMapReloadUnavailable(reason=f"invalid JSON from {url}: {e}")
|
||||
if not isinstance(parsed, dict):
|
||||
return ModelCostMapReloadUnavailable(reason=f"expected a JSON object from {url}, got {type(parsed).__name__}")
|
||||
return ModelCostMapReloaded(model_cost_map=parsed)
|
||||
return ModelCostMapReloaded(
|
||||
model_cost_map=parsed, revision=git_blob_id(response.content), etag=response.headers.get("etag")
|
||||
)
|
||||
|
||||
|
||||
def _next_retry_wait(
|
||||
|
|
@ -328,13 +347,12 @@ async def refetch_model_cost_map(
|
|||
map they already have.
|
||||
"""
|
||||
if os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", "").lower() == "true":
|
||||
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
|
||||
_cost_map_source_info.source = "local"
|
||||
_cost_map_source_info.url = None
|
||||
_cost_map_source_info.is_env_forced = True
|
||||
_cost_map_source_info.fallback_reason = None
|
||||
return ModelCostMapReloaded(
|
||||
model_cost_map=_finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
|
||||
)
|
||||
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision())
|
||||
|
||||
result: Final = await _fetch_remote_model_cost_map_with_retry(
|
||||
url=url,
|
||||
|
|
@ -355,11 +373,12 @@ async def refetch_model_cost_map(
|
|||
backup_model_count=GetModelCostMap._get_backup_model_count(),
|
||||
):
|
||||
return ModelCostMapReloadUnavailable(reason=f"model cost map from {url} failed integrity validation")
|
||||
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
|
||||
_cost_map_source_info.source = "remote"
|
||||
_cost_map_source_info.url = url
|
||||
_cost_map_source_info.is_env_forced = False
|
||||
_cost_map_source_info.fallback_reason = None
|
||||
return ModelCostMapReloaded(model_cost_map=_finalize_model_cost_map(result.model_cost_map))
|
||||
return _finalize_loaded_model_cost_map(result)
|
||||
|
||||
|
||||
class ModelCostMapSourceInfo:
|
||||
|
|
@ -370,13 +389,35 @@ class ModelCostMapSourceInfo:
|
|||
is_env_forced: bool = False
|
||||
fallback_reason: str | None = None
|
||||
loaded_at: "datetime | None" = None
|
||||
source_revision: str | None = None
|
||||
etag: str | None = None
|
||||
|
||||
|
||||
# Module-level singleton tracking the source of the current cost map
|
||||
_cost_map_source_info: Final = ModelCostMapSourceInfo()
|
||||
|
||||
|
||||
def get_model_cost_map_source_info() -> dict:
|
||||
class CostMapProvenance(TypedDict):
|
||||
source_revision: ReadOnly[str | None]
|
||||
etag: ReadOnly[str | None]
|
||||
|
||||
|
||||
class CostMapSourceInfo(CostMapProvenance):
|
||||
source: ReadOnly[str]
|
||||
url: ReadOnly[str | None]
|
||||
is_env_forced: ReadOnly[bool]
|
||||
fallback_reason: ReadOnly[str | None]
|
||||
loaded_at: ReadOnly[str | None]
|
||||
|
||||
|
||||
def get_model_cost_map_provenance() -> CostMapProvenance:
|
||||
return {
|
||||
"source_revision": _cost_map_source_info.source_revision,
|
||||
"etag": _cost_map_source_info.etag,
|
||||
}
|
||||
|
||||
|
||||
def get_model_cost_map_source_info() -> CostMapSourceInfo:
|
||||
"""
|
||||
Return metadata about where the current model cost map was loaded from.
|
||||
|
||||
|
|
@ -385,12 +426,19 @@ def get_model_cost_map_source_info() -> dict:
|
|||
- url: the remote URL attempted (or None for local-only)
|
||||
- is_env_forced: True if LITELLM_LOCAL_MODEL_COST_MAP=True forced local usage
|
||||
- fallback_reason: human-readable reason if remote failed and local was used
|
||||
- loaded_at: ISO 8601 time this process last loaded the map
|
||||
- source_revision: git blob id of the loaded file's bytes
|
||||
- etag: the ETag of the remote fetch (None for the bundled backup)
|
||||
"""
|
||||
loaded_at: Final = _cost_map_source_info.loaded_at
|
||||
return {
|
||||
"source": _cost_map_source_info.source,
|
||||
"url": _cost_map_source_info.url,
|
||||
"is_env_forced": _cost_map_source_info.is_env_forced,
|
||||
"fallback_reason": _cost_map_source_info.fallback_reason,
|
||||
"loaded_at": loaded_at.isoformat() if loaded_at is not None else None,
|
||||
"source_revision": _cost_map_source_info.source_revision,
|
||||
"etag": _cost_map_source_info.etag,
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -466,6 +514,12 @@ def _finalize_model_cost_map(model_cost: dict) -> dict:
|
|||
return _expand_model_aliases(model_cost)
|
||||
|
||||
|
||||
def _finalize_loaded_model_cost_map(loaded: ModelCostMapReloaded) -> ModelCostMapReloaded:
|
||||
_cost_map_source_info.source_revision = loaded.revision
|
||||
_cost_map_source_info.etag = loaded.etag
|
||||
return replace(loaded, model_cost_map=_finalize_model_cost_map(loaded.model_cost_map))
|
||||
|
||||
|
||||
def get_model_cost_map(
|
||||
url: str,
|
||||
timeout: int = 5,
|
||||
|
|
@ -494,7 +548,7 @@ def get_model_cost_map(
|
|||
_cost_map_source_info.url = None
|
||||
_cost_map_source_info.is_env_forced = True
|
||||
_cost_map_source_info.fallback_reason = None
|
||||
return _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
|
||||
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision()).model_cost_map
|
||||
|
||||
_cost_map_source_info.url = url
|
||||
_cost_map_source_info.is_env_forced = False
|
||||
|
|
@ -515,7 +569,7 @@ def get_model_cost_map(
|
|||
)
|
||||
_cost_map_source_info.source = "local"
|
||||
_cost_map_source_info.fallback_reason = f"Remote fetch failed: {result.reason}"
|
||||
return _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
|
||||
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision()).model_cost_map
|
||||
content: Final = result.model_cost_map
|
||||
|
||||
# Validate using cached count (cheap int comparison, no file I/O)
|
||||
|
|
@ -529,8 +583,8 @@ def get_model_cost_map(
|
|||
)
|
||||
_cost_map_source_info.source = "local"
|
||||
_cost_map_source_info.fallback_reason = "Remote data failed integrity validation"
|
||||
return _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map())
|
||||
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision()).model_cost_map
|
||||
|
||||
_cost_map_source_info.source = "remote"
|
||||
_cost_map_source_info.fallback_reason = None
|
||||
return _finalize_model_cost_map(content)
|
||||
return _finalize_loaded_model_cost_map(result).model_cost_map
|
||||
|
|
|
|||
|
|
@ -202,6 +202,7 @@ if TYPE_CHECKING:
|
|||
from mcp.types import EmbeddedResource, ImageContent, TextContent
|
||||
|
||||
from litellm.integrations.otel.logger import OpenTelemetryV2
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
|
||||
try:
|
||||
from litellm_enterprise.enterprise_callbacks.callback_controls import (
|
||||
|
|
@ -448,6 +449,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, \
|
||||
|
|
@ -1190,14 +1198,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={}):
|
||||
"""
|
||||
|
|
@ -2950,6 +2951,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(
|
||||
|
|
@ -2965,6 +2979,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(
|
||||
|
|
@ -4831,31 +4851,83 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[Custom
|
|||
|
||||
Returns ``None`` when V2 is off OR when there's no preset registered for
|
||||
``callback_name`` — callers should then fall through to the legacy path.
|
||||
|
||||
A preset that needs operator credentials it cannot find is allowed to build
|
||||
only when this request has a key/team destination for that backend and another
|
||||
V2 logger is already registered to carry the fan-out. The resulting logger keeps
|
||||
only its credential-gated exporter, while the registered logger owns operator
|
||||
delivery. Without that carrier, a preset that raises or that ends up with nothing
|
||||
but its gated exporter and the default console placeholder returns ``None``, so the
|
||||
caller falls through to the legacy path exactly as before V2 landed.
|
||||
"""
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
|
||||
if not is_otel_v2_enabled():
|
||||
return None
|
||||
from litellm.integrations.otel.logger import OpenTelemetryV2, build_otel_v2_logger
|
||||
from litellm.integrations.otel.plumbing.context import destination_backends
|
||||
from litellm.integrations.otel.presets import PRESET_BY_CALLBACK
|
||||
|
||||
preset_fn: Final = PRESET_BY_CALLBACK.get(callback_name)
|
||||
if preset_fn is None:
|
||||
return None
|
||||
serves_a_destination: Final = callback_name in destination_backends()
|
||||
has_v2_logger: Final = any(isinstance(callback, OpenTelemetryV2) for callback in _in_memory_loggers)
|
||||
carried: Final = serves_a_destination and has_v2_logger
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, OpenTelemetryV2) and getattr(callback, "callback_name", None) == callback_name:
|
||||
if (
|
||||
isinstance(callback, OpenTelemetryV2)
|
||||
and getattr(callback, "callback_name", None) == callback_name
|
||||
and (serves_a_destination or not _exports_nowhere(callback.config))
|
||||
):
|
||||
return callback
|
||||
try:
|
||||
config: Final = preset_fn()
|
||||
built: Final = preset_fn(allow_missing_credentials=carried)
|
||||
except Exception:
|
||||
# If env vars are missing or the preset raises, defer to the legacy path
|
||||
# so customers get the same error story they had before V2 landed.
|
||||
return None
|
||||
gated: Final = _is_credential_gated(built)
|
||||
if gated and not carried and not _has_operator_exporter(built):
|
||||
return None
|
||||
config: Final = _only_the_gated_exporter(built) if gated and carried else built
|
||||
if _exports_nowhere(config):
|
||||
verbose_logger.warning(
|
||||
"OTel V2: no operator credentials for '%s'; only key/team destinations will receive its traces",
|
||||
callback_name,
|
||||
)
|
||||
v2_logger: Final = build_otel_v2_logger(config=config, callback_name=callback_name)
|
||||
_in_memory_loggers.append(v2_logger)
|
||||
return v2_logger
|
||||
|
||||
|
||||
def _exports_nowhere(config: "OpenTelemetryV2Config") -> bool:
|
||||
"""Whether every exporter in ``config`` is waiting on credentials it never got."""
|
||||
return all(_is_gated(spec) for spec in config.exporters)
|
||||
|
||||
|
||||
def _is_credential_gated(config: "OpenTelemetryV2Config") -> bool:
|
||||
"""Whether the preset built without the operator's own credentials for its backend."""
|
||||
return any(_is_gated(spec) for spec in config.exporters)
|
||||
|
||||
|
||||
def _has_operator_exporter(config: "OpenTelemetryV2Config") -> bool:
|
||||
"""Whether the operator configured somewhere real to export, beyond the default console placeholder."""
|
||||
from litellm.integrations.otel.presets.utils import is_unconfigured_placeholder
|
||||
|
||||
return any(not _is_gated(spec) and not is_unconfigured_placeholder(spec) for spec in config.exporters)
|
||||
|
||||
|
||||
def _only_the_gated_exporter(config: "OpenTelemetryV2Config") -> "OpenTelemetryV2Config":
|
||||
return config.model_copy(
|
||||
update={"exporters": [spec for spec in config.exporters if _is_gated(spec)]} # mutable-ok: model_copy update
|
||||
)
|
||||
|
||||
|
||||
def _is_gated(spec: "ExporterSpec") -> bool:
|
||||
return spec.requires_headers and not spec.headers
|
||||
|
||||
|
||||
def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list[CustomLogger]) -> None:
|
||||
"""
|
||||
Auto-initialize ArizePhoenixLogger when Phoenix env vars are detected.
|
||||
|
|
|
|||
|
|
@ -780,6 +780,7 @@ class PromptTokensDetailsResult(TypedDict):
|
|||
image_count: int
|
||||
video_length_seconds: float
|
||||
audio_length_seconds: float
|
||||
query_count: int
|
||||
|
||||
|
||||
def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
|
||||
|
|
@ -828,6 +829,7 @@ def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
|
|||
)
|
||||
or 0.0
|
||||
)
|
||||
query_count: Final = _coerce_token_count(getattr(usage.prompt_tokens_details, "query_count", 0))
|
||||
|
||||
return PromptTokensDetailsResult(
|
||||
cache_hit_tokens=cache_hit_tokens,
|
||||
|
|
@ -841,6 +843,7 @@ def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
|
|||
image_count=image_count,
|
||||
video_length_seconds=float(video_length_seconds),
|
||||
audio_length_seconds=float(audio_length_seconds),
|
||||
query_count=query_count,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -978,6 +981,11 @@ def _calculate_input_cost(
|
|||
prompt_tokens_details["audio_length_seconds"],
|
||||
)
|
||||
|
||||
if prompt_tokens_details["query_count"]:
|
||||
prompt_cost += calculate_cost_component(
|
||||
model_info, "input_cost_per_query", prompt_tokens_details["query_count"]
|
||||
)
|
||||
|
||||
return prompt_cost
|
||||
|
||||
|
||||
|
|
@ -1149,6 +1157,7 @@ def generic_cost_per_token(
|
|||
image_count=0,
|
||||
video_length_seconds=0.0,
|
||||
audio_length_seconds=0.0,
|
||||
query_count=0,
|
||||
)
|
||||
if usage.prompt_tokens_details:
|
||||
prompt_tokens_details = parse_prompt_tokens_details(usage)
|
||||
|
|
|
|||
|
|
@ -183,6 +183,7 @@ class _RemoteSource:
|
|||
class RemoteMedia:
|
||||
url: str
|
||||
fields: Mapping[str, object]
|
||||
part_type: str
|
||||
|
||||
|
||||
_NO_FIELDS: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
|
@ -192,6 +193,10 @@ def inline_every_remote_url(_media: RemoteMedia) -> bool:
|
|||
return True
|
||||
|
||||
|
||||
def inline_remote_image_urls(media: RemoteMedia) -> bool:
|
||||
return media.part_type == "image_url"
|
||||
|
||||
|
||||
def _parse_remote_image(fields: Mapping[str, object]) -> _RemoteImage | None:
|
||||
if fields.get("type") != "image_url":
|
||||
return None
|
||||
|
|
@ -223,11 +228,11 @@ def _parse_remote_part(part: object) -> _RemoteImage | _RemoteFile | _RemoteSour
|
|||
def _remote_media(remote: _RemoteImage | _RemoteFile | _RemoteSource) -> RemoteMedia:
|
||||
match remote:
|
||||
case _RemoteImage(_, image_url, url):
|
||||
return RemoteMedia(url, image_url if image_url is not None else _NO_FIELDS)
|
||||
return RemoteMedia(url, image_url if image_url is not None else _NO_FIELDS, "image_url")
|
||||
case _RemoteFile(_, file, url):
|
||||
return RemoteMedia(url, file)
|
||||
case _RemoteSource(_, source, url):
|
||||
return RemoteMedia(url, source)
|
||||
return RemoteMedia(url, file, "file")
|
||||
case _RemoteSource(part, source, url):
|
||||
return RemoteMedia(url, source, str(part.get("type")))
|
||||
|
||||
|
||||
_PDF_FORMAT: Final = MappingProxyType({"format": "application/pdf"})
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from litellm.types.utils import (
|
|||
Choices,
|
||||
CompletionTokensDetails,
|
||||
CompletionTokensDetailsWrapper,
|
||||
Delta,
|
||||
Function,
|
||||
FunctionCall,
|
||||
ModelResponse,
|
||||
|
|
@ -326,6 +327,18 @@ class ChunkProcessor:
|
|||
return chunk_id
|
||||
return ""
|
||||
|
||||
@staticmethod
|
||||
def _get_role_from_chunks(chunks: Sequence["_BaseChunk"]) -> str:
|
||||
return ChunkProcessor._role_of_choice(next((c["choices"][0] for c in chunks if c.get("choices")), None))
|
||||
|
||||
@staticmethod
|
||||
def _role_of_choice(choice: object) -> str:
|
||||
match choice:
|
||||
case StreamingChoices(delta=Delta(role=str() as role)) | {"delta": {"role": str() as role}} if role:
|
||||
return role
|
||||
case _:
|
||||
return "assistant"
|
||||
|
||||
@staticmethod
|
||||
def _get_model_from_chunks(chunks: Sequence["_BaseChunk"], first_chunk_model: str) -> str:
|
||||
"""
|
||||
|
|
@ -353,8 +366,7 @@ class ChunkProcessor:
|
|||
model: Final = ChunkProcessor._get_model_from_chunks(chunks, first_chunk_model)
|
||||
system_fingerprint: Final = chunk.get("system_fingerprint", None)
|
||||
|
||||
first_chunk_with_choices: Final = next((c for c in chunks if c.get("choices")), chunk)
|
||||
role: Final = first_chunk_with_choices["choices"][0]["delta"]["role"]
|
||||
role: Final = ChunkProcessor._get_role_from_chunks(chunks)
|
||||
finish_reason = "stop"
|
||||
for chunk in chunks:
|
||||
if "choices" in chunk and len(chunk["choices"]) > 0:
|
||||
|
|
|
|||
|
|
@ -13,9 +13,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,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ from collections.abc import Mapping
|
|||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import ModelInfo
|
||||
|
||||
|
|
@ -21,10 +23,27 @@ _EFFORT_DEGRADATION_CHAIN: Final[Mapping[str, tuple[str, ...]]] = MappingProxyTy
|
|||
_THINKING_OFF: Final = "none"
|
||||
|
||||
|
||||
class _ClaudeCodeUserId(BaseModel):
|
||||
"""The JSON Claude Code packs into ``metadata.user_id``; only ``session_id`` is per conversation."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
session_id: str
|
||||
|
||||
|
||||
def prompt_cache_key_from_user_id(user_id: object) -> str | None:
|
||||
if user_id is None:
|
||||
"""The per-session key Claude Code carries inside ``metadata.user_id``, or nothing.
|
||||
|
||||
Anthropic defines ``user_id`` as an opaque end-user id, so a plain string names a person, not
|
||||
a conversation. Keying the provider cache on it pins every parallel session and subagent of that
|
||||
person to one slot, which caches worse than the provider's own prompt-prefix hashing does.
|
||||
"""
|
||||
if not isinstance(user_id, str):
|
||||
return None
|
||||
try:
|
||||
return _ClaudeCodeUserId.model_validate_json(user_id).session_id[:OPENAI_MAX_PROMPT_CACHE_KEY_LENGTH] or None
|
||||
except ValidationError:
|
||||
return None
|
||||
return str(user_id)[:OPENAI_MAX_PROMPT_CACHE_KEY_LENGTH] or None
|
||||
|
||||
|
||||
def litellm_logging_obj_from_kwargs(kwargs: Mapping[str, object]) -> "LiteLLMLoggingObject | None":
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
azure_ad_token: str | None = None,
|
||||
atranscription: bool = False,
|
||||
litellm_params: dict | None = None,
|
||||
custom_llm_provider: str = "azure",
|
||||
) -> TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse]:
|
||||
data: Final = {"model": model, "file": audio_file, **optional_params}
|
||||
|
||||
|
|
@ -53,6 +54,7 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
litellm_params=litellm_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
azure_client: Final = self.get_azure_openai_client(
|
||||
|
|
@ -99,7 +101,7 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=stringified_response,
|
||||
)
|
||||
hidden_params: Final = {"model": model, "custom_llm_provider": "azure"}
|
||||
hidden_params: Final = {"model": model, "custom_llm_provider": custom_llm_provider}
|
||||
final_response: Final[TranscriptionResponse] = convert_to_model_response_object(
|
||||
response_object=stringified_response,
|
||||
model_response_object=model_response,
|
||||
|
|
@ -122,6 +124,7 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
client=None,
|
||||
max_retries=None,
|
||||
litellm_params: dict | None = None,
|
||||
custom_llm_provider: str = "azure",
|
||||
) -> TranscriptionResponse:
|
||||
response = None
|
||||
try:
|
||||
|
|
@ -178,7 +181,7 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
},
|
||||
original_response=stringified_response,
|
||||
)
|
||||
hidden_params: Final = {"model": model, "custom_llm_provider": "azure"}
|
||||
hidden_params: Final = {"model": model, "custom_llm_provider": custom_llm_provider}
|
||||
response = convert_to_model_response_object(
|
||||
_response_headers=headers,
|
||||
response_object=stringified_response,
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from litellm.types.utils import Usage
|
|||
from litellm.utils import get_model_info
|
||||
|
||||
|
||||
def _is_azure_model_router(model: str) -> bool:
|
||||
def is_azure_model_router(model: str) -> bool:
|
||||
"""
|
||||
Check if the model is Azure AI Foundry Model Router.
|
||||
|
||||
|
|
@ -31,6 +31,18 @@ def _is_azure_model_router(model: str) -> bool:
|
|||
return "model-router" in model_lower or "model_router" in model_lower or model_lower == "azure-model-router"
|
||||
|
||||
|
||||
ROUTER_FEE_ENTRY_NAMES: Final = frozenset({"model-router", "model_router"})
|
||||
|
||||
|
||||
def is_router_fee_entry(model: str) -> bool:
|
||||
return model.lower().removeprefix("azure_ai/") in ROUTER_FEE_ENTRY_NAMES
|
||||
|
||||
|
||||
def _router_fee_entry_name(model: str) -> str:
|
||||
entry_name: Final = model.lower().removeprefix("azure_ai/")
|
||||
return entry_name if entry_name in ROUTER_FEE_ENTRY_NAMES else "model_router"
|
||||
|
||||
|
||||
def calculate_azure_model_router_flat_cost(model: str, prompt_tokens: int) -> float:
|
||||
"""
|
||||
Calculate the flat cost for Azure AI Foundry Model Router.
|
||||
|
|
@ -42,20 +54,39 @@ def calculate_azure_model_router_flat_cost(model: str, prompt_tokens: int) -> fl
|
|||
Returns:
|
||||
float: The flat cost in USD, or 0.0 if not applicable
|
||||
"""
|
||||
if not _is_azure_model_router(model):
|
||||
if not is_azure_model_router(model):
|
||||
return 0.0
|
||||
|
||||
# Get the model router pricing from model_prices_and_context_window.json
|
||||
# Use "model_router" as the key (without actual model name suffix)
|
||||
model_info: Final = get_model_info(model="model_router", custom_llm_provider="azure_ai")
|
||||
model_info: Final = get_model_info(model=_router_fee_entry_name(model), custom_llm_provider="azure_ai")
|
||||
router_flat_cost_per_token: Final = model_info.get("input_cost_per_token", 0)
|
||||
|
||||
if router_flat_cost_per_token and router_flat_cost_per_token > 0:
|
||||
return prompt_tokens * router_flat_cost_per_token
|
||||
|
||||
return 0.0
|
||||
|
||||
|
||||
def _response_model_cost(model: str, usage: Usage, service_tier: str | None) -> tuple[float, float]:
|
||||
try:
|
||||
return generic_cost_per_token(
|
||||
model=model, usage=usage, custom_llm_provider="azure_ai", service_tier=service_tier
|
||||
)
|
||||
except Exception as e:
|
||||
if not is_azure_model_router(model):
|
||||
raise
|
||||
verbose_logger.debug(
|
||||
"Azure AI Model Router: model '%s' not in cost map, only the routing fee applies. Error: %s", model, e
|
||||
)
|
||||
return 0.0, 0.0
|
||||
|
||||
|
||||
def _router_fee_name(model: str, request_model: str | None) -> str | None:
|
||||
if is_router_fee_entry(model):
|
||||
return None
|
||||
if is_azure_model_router(model):
|
||||
return model
|
||||
if request_model is not None and is_azure_model_router(request_model):
|
||||
return request_model
|
||||
return None
|
||||
|
||||
|
||||
def cost_per_token(
|
||||
model: str,
|
||||
usage: Usage,
|
||||
|
|
@ -64,68 +95,31 @@ def cost_per_token(
|
|||
service_tier: str | None = None,
|
||||
) -> tuple[float, float]:
|
||||
"""
|
||||
Calculate the cost per token for Azure AI models.
|
||||
Price the response model's own tokens for Azure AI, plus the Model Router fee exactly once when either the
|
||||
priced name or request_model is a Model Router name.
|
||||
|
||||
For Azure AI Foundry Model Router:
|
||||
- Adds a flat cost of $0.14 per million input tokens (from model_prices_and_context_window.json)
|
||||
- Plus the cost of the actual model used (handled by generic_cost_per_token)
|
||||
A response priced as the router entry itself already carries the fee, so nothing is added on top of it. A
|
||||
router deployment name that is missing from the cost map prices at the fee alone.
|
||||
|
||||
completion_cost passes only the priced name: when that name is a routed model it adds the fee itself through
|
||||
AzureModelRouterConfig.calculate_additional_costs as the "Azure Model Router Flat Cost" line of the cost
|
||||
breakdown, and when the name is router-shaped the fee is already in the prompt cost returned here.
|
||||
|
||||
Args:
|
||||
model: str, the model name without provider prefix (from response)
|
||||
usage: LiteLLM Usage block
|
||||
response_time_ms: Optional response time in milliseconds
|
||||
request_model: Optional[str], the original request model name (to detect router usage)
|
||||
request_model: Optional[str], the original request model name; a Model Router name adds the routing fee
|
||||
service_tier: Optional service tier the request was priced on
|
||||
|
||||
Returns:
|
||||
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
|
||||
|
||||
Raises:
|
||||
ValueError: If the model is not found in the cost map and cost cannot be calculated
|
||||
(except for Model Router models where we return just the routing flat cost)
|
||||
ValueError: If a model that is not a Model Router name is missing from the cost map
|
||||
"""
|
||||
prompt_cost = 0.0
|
||||
completion_cost = 0.0
|
||||
|
||||
# Determine if this was a model router request
|
||||
# Check both the response model and the request model
|
||||
is_router_request: Final = _is_azure_model_router(model) or (
|
||||
request_model is not None and _is_azure_model_router(request_model)
|
||||
)
|
||||
|
||||
# Calculate base cost using generic cost calculator
|
||||
# This may raise an exception if the model is not in the cost map
|
||||
try:
|
||||
prompt_cost, completion_cost = generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider="azure_ai",
|
||||
service_tier=service_tier,
|
||||
)
|
||||
except Exception as e:
|
||||
# For Model Router, the model name (e.g., "azure-model-router") may not be in the cost map
|
||||
# because it's a routing service, not an actual model. In this case, we continue
|
||||
# to calculate just the routing flat cost.
|
||||
if not _is_azure_model_router(model):
|
||||
# Re-raise for non-router models - they should have pricing defined
|
||||
raise
|
||||
verbose_logger.debug(
|
||||
"Azure AI Model Router: model '%s' not in cost map, calculating routing flat cost only. Error: %s", model, e
|
||||
)
|
||||
|
||||
# Add flat cost for Azure Model Router
|
||||
# The flat cost is defined in model_prices_and_context_window.json for azure_ai/model_router
|
||||
if is_router_request:
|
||||
# Use the request model for flat cost calculation if available, otherwise use response model
|
||||
router_model_for_calc: Final = request_model if request_model else model
|
||||
router_flat_cost: Final = calculate_azure_model_router_flat_cost(router_model_for_calc, usage.prompt_tokens)
|
||||
|
||||
if router_flat_cost > 0:
|
||||
verbose_logger.debug(
|
||||
f"Azure AI Model Router flat cost: ${router_flat_cost:.6f} "
|
||||
f"({usage.prompt_tokens} tokens × ${router_flat_cost / usage.prompt_tokens:.9f}/token)"
|
||||
)
|
||||
|
||||
# Add flat cost to prompt cost
|
||||
prompt_cost += router_flat_cost
|
||||
|
||||
return prompt_cost, completion_cost
|
||||
prompt_cost, completion_cost = _response_model_cost(model=model, usage=usage, service_tier=service_tier)
|
||||
fee_name: Final = _router_fee_name(model=model, request_model=request_model)
|
||||
if fee_name is None:
|
||||
return prompt_cost, completion_cost
|
||||
return prompt_cost + calculate_azure_model_router_flat_cost(fee_name, usage.prompt_tokens), completion_cost
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -52,6 +52,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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -35,7 +35,7 @@ from .amazon_titan_multimodal_transformation import (
|
|||
)
|
||||
from .amazon_titan_v2_transformation import AmazonTitanV2Config
|
||||
from .cohere_transformation import BedrockCohereEmbeddingConfig
|
||||
from .twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig
|
||||
from .twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig, drop_params_enabled
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -239,7 +239,7 @@ class BedrockEmbedding(BaseAWSLLM):
|
|||
returned_response = AmazonTitanG1Config()._transform_response(response_list=response_list, model=model)
|
||||
elif provider == "twelvelabs":
|
||||
returned_response = TwelveLabsMarengoEmbeddingConfig()._transform_response(
|
||||
response_list=response_list, model=model
|
||||
response_list=response_list, model=model, batch_data=batch_data
|
||||
)
|
||||
elif provider == "nova":
|
||||
returned_response = AmazonNovaEmbeddingConfig()._transform_response(
|
||||
|
|
@ -484,12 +484,13 @@ class BedrockEmbedding(BaseAWSLLM):
|
|||
elif provider == "twelvelabs":
|
||||
batch_data = []
|
||||
for i in input:
|
||||
twelvelabs_request = TwelveLabsMarengoEmbeddingConfig()._transform_request(
|
||||
twelvelabs_request = TwelveLabsMarengoEmbeddingConfig(model=model)._transform_request(
|
||||
input=i,
|
||||
inference_params=inference_params,
|
||||
async_invoke_route=has_async_invoke,
|
||||
model_id=modelId,
|
||||
output_s3_uri=inference_params.get("output_s3_uri"),
|
||||
drop_params=drop_params_enabled(litellm_params),
|
||||
)
|
||||
batch_data.append(twelvelabs_request)
|
||||
elif provider == "nova":
|
||||
|
|
|
|||
|
|
@ -0,0 +1,239 @@
|
|||
"""
|
||||
Request builder for Bedrock TwelveLabs Marengo Embed 3.0, whose payload nests the input under a key named after
|
||||
``inputType`` instead of the flat 2.7 layout.
|
||||
|
||||
Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo-3.html
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from litellm.types.llms.bedrock import (
|
||||
TWELVELABS_MARENGO_3_EMBEDDING_OPTIONS,
|
||||
TWELVELABS_MARENGO_3_EMBEDDING_SCOPES,
|
||||
TWELVELABS_MARENGO_3_EMBEDDING_TYPES,
|
||||
TWELVELABS_MARENGO_3_INPUT_TYPES,
|
||||
TwelveLabsMarengo3AudioRequest,
|
||||
TwelveLabsMarengo3EmbeddingRequest,
|
||||
TwelveLabsMarengo3ImageRequest,
|
||||
TwelveLabsMarengo3MultiInputRequest,
|
||||
TwelveLabsMarengo3NamedMediaSource,
|
||||
TwelveLabsMarengo3RequestBase,
|
||||
TwelveLabsMarengo3Segmentation,
|
||||
TwelveLabsMarengo3TextImageRequest,
|
||||
TwelveLabsMarengo3TextRequest,
|
||||
TwelveLabsMarengo3TimedMediaInput,
|
||||
TwelveLabsMarengo3TimedMediaOptions,
|
||||
TwelveLabsMarengo3VideoRequest,
|
||||
TwelveLabsMediaSource,
|
||||
TwelveLabsS3Location,
|
||||
)
|
||||
from litellm.utils import get_base64_str
|
||||
|
||||
MARENGO_3_MODEL_MARKER: Final = "marengo-embed-3-"
|
||||
S3_URI_PREFIX: Final = "s3://"
|
||||
TIMED_MEDIA_OPTION_FIELDS: Final = MappingProxyType(
|
||||
{
|
||||
"startSec": True,
|
||||
"endSec": True,
|
||||
"segmentation": True,
|
||||
"embeddingOption": True,
|
||||
"embeddingType": True,
|
||||
"embeddingScope": True,
|
||||
}
|
||||
)
|
||||
TIMED_MEDIA_OPTIONS: Final = TypeAdapter(TwelveLabsMarengo3TimedMediaOptions)
|
||||
TIMED_INPUT_TYPES: Final = frozenset({"video", "audio"})
|
||||
MARENGO_2_7_ONLY_PARAMS: Final = ("textTruncate", "lengthSec", "useFixedLengthSec", "minClipSec")
|
||||
MARENGO_2_7_ONLY_FIELDS: Final = MappingProxyType({name: True for name in MARENGO_2_7_ONLY_PARAMS})
|
||||
|
||||
|
||||
def is_marengo_3_model(model: str | None) -> bool:
|
||||
return MARENGO_3_MODEL_MARKER in (model or "")
|
||||
|
||||
|
||||
class Marengo3Params(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
inputType: TWELVELABS_MARENGO_3_INPUT_TYPES | None = None
|
||||
input_type: TWELVELABS_MARENGO_3_INPUT_TYPES | None = None
|
||||
media_source: str | None = None
|
||||
media_sources: Mapping[str, str] | None = None
|
||||
bucketOwner: str | None = None
|
||||
startSec: float | None = None
|
||||
endSec: float | None = None
|
||||
segmentation: TwelveLabsMarengo3Segmentation | None = None
|
||||
embeddingOption: tuple[TWELVELABS_MARENGO_3_EMBEDDING_OPTIONS, ...] | None = None
|
||||
embeddingType: tuple[TWELVELABS_MARENGO_3_EMBEDDING_TYPES, ...] | None = None
|
||||
embeddingScope: tuple[TWELVELABS_MARENGO_3_EMBEDDING_SCOPES, ...] | None = None
|
||||
inferenceId: str | None = None
|
||||
textTruncate: object = None
|
||||
lengthSec: object = None
|
||||
useFixedLengthSec: object = None
|
||||
minClipSec: object = None
|
||||
|
||||
@property
|
||||
def resolved_input_type(self) -> TWELVELABS_MARENGO_3_INPUT_TYPES:
|
||||
return self.inputType or self.input_type or "text"
|
||||
|
||||
def timed_media_options(self) -> TwelveLabsMarengo3TimedMediaOptions:
|
||||
return TIMED_MEDIA_OPTIONS.validate_python(self.given_timed_media_options())
|
||||
|
||||
def given_timed_media_options(self) -> dict[str, object]:
|
||||
return self.model_dump(include=TIMED_MEDIA_OPTION_FIELDS, exclude_none=True)
|
||||
|
||||
def given_2_7_only_params(self) -> dict[str, object]:
|
||||
return self.model_dump(include=MARENGO_2_7_ONLY_FIELDS, exclude_none=True)
|
||||
|
||||
|
||||
def _require_bucket_owner(bucket_owner: str | None) -> str:
|
||||
if bucket_owner is None:
|
||||
raise BedrockError(
|
||||
status_code=400,
|
||||
message="s3:// media requires the 'bucketOwner' parameter, the account id that owns the bucket",
|
||||
)
|
||||
return bucket_owner
|
||||
|
||||
|
||||
def _media_source(media: str, bucket_owner: str | None) -> TwelveLabsMediaSource:
|
||||
if not media.startswith(S3_URI_PREFIX):
|
||||
inline: Final[TwelveLabsMediaSource] = {"base64String": get_base64_str(media)}
|
||||
return inline
|
||||
s3_location: Final[TwelveLabsS3Location] = {"uri": media, "bucketOwner": _require_bucket_owner(bucket_owner)}
|
||||
remote: Final[TwelveLabsMediaSource] = {"s3Location": s3_location}
|
||||
return remote
|
||||
|
||||
|
||||
def _named_media_source(name: str, media: str, bucket_owner: str | None) -> TwelveLabsMarengo3NamedMediaSource:
|
||||
named: Final[TwelveLabsMarengo3NamedMediaSource] = {
|
||||
"name": name,
|
||||
"mediaType": "image",
|
||||
**_media_source(media, bucket_owner),
|
||||
}
|
||||
return named
|
||||
|
||||
|
||||
def _timed_media_input(media: str, params: Marengo3Params) -> TwelveLabsMarengo3TimedMediaInput:
|
||||
timed: Final[TwelveLabsMarengo3TimedMediaInput] = {
|
||||
"mediaSource": _media_source(media, params.bucketOwner),
|
||||
**params.timed_media_options(),
|
||||
}
|
||||
return timed
|
||||
|
||||
|
||||
def _describe(error: ValidationError) -> str:
|
||||
return "; ".join(
|
||||
f"{'.'.join(str(part) for part in problem['loc'])}: {problem['msg']}" for problem in error.errors()
|
||||
)
|
||||
|
||||
|
||||
def _validated_params(inference_params: Mapping[str, object]) -> Marengo3Params:
|
||||
try:
|
||||
return Marengo3Params.model_validate(inference_params)
|
||||
except ValidationError as error:
|
||||
raise BedrockError(status_code=400, message=f"Invalid Marengo 3.0 parameters: {_describe(error)}") from error
|
||||
|
||||
|
||||
def _reject_unless_dropped(given: Mapping[str, object], drop_params: bool, reason: str) -> None:
|
||||
if not given or drop_params:
|
||||
return
|
||||
raise BedrockError(status_code=400, message=f"{reason} {', '.join(given)}; set drop_params to drop them")
|
||||
|
||||
|
||||
def _require(value: str | None, input_type: str, param_name: str) -> str:
|
||||
if value is None:
|
||||
raise BedrockError(status_code=400, message=f"Input type '{input_type}' requires the '{param_name}' parameter")
|
||||
return value
|
||||
|
||||
|
||||
def _require_media_sources(value: Mapping[str, str] | None) -> Mapping[str, str]:
|
||||
if not value:
|
||||
raise BedrockError(
|
||||
status_code=400,
|
||||
message="Input type 'multi_input' requires a non-empty 'media_sources' mapping of name to media",
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def _request_base(inference_id: str | None) -> TwelveLabsMarengo3RequestBase:
|
||||
if inference_id is None:
|
||||
anonymous: Final[TwelveLabsMarengo3RequestBase] = {}
|
||||
return anonymous
|
||||
identified: Final[TwelveLabsMarengo3RequestBase] = {"inferenceId": inference_id}
|
||||
return identified
|
||||
|
||||
|
||||
def build_marengo_3_request(
|
||||
input: str, inference_params: Mapping[str, object], drop_params: bool = False
|
||||
) -> TwelveLabsMarengo3EmbeddingRequest:
|
||||
params: Final = _validated_params(inference_params)
|
||||
base: Final = _request_base(params.inferenceId)
|
||||
input_type: Final = params.resolved_input_type
|
||||
_reject_unless_dropped(
|
||||
params.given_2_7_only_params(), drop_params, "Marengo 3.0 does not accept the Marengo 2.7 parameters"
|
||||
)
|
||||
if input_type not in TIMED_INPUT_TYPES:
|
||||
_reject_unless_dropped(
|
||||
params.given_timed_media_options(), drop_params, f"Input type '{input_type}' does not accept"
|
||||
)
|
||||
match input_type:
|
||||
case "text":
|
||||
text_request: Final[TwelveLabsMarengo3TextRequest] = {
|
||||
**base,
|
||||
"inputType": "text",
|
||||
"text": {"inputText": input},
|
||||
}
|
||||
return text_request
|
||||
case "image":
|
||||
image_request: Final[TwelveLabsMarengo3ImageRequest] = {
|
||||
**base,
|
||||
"inputType": "image",
|
||||
"image": {"mediaSource": _media_source(input, params.bucketOwner)},
|
||||
}
|
||||
return image_request
|
||||
case "video":
|
||||
video_request: Final[TwelveLabsMarengo3VideoRequest] = {
|
||||
**base,
|
||||
"inputType": "video",
|
||||
"video": _timed_media_input(input, params),
|
||||
}
|
||||
return video_request
|
||||
case "audio":
|
||||
audio_request: Final[TwelveLabsMarengo3AudioRequest] = {
|
||||
**base,
|
||||
"inputType": "audio",
|
||||
"audio": _timed_media_input(input, params),
|
||||
}
|
||||
return audio_request
|
||||
case "text_image":
|
||||
text_image_request: Final[TwelveLabsMarengo3TextImageRequest] = {
|
||||
**base,
|
||||
"inputType": "text_image",
|
||||
"text_image": {
|
||||
"inputText": input,
|
||||
"mediaSource": _media_source(
|
||||
_require(params.media_source, input_type, "media_source"), params.bucketOwner
|
||||
),
|
||||
},
|
||||
}
|
||||
return text_image_request
|
||||
case "multi_input":
|
||||
media_sources: Final = tuple(
|
||||
_named_media_source(name, media, params.bucketOwner)
|
||||
for name, media in _require_media_sources(params.media_sources).items()
|
||||
)
|
||||
multi_input_request: Final[TwelveLabsMarengo3MultiInputRequest] = {
|
||||
**base,
|
||||
"inputType": "multi_input",
|
||||
"multi_input": {"inputText": input, "mediaSources": media_sources}
|
||||
if input
|
||||
else {"mediaSources": media_sources},
|
||||
}
|
||||
return multi_input_request
|
||||
case _:
|
||||
assert_never(input_type)
|
||||
|
|
@ -4,19 +4,120 @@ Transformation logic from OpenAI /v1/embeddings format to Bedrock TwelveLabs Mar
|
|||
Why separate file? Make it easy to see how transformation works
|
||||
|
||||
Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo.html
|
||||
Marengo 3.0 docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo-3.html
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, cast
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter
|
||||
from typing_extensions import assert_never
|
||||
|
||||
import litellm
|
||||
from litellm.llms.bedrock.embed.twelvelabs_marengo_3_transformation import (
|
||||
MARENGO_2_7_ONLY_PARAMS,
|
||||
build_marengo_3_request,
|
||||
is_marengo_3_model,
|
||||
)
|
||||
from litellm.types.llms.bedrock import (
|
||||
TWELVELABS_EMBEDDING_INPUT_TYPES,
|
||||
TWELVELABS_MARENGO_3_INPUT_TYPES,
|
||||
TwelveLabsAsyncInvokeRequest,
|
||||
TwelveLabsMarengo3EmbeddingRequest,
|
||||
TwelveLabsMarengoEmbeddingRequest,
|
||||
TwelveLabsOutputDataConfig,
|
||||
TwelveLabsS3Location,
|
||||
TwelveLabsS3OutputDataConfig,
|
||||
)
|
||||
from litellm.types.utils import Embedding, EmbeddingResponse, Usage
|
||||
from litellm.types.utils import Embedding, EmbeddingResponse, PromptTokensDetailsWrapper, Usage
|
||||
|
||||
|
||||
class MarengoEmbeddingItem(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
embedding: tuple[float, ...] | None = None
|
||||
|
||||
|
||||
class MarengoInvokeResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
data: tuple[MarengoEmbeddingItem, ...] = ()
|
||||
embedding: tuple[float, ...] | None = None
|
||||
embeddings: tuple[MarengoEmbeddingItem, ...] = ()
|
||||
|
||||
def vectors(self) -> tuple[tuple[float, ...], ...]:
|
||||
if self.data:
|
||||
return tuple(item.embedding for item in self.data if item.embedding is not None)
|
||||
if self.embedding is not None:
|
||||
return (self.embedding,)
|
||||
return tuple(item.embedding for item in self.embeddings if item.embedding is not None)
|
||||
|
||||
|
||||
class MarengoBilledMultiInput(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
inputText: str | None = None
|
||||
mediaSources: tuple[Mapping[str, object], ...] = ()
|
||||
|
||||
|
||||
class MarengoBilledRequest(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
inputType: TWELVELABS_MARENGO_3_INPUT_TYPES | None = None
|
||||
multi_input: MarengoBilledMultiInput | None = None
|
||||
|
||||
|
||||
INVOKE_RESPONSES: Final = TypeAdapter(tuple[MarengoInvokeResponse, ...])
|
||||
BILLED_REQUESTS: Final = TypeAdapter(tuple[MarengoBilledRequest, ...])
|
||||
|
||||
|
||||
def _billed_units(request: MarengoBilledRequest) -> tuple[int, int]:
|
||||
input_type: Final = request.inputType
|
||||
match input_type:
|
||||
case "text":
|
||||
return (1, 0)
|
||||
case "image":
|
||||
return (0, 1)
|
||||
case "text_image":
|
||||
return (1, 1)
|
||||
case "multi_input":
|
||||
multi_input: Final = request.multi_input or MarengoBilledMultiInput()
|
||||
return (1 if multi_input.inputText else 0, len(multi_input.mediaSources))
|
||||
case "video" | "audio" | None:
|
||||
return (0, 0)
|
||||
case _:
|
||||
assert_never(input_type)
|
||||
|
||||
|
||||
def _billed_usage(batch_data: list[dict] | None) -> Usage:
|
||||
units: Final = tuple(_billed_units(request) for request in BILLED_REQUESTS.validate_python(batch_data or ()))
|
||||
query_count: Final = sum(text_requests for text_requests, _ in units)
|
||||
image_count: Final = sum(images for _, images in units)
|
||||
details: Final = (
|
||||
PromptTokensDetailsWrapper(query_count=query_count or None, image_count=image_count or None)
|
||||
if query_count or image_count
|
||||
else None
|
||||
)
|
||||
return Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0, prompt_tokens_details=details)
|
||||
|
||||
|
||||
MARENGO_SHARED_PARAMS: Final = (
|
||||
"encoding_format",
|
||||
"embeddingOption",
|
||||
"startSec",
|
||||
"input_type",
|
||||
"endSec",
|
||||
"segmentation",
|
||||
"embeddingType",
|
||||
"embeddingScope",
|
||||
"inferenceId",
|
||||
"media_source",
|
||||
"media_sources",
|
||||
)
|
||||
|
||||
|
||||
def drop_params_enabled(litellm_params: Mapping[str, object]) -> bool:
|
||||
return litellm.drop_params is True or litellm_params.get("drop_params") is True
|
||||
|
||||
|
||||
class TwelveLabsMarengoEmbeddingConfig:
|
||||
|
|
@ -26,28 +127,24 @@ class TwelveLabsMarengoEmbeddingConfig:
|
|||
Supports text, image, video, and audio inputs.
|
||||
- InvokeModel: text and image inputs
|
||||
- StartAsyncInvoke: video, audio, image, and text inputs
|
||||
|
||||
Marengo 3.0 (model ids containing "marengo-embed-3") nests the input under a key named after inputType and
|
||||
adds the text_image and multi_input input types; that payload is built by build_marengo_3_request.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
def __init__(self, model: str | None = None) -> None:
|
||||
self.is_marengo_3: Final = is_marengo_3_model(model)
|
||||
|
||||
def get_supported_openai_params(self) -> list[str]:
|
||||
return [
|
||||
"encoding_format",
|
||||
"textTruncate",
|
||||
"embeddingOption",
|
||||
"startSec",
|
||||
"lengthSec",
|
||||
"useFixedLengthSec",
|
||||
"minClipSec",
|
||||
"input_type",
|
||||
]
|
||||
if self.is_marengo_3:
|
||||
return list(MARENGO_SHARED_PARAMS)
|
||||
return [*MARENGO_SHARED_PARAMS, *MARENGO_2_7_ONLY_PARAMS]
|
||||
|
||||
def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict:
|
||||
for k, v in non_default_params.items():
|
||||
if k == "encoding_format":
|
||||
# TwelveLabs doesn't have encoding_format, but we can map it to embeddingOption
|
||||
if v == "float":
|
||||
if v == "float" and not self.is_marengo_3:
|
||||
optional_params["embeddingOption"] = ["visual-text", "visual-image"]
|
||||
elif k == "textTruncate":
|
||||
optional_params["textTruncate"] = v
|
||||
|
|
@ -56,7 +153,19 @@ class TwelveLabsMarengoEmbeddingConfig:
|
|||
elif k == "input_type":
|
||||
# Map input_type to inputType for Bedrock
|
||||
optional_params["inputType"] = v
|
||||
elif k in ["startSec", "lengthSec", "useFixedLengthSec", "minClipSec"]:
|
||||
elif k in (
|
||||
"startSec",
|
||||
"lengthSec",
|
||||
"useFixedLengthSec",
|
||||
"minClipSec",
|
||||
"endSec",
|
||||
"segmentation",
|
||||
"embeddingType",
|
||||
"embeddingScope",
|
||||
"inferenceId",
|
||||
"media_source",
|
||||
"media_sources",
|
||||
):
|
||||
optional_params[k] = v
|
||||
return optional_params
|
||||
|
||||
|
|
@ -77,7 +186,8 @@ class TwelveLabsMarengoEmbeddingConfig:
|
|||
async_invoke_route: bool = False,
|
||||
model_id: str | None = None,
|
||||
output_s3_uri: str | None = None,
|
||||
) -> TwelveLabsMarengoEmbeddingRequest | TwelveLabsAsyncInvokeRequest:
|
||||
drop_params: bool = False,
|
||||
) -> TwelveLabsMarengoEmbeddingRequest | TwelveLabsMarengo3EmbeddingRequest | TwelveLabsAsyncInvokeRequest:
|
||||
"""
|
||||
Transform OpenAI-style input to TwelveLabs Marengo format/async-invoke format.
|
||||
|
||||
|
|
@ -87,20 +197,29 @@ class TwelveLabsMarengoEmbeddingConfig:
|
|||
- Video inputs (async-invoke only)
|
||||
- Audio inputs (async-invoke only)
|
||||
- S3 URLs for all media types (async-invoke only)
|
||||
- Marengo 3.0 only: text_image and multi_input inputs (nested payload)
|
||||
"""
|
||||
# Get input_type or default to "text"
|
||||
input_type: Final = cast(
|
||||
TWELVELABS_EMBEDDING_INPUT_TYPES,
|
||||
inference_params.get("inputType") or inference_params.get("input_type") or "text",
|
||||
)
|
||||
|
||||
# Validate that async-invoke is used for video/audio
|
||||
if input_type in ["video", "audio"] and not async_invoke_route:
|
||||
raise ValueError(
|
||||
f"Input type '{input_type}' requires async_invoke route. "
|
||||
f"Use model format: 'bedrock/async_invoke/model_id'"
|
||||
)
|
||||
|
||||
if self.is_marengo_3:
|
||||
marengo_3_request: Final = build_marengo_3_request(
|
||||
input=input, inference_params=inference_params, drop_params=drop_params
|
||||
)
|
||||
if async_invoke_route and model_id:
|
||||
return self._wrap_async_invoke_request(
|
||||
model_input=marengo_3_request, model_id=model_id, output_s3_uri=output_s3_uri
|
||||
)
|
||||
return marengo_3_request
|
||||
|
||||
transformed_request: Final[TwelveLabsMarengoEmbeddingRequest] = {"inputType": input_type}
|
||||
|
||||
if input_type == "text":
|
||||
|
|
@ -154,7 +273,7 @@ class TwelveLabsMarengoEmbeddingConfig:
|
|||
|
||||
def _wrap_async_invoke_request(
|
||||
self,
|
||||
model_input: TwelveLabsMarengoEmbeddingRequest,
|
||||
model_input: TwelveLabsMarengoEmbeddingRequest | TwelveLabsMarengo3EmbeddingRequest,
|
||||
model_id: str,
|
||||
output_s3_uri: str | None = None,
|
||||
) -> TwelveLabsAsyncInvokeRequest:
|
||||
|
|
@ -188,62 +307,16 @@ class TwelveLabsMarengoEmbeddingConfig:
|
|||
),
|
||||
)
|
||||
|
||||
def _transform_response(self, response_list: list[dict], model: str) -> EmbeddingResponse:
|
||||
"""
|
||||
Transform TwelveLabs response to OpenAI format.
|
||||
Handles the actual TwelveLabs response format: {"data": [{"embedding": [...]}]}
|
||||
"""
|
||||
embeddings: Final[list[Embedding]] = []
|
||||
total_tokens = 0
|
||||
|
||||
for response in response_list:
|
||||
# TwelveLabs response format has a "data" field containing the embeddings
|
||||
if "data" in response and isinstance(response["data"], list):
|
||||
for item in response["data"]:
|
||||
if "embedding" in item:
|
||||
# Single embedding response
|
||||
embedding = Embedding(
|
||||
embedding=item["embedding"],
|
||||
index=len(embeddings),
|
||||
object="embedding",
|
||||
)
|
||||
embeddings.append(embedding)
|
||||
|
||||
# Estimate token count (rough approximation)
|
||||
if "inputTextTokenCount" in item:
|
||||
total_tokens += item["inputTextTokenCount"]
|
||||
else:
|
||||
# Rough estimate: 1 token per 4 characters for text, or use embedding size
|
||||
total_tokens += len(item["embedding"]) // 4
|
||||
elif "embedding" in response:
|
||||
# Direct embedding response (fallback for other formats)
|
||||
embedding = Embedding(
|
||||
embedding=response["embedding"],
|
||||
index=len(embeddings),
|
||||
object="embedding",
|
||||
)
|
||||
embeddings.append(embedding)
|
||||
|
||||
# Estimate token count (rough approximation)
|
||||
if "inputTextTokenCount" in response:
|
||||
total_tokens += response["inputTextTokenCount"]
|
||||
else:
|
||||
# Rough estimate: 1 token per 4 characters for text
|
||||
total_tokens += len(response.get("inputText", "")) // 4
|
||||
elif "embeddings" in response:
|
||||
# Multiple embeddings response (from video/audio)
|
||||
for i, emb in enumerate(response["embeddings"]):
|
||||
embedding = Embedding(
|
||||
embedding=emb["embedding"],
|
||||
index=len(embeddings),
|
||||
object="embedding",
|
||||
)
|
||||
embeddings.append(embedding)
|
||||
total_tokens += len(emb["embedding"]) // 4 # Rough estimate
|
||||
|
||||
usage: Final = Usage(prompt_tokens=total_tokens, total_tokens=total_tokens)
|
||||
|
||||
return EmbeddingResponse(data=embeddings, model=model, usage=usage)
|
||||
def _transform_response(
|
||||
self, response_list: list[dict], model: str, batch_data: list[dict] | None = None
|
||||
) -> EmbeddingResponse:
|
||||
vectors: Final = tuple(
|
||||
vector for response in INVOKE_RESPONSES.validate_python(response_list) for vector in response.vectors()
|
||||
)
|
||||
embeddings: Final = [
|
||||
Embedding(embedding=list(vector), index=index, object="embedding") for index, vector in enumerate(vectors)
|
||||
]
|
||||
return EmbeddingResponse(data=embeddings, model=model, usage=_billed_usage(batch_data))
|
||||
|
||||
def _transform_async_invoke_response(self, response: dict, model: str) -> EmbeddingResponse:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -9826,7 +9826,7 @@ class BaseLLMHTTPHandler:
|
|||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
litellm_params=MappingProxyType(dict(litellm_params, timeout=timeout)),
|
||||
extra_body=extra_body,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
|
|
@ -9871,6 +9871,12 @@ class BaseLLMHTTPHandler:
|
|||
data=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise vector_store_provider_config.get_error_class(
|
||||
error_message="Vector store search exceeded the caller timeout.",
|
||||
status_code=408,
|
||||
headers=httpx.Headers(),
|
||||
) from None
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
|
||||
|
||||
|
|
@ -9955,7 +9961,7 @@ class BaseLLMHTTPHandler:
|
|||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
litellm_params=MappingProxyType(dict(litellm_params, timeout=timeout)),
|
||||
extra_body=extra_body,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
|
|
@ -10000,7 +10006,14 @@ class BaseLLMHTTPHandler:
|
|||
url=url,
|
||||
headers=headers,
|
||||
data=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise vector_store_provider_config.get_error_class(
|
||||
error_message="Vector store search exceeded the caller timeout.",
|
||||
status_code=408,
|
||||
headers=httpx.Headers(),
|
||||
) from None
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
|
||||
|
||||
|
|
@ -10030,6 +10043,8 @@ class BaseLLMHTTPHandler:
|
|||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
vector_store_provider_config.validate_create_vector_store()
|
||||
|
||||
headers: Final = vector_store_provider_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=litellm_params
|
||||
)
|
||||
|
|
@ -10100,6 +10115,8 @@ class BaseLLMHTTPHandler:
|
|||
else:
|
||||
sync_httpx_client = client
|
||||
|
||||
vector_store_provider_config.validate_create_vector_store()
|
||||
|
||||
headers: Final = vector_store_provider_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=litellm_params
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from urllib.parse import unquote
|
||||
|
||||
import httpx
|
||||
from openai.types.responses import EasyInputMessageParam, ResponseInputItemParam
|
||||
from openai.types.responses import EasyInputMessageParam, ResponseInputContentParam, ResponseInputItemParam
|
||||
|
||||
from litellm.llms.fireworks_ai.common_utils import (
|
||||
resolve_fireworks_api_key,
|
||||
|
|
@ -31,6 +31,17 @@ def _session_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object
|
|||
)
|
||||
|
||||
|
||||
_INSTRUCTION_ROLES: Final = frozenset({"system", "developer"})
|
||||
|
||||
|
||||
def _role(item: ResponseInputItemParam) -> str | None:
|
||||
match item:
|
||||
case {"role": str(role)}:
|
||||
return role
|
||||
case _:
|
||||
return None
|
||||
|
||||
|
||||
def _developer_item_as_system(item: ResponseInputItemParam) -> ResponseInputItemParam:
|
||||
if "role" not in item or item["role"] != "developer":
|
||||
return item
|
||||
|
|
@ -43,6 +54,69 @@ def _developer_items_as_system(input: str | ResponseInputParam) -> str | Respons
|
|||
return [_developer_item_as_system(item) for item in input]
|
||||
|
||||
|
||||
def _text_part(part: ResponseInputContentParam) -> str | None:
|
||||
match part:
|
||||
case {"type": "input_text", "text": str(text)}:
|
||||
return text
|
||||
case _:
|
||||
return None
|
||||
|
||||
|
||||
def _text_only_content(item: ResponseInputItemParam) -> str | None:
|
||||
match item:
|
||||
case {"role": "system" | "developer", "content": str(text)}:
|
||||
return text
|
||||
case {"role": "system" | "developer", "content": [*parts]}:
|
||||
texts: Final = tuple(map(_text_part, parts))
|
||||
return None if any(text is None for text in texts) else "\n\n".join(text for text in texts if text)
|
||||
case _:
|
||||
return None
|
||||
|
||||
|
||||
def _leading_instruction_block_length(roles: Sequence[str | None]) -> int:
|
||||
return next((index for index, role in enumerate(roles) if role not in _INSTRUCTION_ROLES), len(roles))
|
||||
|
||||
|
||||
def _closing_instruction_block_start(roles: Sequence[str | None], leading_length: int) -> int:
|
||||
last_conversation_index: Final = next(
|
||||
(index for index in range(len(roles) - 1, leading_length - 1, -1) if roles[index] not in _INSTRUCTION_ROLES),
|
||||
None,
|
||||
)
|
||||
if last_conversation_index is None or roles[last_conversation_index] != "assistant":
|
||||
return len(roles)
|
||||
return last_conversation_index + 1
|
||||
|
||||
|
||||
def _hoisted_indices(roles: Sequence[str | None]) -> tuple[int, ...]:
|
||||
leading_length: Final = _leading_instruction_block_length(roles)
|
||||
closing_start: Final = _closing_instruction_block_start(roles, leading_length)
|
||||
return tuple(
|
||||
index for index, role in enumerate(roles[:closing_start]) if index < leading_length or role == "developer"
|
||||
)
|
||||
|
||||
|
||||
def _with_instruction_items_folded(
|
||||
input: str | ResponseInputParam, instructions: str | None
|
||||
) -> tuple[str | None, str | ResponseInputParam]:
|
||||
if isinstance(input, str):
|
||||
return instructions, input
|
||||
items: Final = tuple(input)
|
||||
folded: Final = MappingProxyType(
|
||||
{
|
||||
index: text
|
||||
for index in _hoisted_indices(tuple(map(_role, items)))
|
||||
if (text := _text_only_content(items[index])) is not None
|
||||
}
|
||||
)
|
||||
joined: Final = "\n\n".join(chunk for chunk in (instructions, *folded.values()) if chunk)
|
||||
return (
|
||||
instructions if not folded else joined or None,
|
||||
[ # mutable-ok: the base class takes the input items as a list
|
||||
_developer_item_as_system(item) for index, item in enumerate(items) if index not in folded
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
|
|
@ -68,9 +142,6 @@ class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
base: Final = (api_base or get_secret_str("FIREWORKS_API_BASE") or FIREWORKS_AI_DEFAULT_API_BASE).rstrip("/")
|
||||
return f"{base}/responses"
|
||||
|
||||
def _validate_input_param(self, input: str | ResponseInputParam) -> str | ResponseInputParam:
|
||||
return _developer_items_as_system(super()._validate_input_param(input))
|
||||
|
||||
def transform_responses_api_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -79,10 +150,25 @@ class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict, # mutable-ok: overrides the base class signature
|
||||
) -> dict: # mutable-ok: overrides the base class signature
|
||||
instructions_param: Final[object] = response_api_optional_request_params.get("instructions")
|
||||
validated_input: Final = self._validate_input_param(input)
|
||||
instructions, folded_input = (
|
||||
_with_instruction_items_folded(validated_input, instructions_param)
|
||||
if isinstance(instructions_param, str | None)
|
||||
else (instructions_param, _developer_items_as_system(validated_input))
|
||||
)
|
||||
instruction_entries: Final = () if instructions is None else (("instructions", instructions),)
|
||||
folded_params: Final = { # mutable-ok: the base class takes the optional params as a dict
|
||||
key: value
|
||||
for key, value in (
|
||||
*((key, value) for key, value in response_api_optional_request_params.items() if key != "instructions"),
|
||||
*instruction_entries,
|
||||
)
|
||||
}
|
||||
return super().transform_responses_api_request(
|
||||
model=resolve_fireworks_resource_name(model),
|
||||
input=input,
|
||||
response_api_optional_request_params=response_api_optional_request_params,
|
||||
input=folded_input,
|
||||
response_api_optional_request_params=folded_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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={}
|
||||
|
|
|
|||
|
|
@ -6545,7 +6545,7 @@ def embedding(
|
|||
client=client,
|
||||
timeout=timeout,
|
||||
aembedding=aembedding,
|
||||
litellm_params={},
|
||||
litellm_params=litellm_params_dict,
|
||||
api_base=api_base,
|
||||
print_verbose=print_verbose,
|
||||
extra_headers=headers,
|
||||
|
|
@ -7805,6 +7805,7 @@ def transcription(
|
|||
azure_ad_token=azure_ad_token,
|
||||
max_retries=max_retries,
|
||||
litellm_params=litellm_params_dict,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
elif custom_llm_provider == "openai" or (custom_llm_provider in litellm.openai_compatible_providers):
|
||||
api_base = (
|
||||
|
|
|
|||
|
|
@ -650,7 +650,10 @@
|
|||
},
|
||||
"twelvelabs.marengo-embed-2-7-v1:0": {
|
||||
"deprecation_date": "2026-11-30",
|
||||
"input_cost_per_token": 7e-05,
|
||||
"input_cost_per_query": 7e-05,
|
||||
"input_cost_per_video_per_second": 0.0007,
|
||||
"input_cost_per_audio_per_second": 0.00014,
|
||||
"input_cost_per_image": 0.0001,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 77,
|
||||
"max_tokens": 77,
|
||||
|
|
@ -662,7 +665,7 @@
|
|||
},
|
||||
"us.twelvelabs.marengo-embed-2-7-v1:0": {
|
||||
"deprecation_date": "2026-11-30",
|
||||
"input_cost_per_token": 7e-05,
|
||||
"input_cost_per_query": 7e-05,
|
||||
"input_cost_per_video_per_second": 0.0007,
|
||||
"input_cost_per_audio_per_second": 0.00014,
|
||||
"input_cost_per_image": 0.0001,
|
||||
|
|
@ -677,7 +680,7 @@
|
|||
},
|
||||
"eu.twelvelabs.marengo-embed-2-7-v1:0": {
|
||||
"deprecation_date": "2026-11-30",
|
||||
"input_cost_per_token": 7e-05,
|
||||
"input_cost_per_query": 7e-05,
|
||||
"input_cost_per_video_per_second": 0.0007,
|
||||
"input_cost_per_audio_per_second": 0.00014,
|
||||
"input_cost_per_image": 0.0001,
|
||||
|
|
@ -690,6 +693,48 @@
|
|||
"supports_embedding_image_input": true,
|
||||
"supports_image_input": true
|
||||
},
|
||||
"twelvelabs.marengo-embed-3-0-v1:0": {
|
||||
"input_cost_per_query": 7e-05,
|
||||
"input_cost_per_video_per_second": 0.0007,
|
||||
"input_cost_per_audio_per_second": 0.00014,
|
||||
"input_cost_per_image": 0.0001,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 500,
|
||||
"max_tokens": 500,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 512,
|
||||
"supports_embedding_image_input": true,
|
||||
"supports_image_input": true
|
||||
},
|
||||
"us.twelvelabs.marengo-embed-3-0-v1:0": {
|
||||
"input_cost_per_query": 7e-05,
|
||||
"input_cost_per_video_per_second": 0.0007,
|
||||
"input_cost_per_audio_per_second": 0.00014,
|
||||
"input_cost_per_image": 0.0001,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 500,
|
||||
"max_tokens": 500,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 512,
|
||||
"supports_embedding_image_input": true,
|
||||
"supports_image_input": true
|
||||
},
|
||||
"eu.twelvelabs.marengo-embed-3-0-v1:0": {
|
||||
"input_cost_per_query": 7e-05,
|
||||
"input_cost_per_video_per_second": 0.0007,
|
||||
"input_cost_per_audio_per_second": 0.00014,
|
||||
"input_cost_per_image": 0.0001,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 500,
|
||||
"max_tokens": 500,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0.0,
|
||||
"output_vector_size": 512,
|
||||
"supports_embedding_image_input": true,
|
||||
"supports_image_input": true
|
||||
},
|
||||
"twelvelabs.pegasus-1-2-v1:0": {
|
||||
"input_cost_per_video_per_second": 0.00049,
|
||||
"output_cost_per_token": 7.5e-06,
|
||||
|
|
@ -3581,6 +3626,79 @@
|
|||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure_ai/gpt-chat-latest": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"deprecation_date": "2026-12-02",
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 272000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure_ai/codex-mini": {
|
||||
"cache_read_input_token_cost": 3.75e-07,
|
||||
"deprecation_date": "2026-11-15",
|
||||
"input_cost_per_token": 1.5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 100000,
|
||||
"max_tokens": 100000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 6e-06,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure_ai/whisper": {
|
||||
"deprecation_date": "2026-12-15",
|
||||
"input_cost_per_second": 0.0001,
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "audio_transcription",
|
||||
"output_cost_per_second": 0.0001,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/"
|
||||
},
|
||||
"azure_ai/gpt-5.5-2026-04-23": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
|
||||
|
|
@ -3984,13 +4102,29 @@
|
|||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure_ai/model_router": {
|
||||
"deprecation_date": "2027-05-20",
|
||||
"input_cost_per_token": 1.4e-07,
|
||||
"output_cost_per_token": 0,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/aoai/",
|
||||
"comment": "Flat cost of $0.14 per M input tokens for Azure AI Foundry Model Router infrastructure. Use pattern: azure_ai/model_router/<deployment-name> where deployment-name is your Azure deployment (e.g., azure-model-router)"
|
||||
},
|
||||
"azure_ai/model-router": {
|
||||
"deprecation_date": "2027-05-20",
|
||||
"input_cost_per_token": 1.4e-07,
|
||||
"output_cost_per_token": 0,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/aoai/",
|
||||
"comment": "Catalog-name twin of azure_ai/model_router: the flat $0.14 per M input tokens is the router's own fee, the routed model is priced on top of it"
|
||||
},
|
||||
"azure/eu/gpt-4o-2024-08-06": {
|
||||
"deprecation_date": "2027-04-14",
|
||||
"cache_read_input_token_cost": 1.375e-06,
|
||||
|
|
@ -10302,6 +10436,18 @@
|
|||
"/v1/ocr"
|
||||
]
|
||||
},
|
||||
"azure_ai/cohere-command-a": {
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 8182,
|
||||
"max_tokens": 8182,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/cohere/",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"azure_ai/doc-intelligence/prebuilt-read": {
|
||||
"litellm_provider": "azure_ai",
|
||||
"ocr_cost_per_page": 0.0015,
|
||||
|
|
@ -10653,6 +10799,41 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure_ai/grok-4-20-reasoning": {
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"deprecation_date": "2027-04-06",
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 262000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/grok/",
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_reasoning": true
|
||||
},
|
||||
"azure_ai/grok-4-20-non-reasoning": {
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"deprecation_date": "2027-04-06",
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 262000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/grok/",
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure_ai/grok-4-fast-non-reasoning": {
|
||||
"deprecation_date": "2026-05-01",
|
||||
"input_cost_per_token": 2e-07,
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase):
|
|||
file_object: LiteLLMBatch | LiteLLMFineTuningJob | ResponsesAPIResponse
|
||||
created_by: str | None = None
|
||||
team_id: str | None = None
|
||||
org_id: str | None = None
|
||||
|
||||
|
||||
class LiteLLM_ManagedVectorStoreTable(LiteLLMPydanticObjectBase):
|
||||
|
|
|
|||
|
|
@ -21,3 +21,6 @@ _mcp_gateway_initialize_instructions: Final[ContextVar[str | None]] = ContextVar
|
|||
# Per-request scoped server name; set in MCP HTTP/SSE handlers when the path
|
||||
# identifies exactly one upstream server. Never populated from client-supplied headers.
|
||||
_mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gateway_server_name", default=None)
|
||||
|
||||
# Set server-side by the /mcp/proxy route. Never populated from client-supplied headers.
|
||||
_mcp_proxy_mode: Final[ContextVar[bool]] = ContextVar("_mcp_proxy_mode", default=False)
|
||||
|
|
|
|||
|
|
@ -15,11 +15,11 @@ import types
|
|||
import uuid
|
||||
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol
|
||||
|
||||
import httpx
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from pydantic import AnyUrl, ConfigDict
|
||||
from pydantic import AnyUrl, ConfigDict, TypeAdapter, ValidationError
|
||||
from starlette.requests import Request as StarletteRequest
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
|
@ -47,6 +47,7 @@ from litellm.proxy._experimental.mcp_server.mcp_context import (
|
|||
_mcp_active_toolset_id,
|
||||
_mcp_gateway_initialize_instructions,
|
||||
_mcp_gateway_server_name,
|
||||
_mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # server-owned request mode
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
|
|
@ -108,9 +109,9 @@ _MAX_STATEFUL_SESSIONS_PER_OWNER: Final = 100
|
|||
# prevents an authenticated client from forcing the proxy to buffer an
|
||||
# arbitrarily large body just to make a routing decision.
|
||||
_MCP_ROUTING_PEEK_MAX_BYTES: Final = 4096
|
||||
# ASGI scope key holding the tracing span of the request carrying an MCP
|
||||
# message, written on the request task and read back by the message handler.
|
||||
# ASGI scope keys carrying OTel request state into a stateful MCP message handler.
|
||||
_MCP_TRANSPORT_SPAN_SCOPE_KEY: Final = "litellm_otel_transport_span"
|
||||
_MCP_DESTINATIONS_SCOPE_KEY: Final = "litellm_otel_request_destinations"
|
||||
|
||||
|
||||
def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None:
|
||||
|
|
@ -328,18 +329,17 @@ def _otel_publish_transport_span_on_scope(scope: Scope) -> None:
|
|||
scope[_MCP_TRANSPORT_SPAN_SCOPE_KEY] = span
|
||||
|
||||
|
||||
def _otel_transport_span_from_message(req_ctx: object) -> object:
|
||||
"""The tracing span of the HTTP request that carried this MCP message.
|
||||
|
||||
Read off that request's ASGI scope, reached through the ``Request`` the
|
||||
streamable-HTTP transport attaches to each message, so it is this message's
|
||||
transport and not whichever request happens to have touched the session last.
|
||||
Returns whatever the scope holds; the otel plumbing validates it."""
|
||||
def _otel_value_from_message_scope(req_ctx: object, key: str) -> object:
|
||||
request: Final = getattr(req_ctx, "request", None)
|
||||
scope: Final = getattr(request, "scope", None)
|
||||
if not isinstance(scope, Mapping):
|
||||
return None
|
||||
return scope.get(_MCP_TRANSPORT_SPAN_SCOPE_KEY)
|
||||
return scope.get(key)
|
||||
|
||||
|
||||
def _otel_transport_span_from_message(req_ctx: object) -> object:
|
||||
"""The tracing span of the HTTP request that carried this MCP message."""
|
||||
return _otel_value_from_message_scope(req_ctx, _MCP_TRANSPORT_SPAN_SCOPE_KEY)
|
||||
|
||||
|
||||
def _otel_set_mcp_transport_span(span: object) -> object:
|
||||
|
|
@ -372,6 +372,44 @@ def _otel_reset_mcp_transport_span(token: object) -> None:
|
|||
return
|
||||
|
||||
|
||||
def _otel_publish_request_destinations_on_scope(scope: Scope) -> None:
|
||||
try:
|
||||
from litellm.integrations.otel.plumbing.context import request_destinations
|
||||
|
||||
scope[_MCP_DESTINATIONS_SCOPE_KEY] = request_destinations()
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
|
||||
def _otel_set_mcp_request_destinations(req_ctx: object) -> object:
|
||||
destinations: Final = _otel_value_from_message_scope(req_ctx, _MCP_DESTINATIONS_SCOPE_KEY)
|
||||
if not isinstance(destinations, tuple):
|
||||
return None
|
||||
try:
|
||||
from litellm.integrations.otel.model.destination import OtelDestination
|
||||
from litellm.integrations.otel.plumbing.context import set_request_destinations
|
||||
|
||||
destination_adapter: Final[TypeAdapter[tuple[OtelDestination, ...]]] = TypeAdapter(
|
||||
tuple[OtelDestination, ...],
|
||||
config=ConfigDict(revalidate_instances="always"),
|
||||
)
|
||||
validated_destinations: Final = destination_adapter.validate_python(destinations, strict=True)
|
||||
return set_request_destinations(validated_destinations)
|
||||
except (ImportError, ValidationError):
|
||||
return None
|
||||
|
||||
|
||||
def _otel_reset_mcp_request_destinations(token: object) -> None:
|
||||
if token is None:
|
||||
return
|
||||
try:
|
||||
from litellm.integrations.otel.plumbing.context import reset_request_destinations
|
||||
|
||||
reset_request_destinations(token)
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
|
||||
def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException:
|
||||
"""Map a ``ProxyException`` to an ``HTTPException`` that preserves its real
|
||||
status code and headers.
|
||||
|
|
@ -500,11 +538,22 @@ if MCP_AVAILABLE:
|
|||
notification_options: NotificationOptions | None = None,
|
||||
experimental_capabilities: dict[str, dict[str, object]] | None = None,
|
||||
) -> InitializationOptions:
|
||||
opts: Final = Server.create_initialization_options(
|
||||
base_options: Final = Server.create_initialization_options(
|
||||
self,
|
||||
notification_options=notification_options,
|
||||
experimental_capabilities=experimental_capabilities or {},
|
||||
)
|
||||
opts: Final = (
|
||||
base_options.model_copy(
|
||||
update={ # mutable-ok: Pydantic update payload
|
||||
"capabilities": base_options.capabilities.model_copy(
|
||||
update={"prompts": None, "resources": None} # mutable-ok: Pydantic update payload
|
||||
)
|
||||
}
|
||||
)
|
||||
if _mcp_proxy_mode.get()
|
||||
else base_options
|
||||
)
|
||||
updates: Final[dict[str, str]] = {}
|
||||
merged: Final = _mcp_gateway_initialize_instructions.get()
|
||||
if merged is not None:
|
||||
|
|
@ -718,12 +767,12 @@ if MCP_AVAILABLE:
|
|||
_stateful_auth_context_cleanup_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await _stateful_auth_context_cleanup_task
|
||||
if _session_manager_cm:
|
||||
await _session_manager_cm.__aexit__(None, None, None)
|
||||
if _session_manager_stateful_cm:
|
||||
await _session_manager_stateful_cm.__aexit__(None, None, None)
|
||||
if _sse_session_manager_cm:
|
||||
await _sse_session_manager_cm.__aexit__(None, None, None)
|
||||
if _session_manager_stateful_cm:
|
||||
await _session_manager_stateful_cm.__aexit__(None, None, None)
|
||||
if _session_manager_cm:
|
||||
await _session_manager_cm.__aexit__(None, None, None)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error during session manager shutdown: %s", e)
|
||||
|
||||
|
|
@ -763,10 +812,12 @@ if MCP_AVAILABLE:
|
|||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
_trace_token = None
|
||||
_transport_token = None
|
||||
_destinations_token = None
|
||||
|
||||
try:
|
||||
_trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx))
|
||||
_transport_token = _otel_set_mcp_transport_span(_otel_transport_span_from_message(req_ctx))
|
||||
_destinations_token = _otel_set_mcp_request_destinations(req_ctx)
|
||||
# Get user authentication from context variable
|
||||
(
|
||||
user_api_key_auth,
|
||||
|
|
@ -783,17 +834,20 @@ if MCP_AVAILABLE:
|
|||
"MCP list_tools - MCP server auth headers: %s",
|
||||
list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None,
|
||||
)
|
||||
from mcp.types import Tool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
get_mcp_proxy_tool_definitions,
|
||||
get_virtual_tool_definitions,
|
||||
)
|
||||
|
||||
if _mcp_proxy_mode.get():
|
||||
return [Tool.model_validate(d) for d in get_mcp_proxy_tool_definitions()] # mutable-ok: MCP SDK list
|
||||
if getattr(
|
||||
getattr(user_api_key_auth, "object_permission", None),
|
||||
"mcp_tool_search_enabled",
|
||||
False,
|
||||
):
|
||||
from mcp.types import Tool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
get_virtual_tool_definitions,
|
||||
)
|
||||
|
||||
return [Tool.model_validate(d) for d in get_virtual_tool_definitions()]
|
||||
|
||||
# Get mcp_servers from context variable
|
||||
|
|
@ -828,6 +882,7 @@ if MCP_AVAILABLE:
|
|||
# This prevents the HTTP stream from failing and allows the client to get a response
|
||||
return []
|
||||
finally:
|
||||
_otel_reset_mcp_request_destinations(_destinations_token)
|
||||
_otel_reset_mcp_transport_span(_transport_token)
|
||||
_otel_reset_mcp_trace_carrier(_trace_token)
|
||||
if _session_reset_token is not None:
|
||||
|
|
@ -866,6 +921,12 @@ if MCP_AVAILABLE:
|
|||
verbose_logger.debug("Host progressToken captured: %s...", str(host_token)[:8])
|
||||
return forward_progress
|
||||
|
||||
def _reject_mcp_proxy_operation() -> NoReturn:
|
||||
from mcp.shared.exceptions import McpError
|
||||
from mcp.types import METHOD_NOT_FOUND, ErrorData
|
||||
|
||||
raise McpError(ErrorData(code=METHOD_NOT_FOUND, message="Operation unavailable on /mcp/proxy"))
|
||||
|
||||
async def _build_virtual_call_logging_obj(
|
||||
name: str,
|
||||
arguments: dict[str, object],
|
||||
|
|
@ -921,16 +982,91 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
AGENT_SEARCH_TOOL_NAME,
|
||||
DEFAULT_AGENT_SEARCH_TOP_K,
|
||||
MCP_PROXY_CALL_TOOL_NAME,
|
||||
MCP_PROXY_TOOL_NAMES,
|
||||
MCP_TOOL_SEARCH_TOOL_NAME,
|
||||
SKILL_SEARCH_TOOL_NAME,
|
||||
VIRTUAL_TOOL_NAMES,
|
||||
coerce_top_k,
|
||||
handle_agent_search,
|
||||
handle_mcp_proxy_tool,
|
||||
handle_mcp_tool_call,
|
||||
handle_mcp_tool_search,
|
||||
handle_skill_search,
|
||||
)
|
||||
|
||||
if _mcp_proxy_mode.get() and name not in MCP_PROXY_TOOL_NAMES:
|
||||
return CallToolResult(
|
||||
content=[ # mutable-ok: MCP result content
|
||||
TextContent(type="text", text=f"Tool {name} is unavailable on /mcp/proxy")
|
||||
],
|
||||
isError=True,
|
||||
)
|
||||
|
||||
if _mcp_proxy_mode.get() and name in MCP_PROXY_TOOL_NAMES:
|
||||
assert user_api_key_auth is not None
|
||||
proxy_call_start: Final = datetime.now() # noqa: DTZ005 # logging pipeline uses naive datetimes
|
||||
proxy_logging_obj: Final = (
|
||||
await _build_virtual_call_logging_obj(
|
||||
name=name,
|
||||
arguments=arguments or {}, # mutable-ok: logging pipeline payload
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
if name == MCP_PROXY_CALL_TOOL_NAME
|
||||
else None
|
||||
)
|
||||
try:
|
||||
proxy_result: Final = await handle_mcp_proxy_tool(
|
||||
name=name,
|
||||
arguments=arguments or {}, # mutable-ok: proxy handler payload
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
client_ip=client_ip,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as exc:
|
||||
if proxy_logging_obj is not None:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj as request_logging_obj
|
||||
|
||||
failure_end: Final = datetime.now() # noqa: DTZ005 # matches the logging pipeline start time
|
||||
failure_traceback: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG)
|
||||
try:
|
||||
proxy_logging_obj.failure_handler(exc, failure_traceback, proxy_call_start, failure_end)
|
||||
await proxy_logging_obj.async_failure_handler(
|
||||
exc, failure_traceback, proxy_call_start, failure_end
|
||||
)
|
||||
if not isinstance(exc, MCPUpstreamAuthError):
|
||||
await request_logging_obj.post_call_failure_hook(
|
||||
request_data={ # mutable-ok: failure hook mutates its request payload
|
||||
"name": name,
|
||||
"arguments": arguments,
|
||||
"litellm_logging_obj": proxy_logging_obj,
|
||||
},
|
||||
original_exception=exc,
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
route="/mcp/call_tool",
|
||||
traceback_str=failure_traceback,
|
||||
)
|
||||
except Exception:
|
||||
verbose_logger.exception("Error logging failed MCP proxy tool call")
|
||||
raise
|
||||
if proxy_logging_obj is not None:
|
||||
return await _fire_mcp_tool_call_logging(
|
||||
logging_obj=proxy_logging_obj,
|
||||
result=proxy_result,
|
||||
start_time=proxy_call_start,
|
||||
end_time=datetime.now(), # noqa: DTZ005 # matches the logging pipeline start time
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
request_data=types.MappingProxyType({"name": name, "arguments": arguments}),
|
||||
)
|
||||
return proxy_result
|
||||
|
||||
if name not in VIRTUAL_TOOL_NAMES:
|
||||
return None
|
||||
|
||||
|
|
@ -1021,10 +1157,12 @@ if MCP_AVAILABLE:
|
|||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
_trace_token = None
|
||||
_transport_token = None
|
||||
_destinations_token = None
|
||||
|
||||
try:
|
||||
_trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx))
|
||||
_transport_token = _otel_set_mcp_transport_span(_otel_transport_span_from_message(req_ctx))
|
||||
_destinations_token = _otel_set_mcp_request_destinations(req_ctx)
|
||||
# Validate arguments
|
||||
(
|
||||
user_api_key_auth,
|
||||
|
|
@ -1163,6 +1301,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
return response
|
||||
finally:
|
||||
_otel_reset_mcp_request_destinations(_destinations_token)
|
||||
_otel_reset_mcp_transport_span(_transport_token)
|
||||
_otel_reset_mcp_trace_carrier(_trace_token)
|
||||
if _session_reset_token is not None:
|
||||
|
|
@ -1173,6 +1312,8 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
List all available prompts
|
||||
"""
|
||||
if _mcp_proxy_mode.get():
|
||||
_reject_mcp_proxy_operation()
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx: Final = request_ctx.get(None)
|
||||
|
|
@ -1230,8 +1371,8 @@ if MCP_AVAILABLE:
|
|||
Returns:
|
||||
GetPromptResult: Getting prompt execution results
|
||||
"""
|
||||
|
||||
# Validate arguments
|
||||
if _mcp_proxy_mode.get():
|
||||
_reject_mcp_proxy_operation()
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx: Final = request_ctx.get(None)
|
||||
|
|
@ -1268,6 +1409,8 @@ if MCP_AVAILABLE:
|
|||
@server.list_resources()
|
||||
async def list_resources() -> list[Resource]:
|
||||
"""List all available resources."""
|
||||
if _mcp_proxy_mode.get():
|
||||
_reject_mcp_proxy_operation()
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx: Final = request_ctx.get(None)
|
||||
|
|
@ -1312,6 +1455,8 @@ if MCP_AVAILABLE:
|
|||
@server.list_resource_templates()
|
||||
async def list_resource_templates() -> list[ResourceTemplate]:
|
||||
"""List all available resource templates."""
|
||||
if _mcp_proxy_mode.get():
|
||||
_reject_mcp_proxy_operation()
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx: Final = request_ctx.get(None)
|
||||
|
|
@ -1357,6 +1502,8 @@ if MCP_AVAILABLE:
|
|||
|
||||
@server.read_resource()
|
||||
async def read_resource(url: AnyUrl) -> list[ReadResourceContents]:
|
||||
if _mcp_proxy_mode.get():
|
||||
_reject_mcp_proxy_operation()
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx: Final = request_ctx.get(None)
|
||||
|
|
@ -1955,6 +2102,7 @@ if MCP_AVAILABLE:
|
|||
litellm_trace_id: str | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -2134,9 +2282,14 @@ if MCP_AVAILABLE:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
# Apply display-name/description overrides last so that
|
||||
# permission filtering always works against original names.
|
||||
filtered_tools = apply_tool_overrides(filtered_tools, server)
|
||||
if mcp_proxy_mode:
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
|
||||
|
||||
filtered_tools = [ # mutable-ok: MCP tool pipeline
|
||||
with_mcp_proxy_identity(tool, server.server_id) for tool in filtered_tools
|
||||
]
|
||||
else:
|
||||
filtered_tools = apply_tool_overrides(filtered_tools, server)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Successfully fetched %s tools from server %s, %s after filtering",
|
||||
|
|
@ -2448,6 +2601,7 @@ if MCP_AVAILABLE:
|
|||
log_list_tools_to_spendlogs: bool = False,
|
||||
list_tools_log_source: str | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
List all available MCP tools.
|
||||
|
|
@ -2477,6 +2631,7 @@ if MCP_AVAILABLE:
|
|||
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
|
||||
list_tools_log_source=list_tools_log_source,
|
||||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
)
|
||||
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
|
||||
return listing
|
||||
|
|
@ -3376,7 +3531,9 @@ if MCP_AVAILABLE:
|
|||
server_name: str | None,
|
||||
session_id: str | None = None,
|
||||
) -> StandardLoggingMCPToolCall:
|
||||
mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
|
||||
mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name(
|
||||
add_server_prefix_to_name(name, server_name) if server_name else name
|
||||
)
|
||||
namespaced_tool_name: Final = f"{server_name}/{name}" if server_name else name
|
||||
if mcp_server:
|
||||
mcp_info: Final = mcp_server.mcp_info or {}
|
||||
|
|
@ -4493,6 +4650,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def _dispatch() -> None:
|
||||
_otel_publish_transport_span_on_scope(scope)
|
||||
_otel_publish_request_destinations_on_scope(scope)
|
||||
auth_user: Final = _set_or_update_auth_context(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -30,6 +31,12 @@ if TYPE_CHECKING:
|
|||
MCP_TOOL_SEARCH_SETTINGS_KEY: Final[str] = "mcp_tool_search"
|
||||
MCP_TOOL_SEARCH_TOOL_NAME: Final[str] = "mcp_tool_search"
|
||||
MCP_TOOL_CALL_TOOL_NAME: Final[str] = "mcp_tool_call"
|
||||
MCP_PROXY_SEARCH_TOOL_NAME: Final[str] = "search_tools"
|
||||
MCP_PROXY_SCHEMA_TOOL_NAME: Final[str] = "get_tool_schema"
|
||||
MCP_PROXY_CALL_TOOL_NAME: Final[str] = "call_tool"
|
||||
MCP_PROXY_TOOL_NAMES: Final = frozenset(
|
||||
(MCP_PROXY_SEARCH_TOOL_NAME, MCP_PROXY_SCHEMA_TOOL_NAME, MCP_PROXY_CALL_TOOL_NAME)
|
||||
)
|
||||
AGENT_SEARCH_TOOL_NAME: Final[str] = "agent_search"
|
||||
SKILL_SEARCH_TOOL_NAME: Final[str] = "skill_search"
|
||||
VIRTUAL_TOOL_NAMES: Final = frozenset(
|
||||
|
|
@ -51,6 +58,29 @@ class ToolSearchResult(TypedDict, total=False):
|
|||
score: ReadOnly[float]
|
||||
|
||||
|
||||
class MCPProxySearchResult(TypedDict, total=False):
|
||||
tool_id: Required[ReadOnly[str]]
|
||||
name: Required[ReadOnly[str]]
|
||||
description: Required[ReadOnly[str]]
|
||||
score: ReadOnly[float]
|
||||
|
||||
|
||||
class MCPProxySchemaResult(MCPProxySearchResult, total=False):
|
||||
inputSchema: Required[ReadOnly[Mapping[str, object]]]
|
||||
outputSchema: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class MCPProxyToolIdentity(TypedDict):
|
||||
server_id: ReadOnly[str]
|
||||
tool_name: ReadOnly[str]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MCPToolSearchHit:
|
||||
tool: Tool
|
||||
score: float | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SemanticToolRanker:
|
||||
embed: Embedder
|
||||
|
|
@ -76,6 +106,55 @@ def _scored_result(tool: Tool, score: float) -> ToolSearchResult:
|
|||
return {"name": tool.name, "description": tool.description or "", "inputSchema": tool.inputSchema, "score": score}
|
||||
|
||||
|
||||
_MCP_PROXY_IDENTITY_META_KEY: Final[str] = "litellm.ai/proxy_tool_identity"
|
||||
|
||||
|
||||
def with_mcp_proxy_identity(tool: Tool, server_id: str) -> Tool:
|
||||
identity: Final[MCPProxyToolIdentity] = {"server_id": server_id, "tool_name": tool.name}
|
||||
return tool.model_copy( # mutable-ok: Pydantic requires mutable update and metadata mappings
|
||||
update={ # mutable-ok: Pydantic update payload
|
||||
"meta": {**(tool.meta or {}), _MCP_PROXY_IDENTITY_META_KEY: identity} # mutable-ok: metadata mapping
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _mcp_proxy_identity(tool: Tool) -> MCPProxyToolIdentity:
|
||||
identity: Final = (tool.meta or {}).get(_MCP_PROXY_IDENTITY_META_KEY) # mutable-ok: absent metadata default
|
||||
if not isinstance(identity, Mapping):
|
||||
raise TypeError("MCP proxy tool identity is missing")
|
||||
server_id: Final = identity.get("server_id")
|
||||
tool_name: Final = identity.get("tool_name")
|
||||
if not isinstance(server_id, str) or not isinstance(tool_name, str):
|
||||
raise TypeError("MCP proxy tool identity is invalid")
|
||||
return {"server_id": server_id, "tool_name": tool_name} # mutable-ok: TypedDict identity payload
|
||||
|
||||
|
||||
def mcp_proxy_tool_id(tool: Tool) -> str:
|
||||
identity: Final = _mcp_proxy_identity(tool)
|
||||
return hashlib.sha256(f"{identity['server_id']}\0{identity['tool_name']}".encode()).hexdigest()[:32]
|
||||
|
||||
|
||||
def _proxy_search_result(hit: MCPToolSearchHit) -> MCPProxySearchResult:
|
||||
base: Final[MCPProxySearchResult] = {
|
||||
"tool_id": mcp_proxy_tool_id(hit.tool),
|
||||
"name": hit.tool.name,
|
||||
"description": hit.tool.description or "",
|
||||
}
|
||||
return {**base, "score": hit.score} if hit.score is not None else base # mutable-ok: wire result payload
|
||||
|
||||
|
||||
def _proxy_schema_result(tool: Tool) -> MCPProxySchemaResult:
|
||||
base: Final[MCPProxySchemaResult] = {
|
||||
"tool_id": mcp_proxy_tool_id(tool),
|
||||
"name": tool.name,
|
||||
"description": tool.description or "",
|
||||
"inputSchema": tool.inputSchema,
|
||||
}
|
||||
if tool.outputSchema is None:
|
||||
return base
|
||||
return {**base, "outputSchema": tool.outputSchema} # mutable-ok: wire schema payload
|
||||
|
||||
|
||||
def _tool_text(tool: Tool) -> str:
|
||||
return "\n".join(part for part in (tool.name, tool.description or "") if part)
|
||||
|
||||
|
|
@ -107,6 +186,38 @@ def search_tools(query: str, tools: Sequence[Tool], top_k: int = 5) -> tuple[Too
|
|||
return tuple(_tool_result(tool) for _, tool in _top_hits(tools, scores, minimum=1.0, limit=top_k))
|
||||
|
||||
|
||||
async def rank_mcp_tools(
|
||||
query: str,
|
||||
tools: Sequence[Tool],
|
||||
top_k: int,
|
||||
settings: MCPToolSearchSettings,
|
||||
ranker: SemanticToolRanker | None,
|
||||
) -> tuple[MCPToolSearchHit, ...] | EmbeddingFailed:
|
||||
core, rest = _split_core_tools(tools, settings.core_tools)
|
||||
core_hits: Final = tuple(MCPToolSearchHit(tool) for tool in core)
|
||||
if not query:
|
||||
return core_hits
|
||||
limit: Final = min(top_k, settings.top_k)
|
||||
if ranker is None:
|
||||
scores: Final = tuple(_keyword_score(query, tool) for tool in rest)
|
||||
return (
|
||||
*core_hits,
|
||||
*(MCPToolSearchHit(tool) for _, tool in _top_hits(rest, scores, minimum=1.0, limit=limit)),
|
||||
)
|
||||
semantic_scores: Final = await ranker.index.scores(
|
||||
query, tuple(_tool_text(tool) for tool in rest), ranker.embed, ranker.embedding_model
|
||||
)
|
||||
if isinstance(semantic_scores, EmbeddingFailed):
|
||||
return semantic_scores
|
||||
return (
|
||||
*core_hits,
|
||||
*(
|
||||
MCPToolSearchHit(tool, score)
|
||||
for score, tool in _top_hits(rest, semantic_scores, settings.similarity_threshold, limit)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def search_mcp_tools(
|
||||
query: str,
|
||||
tools: Sequence[Tool],
|
||||
|
|
@ -114,21 +225,12 @@ async def search_mcp_tools(
|
|||
settings: MCPToolSearchSettings,
|
||||
ranker: SemanticToolRanker | None,
|
||||
) -> tuple[ToolSearchResult, ...] | EmbeddingFailed:
|
||||
"""Core tools the caller can access come first, then up to `top_k` ranked matches from the remaining tools."""
|
||||
core, rest = _split_core_tools(tools, settings.core_tools)
|
||||
limit: Final = min(top_k, settings.top_k)
|
||||
core_results: Final = tuple(_tool_result(tool) for tool in core)
|
||||
if ranker is None:
|
||||
return (*core_results, *search_tools(query, rest, limit))
|
||||
if not query:
|
||||
return core_results
|
||||
scores: Final = await ranker.index.scores(
|
||||
query, tuple(_tool_text(tool) for tool in rest), ranker.embed, ranker.embedding_model
|
||||
hits: Final = await rank_mcp_tools(query, tools, top_k, settings, ranker)
|
||||
if isinstance(hits, EmbeddingFailed):
|
||||
return hits
|
||||
return tuple(
|
||||
_scored_result(hit.tool, hit.score) if hit.score is not None else _tool_result(hit.tool) for hit in hits
|
||||
)
|
||||
if isinstance(scores, EmbeddingFailed):
|
||||
return scores
|
||||
hits: Final = _top_hits(rest, scores, minimum=settings.similarity_threshold, limit=limit)
|
||||
return (*core_results, *(_scored_result(tool, score) for score, tool in hits))
|
||||
|
||||
|
||||
class _ToolParamSchema(TypedDict, total=False):
|
||||
|
|
@ -223,10 +325,48 @@ _SKILL_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
|
|||
}
|
||||
|
||||
|
||||
_MCP_PROXY_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
|
||||
"name": MCP_PROXY_SEARCH_TOOL_NAME,
|
||||
"description": "Search accessible MCP tools by describing what you need. Returns opaque tool IDs.",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string", "description": "What the tool should do."}},
|
||||
"required": _json_array("query"),
|
||||
},
|
||||
}
|
||||
|
||||
_MCP_PROXY_SCHEMA_DEFINITION: Final[VirtualToolDefinition] = {
|
||||
"name": MCP_PROXY_SCHEMA_TOOL_NAME,
|
||||
"description": "Return the complete schema for an accessible MCP tool ID.",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {"tool_id": {"type": "string", "description": "Opaque ID from search_tools."}},
|
||||
"required": _json_array("tool_id"),
|
||||
},
|
||||
}
|
||||
|
||||
_MCP_PROXY_CALL_DEFINITION: Final[VirtualToolDefinition] = {
|
||||
"name": MCP_PROXY_CALL_TOOL_NAME,
|
||||
"description": "Call an accessible MCP tool by opaque ID with schema-valid arguments.",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"tool_id": {"type": "string", "description": "Opaque ID from search_tools."},
|
||||
"arguments": {"type": "object", "description": "Arguments validated against the selected tool schema."},
|
||||
},
|
||||
"required": _json_array("tool_id"),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_virtual_tool_definitions() -> tuple[VirtualToolDefinition, ...]:
|
||||
return (_MCP_TOOL_SEARCH_DEFINITION, _MCP_TOOL_CALL_DEFINITION, _AGENT_SEARCH_DEFINITION, _SKILL_SEARCH_DEFINITION)
|
||||
|
||||
|
||||
def get_mcp_proxy_tool_definitions() -> tuple[VirtualToolDefinition, ...]:
|
||||
return (_MCP_PROXY_SEARCH_DEFINITION, _MCP_PROXY_SCHEMA_DEFINITION, _MCP_PROXY_CALL_DEFINITION)
|
||||
|
||||
|
||||
def _text_tool_result(text: str, is_error: bool) -> CallToolResult:
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
|
|
@ -314,7 +454,9 @@ async def handle_mcp_tool_search(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
) -> CallToolResult:
|
||||
from litellm.proxy._experimental.mcp_server.server import _list_mcp_tools
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_list_mcp_tools, # pyright: ignore[reportPrivateUsage] # shared catalog owner
|
||||
)
|
||||
from litellm.proxy.proxy_server import llm_router, proxy_logging_obj
|
||||
|
||||
settings: Final = mcp_tool_search_settings()
|
||||
|
|
@ -351,6 +493,97 @@ async def handle_mcp_tool_search(
|
|||
return _text_tool_result(json.dumps(results), is_error=False)
|
||||
|
||||
|
||||
async def handle_mcp_proxy_tool(
|
||||
name: str,
|
||||
arguments: dict[str, object], # mutable-ok: MCP dispatcher passes mutable call arguments
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
client_ip: str | None = None,
|
||||
mcp_servers: list[str] | None = None, # mutable-ok: preserve MCP scope container for existing resolver
|
||||
mcp_auth_header: str | None = None,
|
||||
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, # mutable-ok: preserve forwarded headers
|
||||
oauth2_headers: dict[str, str] | None = None, # mutable-ok: preserve forwarded headers
|
||||
raw_headers: dict[str, str] | None = None, # mutable-ok: preserve request headers
|
||||
litellm_logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> CallToolResult:
|
||||
from fastapi import HTTPException
|
||||
from jsonschema import ValidationError as JsonSchemaValidationError
|
||||
from jsonschema import validate
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server.server import ( # pyright: ignore[reportPrivateUsage] # shared catalog owner
|
||||
_list_mcp_tools, # pyright: ignore[reportPrivateUsage] # shared catalog owner
|
||||
)
|
||||
|
||||
listing: Final = await _list_mcp_tools(
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=client_ip,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_proxy_mode=True,
|
||||
)
|
||||
tools_by_id: Final = {mcp_proxy_tool_id(tool): tool for tool in listing.tools} # mutable-ok: lookup index
|
||||
|
||||
if name == MCP_PROXY_SEARCH_TOOL_NAME:
|
||||
llm_router: Final = proxy_server.llm_router
|
||||
proxy_logging_obj: Final = proxy_server.proxy_logging_obj
|
||||
settings: Final = mcp_tool_search_settings()
|
||||
if isinstance(settings, ValidationError):
|
||||
return _text_tool_result(str(settings), is_error=True)
|
||||
if settings.embedding_model is not None and llm_router is None:
|
||||
return _text_tool_result(
|
||||
f"litellm_settings.{MCP_TOOL_SEARCH_SETTINGS_KEY}.embedding_model needs a model_list so it can be called",
|
||||
is_error=True,
|
||||
)
|
||||
ranker: Final = (
|
||||
SemanticToolRanker(
|
||||
embed=router_embedder(llm_router, settings.embedding_model, user_api_key_dict, proxy_logging_obj),
|
||||
embedding_model=settings.embedding_model,
|
||||
index=global_mcp_tool_search_index,
|
||||
)
|
||||
if settings.embedding_model is not None and llm_router is not None
|
||||
else None
|
||||
)
|
||||
results: Final = await rank_mcp_tools(str(arguments.get("query", "")), listing.tools, 5, settings, ranker)
|
||||
if isinstance(results, EmbeddingFailed):
|
||||
return _text_tool_result(results.reason, is_error=True)
|
||||
return _text_tool_result(json.dumps(tuple(_proxy_search_result(hit) for hit in results)), is_error=False)
|
||||
|
||||
tool_id: Final = arguments.get("tool_id")
|
||||
tool: Final = tools_by_id.get(tool_id) if isinstance(tool_id, str) else None
|
||||
if tool is None:
|
||||
return _text_tool_result("Unknown or unauthorized tool_id", is_error=True)
|
||||
|
||||
if name == MCP_PROXY_SCHEMA_TOOL_NAME:
|
||||
return _text_tool_result(json.dumps(_proxy_schema_result(tool)), is_error=False)
|
||||
if name != MCP_PROXY_CALL_TOOL_NAME:
|
||||
raise HTTPException(status_code=400, detail=f"Unknown MCP proxy tool: {name}")
|
||||
|
||||
tool_arguments: Final = arguments.get("arguments", {}) # mutable-ok: JSON Schema validator consumes mapping
|
||||
if not isinstance(tool_arguments, dict):
|
||||
return _text_tool_result("arguments must be an object", is_error=True)
|
||||
try:
|
||||
validate(instance=tool_arguments, schema=tool.inputSchema)
|
||||
except JsonSchemaValidationError as exc:
|
||||
return _text_tool_result(f"Invalid arguments: {exc.message}", is_error=True)
|
||||
|
||||
return await handle_mcp_tool_call(
|
||||
tool_name=_mcp_proxy_identity(tool)["tool_name"],
|
||||
arguments=tool_arguments,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
requested_server_id=_mcp_proxy_identity(tool)["server_id"],
|
||||
client_ip=client_ip,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
async def handle_mcp_tool_call(
|
||||
tool_name: str,
|
||||
arguments: dict[str, Any],
|
||||
|
|
@ -362,6 +595,7 @@ async def handle_mcp_tool_call(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
litellm_logging_obj: LiteLLMLoggingObj | None = None,
|
||||
requested_server_id: str | None = None,
|
||||
) -> CallToolResult:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_allowed_mcp_servers,
|
||||
|
|
@ -400,4 +634,5 @@ async def handle_mcp_tool_call(
|
|||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
requested_server_id=requested_server_id,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -132,10 +132,18 @@ async def update_mcp_toolset(
|
|||
data: UpdateMCPToolsetRequest,
|
||||
touched_by: str,
|
||||
) -> MCPToolset | None:
|
||||
data_dict: Final = data.model_dump(exclude_none=True, exclude={"toolset_id"})
|
||||
if "tools" in data_dict:
|
||||
data_dict["tools"] = json.dumps(data_dict["tools"])
|
||||
data_dict["updated_by"] = touched_by
|
||||
"""A partial update: absent keeps, null clears. A toolset always has a name and a
|
||||
tool list, so a null ``toolset_name`` or ``tools`` is a no-op rather than a clear;
|
||||
emptying the tool selection is an explicit ``[]``, which cannot be mistaken for a
|
||||
caller that left the field out."""
|
||||
data_dict: Final = dict( # mutable-ok: Prisma requires a plain dict for JSON query serialization
|
||||
(
|
||||
(field, json.dumps(value) if field == "tools" else value)
|
||||
for field, value in data.model_dump(exclude_unset=True).items()
|
||||
if field != "toolset_id" and (field not in ("toolset_name", "tools") or value is not None)
|
||||
),
|
||||
updated_by=touched_by,
|
||||
)
|
||||
try:
|
||||
row: Final = await _toolset_table(prisma_client).update(
|
||||
where={"toolset_id": data.toolset_id},
|
||||
|
|
|
|||
|
|
@ -17027,6 +17027,134 @@
|
|||
"mcp_app"
|
||||
]
|
||||
}
|
||||
},
|
||||
"/mcp/proxy": {
|
||||
"delete": {
|
||||
"description": "Serve the fixed three-tool MCP proxy surface.",
|
||||
"operationId": "proxy_mcp_route_mcp_proxy_delete",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"summary": "Proxy Mcp Route",
|
||||
"tags": [
|
||||
"mcp_app"
|
||||
]
|
||||
},
|
||||
"get": {
|
||||
"description": "Serve the fixed three-tool MCP proxy surface.",
|
||||
"operationId": "proxy_mcp_route_mcp_proxy_get",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"summary": "Proxy Mcp Route",
|
||||
"tags": [
|
||||
"mcp_app"
|
||||
]
|
||||
},
|
||||
"head": {
|
||||
"description": "Serve the fixed three-tool MCP proxy surface.",
|
||||
"operationId": "proxy_mcp_route_mcp_proxy_head",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"summary": "Proxy Mcp Route",
|
||||
"tags": [
|
||||
"mcp_app"
|
||||
]
|
||||
},
|
||||
"options": {
|
||||
"description": "Serve the fixed three-tool MCP proxy surface.",
|
||||
"operationId": "proxy_mcp_route_mcp_proxy_options",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"summary": "Proxy Mcp Route",
|
||||
"tags": [
|
||||
"mcp_app"
|
||||
]
|
||||
},
|
||||
"patch": {
|
||||
"description": "Serve the fixed three-tool MCP proxy surface.",
|
||||
"operationId": "proxy_mcp_route_mcp_proxy_patch",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"summary": "Proxy Mcp Route",
|
||||
"tags": [
|
||||
"mcp_app"
|
||||
]
|
||||
},
|
||||
"post": {
|
||||
"description": "Serve the fixed three-tool MCP proxy surface.",
|
||||
"operationId": "proxy_mcp_route_mcp_proxy_post",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"summary": "Proxy Mcp Route",
|
||||
"tags": [
|
||||
"mcp_app"
|
||||
]
|
||||
},
|
||||
"put": {
|
||||
"description": "Serve the fixed three-tool MCP proxy surface.",
|
||||
"operationId": "proxy_mcp_route_mcp_proxy_put",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"summary": "Proxy Mcp Route",
|
||||
"tags": [
|
||||
"mcp_app"
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
|
|
|
|||
|
|
@ -502,6 +502,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
mcp_inference_routes = [
|
||||
"/mcp",
|
||||
"/mcp/",
|
||||
"/mcp/proxy",
|
||||
"/mcp/{subpath}",
|
||||
"/mcp/tools",
|
||||
"/mcp/tools/list",
|
||||
|
|
@ -835,6 +836,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",
|
||||
|
|
@ -3584,6 +3593,7 @@ 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,
|
||||
|
|
|
|||
|
|
@ -25,6 +25,11 @@ from litellm.proxy.common_request_processing import (
|
|||
proxy_exception_from_http_exception,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
|
@ -243,9 +248,9 @@ async def anthropic_response(
|
|||
return _anthropic_error_json_response(
|
||||
ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
headers=headers,
|
||||
),
|
||||
request,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -2836,6 +2836,43 @@ async def _authorize_authenticated_request(
|
|||
return None
|
||||
|
||||
|
||||
def _seed_request_destinations(user_api_key_dict: UserAPIKeyAuth, request: Request | None = None) -> None:
|
||||
"""Anchor the OTLP destinations this key or team overrides its traces to.
|
||||
|
||||
Called inside the ``auth`` phase span so that span reaches the tenant's account
|
||||
as well, and on the request task so the ``ContextVar`` is inherited by the logging
|
||||
tasks that close the LLM span. Best-effort: trace routing must never fail auth.
|
||||
|
||||
``request`` carries the headers, so a backend this request disabled with
|
||||
``x-litellm-disable-callbacks`` resolves to no destination.
|
||||
|
||||
Only destinations the published fan-out can build are anchored. Anchoring one is
|
||||
what tells the operator's exporter to hold that backend's spans back under
|
||||
``override``, so an unbuildable one would leave the span with nowhere to go.
|
||||
|
||||
The ``postgres`` spans under ``auth`` close before this runs, because they are the
|
||||
reads that resolve the identity being read here. They never reach the tenant's
|
||||
account, and they are never withheld from the operator's backend, whichever mode
|
||||
is set.
|
||||
"""
|
||||
try:
|
||||
from litellm.integrations.otel.logger import fan_out_provider
|
||||
from litellm.integrations.otel.plumbing.context import set_request_destinations
|
||||
from litellm.integrations.otel.plumbing.providers import deliverable_destinations
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
resolve_tenant_otel_destinations,
|
||||
)
|
||||
|
||||
set_request_destinations(
|
||||
deliverable_destinations(
|
||||
resolve_tenant_otel_destinations(user_api_key_dict, _safe_get_request_headers(request)),
|
||||
fan_out_provider(),
|
||||
)
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 # telemetry routing is best-effort and must never break authentication
|
||||
verbose_proxy_logger.debug("OTel V2: tenant destination resolution failed: %s", exc)
|
||||
|
||||
|
||||
@tracer.wrap()
|
||||
async def user_api_key_auth(
|
||||
request: Request,
|
||||
|
|
@ -2883,6 +2920,7 @@ async def user_api_key_auth(
|
|||
raise body_parse_exception
|
||||
raise
|
||||
user_api_key_auth_obj.budget_reservation = None
|
||||
_seed_request_destinations(user_api_key_auth_obj, request)
|
||||
|
||||
# A body that never parsed is authenticated (so the trace carries identity
|
||||
# and this ``auth`` span) but not authorized: there is no model to check it
|
||||
|
|
|
|||
|
|
@ -54,6 +54,12 @@ from litellm.proxy.common_utils.callback_utils import (
|
|||
get_logging_caching_headers,
|
||||
get_remaining_tokens_and_requests_from_request_data,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
attribute_of,
|
||||
error_status_code,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
from litellm.proxy.common_utils.sse_keepalive import (
|
||||
SSE_COMMENT_PING_BYTES,
|
||||
coerce_keepalive_interval,
|
||||
|
|
@ -464,46 +470,6 @@ def _stream_usage_tracking_updates(
|
|||
}
|
||||
|
||||
|
||||
def _getattr_object(value: object, name: str, default: object = None) -> object:
|
||||
return getattr(value, name, default)
|
||||
|
||||
|
||||
_OPENAI_ERROR_TYPE_BY_STATUS: Final[Mapping[int, str]] = MappingProxyType(
|
||||
{
|
||||
status.HTTP_401_UNAUTHORIZED: "authentication_error",
|
||||
status.HTTP_403_FORBIDDEN: "permission_error",
|
||||
status.HTTP_429_TOO_MANY_REQUESTS: "rate_limit_error",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _error_status_code(exc: object, default: int) -> int:
|
||||
"""The HTTP status an exception carries, or ``default`` when it carries none."""
|
||||
carried: Final = _getattr_object(exc, "status_code")
|
||||
return carried if isinstance(carried, int) and not isinstance(carried, bool) else default
|
||||
|
||||
|
||||
def _openai_error_type(exc: object, status_code: int) -> str:
|
||||
"""OpenAI types ``error.type`` as a required string, so an exception carrying none
|
||||
falls back to the type its status code stands for."""
|
||||
carried: Final = _getattr_object(exc, "type")
|
||||
if isinstance(carried, str):
|
||||
return carried
|
||||
mapped: Final = _OPENAI_ERROR_TYPE_BY_STATUS.get(status_code)
|
||||
if mapped is not None:
|
||||
return mapped
|
||||
if status_code < status.HTTP_500_INTERNAL_SERVER_ERROR:
|
||||
return "invalid_request_error"
|
||||
return "internal_server_error"
|
||||
|
||||
|
||||
def _openai_error_param(exc: object) -> str | None:
|
||||
"""OpenAI types ``error.param`` as nullable, so an exception carrying none
|
||||
serializes as JSON ``null``."""
|
||||
carried: Final = _getattr_object(exc, "param")
|
||||
return carried if isinstance(carried, str) else None
|
||||
|
||||
|
||||
class _UpstreamHttpResponse(Protocol):
|
||||
@property
|
||||
def status_code(self) -> int: ...
|
||||
|
|
@ -573,15 +539,15 @@ def serialize_http_exception_detail(
|
|||
|
||||
|
||||
def proxy_exception_from_http_exception(exc: HTTPException, headers: dict[str, str]) -> ProxyException:
|
||||
raw_detail: Final = _getattr_object(exc, "detail", str(exc))
|
||||
raw_detail: Final = attribute_of(exc, "detail", str(exc))
|
||||
message, structured_fields = serialize_http_exception_detail(raw_detail)
|
||||
existing_fields: Final = getattr(exc, "provider_specific_fields", None) or {}
|
||||
merged_fields: Final = {**existing_fields, **structured_fields} if structured_fields else (existing_fields or None)
|
||||
error_status: Final = _error_status_code(exc, status.HTTP_400_BAD_REQUEST)
|
||||
error_status: Final = error_status_code(exc, status.HTTP_400_BAD_REQUEST)
|
||||
return ProxyException(
|
||||
message=message,
|
||||
type=_openai_error_type(exc, error_status),
|
||||
param=_openai_error_param(exc),
|
||||
type=openai_error_type(exc, error_status),
|
||||
param=openai_error_param(exc),
|
||||
code=error_status,
|
||||
provider_specific_fields=merged_fields,
|
||||
headers=headers,
|
||||
|
|
@ -865,8 +831,8 @@ def sse_error_payload(exc: BaseException) -> tuple[int, Mapping[str, object]]:
|
|||
are byte-identical.
|
||||
"""
|
||||
# Preserve status code from HTTPException (e.g. guardrail blocks)
|
||||
error_status: Final = _error_status_code(exc, status.HTTP_500_INTERNAL_SERVER_ERROR)
|
||||
raw_detail: Final = _getattr_object(exc, "detail", "Error processing stream start")
|
||||
error_status: Final = error_status_code(exc, status.HTTP_500_INTERNAL_SERVER_ERROR)
|
||||
raw_detail: Final = attribute_of(exc, "detail", "Error processing stream start")
|
||||
message, structured_fields = serialize_http_exception_detail(raw_detail)
|
||||
|
||||
existing_fields: Final = getattr(exc, "provider_specific_fields", None) or {}
|
||||
|
|
@ -874,8 +840,8 @@ def sse_error_payload(exc: BaseException) -> tuple[int, Mapping[str, object]]:
|
|||
|
||||
error_obj: Final = {
|
||||
"message": message,
|
||||
"type": _openai_error_type(exc, error_status),
|
||||
"param": _openai_error_param(exc),
|
||||
"type": openai_error_type(exc, error_status),
|
||||
"param": openai_error_param(exc),
|
||||
"code": str(error_status),
|
||||
}
|
||||
if not merged_fields:
|
||||
|
|
@ -2815,10 +2781,10 @@ class ProxyBaseLLMRequestProcessing:
|
|||
``ResponsesAPIResponse`` directly. Handle both shapes so the
|
||||
container-ownership recording path can walk ``.output`` either way.
|
||||
"""
|
||||
completed: Final = _getattr_object(stream_response, "completed_response")
|
||||
completed: Final = attribute_of(stream_response, "completed_response")
|
||||
if completed is None:
|
||||
return None
|
||||
response_obj: Final = _getattr_object(completed, "response")
|
||||
response_obj: Final = attribute_of(completed, "response")
|
||||
if response_obj is not None:
|
||||
return response_obj
|
||||
return completed
|
||||
|
|
@ -3468,7 +3434,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
headers = getattr(e, "headers", None) or {}
|
||||
if not headers:
|
||||
# Try to get headers from e.response.headers (httpx.Response)
|
||||
_response: Final = _getattr_object(e, "response")
|
||||
_response: Final = attribute_of(e, "response")
|
||||
if _response is not None:
|
||||
_response_headers: Final = getattr(_response, "headers", None)
|
||||
if _response_headers:
|
||||
|
|
@ -3543,8 +3509,8 @@ class ProxyBaseLLMRequestProcessing:
|
|||
_code = status.HTTP_500_INTERNAL_SERVER_ERROR
|
||||
raise ProxyException(
|
||||
message=redact_internal_details_from_client_message(getattr(e, "message", error_msg)),
|
||||
type=_openai_error_type(e, _code),
|
||||
param=_openai_error_param(e),
|
||||
type=openai_error_type(e, _code),
|
||||
param=openai_error_param(e),
|
||||
openai_code=getattr(e, "code", None),
|
||||
code=_code,
|
||||
provider_specific_fields=getattr(e, "provider_specific_fields", None),
|
||||
|
|
@ -3754,11 +3720,11 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
if isinstance(e, HTTPException):
|
||||
raise e
|
||||
stream_error_status: Final = _error_status_code(e, status.HTTP_500_INTERNAL_SERVER_ERROR)
|
||||
stream_error_status: Final = error_status_code(e, status.HTTP_500_INTERNAL_SERVER_ERROR)
|
||||
proxy_exception: Final = ProxyException(
|
||||
message=redact_internal_details_from_client_message(getattr(e, "message", str(e))),
|
||||
type=_openai_error_type(e, stream_error_status),
|
||||
param=_openai_error_param(e),
|
||||
type=openai_error_type(e, stream_error_status),
|
||||
param=openai_error_param(e),
|
||||
code=stream_error_status,
|
||||
)
|
||||
stream_completed = True
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ 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
|
||||
|
|
@ -509,6 +510,7 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset(
|
|||
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",
|
||||
|
|
|
|||
52
litellm/proxy/common_utils/openai_error_payload.py
Normal file
52
litellm/proxy/common_utils/openai_error_payload.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
"""Shapes the ``error`` object the proxy answers with so it matches OpenAI's contract:
|
||||
``type`` is a required string and ``param`` is nullable, neither of which the literal
|
||||
string ``"None"`` satisfies."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from fastapi import status
|
||||
|
||||
_OPENAI_ERROR_TYPE_BY_STATUS: Final[Mapping[int, str]] = MappingProxyType(
|
||||
{
|
||||
status.HTTP_401_UNAUTHORIZED: "authentication_error",
|
||||
status.HTTP_403_FORBIDDEN: "permission_error",
|
||||
status.HTTP_429_TOO_MANY_REQUESTS: "rate_limit_error",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def attribute_of(value: object, name: str, default: object = None) -> object:
|
||||
return getattr(value, name, default)
|
||||
|
||||
|
||||
def error_status_code(exc: object, default: int) -> int:
|
||||
"""The HTTP status an exception carries as ``status_code`` or, the way ``ProxyException``
|
||||
stores it, as a stringified ``code``; ``default`` when it carries neither."""
|
||||
carried: Final = attribute_of(exc, "status_code")
|
||||
if isinstance(carried, int) and not isinstance(carried, bool):
|
||||
return carried
|
||||
stringified: Final = attribute_of(exc, "code")
|
||||
return int(stringified) if isinstance(stringified, str) and stringified.isdecimal() else default
|
||||
|
||||
|
||||
def openai_error_type(exc: object, status_code: int) -> str:
|
||||
"""OpenAI types ``error.type`` as a required string, so an exception carrying none
|
||||
falls back to the type its status code stands for."""
|
||||
carried: Final = attribute_of(exc, "type")
|
||||
if isinstance(carried, str):
|
||||
return carried
|
||||
mapped: Final = _OPENAI_ERROR_TYPE_BY_STATUS.get(status_code)
|
||||
if mapped is not None:
|
||||
return mapped
|
||||
if status_code < status.HTTP_500_INTERNAL_SERVER_ERROR:
|
||||
return "invalid_request_error"
|
||||
return "internal_server_error"
|
||||
|
||||
|
||||
def openai_error_param(exc: object) -> str | None:
|
||||
"""OpenAI types ``error.param`` as nullable, so an exception carrying none
|
||||
serializes as JSON ``null``."""
|
||||
carried: Final = attribute_of(exc, "param")
|
||||
return carried if isinstance(carried, str) else None
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -20,6 +20,11 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
coerce_numeric_form_fields,
|
||||
numeric_form_fields,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.types.images.main import ImageEditRequestParams
|
||||
from litellm.types.llms.openai import ChatCompletionUserMessage
|
||||
|
|
@ -200,18 +205,18 @@ async def image_generation(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
openai_code=getattr(e, "code", None),
|
||||
code=getattr(e, "status_code", 500),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from pydantic import TypeAdapter
|
||||
from pydantic import ValidationError as PydanticValidationError
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
|
|
@ -24,9 +25,11 @@ from litellm.constants import (
|
|||
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,
|
||||
X_LITELLM_DISABLE_CALLBACKS,
|
||||
)
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
||||
|
|
@ -158,6 +161,7 @@ from litellm.types.utils import (
|
|||
CustomPricingLiteLLMParams,
|
||||
LlmProviders,
|
||||
ProviderSpecificHeader,
|
||||
StandardCallbackDynamicParams,
|
||||
StandardLoggingUserAPIKeyMetadata,
|
||||
SupportedCacheControls,
|
||||
)
|
||||
|
|
@ -171,6 +175,7 @@ _ENABLE_TEAM_STALE_ALIAS_BYPASS: bool | None = None
|
|||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.otel.model.destination import OtelDestination
|
||||
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
|
||||
from litellm.types.proxy.policy_engine import Policy, PolicyMatchContext
|
||||
|
||||
|
|
@ -290,6 +295,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",
|
||||
|
|
@ -973,6 +979,145 @@ def _get_dynamic_logging_metadata(
|
|||
return callback_settings_obj
|
||||
|
||||
|
||||
_TENANT_OTEL_PARAMS: Final = TypeAdapter(StandardCallbackDynamicParams)
|
||||
|
||||
|
||||
def _tenant_otel_params(callback_vars: Mapping[str, str]) -> StandardCallbackDynamicParams:
|
||||
try:
|
||||
return _TENANT_OTEL_PARAMS.validate_python(callback_vars)
|
||||
except PydanticValidationError:
|
||||
return StandardCallbackDynamicParams()
|
||||
|
||||
|
||||
_NO_REQUEST_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _dynamically_disabled_backends(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_headers: Mapping[str, str] | None,
|
||||
) -> frozenset[str]:
|
||||
"""The callbacks this request turned off, read the way dispatch reads them.
|
||||
|
||||
Same sources, precedence, and premium gate ``EnterpriseCallbackControls`` applies
|
||||
before it skips a callback: the ``x-litellm-disable-callbacks`` header wins over the
|
||||
key's stored list, team settings are not a source, and a non-premium proxy honours
|
||||
neither. A destination has to agree with that decision, or a backend the key turned
|
||||
off would still be exported to, now through the fan-out instead of the callback.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
|
||||
if litellm.allow_dynamic_callback_disabling is not True or not premium_user:
|
||||
return frozenset()
|
||||
header: Final = (request_headers if request_headers is not None else _NO_REQUEST_HEADERS).get(
|
||||
X_LITELLM_DISABLE_CALLBACKS
|
||||
)
|
||||
if header is not None:
|
||||
return frozenset(name.strip().lower() for name in header.split(","))
|
||||
metadata: Final = user_api_key_dict.metadata
|
||||
disabled: Final = metadata.get("litellm_disabled_callbacks") if metadata else None
|
||||
if not isinstance(disabled, list):
|
||||
return frozenset()
|
||||
return frozenset(name.lower() for name in disabled if isinstance(name, str))
|
||||
|
||||
|
||||
def resolve_tenant_otel_destinations(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_headers: Mapping[str, str] | None = None,
|
||||
) -> "tuple[OtelDestination, ...]":
|
||||
"""The OTLP destinations this request's key or team config overrides its traces to.
|
||||
|
||||
Key settings win over team settings outright, the same precedence
|
||||
``_get_dynamic_logging_metadata`` applies, so one caller never exports the same
|
||||
backend to two accounts. An empty key-level list counts as configured, since that
|
||||
is what disabling a key's callbacks writes. Returns empty when OTEL V2 is off, when
|
||||
neither level named a destination-capable backend, or when the config is
|
||||
incomplete, and the request then keeps the operator's own exporters.
|
||||
|
||||
Two entries naming the same backend merge their ``callback_vars`` last-wins, the
|
||||
way ``convert_key_logging_metadata_to_callback`` merges them, so the destination
|
||||
and the per-request tracer routing cannot read one config two ways.
|
||||
|
||||
A ``failure``-only entry is skipped: a destination is resolved during auth, before
|
||||
the request has an outcome, so honouring the filter would mean holding every span
|
||||
back until the call finishes. Those entries keep today's behaviour instead, where
|
||||
the tenant's credentials reach the backend through per-request tracer routing and
|
||||
the operator's exporter is left alone. Its ``callback_vars`` still take part in the
|
||||
merge for a backend another entry made eligible, so the destination carries the
|
||||
same credentials the runtime parser resolves for that request.
|
||||
|
||||
A backend the request disabled dynamically, through the key's
|
||||
``litellm_disabled_callbacks`` or the ``x-litellm-disable-callbacks`` header in
|
||||
``request_headers``, resolves to no destination, so the fan-out never carries the
|
||||
request tree to that account and the operator's exporter is never suppressed for
|
||||
it. That leaves the request exactly where it stood before destinations existed:
|
||||
the OTel V2 logger itself is not on the disable list's class registry, so its own
|
||||
span still routes to the tenant's credentials the way it did then.
|
||||
"""
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
from litellm.integrations.otel.presets.destinations import destination_for
|
||||
|
||||
if not is_otel_v2_enabled():
|
||||
return ()
|
||||
key_entries: Final = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict)
|
||||
entries: Final = (
|
||||
key_entries
|
||||
if key_entries is not None
|
||||
else KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict)
|
||||
)
|
||||
if not entries:
|
||||
return ()
|
||||
disabled: Final = _dynamically_disabled_backends(user_api_key_dict, request_headers)
|
||||
callbacks: Final = tuple(
|
||||
callback
|
||||
for item in entries
|
||||
if (callback := _get_validated_callback_metadata(item=item, source="otel-destination")) is not None
|
||||
if callback.callback_name.lower() not in disabled
|
||||
)
|
||||
return tuple(
|
||||
destination
|
||||
for name in dict.fromkeys(
|
||||
callback.callback_name for callback in callbacks if callback.callback_type != "failure"
|
||||
)
|
||||
if (
|
||||
destination := destination_for(
|
||||
name,
|
||||
_tenant_otel_params(
|
||||
MappingProxyType(
|
||||
{
|
||||
var: value
|
||||
for callback in callbacks
|
||||
if callback.callback_name == name
|
||||
for var, value in callback.callback_vars.items()
|
||||
}
|
||||
)
|
||||
),
|
||||
_tenant_service_name(user_api_key_dict),
|
||||
)
|
||||
)
|
||||
is not None
|
||||
)
|
||||
|
||||
|
||||
def _tenant_service_name(user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
"""The ``service.name`` this key or team configured, the key winning over its team.
|
||||
|
||||
Same fields and same precedence the request-metadata build applies, read straight
|
||||
off the auth object because destinations resolve during auth, before that metadata
|
||||
is assembled.
|
||||
"""
|
||||
sources: Final = (user_api_key_dict.metadata, user_api_key_dict.team_metadata)
|
||||
return next(
|
||||
(
|
||||
stripped
|
||||
for source in sources
|
||||
if source
|
||||
for field in OTEL_SERVICE_NAME_METADATA_KEYS
|
||||
if isinstance(value := source.get(field), str) and (stripped := value.strip())
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def clean_headers(
|
||||
headers: Headers,
|
||||
litellm_key_header_name: str | None = None,
|
||||
|
|
@ -3078,10 +3223,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(
|
||||
|
|
@ -3098,10 +3242,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))
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.proxy._types import CommonProxyErrors
|
|||
from litellm.proxy.spend_tracking.key_metadata_recovery import (
|
||||
attach_user_emails,
|
||||
recover_double_hashed_key_metadata,
|
||||
recover_key_metadata_from_spend_logs,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
|
@ -433,9 +434,29 @@ def update_breakdown_metrics(
|
|||
return breakdown
|
||||
|
||||
|
||||
def _spend_logs_window(dates: AbstractSet[str | None]) -> tuple[datetime, datetime] | None:
|
||||
parsed: Final = sorted(day for day in (_parse_spend_date(raw) for raw in dates) if day is not None)
|
||||
if not parsed:
|
||||
return None
|
||||
return (parsed[0] - timedelta(days=1), parsed[-1] + timedelta(days=2))
|
||||
|
||||
|
||||
def _parse_spend_date(raw: str | None) -> datetime | None:
|
||||
if not isinstance(raw, str):
|
||||
return None
|
||||
try:
|
||||
return datetime.fromisoformat(raw)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
_EMPTY_KEY_METADATA: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType({})
|
||||
|
||||
|
||||
async def get_api_key_metadata(
|
||||
prisma_client: PrismaClient,
|
||||
api_keys: AbstractSet[str],
|
||||
spend_logs_window: tuple[datetime, datetime] | None = None,
|
||||
) -> Mapping[str, _KeyMetadataDict]:
|
||||
"""Get api key metadata, falling back to deleted keys table for keys not found in active table.
|
||||
|
||||
|
|
@ -481,11 +502,17 @@ async def get_api_key_metadata(
|
|||
)
|
||||
|
||||
still_missing: Final = api_keys - frozenset(result)
|
||||
combined: Final = (
|
||||
result
|
||||
if not still_missing
|
||||
else MappingProxyType({**result, **(await recover_double_hashed_key_metadata(prisma_client, still_missing))})
|
||||
from_reverse_hash: Final = (
|
||||
await recover_double_hashed_key_metadata(prisma_client, still_missing) if still_missing else _EMPTY_KEY_METADATA
|
||||
)
|
||||
after_token_recovery: Final = MappingProxyType({**result, **from_reverse_hash})
|
||||
unresolved: Final = api_keys - frozenset(after_token_recovery)
|
||||
from_spend_logs: Final = (
|
||||
await recover_key_metadata_from_spend_logs(prisma_client, unresolved, spend_logs_window)
|
||||
if unresolved and spend_logs_window is not None
|
||||
else _EMPTY_KEY_METADATA
|
||||
)
|
||||
combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs})
|
||||
return await attach_user_emails(prisma_client, combined)
|
||||
|
||||
|
||||
|
|
@ -898,7 +925,9 @@ async def _aggregate_spend_records(
|
|||
|
||||
api_key_metadata: dict[str, _KeyMetadataDict] = {}
|
||||
if api_keys:
|
||||
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
|
||||
api_key_metadata = await get_api_key_metadata(
|
||||
prisma_client, api_keys, _spend_logs_window(frozenset(record.date for record in records))
|
||||
)
|
||||
|
||||
return await asyncio.to_thread(
|
||||
_aggregate_spend_records_sync,
|
||||
|
|
@ -1094,7 +1123,9 @@ async def _aggregate_grouping_sets_records(
|
|||
|
||||
api_key_metadata: dict[str, _KeyMetadataDict] = {}
|
||||
if api_keys:
|
||||
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
|
||||
api_key_metadata = await get_api_key_metadata(
|
||||
prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records))
|
||||
)
|
||||
|
||||
return await asyncio.to_thread(
|
||||
_aggregate_grouping_sets_records_sync,
|
||||
|
|
@ -1357,7 +1388,9 @@ async def get_daily_activity_aggregated(
|
|||
r.api_key for r in entity_records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY
|
||||
)
|
||||
entity_key_metadata: Final = (
|
||||
await get_api_key_metadata(prisma_client, entity_api_keys)
|
||||
await get_api_key_metadata(
|
||||
prisma_client, entity_api_keys, _spend_logs_window(frozenset(r.date for r in entity_records))
|
||||
)
|
||||
if entity_api_keys
|
||||
else {} # mutable-ok: matches the helper's dict return
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2673,6 +2673,8 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
Updates the MCP Server in the db.
|
||||
|
||||
Partial update: a field left out of the payload keeps its stored value, and a field sent as null is cleared.
|
||||
|
||||
Parameters:
|
||||
- payload: UpdateMCPServerRequest - Required. The updated mcp server data.
|
||||
```
|
||||
|
|
@ -3098,6 +3100,8 @@ if MCP_AVAILABLE:
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
litellm_changed_by: str | None = Header(None),
|
||||
):
|
||||
"""Partial update: a field left out keeps its stored value, and a field sent as null is cleared, except
|
||||
``toolset_name`` and ``tools``, which a toolset always has; empty the tool selection with an explicit []."""
|
||||
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
|
||||
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -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 /
|
||||
|
|
|
|||
|
|
@ -45,6 +45,11 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
get_custom_llm_provider_from_request_headers,
|
||||
get_custom_llm_provider_from_request_query,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.batch_file_validation import (
|
||||
check_batch_file_upload,
|
||||
raise_batch_file_validation_failure,
|
||||
|
|
@ -296,22 +301,22 @@ async def route_create_file(
|
|||
if managed_files_obj is None:
|
||||
raise ProxyException(
|
||||
message="Managed files hook not found",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
if llm_router is None:
|
||||
raise ProxyException(
|
||||
message="LLM Router not found",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
if not isinstance(managed_files_obj, BaseFileEndpoints):
|
||||
raise ProxyException(
|
||||
message="Managed files hook is not a BaseFileEndpoints",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
# Managed files internally calls llm_router.acreate_file() which includes loadbalancing
|
||||
|
|
@ -713,17 +718,17 @@ async def create_file(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e.detail)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
finally:
|
||||
for spool in spools:
|
||||
|
|
@ -812,22 +817,22 @@ async def get_file_content(
|
|||
if managed_files_obj is None:
|
||||
raise ProxyException(
|
||||
message="Managed files hook not found",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
if llm_router is None:
|
||||
raise ProxyException(
|
||||
message="LLM Router not found",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
if not isinstance(managed_files_obj, BaseFileEndpoints):
|
||||
raise ProxyException(
|
||||
message="Managed files hook is not a BaseFileEndpoints",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
|
||||
|
|
@ -1021,17 +1026,17 @@ async def get_file_content(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e.detail)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1151,15 +1156,15 @@ async def get_file(
|
|||
if managed_files_obj is None:
|
||||
raise ProxyException(
|
||||
message="Managed files hook not found",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
if not isinstance(managed_files_obj, BaseFileEndpoints):
|
||||
raise ProxyException(
|
||||
message="Managed files hook is not a BaseFileEndpoints",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
response = await managed_files_obj.afile_retrieve(
|
||||
|
|
@ -1215,17 +1220,17 @@ async def get_file(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e.detail)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1355,22 +1360,22 @@ async def delete_file(
|
|||
if managed_files_obj is None:
|
||||
raise ProxyException(
|
||||
message="Managed files hook not found",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
if llm_router is None:
|
||||
raise ProxyException(
|
||||
message="LLM Router not found",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
if not isinstance(managed_files_obj, BaseFileEndpoints):
|
||||
raise ProxyException(
|
||||
message="Managed files hook is not a BaseFileEndpoints",
|
||||
type="None",
|
||||
param="None",
|
||||
type=ProxyErrorTypes.internal_server_error.value,
|
||||
param=None,
|
||||
code=500,
|
||||
)
|
||||
|
||||
|
|
@ -1427,17 +1432,17 @@ async def delete_file(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e.detail)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1629,15 +1634,15 @@ async def list_files(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e.detail)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -78,6 +78,11 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
_read_request_body,
|
||||
_safe_get_request_headers,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
from litellm.proxy.common_utils.sse_keepalive import (
|
||||
wrap_passthrough_sse_bytes_with_keepalive_pings,
|
||||
)
|
||||
|
|
@ -311,9 +316,9 @@ async def chat_completion_pass_through_endpoint(
|
|||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1728,18 +1733,18 @@ async def pass_through_request(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(getattr(e, "detail", str(e)))),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
headers=custom_headers,
|
||||
)
|
||||
else:
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
headers=custom_headers,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -5,12 +5,16 @@ 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 Mapping, 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,
|
||||
|
|
@ -21,6 +25,7 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
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,
|
||||
)
|
||||
|
|
@ -29,7 +34,14 @@ from litellm.types.proxy.policy_engine.pipeline_types import (
|
|||
PipelineStep,
|
||||
PipelineStepResult,
|
||||
)
|
||||
from litellm.types.utils import StandardLoggingGuardrailInformation
|
||||
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
|
||||
|
|
@ -37,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."""
|
||||
|
||||
|
|
@ -49,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.
|
||||
|
|
@ -65,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
|
||||
|
|
@ -89,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
|
||||
|
|
@ -114,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
|
||||
|
|
@ -129,6 +268,7 @@ class PipelineExecutor:
|
|||
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":
|
||||
|
|
@ -137,6 +277,7 @@ class PipelineExecutor:
|
|||
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
|
||||
|
|
@ -144,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,
|
||||
|
|
@ -152,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,
|
||||
|
|
@ -175,29 +363,13 @@ 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)
|
||||
|
||||
# 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]
|
||||
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 = "apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks
|
||||
if use_unified:
|
||||
use_unified: Final = PipelineExecutor.supports_unified_execution(callback)
|
||||
if use_unified and streaming_chunks is None:
|
||||
hook_input["guardrail_to_apply"] = callback
|
||||
target = UnifiedLLMGuardrails()
|
||||
|
||||
|
|
@ -213,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,
|
||||
|
|
@ -226,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):
|
||||
|
|
@ -246,6 +444,12 @@ class PipelineExecutor:
|
|||
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:
|
||||
"""Look up an initialized guardrail callback by name from litellm.callbacks."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -2582,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
|
||||
|
|
@ -2600,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(
|
||||
|
|
@ -2663,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",
|
||||
|
|
@ -2685,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,
|
||||
|
|
@ -5517,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",
|
||||
|
|
@ -17745,6 +17750,7 @@ async def reload_model_cost_map(
|
|||
# Immediately reload the model cost map in the current pod
|
||||
from litellm.litellm_core_utils.get_model_cost_map import (
|
||||
ModelCostMapReloadUnavailable,
|
||||
get_model_cost_map_provenance,
|
||||
refetch_model_cost_map,
|
||||
)
|
||||
|
||||
|
|
@ -17758,6 +17764,7 @@ async def reload_model_cost_map(
|
|||
models_count = _swap_in_model_cost_map(reload_result.model_cost_map)
|
||||
current_time = utc_now()
|
||||
proxy_config.model_cost_map_loaded_at = current_time
|
||||
provenance: Final = get_model_cost_map_provenance()
|
||||
|
||||
# Publish a new revision so every other pod reloads on its next poll; this pod has
|
||||
# already served it, so adopt it here rather than reloading again a tick later
|
||||
|
|
@ -17772,6 +17779,7 @@ async def reload_model_cost_map(
|
|||
"status": "success",
|
||||
"models_count": models_count,
|
||||
"timestamp": current_time.isoformat(),
|
||||
**provenance,
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
|
|
@ -17892,12 +17900,17 @@ async def get_model_cost_map_reload_status(
|
|||
|
||||
try:
|
||||
global prisma_client
|
||||
from litellm.litellm_core_utils.get_model_cost_map import (
|
||||
get_model_cost_map_provenance,
|
||||
)
|
||||
|
||||
provenance: Final = get_model_cost_map_provenance()
|
||||
if prisma_client is None:
|
||||
verbose_proxy_logger.info("No database connection, returning not scheduled")
|
||||
return reload_schedule_status(None)
|
||||
return {**reload_schedule_status(None), **provenance}
|
||||
|
||||
return reload_schedule_status(await read_reload_schedule(prisma_client, MODEL_COST_MAP_RELOAD_PARAM_NAME))
|
||||
schedule: Final = await read_reload_schedule(prisma_client, MODEL_COST_MAP_RELOAD_PARAM_NAME)
|
||||
return {**reload_schedule_status(schedule), **provenance}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Failed to get model cost map reload status: %s", e)
|
||||
raise HTTPException(
|
||||
|
|
@ -17925,6 +17938,9 @@ async def get_model_cost_map_source(
|
|||
- url: the remote URL that was attempted (null when env-forced local)
|
||||
- is_env_forced: true if LITELLM_LOCAL_MODEL_COST_MAP=True forced local usage
|
||||
- fallback_reason: human-readable reason why remote failed (null on success)
|
||||
- loaded_at: when this pod last loaded the map
|
||||
- source_revision: git blob id of the loaded file, what git rev-parse <commit>:<path> prints for it
|
||||
- etag: the ETag of the remote fetch (null for the bundled backup)
|
||||
- model_count: number of models in the currently loaded cost map
|
||||
"""
|
||||
# Read-only source info — admin viewers can read.
|
||||
|
|
@ -18479,6 +18495,31 @@ async def _stream_mcp_asgi_response(handle_fn, scope: dict, receive) -> "Streami
|
|||
########################################################
|
||||
|
||||
|
||||
@app.api_route(
|
||||
"/mcp/proxy",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], # mutable-ok: FastAPI route methods
|
||||
)
|
||||
async def proxy_mcp_route(request: Request) -> Response:
|
||||
"""Serve the fixed three-tool MCP proxy surface."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import ( # pyright: ignore[reportPrivateUsage] # route-owned mode
|
||||
_mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # route-owned mode
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.server import handle_streamable_http_mcp
|
||||
from litellm.proxy._experimental.mcp_server.utils import is_mcp_available
|
||||
|
||||
if not is_mcp_available():
|
||||
raise HTTPException(status_code=404, detail="Not Found")
|
||||
|
||||
token: Final = _mcp_proxy_mode.set(True)
|
||||
try:
|
||||
scope: Final = dict(request.scope) # mutable-ok: ASGI scope rewrite
|
||||
scope["_original_path"] = scope.get("path", "")
|
||||
scope["path"] = BASE_MCP_ROUTE
|
||||
return await _stream_mcp_asgi_response(handle_streamable_http_mcp, scope, request.receive)
|
||||
finally:
|
||||
_mcp_proxy_mode.reset(token)
|
||||
|
||||
|
||||
@app.api_route(
|
||||
BASE_MCP_ROUTE,
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"],
|
||||
|
|
|
|||
|
|
@ -10,7 +10,12 @@
|
|||
"REASONING": ["claude-opus-5"]
|
||||
},
|
||||
"tier_model_configs": {
|
||||
"REASONING": [{ "model_name": "claude-opus-5", "litellm_params": { "reasoning_effort": "high" } }]
|
||||
"REASONING": [
|
||||
{
|
||||
"model_name": "claude-opus-5",
|
||||
"litellm_params": { "reasoning_effort": "high" }
|
||||
}
|
||||
]
|
||||
},
|
||||
"classifier_type": "heuristic_v2",
|
||||
"escalation_keywords": ["LITELLM ESCALATE"],
|
||||
|
|
@ -23,16 +28,21 @@
|
|||
},
|
||||
"anthropic_family": {
|
||||
"label": "Anthropic Family",
|
||||
"description": "Routes across the Claude model family: Haiku for simple queries, Sonnet for medium, Opus for complex, Opus at high thinking for reasoning.",
|
||||
"description": "Routes across the Claude model family: Haiku for simple queries, Sonnet for medium, Opus for complex, Fable 5.1 at high thinking for reasoning.",
|
||||
"complexity_router_config": {
|
||||
"tiers": {
|
||||
"SIMPLE": ["claude-haiku-4-5"],
|
||||
"MEDIUM": ["claude-sonnet-5"],
|
||||
"COMPLEX": ["claude-opus-5"],
|
||||
"REASONING": ["claude-opus-5"]
|
||||
"REASONING": ["claude-fable-5-1"]
|
||||
},
|
||||
"tier_model_configs": {
|
||||
"REASONING": [{ "model_name": "claude-opus-5", "litellm_params": { "reasoning_effort": "high" } }]
|
||||
"REASONING": [
|
||||
{
|
||||
"model_name": "claude-fable-5-1",
|
||||
"litellm_params": { "reasoning_effort": "high" }
|
||||
}
|
||||
]
|
||||
},
|
||||
"classifier_type": "heuristic",
|
||||
"escalation_keywords": ["LITELLM ESCALATE"],
|
||||
|
|
@ -73,8 +83,18 @@
|
|||
"REASONING": ["claude-opus-5"]
|
||||
},
|
||||
"tier_model_configs": {
|
||||
"MEDIUM": [{ "model_name": "muse-spark-1.2", "litellm_params": { "reasoning_effort": "xhigh" } }],
|
||||
"COMPLEX": [{ "model_name": "kimi-k3", "litellm_params": { "reasoning_effort": "max" } }]
|
||||
"MEDIUM": [
|
||||
{
|
||||
"model_name": "muse-spark-1.2",
|
||||
"litellm_params": { "reasoning_effort": "xhigh" }
|
||||
}
|
||||
],
|
||||
"COMPLEX": [
|
||||
{
|
||||
"model_name": "kimi-k3",
|
||||
"litellm_params": { "reasoning_effort": "max" }
|
||||
}
|
||||
]
|
||||
},
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {
|
||||
|
|
@ -93,16 +113,21 @@
|
|||
},
|
||||
"openai_family": {
|
||||
"label": "OpenAI Family",
|
||||
"description": "Routes across the GPT model family: Luna for simple queries, Terra for medium, Sol for complex, Sol at xhigh thinking for reasoning.",
|
||||
"description": "Routes across the GPT model family: Luna for simple queries, Terra for medium, Sol for complex, Astra at xhigh thinking for reasoning.",
|
||||
"complexity_router_config": {
|
||||
"tiers": {
|
||||
"SIMPLE": ["gpt-5.6-luna"],
|
||||
"MEDIUM": ["gpt-5.6-terra"],
|
||||
"COMPLEX": ["gpt-5.6-sol"],
|
||||
"REASONING": ["gpt-5.6-sol"]
|
||||
"REASONING": ["gpt-6-astra"]
|
||||
},
|
||||
"tier_model_configs": {
|
||||
"REASONING": [{ "model_name": "gpt-5.6-sol", "litellm_params": { "reasoning_effort": "xhigh" } }]
|
||||
"REASONING": [
|
||||
{
|
||||
"model_name": "gpt-6-astra",
|
||||
"litellm_params": { "reasoning_effort": "xhigh" }
|
||||
}
|
||||
]
|
||||
},
|
||||
"classifier_type": "heuristic",
|
||||
"escalation_keywords": ["LITELLM ESCALATE"],
|
||||
|
|
|
|||
|
|
@ -17,6 +17,11 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
|||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
from litellm.types.realtime import (
|
||||
RealtimeClientSecretRequest,
|
||||
RealtimeClientSecretResponse,
|
||||
|
|
@ -304,15 +309,15 @@ async def create_realtime_client_secret(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", http_status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, http_status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, http_status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
if upstream_resp.status_code != 200:
|
||||
|
|
@ -495,15 +500,15 @@ async def proxy_realtime_calls(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", http_status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, http_status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, http_status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
return Response(
|
||||
|
|
@ -608,15 +613,15 @@ async def create_realtime_transcription_session(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "detail", getattr(e, "message", str(e))),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", http_status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, http_status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, http_status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
||||
if upstream_resp.status_code != 200:
|
||||
|
|
|
|||
|
|
@ -11,6 +11,11 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
openai_error_param,
|
||||
openai_error_type,
|
||||
)
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
|
@ -112,15 +117,15 @@ async def rerank(
|
|||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
|
||||
type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
|
||||
)
|
||||
else:
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
type=openai_error_type(e, error_status_code(e, 500)),
|
||||
param=openai_error_param(e),
|
||||
code=error_status_code(e, 500),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue