Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_fix_guardrail_event_hook_resync

# Conflicts:
#	tests/test_litellm/integrations/test_custom_guardrail.py
This commit is contained in:
mateo-berri 2026-09-02 11:57:22 -07:00
commit 70fbf7f4a3
55 changed files with 4433 additions and 359 deletions

View file

@ -46,6 +46,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3.13
# Stage 2 — copy source and install the project + workspace members.
@ -57,6 +58,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -18,7 +18,7 @@ type: application
# This is the chart version. This version number should be incremented each time you make changes
# to the chart and its templates, including the app version.
# Versions are expected to follow Semantic Versioning (https://semver.org/)
version: 1.1.2
version: 1.1.3
# This is the version number of the application being deployed. This version number should be
# incremented each time you make changes to the application. Versions are not expected to

View file

@ -1,4 +1,4 @@
suite: "hpa with behavior"
suite: "hpa"
templates:
- hpa.yaml
tests:
@ -23,14 +23,44 @@ tests:
- equal: { path: spec.behavior.scaleUp.stabilizationWindowSeconds, value: 60 }
- equal: { path: spec.behavior.scaleDown.stabilizationWindowSeconds, value: 90 }
---
suite: "hpa without behavior"
templates:
- hpa.yaml
tests:
- it: "does not render behavior when not set"
set:
autoscaling.enabled: true
asserts:
- isKind: { of: HorizontalPodAutoscaler }
- isNull: { path: spec.behavior }
- it: "scales on cpu at the documented 60 percent by default"
set:
autoscaling.enabled: true
asserts:
- isKind: { of: HorizontalPodAutoscaler }
- equal: { path: "spec.metrics[0].resource.name", value: cpu }
- equal: { path: "spec.metrics[0].resource.target.type", value: Utilization }
- equal: { path: "spec.metrics[0].resource.target.averageUtilization", value: 60 }
- it: "does not scale on memory by default"
set:
autoscaling.enabled: true
asserts:
- lengthEqual: { path: spec.metrics, count: 1 }
- it: "honours an explicit cpu target override"
set:
autoscaling.enabled: true
autoscaling.targetCPUUtilizationPercentage: 75
asserts:
- equal: { path: "spec.metrics[0].resource.target.averageUtilization", value: 75 }
- it: "renders a memory metric only when a memory target is set"
set:
autoscaling.enabled: true
autoscaling.targetMemoryUtilizationPercentage: 80
asserts:
- lengthEqual: { path: spec.metrics, count: 2 }
- equal: { path: "spec.metrics[1].resource.name", value: memory }
- equal: { path: "spec.metrics[1].resource.target.averageUtilization", value: 80 }
- it: "renders no hpa when autoscaling is disabled"
asserts:
- hasDocuments: { count: 0 }

View file

@ -200,7 +200,16 @@ autoscaling:
enabled: false
minReplicas: 1
maxReplicas: 100
targetCPUUtilizationPercentage: 80
# 60 is the documented recommendation. See "Recommended Machine Specifications"
# in https://docs.litellm.ai/docs/proxy/prod. A new replica clears the startupProbe
# above only after up to failureThreshold x periodSeconds = 300 seconds, so a target
# high enough to trip near saturation adds capacity minutes after it was needed.
targetCPUUtilizationPercentage: 60
# Deliberately left unset rather than given a value. The prisma query engine's
# resident memory is a high-water mark that ratchets to the pod's worst-ever write
# and is never returned, so a memory target reads the largest write a pod ever did
# rather than what it is doing now, and replicas ratchet up without scaling back in.
# Memory is a floor to provision under 'resources', not a signal to scale on.
# targetMemoryUtilizationPercentage: 80
# behavior: {}

View file

@ -9,6 +9,38 @@ DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT"
AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
ROUTER_MAX_FALLBACKS: Final = int(os.getenv("ROUTER_MAX_FALLBACKS", 5))
ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS: Final = 2000
RUNTIME_UPDATABLE_ROUTER_SETTINGS: Final[frozenset[str]] = frozenset(
{
"routing_strategy_args",
"routing_strategy",
"routing_groups",
"allowed_fails",
"cooldown_time",
"num_retries",
"timeout",
"max_retries",
"retry_after",
"fallbacks",
"context_window_fallbacks",
"retry_policy",
"model_group_retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_tag_filtering",
"tag_routing_prefix",
"optional_pre_call_checks",
}
)
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
{
"model_list",
"search_tools",
"assistants_config",
"router_general_settings",
"ignore_invalid_deployments",
"fallback_access_check",
}
)
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))

View file

@ -1,4 +1,5 @@
import contextvars
import copy
import hashlib
import os
import secrets
@ -41,6 +42,7 @@ except ImportError:
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
dc: Final = DualCache()
@ -872,6 +874,69 @@ class CustomGuardrail(CustomLogger):
return result
async def async_logging_hook(
self,
kwargs: dict, # mutable-ok: CustomLogger.async_logging_hook contract
result: object,
call_type: str,
) -> tuple[dict, object]: # mutable-ok: CustomLogger.async_logging_hook contract
"""logging_only: run apply_guardrail on copies of the logged request/response and record the verdict."""
from litellm.llms import get_guardrail_translation_mapping
if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks:
return kwargs, result
try:
translation: Final = get_guardrail_translation_mapping(CallTypes(call_type))()
except ValueError:
verbose_logger.debug(
"Guardrail %s: no guardrail translation for call_type=%s, skipping logging_only scan",
self.guardrail_name,
call_type,
)
return kwargs, result
litellm_params: Final = kwargs.get("litellm_params") or {}
scratch_metadata: Final = {
key: value
for key, value in (litellm_params.get("metadata") or {}).items()
if key != "standard_logging_guardrail_information"
}
try:
await self._scan_logged_call(kwargs, result, translation, scratch_metadata)
except Exception as e:
verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e)
recorded: Final = scratch_metadata.get("standard_logging_guardrail_information")
standard_logging_object: Final = kwargs.get("standard_logging_object")
if not recorded or not isinstance(standard_logging_object, dict):
return kwargs, result
entries: Final = recorded if isinstance(recorded, list) else [recorded]
existing: Final = standard_logging_object.get("guardrail_information") or []
return {
**kwargs,
"standard_logging_object": {**standard_logging_object, "guardrail_information": [*existing, *entries]},
}, result
async def _scan_logged_call(
self,
kwargs: dict, # mutable-ok: CustomLogger.async_logging_hook contract
result: object,
translation: "BaseTranslation",
scratch_metadata: dict, # mutable-ok: apply_guardrail records its verdict into request metadata
) -> None:
optional_params: Final = kwargs.get("optional_params") or {}
scratch_input: Final = copy.deepcopy(kwargs.get("messages") or kwargs.get("input"))
scratch_request: Final = {
"model": kwargs.get("model"),
"messages": scratch_input,
"input": scratch_input,
"tools": copy.deepcopy(optional_params.get("tools")),
"litellm_call_id": kwargs.get("litellm_call_id"),
"metadata": scratch_metadata,
}
await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self)
await translation.process_output_response(
response=copy.deepcopy(result), guardrail_to_apply=self, request_data=scratch_request
)
def supports_scan_only_tool_results(self) -> bool:
"""Whether this guardrail can scan tool-result content.

View file

@ -36,6 +36,8 @@ from litellm.types.utils import (
from litellm.utils import print_verbose, token_counter
if TYPE_CHECKING:
from openai.types.completion_usage import CompletionUsage
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.litellm_core_utils.streaming_chunk_builder_utils import (
UsagePerChunk,
@ -794,7 +796,7 @@ class ChunkProcessor:
@staticmethod
def _extract_usage_chunk(chunk: "_UsageBearingChunk | ModelResponse | ModelResponseStream") -> Usage | None:
usage_chunk: Usage | None = None
usage_chunk: Usage | CompletionUsage | None = None
if hasattr(chunk, "usage") and chunk.usage is not None:
usage_chunk = chunk.usage
elif "usage" in chunk:
@ -806,7 +808,9 @@ class ChunkProcessor:
if isinstance(usage_chunk, dict):
return Usage(**usage_chunk)
return usage_chunk
if usage_chunk is None or isinstance(usage_chunk, Usage):
return usage_chunk
return Usage(**usage_chunk.model_dump())
def _calculate_usage_per_chunk(
self,

View file

@ -1378,31 +1378,38 @@ def process_anthropic_headers(headers: httpx.Headers | dict) -> dict:
return additional_headers
def _anthropic_model_entry(model: ModelInfoResponse, created_at: str) -> Mapping[str, object]:
def _anthropic_model_entry(
model: ModelInfoResponse, created_at: str, display_names: Mapping[str, str]
) -> Mapping[str, object]:
return { # mutable-ok: JSON response body, serialized by the route and never mutated
"type": "model",
"id": model["id"],
"display_name": model["id"],
"display_name": display_names.get(model["id"], model["id"]),
"created_at": created_at,
"max_input_tokens": model.get("max_input_tokens"),
"max_tokens": model.get("max_output_tokens"),
}
def create_anthropic_model_list_response(models: Sequence[ModelInfoResponse]) -> Mapping[str, object]:
def create_anthropic_model_list_response(
models: Sequence[ModelInfoResponse],
display_names: Mapping[str, str] = MappingProxyType({}),
) -> Mapping[str, object]:
"""Build the Anthropic-native /v1/models envelope.
Clients that send an anthropic-version header parse the Anthropic Models API
shape (type/display_name/created_at plus has_more/first_id/last_id) and filter
the list themselves, so every model is returned here. The token limits carry
over from the OpenAI-shaped listing, named as the Messages API names them, and
are always present because the vendor shape declares them nullable, not optional
are always present because the vendor shape declares them nullable, not optional.
display_names maps a listed model id to a configured human-readable name; ids
without an entry fall back to the id itself, matching the vendor behavior
"""
created_at: Final = (
datetime.fromtimestamp(DEFAULT_MODEL_CREATED_AT_TIME, tz=timezone.utc).isoformat().replace("+00:00", "Z")
)
data: Final = [ # mutable-ok: JSON response body, serialized by the route and never mutated
_anthropic_model_entry(model, created_at) for model in models
_anthropic_model_entry(model, created_at, display_names) for model in models
]
return { # mutable-ok: JSON response body, serialized by the route and never mutated
"data": data,

View file

@ -1442,7 +1442,7 @@ class BaseAWSLLM:
@tracer.wrap()
def get_request_headers(
self,
credentials: Credentials,
credentials: Credentials | None,
aws_region_name: str,
extra_headers: dict | None,
endpoint_url: str,
@ -1469,9 +1469,13 @@ class BaseAWSLLM:
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.exceptions import NoCredentialsError
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
if credentials is None:
raise NoCredentialsError()
# Filter headers for AWS signature calculation
# AWS SigV4 only includes specific headers in signature calculation
aws_signature_headers: Final = self._filter_headers_for_aws_signature(headers)

View file

@ -1,4 +1,6 @@
import json
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final
import httpx
@ -24,6 +26,22 @@ from ..common_utils import BedrockError, _get_all_bedrock_regions
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
def _sigv4_principal(credentials: Credentials | None) -> Mapping[str, str]:
if credentials is None:
return MappingProxyType({})
return MappingProxyType(
{
key: value
for key, value in (
("aws_access_key_id", credentials.access_key),
("aws_secret_access_key", credentials.secret_key),
("aws_session_token", credentials.token),
)
if value is not None
}
)
def make_sync_call(
client: HTTPHandler | None,
api_base: str,
@ -95,7 +113,7 @@ class BedrockConverseLLM(BaseAWSLLM):
stream,
optional_params: dict,
litellm_params: dict,
credentials: Credentials,
credentials: Credentials | None,
logger_fn=None,
headers={},
client: AsyncHTTPHandler | None = None,
@ -167,7 +185,7 @@ class BedrockConverseLLM(BaseAWSLLM):
stream,
optional_params: dict,
litellm_params: dict,
credentials: Credentials,
credentials: Credentials | None,
logger_fn=None,
headers: dict = {},
client: AsyncHTTPHandler | None = None,
@ -331,7 +349,7 @@ class BedrockConverseLLM(BaseAWSLLM):
litellm_params["aws_region_name"] = aws_region_name # [DO NOT DELETE] important for async calls
credentials: Final[Credentials] = self.get_credentials(
credentials: Final[Credentials | None] = self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
@ -368,19 +386,13 @@ class BedrockConverseLLM(BaseAWSLLM):
# The Rust core owns the whole call for the subset it accepts. Ask
# before transforming so whichever path runs emits pre_call once, and
# hand down the credentials, region and endpoint this handler already
# resolved so both paths sign as the same principal.
# resolved so both paths sign as the same principal. Bearer-token auth
# resolves no SigV4 principal at all, and each path reads that token
# itself.
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
**optional_params,
**{ # mutable-ok: merged into its mutable parent above
key: value
for key, value in (
("aws_access_key_id", credentials.access_key),
("aws_secret_access_key", credentials.secret_key),
("aws_session_token", credentials.token),
("aws_region_name", aws_region_name),
)
if value is not None
},
**_sigv4_principal(credentials),
"aws_region_name": aws_region_name,
}
serves_via_rust: Final = rust_chat_completions_accepts(
model=model,

View file

@ -0,0 +1,90 @@
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final
from pydantic import TypeAdapter, ValidationError
from litellm.utils import get_model_info
PARALLEL_AI_DEFAULT_RESULTS: Final = 10
PARALLEL_AI_ADDITIONAL_RESULT_COST: Final = 0.001
PARALLEL_AI_USAGE_PARAM: Final = "_parallel_ai_usage"
PARALLEL_AI_STANDARD_SEARCH_MODEL: Final = "parallel_ai/search"
PARALLEL_AI_FAST_SEARCH_MODEL: Final = "parallel_ai/search-fast"
PARALLEL_AI_TURBO_SEARCH_MODEL: Final = "parallel_ai/search-turbo"
PARALLEL_AI_PRICING_MODEL_BY_MODE: Final[Mapping[str, str]] = MappingProxyType(
{
"fast": PARALLEL_AI_FAST_SEARCH_MODEL,
"turbo": PARALLEL_AI_TURBO_SEARCH_MODEL,
}
)
ADVANCED_SETTINGS_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
def _non_negative_int(value: object) -> int | None:
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
return None
return value
def _usage_count(usage: Sequence[Mapping[str, object]], sku: str) -> int | None:
counts: Final = tuple(
count
for item in usage
if item.get("name") == sku
if (count := _non_negative_int(item.get("count"))) is not None
)
return sum(counts) if counts else None
def _effective_mode(optional_params: Mapping[str, object]) -> str:
mode: Final = optional_params.get("mode")
if isinstance(mode, str):
return mode
processor: Final = optional_params.get("processor")
if processor == "pro":
return "advanced"
return "basic"
def _effective_max_results(optional_params: Mapping[str, object]) -> int:
try:
advanced_settings: Final = ADVANCED_SETTINGS_ADAPTER.validate_python(optional_params.get("advanced_settings"))
advanced_max_results: Final = _non_negative_int(advanced_settings.get("max_results"))
if advanced_max_results is not None:
return advanced_max_results
except ValidationError:
pass
max_results: Final = _non_negative_int(optional_params.get("max_results"))
return max_results if max_results is not None else PARALLEL_AI_DEFAULT_RESULTS
def _request_cost(mode: str) -> float:
pricing_model: Final = PARALLEL_AI_PRICING_MODEL_BY_MODE.get(mode, PARALLEL_AI_STANDARD_SEARCH_MODEL)
model_info: Final = get_model_info(model=pricing_model, custom_llm_provider="parallel_ai")
return float(model_info.get("input_cost_per_query") or 0.0)
def _additional_results(
optional_params: Mapping[str, object],
usage: Sequence[Mapping[str, object]] | None,
) -> int:
usage_count: Final = _usage_count(usage, "sku_search_additional_results") if usage is not None else None
if usage_count is not None:
return usage_count
if usage is not None:
return 0
return max(_effective_max_results(optional_params) - PARALLEL_AI_DEFAULT_RESULTS, 0)
def parallel_ai_search_cost(
optional_params: Mapping[str, object],
usage: Sequence[Mapping[str, object]] | None,
) -> float:
request_cost: Final = _request_cost(_effective_mode(optional_params))
request_count_from_usage: Final = _usage_count(usage, "sku_search") if usage is not None else None
request_count: Final = request_count_from_usage if request_count_from_usage is not None else 1
additional_results: Final = _additional_results(optional_params, usage)
return request_count * request_cost + additional_results * PARALLEL_AI_ADDITIONAL_RESULT_COST

View file

@ -4,9 +4,13 @@ Calls Parallel AI's /v1/search endpoint to search the web.
Parallel AI API Reference: https://docs.parallel.ai/api-reference/search/search
"""
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final, TypedDict
import httpx
from pydantic import BaseModel, ConfigDict
from typing_extensions import ReadOnly
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.search.transformation import (
@ -14,9 +18,29 @@ from litellm.llms.base_llm.search.transformation import (
SearchResponse,
SearchResult,
)
from litellm.llms.parallel_ai.search.cost_calculator import PARALLEL_AI_USAGE_PARAM
from litellm.secret_managers.main import get_secret_str
class _ParallelAIV1SearchResult(BaseModel):
model_config = ConfigDict(extra="ignore")
url: str | None = None
title: str | None = None
publish_date: str | None = None
excerpts: Sequence[str] | None = None
class _ParallelAIV1SearchResponse(BaseModel):
model_config = ConfigDict(extra="ignore")
search_id: str | None = None
session_id: str | None = None
results: Sequence[_ParallelAIV1SearchResult] = ()
usage: Sequence[Mapping[str, object]] | None = None
warnings: Sequence[Mapping[str, object]] | None = None
class _ParallelAISourcePolicy(TypedDict, total=False):
include_domains: list[str]
exclude_domains: list[str]
@ -27,10 +51,16 @@ class _ParallelAIExcerptSettings(TypedDict, total=False):
max_chars_per_result: int
class _ParallelAIFetchPolicy(TypedDict, total=False):
max_age_seconds: ReadOnly[int]
timeout_seconds: ReadOnly[float]
disable_cache_fallback: ReadOnly[bool]
class _ParallelAIAdvancedSettings(TypedDict, total=False):
source_policy: _ParallelAISourcePolicy
excerpt_settings: _ParallelAIExcerptSettings
fetch_policy: dict
fetch_policy: _ParallelAIFetchPolicy
location: str
max_results: int
@ -43,14 +73,14 @@ class ParallelAISearchRequest(TypedDict, total=False):
search_queries: list[str] # Required - at least one keyword search query
objective: str # Optional - natural-language description of search goal
mode: str # Optional - 'turbo', 'basic', or 'advanced' (default 'advanced')
mode: str # Optional - 'turbo', 'fast', 'basic', or 'advanced' (default 'advanced')
max_chars_total: int # Optional - upper bound on total excerpt characters
session_id: str # Optional - tracks calls across search/extract requests
client_model: str # Optional - model consuming the results
advanced_settings: _ParallelAIAdvancedSettings
LEGACY_PROCESSOR_TO_MODE: Final = {"base": "basic", "pro": "advanced"}
LEGACY_PROCESSOR_TO_MODE: Final = MappingProxyType({"base": "basic", "pro": "advanced"})
class ParallelAISearchConfig(BaseSearchConfig):
@ -67,16 +97,16 @@ class ParallelAISearchConfig(BaseSearchConfig):
api_base: str | None = None,
**kwargs,
) -> dict:
api_key = self.resolve_server_api_key(
resolved_api_key: Final = self.resolve_server_api_key(
caller_api_key=api_key,
caller_api_base=api_base,
key_env_vars=("PARALLEL_AI_API_KEY", "PARALLEL_API_KEY"),
base_env_var="PARALLEL_AI_API_BASE",
default_api_base=self.PARALLEL_AI_API_BASE,
)
if not api_key:
if not resolved_api_key:
raise ValueError("PARALLEL_API_KEY is not set. Set `PARALLEL_API_KEY` environment variable.")
headers["x-api-key"] = api_key
headers["x-api-key"] = resolved_api_key
headers["Content-Type"] = "application/json"
return headers
@ -87,13 +117,12 @@ class ParallelAISearchConfig(BaseSearchConfig):
data: dict | list[dict] | None = None,
**kwargs,
) -> str:
api_base = api_base or get_secret_str("PARALLEL_AI_API_BASE") or self.PARALLEL_AI_API_BASE
resolved_api_base: Final = api_base or get_secret_str("PARALLEL_AI_API_BASE") or self.PARALLEL_AI_API_BASE
api_base = api_base.rstrip("/")
if not api_base.endswith("/v1/search"):
api_base = f"{api_base.removesuffix('/v1')}/v1/search"
return api_base
trimmed: Final = resolved_api_base.rstrip("/")
if trimmed.endswith("/v1/search"):
return trimmed
return f"{trimmed.removesuffix('/v1')}/v1/search"
def transform_search_request(
self,
@ -109,14 +138,17 @@ class ParallelAISearchConfig(BaseSearchConfig):
- If string: maps to `search_queries` (single item) and `objective`
- If list: maps to `search_queries` (keyword queries)
optional_params: Optional parameters for the request
- mode: Search mode ('turbo', 'basic', 'advanced'); defaults to 'basic'
- mode: Search mode ('turbo', 'fast', 'basic', 'advanced'); defaults to 'basic'
- processor: Legacy v1beta param; 'base' maps to mode 'basic', 'pro' to 'advanced'
- max_results: Maximum number of search results -> `advanced_settings.max_results`
- search_domain_filter: Domains to include -> `advanced_settings.source_policy.include_domains`
- search_domain_filter / include_domains: Domains to include -> `advanced_settings.source_policy.include_domains`
- exclude_domains: Domains to exclude -> `advanced_settings.source_policy.exclude_domains`
- country: ISO 3166-1 alpha-2 code -> `advanced_settings.location`
- after_date: RFC 3339 date (YYYY-MM-DD) -> `advanced_settings.source_policy.after_date`
- country / location: ISO 3166-1 alpha-2 code -> `advanced_settings.location`
- max_chars_per_result: -> `advanced_settings.excerpt_settings.max_chars_per_result`
- Any other params are passed through to the request body as-is
- fetch_policy: Cache vs live-fetch policy -> `advanced_settings.fetch_policy`
- Any other params (objective, max_chars_total, session_id, client_model, ...)
are passed through to the request body as-is
Returns:
Dict with request data following the v1 search request spec
@ -137,7 +169,7 @@ class ParallelAISearchConfig(BaseSearchConfig):
mode = LEGACY_PROCESSOR_TO_MODE.get(processor, processor)
# the v1 API defaults to 'advanced' when mode is omitted; default to 'basic'
# instead to keep v1beta's default tier (processor 'base') and litellm's
# $0.004/query cost map entry for `parallel_ai/search` accurate
# cost map entry for `parallel_ai/search` accurate
request_data["mode"] = mode or "basic"
advanced_settings: Final[_ParallelAIAdvancedSettings] = {}
@ -148,17 +180,29 @@ class ParallelAISearchConfig(BaseSearchConfig):
if "country" in params:
advanced_settings["location"] = params.pop("country")
if "location" in params:
advanced_settings["location"] = params.pop("location")
if "max_chars_per_result" in params:
advanced_settings["excerpt_settings"] = {"max_chars_per_result": params.pop("max_chars_per_result")}
if "fetch_policy" in params:
advanced_settings["fetch_policy"] = params.pop("fetch_policy")
source_policy: Final[_ParallelAISourcePolicy] = {}
if "search_domain_filter" in params:
source_policy["include_domains"] = params.pop("search_domain_filter")
if "include_domains" in params:
source_policy["include_domains"] = params.pop("include_domains")
if "exclude_domains" in params:
source_policy["exclude_domains"] = params.pop("exclude_domains")
if "after_date" in params:
source_policy["after_date"] = params.pop("after_date")
if source_policy:
advanced_settings["source_policy"] = source_policy
@ -170,9 +214,11 @@ class ParallelAISearchConfig(BaseSearchConfig):
# unified-spec param with no v1 equivalent
params.pop("max_tokens_per_page", None)
result_data: Final[dict] = dict(request_data)
result_data.update(params)
return result_data
# reserved for the provider's own reported usage, which prices the request;
# a caller-supplied value would otherwise set its own cost
params.pop(PARALLEL_AI_USAGE_PARAM, None)
return {**request_data, **params}
def transform_search_response(
self,
@ -186,26 +232,49 @@ class ParallelAISearchConfig(BaseSearchConfig):
Parallel AI -> LiteLLM mappings:
- results[].title -> SearchResult.title
- results[].url -> SearchResult.url
- results[].excerpts (array) -> SearchResult.snippet (joined string)
- results[].excerpts (array) -> SearchResult.snippet (joined string); the raw
array is preserved as an extra `excerpts` field on each result
- results[].publish_date -> SearchResult.date
- search_id / session_id / warnings are preserved as extra fields on the
response; usage is preserved as `parallel_usage` (the `usage` name is
reserved for LiteLLM's token-usage object)
"""
response_json: Final = raw_response.json()
parsed: Final = _ParallelAIV1SearchResponse.model_validate(raw_response.json())
results: Final = []
for result in response_json.get("results", []):
excerpts = result.get("excerpts") or []
snippet = " ... ".join(excerpts) if excerpts else ""
# written unconditionally: leaving a caller-supplied value in place when the
# provider reports no usage would let the caller price its own request
logging_obj.optional_params = {
**logging_obj.optional_params,
PARALLEL_AI_USAGE_PARAM: parsed.usage,
}
search_result = SearchResult(
title=result.get("title") or "",
url=result.get("url") or "",
snippet=snippet,
date=result.get("publish_date"),
last_updated=None,
results: Final = tuple(
SearchResult.model_validate(
MappingProxyType(
{
"title": result.title or "",
"url": result.url or "",
"snippet": " ... ".join(result.excerpts or ()),
"date": result.publish_date,
"last_updated": None,
"excerpts": result.excerpts or (),
}
)
)
results.append(search_result)
return SearchResponse(
results=results,
object="search",
for result in parsed.results
)
extra_fields: Final = MappingProxyType(
{
key: value
for key, value in (
("search_id", parsed.search_id),
("session_id", parsed.session_id),
("parallel_usage", parsed.usage),
("warnings", parsed.warnings),
)
if value is not None
}
)
return SearchResponse.model_validate(MappingProxyType({"results": results, "object": "search", **extra_fields}))

View file

@ -949,7 +949,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
# For Gemini 3+ models, use thinkingLevel instead of thinkingBudget
if model and VertexGeminiConfig._is_gemini_3_or_newer(model):
if thinking_enabled:
if thinking_budget is None or thinking_budget == 0:
if thinking_budget == 0:
params["includeThoughts"] = False
else:
params["includeThoughts"] = True

View file

@ -177,8 +177,9 @@ class VertexAIDeepSeekOCRConfig(BaseOCRConfig):
content_item = {"type": "image_url", "image_url": document_url}
# Build DeepSeek OCR request
provider_model: Final = model if model.startswith("deepseek-ai/") else f"deepseek-ai/{model}"
data: Final = {
"model": "deepseek-ai/" + model,
"model": provider_model,
"messages": [{"role": "user", "content": [content_item]}],
}

View file

@ -8637,6 +8637,16 @@ def _set_stream_builder_response_cost(response: ModelResponse, logging_obj: Opti
hidden_params["response_cost"] = response_cost
def _stamp_streaming_usage_cost(usage: Usage, response: ModelResponse, logging_obj: Optional["Logging"]) -> None:
if logging_obj is None:
return
if isinstance(getattr(usage, "cost", None), (int, float)):
return
computed_cost: Final = logging_obj._response_cost_calculator(result=response)
if isinstance(computed_cost, (int, float)) and computed_cost > 0:
setattr(usage, "cost", computed_cost)
def stream_chunk_builder(
chunks: list,
messages: list | None = None,
@ -8731,12 +8741,7 @@ def stream_chunk_builder(
)
break
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
setattr(
usage,
"cost",
logging_obj._response_cost_calculator(result=response),
)
_stamp_streaming_usage_cost(usage, response, logging_obj)
_set_stream_builder_response_cost(response, logging_obj)
processor.apply_provider_assembled_streaming_metadata(response, chunks, logging_obj)
@ -8915,10 +8920,7 @@ def stream_chunk_builder(
)
break
# Add cost to usage object if include_cost_in_streaming_usage is True
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
setattr(usage, "cost", logging_obj._response_cost_calculator(result=response))
_stamp_streaming_usage_cost(usage, response, logging_obj)
_set_stream_builder_response_cost(response, logging_obj)
processor.apply_provider_assembled_streaming_metadata(response, chunks, logging_obj)

File diff suppressed because it is too large Load diff

View file

@ -10,13 +10,36 @@ legacy internal names with `general_settings.use_team_public_model_name: false`.
from __future__ import annotations
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, cast
if TYPE_CHECKING:
from litellm.router import Router
def configured_display_names(
entries: Sequence[tuple[str, str]],
llm_router: Router | None,
) -> Mapping[str, str]:
"""response_id -> configured `model_info.display_name` for the listing entries
that have one.
Metadata is looked up by each entry's internal lookup id (so team-scoped rows
resolve), while the returned map is keyed by the public response id the
Anthropic-shaped listing is built from. Entries without a configured name are
omitted so the listing falls back to the id itself.
"""
if llm_router is None:
return MappingProxyType({})
resolved: Final = (
(response_id, llm_router.get_configured_display_name(lookup_id)) for response_id, lookup_id in entries
)
return MappingProxyType(
{response_id: display_name for response_id, display_name in resolved if display_name is not None}
)
class TeamModelNameTranslator:
"""Translates internal team routing keys to their public names for the model
listing/retrieve responses. Stateless; the live router and general_settings

View file

@ -38,6 +38,7 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_checks import (
_delete_cache_access_object, # pyright: ignore[reportPrivateUsage] # the access-group endpoints reach for this same cache primitive
)
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
from litellm.repositories.table_repositories import AccessGroupRepository
@ -72,8 +73,9 @@ _REPOINT_KEY_SQL: Final = (
def _raw_executor(prisma_client: object) -> _RawExecutor:
"""Narrow the untyped Prisma client down to the raw-query call this module makes."""
return AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client
"""Narrow the untyped Prisma client down to the raw-query call this module makes, pinned to the writer."""
db: Final = AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client
return WriterPinnedClient(db).db # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin
async def _invalidate_access_group_cache(access_group_id: str) -> None:

View file

@ -39,7 +39,7 @@ from typing import (
import anyio
import websockets
import websockets.exceptions
from pydantic import BaseModel, Json, JsonValue, ValidationError
from pydantic import BaseModel, Json, JsonValue, TypeAdapter, ValidationError
from typing_extensions import NotRequired, ReadOnly, assert_never
from litellm._uuid import uuid
@ -60,6 +60,7 @@ from litellm.constants import (
LITELLM_SETTINGS_SAFE_DB_OVERRIDES,
LITELLM_UI_ALLOW_HEADERS,
LITELLM_UI_SESSION_DURATION,
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
)
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
@ -253,6 +254,7 @@ from litellm.constants import (
PROXY_BUDGET_RESCHEDULER_MAX_TIME,
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS,
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG,
USER_SPEND_ALERTS_JOB_ID,
WEEKLY_SPEND_REPORT_JOB_ID,
)
@ -352,7 +354,10 @@ from litellm.proxy.common_utils.load_config_utils import (
get_file_contents_from_s3,
)
from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations
from litellm.proxy.common_utils.model_listing_utils import TeamModelNameTranslator
from litellm.proxy.common_utils.model_listing_utils import (
TeamModelNameTranslator,
configured_display_names,
)
from litellm.proxy.common_utils.openai_endpoint_utils import (
remove_sensitive_info_from_deployment,
)
@ -5710,13 +5715,9 @@ class ProxyConfig:
router_settings: Final = config.get("router_settings", None)
if router_settings and isinstance(router_settings, dict):
# model list and search_tools already set
exclude_args: Final = {
"model_list",
"search_tools",
}
available_args: Final = [x for x in litellm.Router.get_valid_args() if x not in exclude_args]
available_args: Final = [
x for x in litellm.Router.get_valid_args() if x not in ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG
]
for k, v in router_settings.items():
if k in available_args:
@ -10223,7 +10224,8 @@ async def model_list(
# The internal routing key drives the metadata/fallback lookup, while the
# public name is what the client sees as the model id.
model_data = []
for response_id, lookup_id in TeamModelNameTranslator.listing_entries(all_models, llm_router, settings):
admin_entries: Final = TeamModelNameTranslator.listing_entries(all_models, llm_router, settings)
for response_id, lookup_id in admin_entries:
model_info = create_model_info_response(
model_id=lookup_id,
provider="openai",
@ -10236,7 +10238,10 @@ async def model_list(
if wants_anthropic_format:
admin_listing: Final = cast(Sequence[ModelInfoResponse], model_data) # cast-ok: rows built above
return create_anthropic_model_list_response(admin_listing)
return create_anthropic_model_list_response(
admin_listing,
display_names=configured_display_names(admin_entries, llm_router),
)
return dict(
data=model_data,
@ -10267,7 +10272,8 @@ async def model_list(
# The internal routing key drives the metadata/fallback lookup, while the
# public name is what the client sees as the model id.
model_data = []
for response_id, lookup_id in TeamModelNameTranslator.listing_entries(all_models, llm_router, settings):
entries: Final = TeamModelNameTranslator.listing_entries(all_models, llm_router, settings)
for response_id, lookup_id in entries:
model_info = create_model_info_response(
model_id=lookup_id,
provider="openai",
@ -10280,7 +10286,10 @@ async def model_list(
if wants_anthropic_format:
listing: Final = cast(Sequence[ModelInfoResponse], model_data) # cast-ok: rows built above
return create_anthropic_model_list_response(listing)
return create_anthropic_model_list_response(
listing,
display_names=configured_display_names(entries, llm_router),
)
return dict(
data=model_data,
@ -16207,6 +16216,7 @@ async def invitation_delete(
)
async def update_config(
config_info: ConfigYAML,
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
@ -16222,6 +16232,26 @@ async def update_config(
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(status_code=403, detail="Only proxy admins can update config")
request_body: Final[Mapping[str, JsonValue]] = TypeAdapter(Mapping[str, JsonValue]).validate_python(
await request.json()
)
raw_router_settings: Final = request_body.get("router_settings")
if isinstance(raw_router_settings, dict):
supported_router_settings: Final = RUNTIME_UPDATABLE_ROUTER_SETTINGS | (
frozenset(litellm.Router.get_valid_args()) - ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG
)
unsupported_router_settings: Final = sorted(set(raw_router_settings) - supported_router_settings)
if unsupported_router_settings:
raise HTTPException(
status_code=400,
detail={
"error": (
f"Unsupported router settings: {', '.join(unsupported_router_settings)} "
"are not valid router settings"
)
},
)
if prisma_client is None:
raise Exception("No DB Connected")
@ -16323,11 +16353,19 @@ async def update_config(
)
# router_settings: merge existing + request, request wins.
if config_info.router_settings is not None:
if isinstance(raw_router_settings, dict):
existing = await _read_section("router_settings")
before_router_settings: Final = copy.deepcopy(existing)
updates = config_info.router_settings.dict(exclude_none=True)
new_router_settings: Final = {**existing, **updates}
typed_router_settings: Final = (
config_info.router_settings.dict(exclude_none=True) if config_info.router_settings is not None else {}
)
raw_router_settings_without_none: Final = {
key: value
for key, value in raw_router_settings.items()
if key not in typed_router_settings and value is not None
}
router_settings_updates: Final = {**typed_router_settings, **raw_router_settings_without_none}
new_router_settings: Final = {**existing, **router_settings_updates}
await _upsert_section("router_settings", new_router_settings)
asyncio.create_task(
create_config_audit_log(

View file

@ -6,6 +6,7 @@ from typing import Any, Final, Literal
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.llms.bedrock.rerank.handler import BedrockRerankHandler
@ -43,10 +44,23 @@ async def arerank(
"""
Async: Reranks a list of documents based on their relevance to the query
"""
_custom_llm_provider: str | None = (
None # rebind-ok: set by the declared-provider guard or the get_llm_provider unpack; read in the except
)
try:
loop: Final = asyncio.get_event_loop()
kwargs["arerank"] = True
declared_provider: Final = declared_authenticating_provider(model, custom_llm_provider)
if declared_provider is not None:
_custom_llm_provider = declared_provider # rebind-ok: see pre-declaration above
else:
_, _custom_llm_provider, _, _ = litellm.get_llm_provider( # rebind-ok: see pre-declaration above
model=model,
custom_llm_provider=custom_llm_provider,
api_base=kwargs.get("api_base", None),
)
func: Final = partial(
rerank,
model,
@ -70,7 +84,11 @@ async def arerank(
response = init_response
return response
except Exception as e:
raise e
raise exception_type(
model=model,
custom_llm_provider=_custom_llm_provider or custom_llm_provider,
original_exception=e,
)
@client
@ -115,6 +133,7 @@ def rerank(
model_info: Final = kwargs.get("model_info", None)
user: Final = kwargs.get("user", None)
client: Final = kwargs.get("client", None)
_custom_llm_provider: str | None = None # rebind-ok: set by the get_llm_provider unpack; read in the except
try:
_is_async: Final = kwargs.pop("arerank", False) is True
optional_params: Final = GenericLiteLLMParams(**kwargs)
@ -127,7 +146,7 @@ def rerank(
(
model,
_custom_llm_provider,
_custom_llm_provider, # rebind-ok: see pre-declaration above
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
@ -538,4 +557,8 @@ def rerank(
return response
except Exception as e:
verbose_logger.error("Error in rerank: %s", e)
raise exception_type(model=model, custom_llm_provider=custom_llm_provider, original_exception=e)
raise exception_type(
model=model,
custom_llm_provider=_custom_llm_provider or custom_llm_provider,
original_exception=e,
)

View file

@ -1169,16 +1169,6 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
def _emit_response_completed_event(self, litellm_model_response: ModelResponse) -> ResponseCompletedEvent | None:
if litellm_model_response:
# Add cost to usage object if include_cost_in_streaming_usage is True
if litellm.include_cost_in_streaming_usage and self.litellm_logging_obj is not None:
usage: Final[object] = getattr(litellm_model_response, "usage", None)
if usage is not None:
setattr(
usage,
"cost",
self.litellm_logging_obj._response_cost_calculator(result=litellm_model_response),
)
# Transform the response
responses_api_response: Final = (
LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(

View file

@ -407,23 +407,7 @@ class BaseResponsesAPIStreamingIterator:
openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED,
):
self.completed_response = openai_responses_api_chunk
# Add cost to usage object if include_cost_in_streaming_usage is True
if litellm.include_cost_in_streaming_usage and self.logging_obj is not None:
response_obj: Final[ResponsesAPIResponse | None] = getattr(
openai_responses_api_chunk, "response", None
)
if response_obj:
usage_obj: Final[ResponseAPIUsage | None] = getattr(response_obj, "usage", None)
if usage_obj is not None:
try:
cost: Final[float | None] = self.logging_obj._response_cost_calculator(
result=response_obj
)
if cost is not None:
setattr(usage_obj, "cost", cost)
except Exception:
# Best-effort usage cost annotation should not break stream replay.
pass
_stamp_responses_usage_cost(getattr(openai_responses_api_chunk, "response", None), self.logging_obj)
if _chunk_type == openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED:
self._handle_logging_failed_response()
@ -1274,6 +1258,24 @@ def _add_text_like_part_events(
)
def _stamp_responses_usage_cost(
response_obj: ResponsesAPIResponse | None, logging_obj: LiteLLMLoggingObj | None
) -> None:
if response_obj is None or logging_obj is None:
return
usage_obj: Final[ResponseAPIUsage | None] = getattr(response_obj, "usage", None)
if usage_obj is None:
return
if isinstance(getattr(usage_obj, "cost", None), (int, float)):
return
try:
cost: Final[float | None] = logging_obj._response_cost_calculator(result=response_obj)
except Exception:
return
if isinstance(cost, (int, float)) and cost > 0:
setattr(usage_obj, "cost", cost)
def build_synthetic_response_events(
*,
transformed: ResponsesAPIResponse,
@ -1281,15 +1283,7 @@ def build_synthetic_response_events(
chunk_size: int,
) -> list[ResponsesAPIStreamingResponse]:
openai_types: Final = _get_openai_response_types()
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
usage_obj: Final = transformed.usage if hasattr(transformed, "usage") else None
if usage_obj is not None:
try:
cost: Final[float | None] = logging_obj._response_cost_calculator(result=transformed)
if cost is not None:
setattr(usage_obj, "cost", cost)
except Exception:
pass
_stamp_responses_usage_cost(transformed, logging_obj)
events: Final[list[ResponsesAPIStreamingResponse]] = [
_build_response_status_event(openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED, transformed),

View file

@ -50,6 +50,7 @@ from litellm.constants import (
DEFAULT_HEALTH_CHECK_INTERVAL,
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
DEFAULT_MAX_LRU_CACHE_SIZE,
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
)
from litellm.integrations.custom_logger import CustomLogger
@ -354,6 +355,13 @@ _PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT")
_ALIAS_PARAMS_NEVER_FORWARDED: Final = frozenset({"model", "api_base", "api_key", "api_version"})
_ALIAS_MARKER_FORWARDED_PARAMS_KWARG: Final = "_alias_marker_forwarded_params"
_RUNTIME_TOGGLEABLE_PRE_CALL_CHECKS: Final[Mapping[str, type[CustomLogger]]] = MappingProxyType(
{
"prompt_caching": PromptCachingDeploymentCheck,
"enforce_model_rate_limits": ModelRateLimitingCheck,
}
)
def _stream_chunks_have_generated_content(chunks: Sequence[ModelResponseStream]) -> bool:
for chunk in chunks:
@ -2072,11 +2080,39 @@ class Router:
if _callback is None:
continue
if self.optional_callbacks is not None and any(
isinstance(callback, type(_callback)) for callback in self.optional_callbacks
):
continue
if self.optional_callbacks is None:
self.optional_callbacks = []
self.optional_callbacks.append(_callback)
litellm.logging_callback_manager.add_litellm_callback(_callback)
def set_optional_pre_call_checks(self, optional_pre_call_checks: OptionalPreCallChecks | None) -> None:
if optional_pre_call_checks is None:
return
requested: Final = frozenset(optional_pre_call_checks)
for name, callback_cls in _RUNTIME_TOGGLEABLE_PRE_CALL_CHECKS.items():
if name not in requested:
self._remove_optional_callbacks_of_type(callback_cls)
self.add_optional_pre_call_checks(optional_pre_call_checks)
def _remove_optional_callbacks_of_type(self, callback_cls: type[CustomLogger]) -> None:
if self.optional_callbacks is None or not any(type(cb) is callback_cls for cb in self.optional_callbacks):
return
self.optional_callbacks = [cb for cb in self.optional_callbacks if type(cb) is not callback_cls]
if any(
router is not self and any(type(cb) is callback_cls for cb in (router.optional_callbacks or []))
for router in tuple(_live_routers)
):
return
for cb in tuple(litellm.callbacks):
if type(cb) is callback_cls:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, cb, require_self=False
)
def print_deployment(self, deployment: dict):
"""
returns a copy of the deployment with the api key masked
@ -9746,6 +9782,26 @@ class Router:
coerce_token_limit(model_info.get("max_output_tokens")),
)
def get_configured_display_name(self, model_name: str) -> "str | None":
"""
Return the display_name explicitly configured in a concrete deployment's
model_info for model_name, via O(1) index lookup.
Returns None for wildcard-expanded or unknown names, and treats a
non-string or empty configured value as absent rather than failing the
listing. Like get_configured_token_limits, this never triggers pattern
matching or deep copies, so it is safe to call per listed model on the
/v1/models hot path.
"""
deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name)
if deployment is None:
return None
display_name: Final = deployment.model_info.get("display_name")
if isinstance(display_name, str) and display_name.strip():
return display_name
return None
def get_deployment_credentials_with_provider(
self, model_id: str, team_id: str | None = None
) -> dict[str, Any] | None:
@ -11331,27 +11387,6 @@ class Router:
"""
Update the router settings.
"""
# only the following settings are allowed to be configured
_allowed_settings: Final = [
"routing_strategy_args",
"routing_strategy",
"routing_groups",
"allowed_fails",
"cooldown_time",
"num_retries",
"timeout",
"max_retries",
"retry_after",
"fallbacks",
"context_window_fallbacks",
"retry_policy",
"model_group_retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_tag_filtering",
"tag_routing_prefix",
]
_int_settings: Final = [
"timeout",
"num_retries",
@ -11364,13 +11399,15 @@ class Router:
rebuild_routing_groups = False
relink_lar1_from_args = False
for var in kwargs:
if var in _allowed_settings:
if var in RUNTIME_UPDATABLE_ROUTER_SETTINGS:
if var in _int_settings:
_casted_value = int(kwargs[var])
setattr(self, var, _casted_value)
elif var == "routing_groups":
self._routing_groups_input = kwargs[var]
rebuild_routing_groups = True
elif var == "optional_pre_call_checks":
self.set_optional_pre_call_checks(kwargs[var])
elif var == "retry_policy":
value = kwargs[var]
if isinstance(value, dict):

View file

@ -9,6 +9,7 @@ import random
import traceback
from collections.abc import Callable
from functools import partial
from types import MappingProxyType
from typing import Any, Final
from litellm._logging import verbose_router_logger
@ -214,6 +215,15 @@ class SearchAPIRouter:
api_key, api_base = SearchAPIRouter._resolve_search_provider_credentials(
tool_litellm_params=litellm_params,
)
protected_params: Final = frozenset(("search_provider", "api_key", "api_base"))
search_params: Final = MappingProxyType(
{
key: value
for params in (litellm_params, kwargs)
for key, value in params.items()
if key not in protected_params and value is not None
}
)
verbose_router_logger.debug("Selected search tool with provider: %s", search_provider)
@ -222,7 +232,7 @@ class SearchAPIRouter:
search_provider=search_provider,
api_key=api_key,
api_base=api_base,
**kwargs,
**search_params,
)
return response

View file

@ -2,16 +2,37 @@
Cost calculation for search providers.
"""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from pydantic import TypeAdapter, ValidationError
from litellm.utils import get_model_info
PROVIDER_USAGE_ADAPTER: Final[TypeAdapter[tuple[Mapping[str, object], ...]]] = TypeAdapter(
tuple[Mapping[str, object], ...]
)
EMPTY_OPTIONAL_PARAMS: Final[Mapping[str, object]] = MappingProxyType({})
def _provider_usage(
optional_params: Mapping[str, object] | None,
usage_param: str,
) -> tuple[Mapping[str, object], ...] | None:
params: Final = optional_params if optional_params is not None else EMPTY_OPTIONAL_PARAMS
raw_usage: Final[object] = params.get(usage_param)
try:
return PROVIDER_USAGE_ADAPTER.validate_python(raw_usage)
except ValidationError:
return None
def search_provider_cost_per_query(
model: str,
custom_llm_provider: str | None = None,
number_of_queries: int = 1,
optional_params: dict | None = None,
optional_params: Mapping[str, object] | None = None,
) -> tuple[float, float]:
"""
Calculate cost for search-only providers.
@ -28,6 +49,18 @@ def search_provider_cost_per_query(
Returns:
Tuple of (input_cost, output_cost) where output_cost is always 0.0
"""
if custom_llm_provider == "parallel_ai":
from litellm.llms.parallel_ai.search.cost_calculator import (
PARALLEL_AI_USAGE_PARAM,
parallel_ai_search_cost,
)
input_cost: Final = parallel_ai_search_cost(
optional_params=optional_params if optional_params is not None else EMPTY_OPTIONAL_PARAMS,
usage=_provider_usage(optional_params, PARALLEL_AI_USAGE_PARAM),
)
return (input_cost, 0.0)
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
# Check for tiered pricing (e.g., Exa AI based on max_results)

View file

@ -106,6 +106,20 @@ class RetryPolicy(BaseModel):
InternalServerErrorRetries: int | None = None
OptionalPreCallChecks = list[
Literal[
"prompt_caching",
"router_budget_limiting",
"responses_api_deployment_check",
"deployment_affinity",
"session_affinity",
"forward_client_headers_by_model_group",
"enforce_model_rate_limits",
"encrypted_content_affinity",
]
]
class UpdateRouterConfig(BaseModel):
"""
Set of params that you can modify via `router.update_settings()`.
@ -128,6 +142,7 @@ class UpdateRouterConfig(BaseModel):
model_group_alias: dict[str, str | dict] | None = {}
enable_tag_filtering: bool | None = None
tag_routing_prefix: str | None = None
optional_pre_call_checks: OptionalPreCallChecks | None = None
model_config = ConfigDict(protected_namespaces=())
@ -869,20 +884,6 @@ class FallbackAccessCheck(Protocol):
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
OptionalPreCallChecks = list[
Literal[
"prompt_caching",
"router_budget_limiting",
"responses_api_deployment_check",
"deployment_affinity",
"session_affinity",
"forward_client_headers_by_model_group",
"enforce_model_rate_limits",
"encrypted_content_affinity",
]
]
class LiteLLM_RouterFileObject(TypedDict, total=False):
"""
Tracking the litellm params hash, used for mapping the file id to the right model

File diff suppressed because it is too large Load diff

View file

@ -24,25 +24,10 @@ from junit_properties import (
)
class FakeMarker:
def __init__(self, name: str, *args: object) -> None:
self.name = name
self.args = args
class FakeItem:
"""The three attributes junit_properties reads off a pytest Item."""
def __init__(
self, nodeid: str, location: tuple[str, int | None, str], markers: tuple[FakeMarker, ...] = ()
) -> None:
self.nodeid = nodeid
self.location = location
self.user_properties: list[tuple[str, str]] = []
self._markers = markers
def iter_markers(self, name: str):
return (marker for marker in self._markers if marker.name == name)
def collected_item(request: pytest.FixtureRequest, name: str) -> pytest.Item:
"""The Item pytest collected for test ``name`` in this file: the real nodeid,
location and marker machinery the collection hook reads, as pytest built it."""
return next(item for item in request.session.items if item.path == request.path and item.name == name)
def repo_root() -> Path | None:
@ -109,22 +94,22 @@ class TestSourceFromLocation:
class TestResultProperties:
def test_every_test_carries_package_covers_and_source(self) -> None:
item = FakeItem(
"logging/test_x.py::TestFoo::test_bar",
("logging/test_x.py", 40, "TestFoo.test_bar"),
(FakeMarker("covers", "LOG-1", "LOG-2"),),
)
assert result_properties(item) == (
("package", "logging"),
def test_every_test_carries_package_covers_and_source(self, request: pytest.FixtureRequest) -> None:
"""Read off this test's own collected Item, so the nodeid and location are
whatever pytest reports for the launch shape in use, and the marker is added
at run time so the coverage registry's collect-only pass never sees it."""
test = type(self).test_every_test_carries_package_covers_and_source
request.applymarker(pytest.mark.covers("LOG-1", "LOG-2"))
assert result_properties(collected_item(request, test.__name__)) == (
("package", "root"),
("covers", "LOG-1,LOG-2"),
("source", "tests/e2e/logging/test_x.py:41"),
("source", f"tests/e2e/test_junit_properties.py:{test.__code__.co_firstlineno}"),
)
def test_attach_is_idempotent(self) -> None:
def test_attach_is_idempotent(self, request: pytest.FixtureRequest) -> None:
"""Collection can run the hook more than once; a second pass must not
double the <property> entries in the report."""
item = FakeItem("logging/test_x.py::test_bar", ("logging/test_x.py", 40, "test_bar"))
item = collected_item(request, type(self).test_attach_is_idempotent.__name__)
attach_result_properties(item)
attach_result_properties(item)
assert [name for name, _ in item.user_properties] == ["package", "covers", "source"]

View file

@ -5,9 +5,11 @@ Note: Vertex AI OCR automatically converts URLs to base64 data URIs since
the Vertex AI endpoint doesn't have internet access.
"""
import os
import json
import os
import tempfile
from typing import Final
import pytest
from base_ocr_unit_tests import BaseOCRTest
@ -139,3 +141,19 @@ def test_vertex_ai_ocr_routing():
assert isinstance(
deepseek_variant, VertexAIDeepSeekOCRConfig
), "DeepSeek variant should route to VertexAIDeepSeekOCRConfig"
@pytest.mark.parametrize("model", ("deepseek-ocr-maas", "deepseek-ai/deepseek-ocr-maas"))
def test_deepseek_request_uses_single_provider_namespace(model: str) -> None:
from litellm.llms.vertex_ai.ocr.deepseek_transformation import (
VertexAIDeepSeekOCRConfig,
)
request: Final = VertexAIDeepSeekOCRConfig().transform_ocr_request(
model=model,
document={"type": "image_url", "image_url": "data:image/png;base64,AA=="},
optional_params={},
headers={},
)
assert request.data["model"] == "deepseek-ai/deepseek-ocr-maas"

View file

@ -3076,7 +3076,9 @@ async def test_update_config_success_callback_normalization():
admin_user = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-test"
)
await proxy_server.update_config(config_update, user_api_key_dict=admin_user)
request = MagicMock()
request.json = AsyncMock(return_value={"litellm_settings": {"success_callback": ["SQS", "sQs"]}})
await proxy_server.update_config(config_update, request=request, user_api_key_dict=admin_user)
assert (
"litellm_settings" in upserted

View file

@ -2294,6 +2294,7 @@ def search_tools():
"search_provider": "perplexity",
"api_key": "test-api-key",
"api_base": "https://api.perplexity.ai",
"mode": "turbo",
},
},
{
@ -2302,6 +2303,7 @@ def search_tools():
"search_provider": "perplexity",
"api_key": "test-api-key-2",
"api_base": "https://api.perplexity.ai",
"mode": "turbo",
},
},
]
@ -2393,6 +2395,7 @@ async def test_asearch_with_fallbacks_helper(search_tools):
assert "search_provider" in kwargs
assert kwargs["search_provider"] == "perplexity"
assert "api_key" in kwargs
assert kwargs["mode"] == "turbo"
assert kwargs["query"] == "helper test query"
return mock_response

View file

@ -2319,3 +2319,202 @@ class TestUpdateInMemoryLitellmParams:
assert guardrail.event_hook is GuardrailEventHooks.during_call
assert getattr(guardrail, "api_base", None) == "https://guardrail.example.com"
class _ApplyOnlyObserver(CustomGuardrail):
"""Overrides only apply_guardrail, like panw_prisma_airs; inherits async_logging_hook."""
def __init__(self, block: bool = False):
from litellm.types.guardrails import GuardrailEventHooks
super().__init__(guardrail_name="apply-only-observer", event_hook=GuardrailEventHooks.logging_only)
self.block = block
self.calls: list = []
@log_guardrail_information
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
from fastapi import HTTPException
self.calls.append((input_type, list(inputs.get("texts") or [])))
if self.block:
raise HTTPException(status_code=400, detail={"error": "flagged"})
return GenericGuardrailAPIInputs(texts=["[MASKED]" for _ in inputs.get("texts") or []])
def _logged_call(messages: list | str) -> tuple[dict, object]:
from litellm.types.utils import Choices, Message, ModelResponse
response = ModelResponse(choices=[Choices(message=Message(role="assistant", content="general kenobi"))])
kwargs = {
"model": "gpt-5.4-mini",
"messages": messages,
"litellm_call_id": "call-1",
"litellm_params": {"metadata": {"user_api_key_user_id": "u1"}},
"optional_params": {},
"standard_logging_object": {"guardrail_information": None},
}
return kwargs, response
class TestLoggingOnlyApplyGuardrail:
"""LIT-4876 regression: a guardrail in mode logging_only that implements only
apply_guardrail must still run against the logged request and response and
record guardrail_information, instead of inheriting the CustomLogger no-op."""
@pytest.mark.asyncio
async def test_runs_apply_guardrail_observe_only_and_records_verdict(self):
guardrail = _ApplyOnlyObserver()
messages = [{"role": "user", "content": "hello there"}]
kwargs, response = _logged_call(messages)
out_kwargs, out_response = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
assert guardrail.calls == [("request", ["hello there"]), ("response", ["general kenobi"])]
assert out_kwargs["messages"] == [{"role": "user", "content": "hello there"}]
assert out_response.choices[0].message.content == "general kenobi"
entries = out_kwargs["standard_logging_object"]["guardrail_information"]
assert [e["guardrail_name"] for e in entries] == ["apply-only-observer", "apply-only-observer"]
assert {e["guardrail_mode"] for e in entries} == {"logging_only"}
assert {e["guardrail_status"] for e in entries} == {"success"}
assert "standard_logging_guardrail_information" not in kwargs["litellm_params"]["metadata"]
assert kwargs["standard_logging_object"] == {"guardrail_information": None}
@pytest.mark.asyncio
async def test_appends_to_pre_call_verdicts_without_duplicating_them(self):
guardrail = _ApplyOnlyObserver()
kwargs, response = _logged_call([{"role": "user", "content": "hello there"}])
pre_call_entry = {"guardrail_name": "pii-blocker", "guardrail_mode": "pre_call", "guardrail_status": "success"}
kwargs["litellm_params"]["metadata"]["standard_logging_guardrail_information"] = [pre_call_entry]
kwargs["standard_logging_object"]["guardrail_information"] = [pre_call_entry]
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
entries = out_kwargs["standard_logging_object"]["guardrail_information"]
assert [e["guardrail_name"] for e in entries] == ["pii-blocker", "apply-only-observer", "apply-only-observer"]
assert kwargs["litellm_params"]["metadata"]["standard_logging_guardrail_information"] == [pre_call_entry]
@pytest.mark.asyncio
async def test_request_copy_failure_is_swallowed(self):
import threading
guardrail = _ApplyOnlyObserver()
kwargs, response = _logged_call([{"role": "user", "content": "hello there", "lock": threading.Lock()}])
out_kwargs, out_response = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
assert guardrail.calls == []
assert out_kwargs is kwargs
assert out_response is response
@pytest.mark.asyncio
async def test_block_verdict_is_recorded_without_raising(self):
guardrail = _ApplyOnlyObserver(block=True)
kwargs, response = _logged_call([{"role": "user", "content": "flagged content"}])
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
assert guardrail.calls == [("request", ["flagged content"])]
entries = out_kwargs["standard_logging_object"]["guardrail_information"]
assert [e["guardrail_status"] for e in entries] == ["guardrail_intervened"]
@pytest.mark.asyncio
async def test_call_type_without_translation_is_skipped(self):
guardrail = _ApplyOnlyObserver()
kwargs, response = _logged_call([{"role": "user", "content": "hello there"}])
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.amoderation.value)
assert guardrail.calls == []
assert out_kwargs["standard_logging_object"]["guardrail_information"] is None
@pytest.mark.asyncio
async def test_aembedding_scans_logged_input(self):
from litellm.types.utils import EmbeddingResponse
guardrail = _ApplyOnlyObserver()
kwargs, _ = _logged_call("hello there")
response = EmbeddingResponse(data=[{"embedding": [0.1], "index": 0, "object": "embedding"}])
out_kwargs, out_response = await guardrail.async_logging_hook(kwargs, response, CallTypes.aembedding.value)
assert guardrail.calls == [("request", ["hello there"])]
assert out_kwargs["messages"] == "hello there"
assert out_response is response
entries = out_kwargs["standard_logging_object"]["guardrail_information"]
assert [e["guardrail_status"] for e in entries] == ["success"]
@pytest.mark.asyncio
async def test_native_lifecycle_hook_guardrail_is_left_alone(self):
class _NativeHooks(_ApplyOnlyObserver):
use_native_lifecycle_hooks = True
guardrail = _NativeHooks()
kwargs, response = _logged_call([{"role": "user", "content": "hello there"}])
out_kwargs, out_response = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
assert guardrail.calls == []
assert out_kwargs is kwargs
assert out_response is response
@pytest.mark.asyncio
async def test_aresponses_scans_logged_messages_when_input_is_cleared(self):
from litellm.types.llms.openai import ResponsesAPIResponse
guardrail = _ApplyOnlyObserver()
kwargs, _ = _logged_call([{"role": "user", "content": "hello there"}])
kwargs["input"] = None
response = ResponsesAPIResponse(
id="resp_1",
created_at=1,
model="gpt-5.4-mini",
object="response",
status="completed",
output=[
{
"type": "message",
"id": "msg_1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "general kenobi"}],
}
],
)
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.aresponses.value)
assert guardrail.calls == [("request", ["hello there"]), ("response", ["general kenobi"])]
entries = out_kwargs["standard_logging_object"]["guardrail_information"]
assert [e["guardrail_status"] for e in entries] == ["success", "success"]
@pytest.mark.asyncio
async def test_async_success_handler_records_verdict_in_standard_logging_object(self):
import datetime as dt
from litellm.litellm_core_utils.litellm_logging import Logging
guardrail = _ApplyOnlyObserver()
guardrail.default_on = True
messages = [{"role": "user", "content": "hello there"}]
_, response = _logged_call(messages)
logging_obj = Logging(
model="gpt-5.4-mini",
messages=messages,
stream=False,
call_type=CallTypes.acompletion.value,
start_time=dt.datetime.now(),
litellm_call_id="call-1",
function_id="fn-1",
dynamic_async_success_callbacks=[guardrail],
)
logging_obj.update_environment_variables(
litellm_params={"metadata": {}}, optional_params={}, model="gpt-5.4-mini", custom_llm_provider="openai"
)
await logging_obj.async_success_handler(
result=response, start_time=dt.datetime.now(), end_time=dt.datetime.now()
)
assert guardrail.calls == [("request", ["hello there"]), ("response", ["general kenobi"])]
entries = logging_obj.model_call_details["standard_logging_object"]["guardrail_information"]
assert [e["guardrail_status"] for e in entries] == ["success", "success"]

View file

@ -1522,7 +1522,7 @@ def test_gpt_5_6_alias_prices_match_sol(local_model_cost_map):
sol = litellm.model_cost["gpt-5.6-sol"]
cost_fields = sorted(field for field in sol if "cost" in field)
assert len(cost_fields) == 23
assert len(cost_fields) == 27
for field in cost_fields:
assert alias.get(field) == sol.get(field), field
@ -4039,8 +4039,8 @@ def test_fast_service_tier_matches_priority_above_the_context_threshold(_local_m
)
assert fast == priority
assert fast[0] == pytest.approx(300_000 * 8e-06, rel=1e-9)
assert fast[1] == pytest.approx(1_000 * 3e-05, rel=1e-9)
assert fast[0] == pytest.approx(300_000 * 1.6e-05, rel=1e-9)
assert fast[1] == pytest.approx(1_000 * 6e-05, rel=1e-9)
def test_priority_reasoning_tokens_bill_at_the_priority_output_rate(_local_model_cost_map):
@ -4200,6 +4200,86 @@ def test_generic_cost_per_token_gemini_37_flash(_local_model_cost_map):
assert completion_cost == pytest.approx(0.001875)
GEMINI_38_FLASH_LAUNCH_PRICING = [
("gemini-3.8-flash", 7.5e-07, 3.75e-06, 7.5e-08),
("gemini/gemini-3.8-flash", 7.5e-07, 3.75e-06, 7.5e-08),
("vertex_ai/gemini-3.8-flash", 7.5e-07, 3.75e-06, 7.5e-08),
]
@pytest.mark.parametrize("model,input_cost,output_cost,cache_read_cost", GEMINI_38_FLASH_LAUNCH_PRICING)
def test_gemini_38_flash_launch_pricing(model, input_cost, output_cost, cache_read_cost, _local_model_cost_map):
model_cost_map = litellm.model_cost[model]
assert model_cost_map["input_cost_per_token"] == input_cost
assert model_cost_map["output_cost_per_token"] == output_cost
assert model_cost_map["output_cost_per_reasoning_token"] == output_cost
assert model_cost_map["cache_read_input_token_cost"] == cache_read_cost
assert model_cost_map["mode"] == "chat"
assert model_cost_map["supports_reasoning"] is True
assert model_cost_map["supports_function_calling"] is True
assert model_cost_map["max_input_tokens"] == 1048576
GEMINI_38_FLASH_FIELDS_SHARED_WITH_37_FLASH = (
"input_cost_per_token",
"output_cost_per_token",
"output_cost_per_reasoning_token",
"cache_read_input_token_cost",
"input_cost_per_token_batches",
"output_cost_per_token_batches",
"input_cost_per_token_flex",
"output_cost_per_token_flex",
"cache_read_input_token_cost_flex",
"input_cost_per_token_priority",
"output_cost_per_token_priority",
"cache_read_input_token_cost_priority",
"search_context_cost_per_query",
"google_maps_grounding_cost_per_query",
"prompt_cache_min_tokens",
"max_input_tokens",
"max_output_tokens",
"supports_reasoning",
"supports_function_calling",
"supports_prompt_caching",
"supports_vision",
"supports_pdf_input",
"supports_audio_input",
"supports_video_input",
"supports_response_schema",
"supports_tool_choice",
"supports_web_search",
"supports_url_context",
)
@pytest.mark.parametrize("prefix", ["", "gemini/", "vertex_ai/"])
def test_gemini_38_flash_matches_37_flash_promotional_pricing(prefix, _local_model_cost_map):
new_model = litellm.model_cost[f"{prefix}gemini-3.8-flash"]
old_model = litellm.model_cost[f"{prefix}gemini-3.7-flash"]
for field in GEMINI_38_FLASH_FIELDS_SHARED_WITH_37_FLASH:
assert new_model[field] == old_model[field], field
def test_generic_cost_per_token_gemini_38_flash(_local_model_cost_map):
usage = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500,
completion_tokens_details=CompletionTokensDetailsWrapper(
reasoning_tokens=200,
text_tokens=300,
),
prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=1000),
)
prompt_cost, completion_cost = generic_cost_per_token(
model="gemini-3.8-flash",
usage=usage,
custom_llm_provider="gemini",
)
assert prompt_cost == pytest.approx(0.00075)
assert completion_cost == pytest.approx(0.001875)
def test_grok_46_launch_pricing(_local_model_cost_map):
model_cost_map = litellm.model_cost["xai/grok-4.6"]
assert model_cost_map["input_cost_per_token"] == 2e-06

View file

@ -592,6 +592,59 @@ def test_stream_chunk_builder_litellm_usage_chunks():
assert usage.total_tokens == 77
def test_calculate_usage_honors_openai_sdk_completion_usage_chunks():
from openai.types.completion_usage import CompletionUsage
content_chunk = ModelResponseStream(
id="chatcmpl-sdk-usage-1",
created=1745513206,
model="mantle-claude",
object="chat.completion.chunk",
system_fingerprint=None,
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(
provider_specific_fields=None,
content="ok",
role=None,
function_call=None,
tool_calls=None,
audio=None,
),
logprobs=None,
)
],
provider_specific_fields=None,
stream_options={"include_usage": True},
)
usage_chunk = ModelResponseStream(
id="chatcmpl-sdk-usage-1",
created=1745513207,
model="mantle-claude",
object="chat.completion.chunk",
system_fingerprint=None,
choices=[],
provider_specific_fields=None,
stream_options={"include_usage": True},
)
usage_chunk.usage = CompletionUsage(
prompt_tokens=20, completion_tokens=60, total_tokens=80, cost=0.000704
)
assert type(usage_chunk.usage) is CompletionUsage
chunks = [content_chunk, usage_chunk]
usage = ChunkProcessor(chunks=chunks).calculate_usage(
chunks=chunks, model="mantle-claude", completion_output=""
)
assert usage.prompt_tokens == 20
assert usage.completion_tokens == 60
assert usage.total_tokens == 80
assert getattr(usage, "cost", None) == pytest.approx(0.000704)
def test_get_model_from_chunks_azure_model_router():
"""
Test that _get_model_from_chunks finds the actual model from Azure Model Router chunks.

View file

@ -96,10 +96,8 @@ def _completion_kwargs(**overrides):
return kwargs
def _run(**overrides):
with patch.object(
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
):
def _run(*, credentials: Credentials | None = RESOLVED_CREDENTIALS, **overrides):
with patch.object(BedrockConverseLLM, "get_credentials", return_value=credentials):
return BedrockConverseLLM().completion(**_completion_kwargs(**overrides))
@ -360,7 +358,7 @@ async def test_async_completion_logs_pre_call_by_default():
def _sync_client_returning_converse_response():
client = MagicMock()
client.post = lambda **_kwargs: httpx.Response(
client.post.side_effect = lambda **_kwargs: httpx.Response(
200,
json=CONVERSE_RESPONSE,
request=httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com"),
@ -487,3 +485,31 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines():
assert response.choices[0].message.content == "hi"
assert len(calls["post_call"]) == 1
assert "hi" in calls["post_call"][0]["original_response"]
def test_bearer_token_auth_serves_when_boto3_resolves_no_sigv4_credentials(monkeypatch):
"""With only `AWS_BEARER_TOKEN_BEDROCK` configured boto3 resolves no
credentials at all. Preparing the Rust handoff must not dereference that
None: the bearer token signs the request on its own."""
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "bedrock-bearer-token")
client = _sync_client_returning_converse_response()
response = _run(credentials=None, litellm_params={}, client=client)
assert response.choices[0].message.content == "hi"
sent_headers = client.post.call_args.kwargs["headers"]
assert sent_headers["Authorization"] == "Bearer bedrock-bearer-token"
def test_the_rust_opt_in_needs_no_sigv4_principal():
"""The core resolves the bearer token itself, so a bearer-only deployment
keeps its opt-in and the gate sees no aws_* credential keys to sign with."""
seen = _inject()
response = _run(credentials=None, api_key="bedrock-bearer-token")
assert response.choices[0].message.content == "hello from rust"
params = seen["call"][0]["optional_params"]
assert not {"aws_access_key_id", "aws_secret_access_key", "aws_session_token"} & params.keys()
assert params["aws_region_name"] == "us-east-1"
assert seen["call"][0]["api_key"] == "bedrock-bearer-token"

View file

@ -15,6 +15,7 @@ from unittest.mock import MagicMock, patch
from botocore.awsrequest import AWSPreparedRequest, AWSRequest
from botocore.auth import SigV4Auth
from botocore.credentials import Credentials
from botocore.exceptions import NoCredentialsError
import litellm
from litellm.llms.bedrock.base_aws_llm import (
@ -801,6 +802,23 @@ def test_get_request_headers_with_sigv4():
assert result == mock_request.prepare.return_value
def test_get_request_headers_without_credentials_or_bearer_token_raises_no_credentials():
"""Bearer-token auth needs no SigV4 principal, so `credentials` may be None.
Reaching the SigV4 branch with neither must fail the way botocore always
has instead of signing with a missing principal."""
llm = BaseAWSLLM()
with patch.dict(os.environ, {}, clear=True), pytest.raises(NoCredentialsError):
llm.get_request_headers(
credentials=None,
aws_region_name="us-west-2",
extra_headers=None,
endpoint_url="https://api.example.com",
data='{"prompt": "test"}',
headers={"Content-Type": "application/json"},
)
def test_sigv4_matches_rust_golden_vector():
request = AWSRequest(
method="POST",

View file

@ -2,6 +2,7 @@
Tests for Parallel AI Search API integration (v1 endpoint).
"""
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -30,13 +31,41 @@ MOCK_V1_RESPONSE = {
}
def _mock_response():
def _mock_response(payload=None):
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = MOCK_V1_RESPONSE
mock_response.json.return_value = payload if payload is not None else MOCK_V1_RESPONSE
return mock_response
@pytest.fixture
def httpx_transport(monkeypatch):
monkeypatch.setattr( # test-quality-ok: respx needs HTTPX enabled to fake the provider HTTP boundary.
litellm,
"disable_aiohttp_transport",
True,
)
litellm.in_memory_llm_clients_cache.flush_cache()
yield
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.fixture
def bundled_cost_map(monkeypatch):
"""Price lookups against the bundled cost map.
litellm caches model-info lookups, so swapping ``model_cost`` only takes
effect once those caches are invalidated -- on the way in and back out.
"""
from litellm.utils import _invalidate_model_cost_lowercase_map
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
_invalidate_model_cost_lowercase_map()
yield
monkeypatch.undo()
_invalidate_model_cost_lowercase_map()
class TestParallelAISearch:
@pytest.fixture(autouse=True)
def _set_api_key(self, monkeypatch):
@ -135,9 +164,7 @@ class TestParallelAISearch:
json_data = mock_post.call_args.kwargs.get("json")
assert json_data["mode"] == "basic"
@pytest.mark.parametrize(
"processor,expected_mode", [("base", "basic"), ("pro", "advanced")]
)
@pytest.mark.parametrize("processor,expected_mode", [("base", "basic"), ("pro", "advanced")])
@pytest.mark.asyncio
async def test_legacy_processor_maps_to_mode(self, processor, expected_mode):
with patch(
@ -222,9 +249,7 @@ class TestParallelAISearch:
"arxiv.org",
"nature.com",
]
assert advanced_settings["source_policy"]["exclude_domains"] == [
"reddit.com"
]
assert advanced_settings["source_policy"]["exclude_domains"] == ["reddit.com"]
assert advanced_settings["excerpt_settings"]["max_chars_per_result"] == 1500
assert "max_results" not in json_data
@ -306,10 +331,7 @@ class TestParallelAISearch:
)
call_args = mock_post.call_args
assert (
call_args.kwargs["url"]
== "https://proxy.internal.example.com/v1/search"
)
assert call_args.kwargs["url"] == "https://proxy.internal.example.com/v1/search"
@pytest.mark.asyncio
async def test_caller_api_base_without_key_is_refused(self, monkeypatch):
@ -338,3 +360,147 @@ class TestParallelAISearch:
query="AI developments",
search_provider="parallel_ai",
)
@pytest.mark.asyncio
async def test_flat_source_and_fetch_params_nest_under_advanced_settings(self, respx_mock, httpx_transport):
route = respx_mock.post("https://api.parallel.ai/v1/search").respond(json=MOCK_V1_RESPONSE)
await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
objective="find peer-reviewed AI research",
include_domains=["arxiv.org"],
after_date="2026-01-01",
location="gb",
fetch_policy={"max_age_seconds": 600, "disable_cache_fallback": True},
client_model="claude-fable-5",
)
json_data = json.loads(route.calls[0].request.content)
assert json_data["objective"] == "find peer-reviewed AI research"
assert json_data["client_model"] == "claude-fable-5"
advanced_settings = json_data["advanced_settings"]
assert advanced_settings["location"] == "gb"
assert advanced_settings["fetch_policy"] == {
"max_age_seconds": 600,
"disable_cache_fallback": True,
}
assert advanced_settings["source_policy"]["include_domains"] == ["arxiv.org"]
assert advanced_settings["source_policy"]["after_date"] == "2026-01-01"
assert "include_domains" not in json_data
assert "after_date" not in json_data
assert "location" not in json_data
assert "fetch_policy" not in json_data
@pytest.mark.asyncio
async def test_response_preserves_raw_parallel_fields(self, respx_mock, httpx_transport):
respx_mock.post("https://api.parallel.ai/v1/search").respond(json=MOCK_V1_RESPONSE)
response = await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
)
dumped = response.model_dump()
assert dumped["search_id"] == "search_abc123"
assert dumped["session_id"] == "session_xyz"
assert dumped["parallel_usage"] == [{"name": "search_advanced", "count": 1}]
first = response.results[0].model_dump()
assert first["excerpts"] == ["First excerpt.", "Second excerpt."]
@pytest.mark.asyncio
async def test_response_normalizes_null_result_fields(self, respx_mock, httpx_transport):
response_payload = {
**MOCK_V1_RESPONSE,
"results": [{"url": None, "title": None, "publish_date": None, "excerpts": None}],
}
respx_mock.post("https://api.parallel.ai/v1/search").respond(json=response_payload)
response = await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
)
assert len(response.results) == 1
result = response.results[0]
assert result.url == ""
assert result.title == ""
assert result.snippet == ""
assert result.date is None
assert result.model_dump()["excerpts"] == ()
@pytest.mark.parametrize(
"mode,usage,max_results,expected_cost",
[
("turbo", [{"name": "sku_search", "count": 1}], None, 0.001),
("fast", [{"name": "sku_search", "count": 1}], None, 0.001),
("basic", [{"name": "sku_search", "count": 1}], None, 0.005),
("advanced", [{"name": "sku_search", "count": 1}], None, 0.005),
(
"basic",
[
{"name": "sku_search", "count": 1},
{"name": "sku_search_additional_results", "count": 2},
],
20,
0.007,
),
("basic", None, 20, 0.015),
],
)
@pytest.mark.asyncio
async def test_search_cost_uses_mode_and_provider_usage(
self, mode, usage, max_results, expected_cost, bundled_cost_map, respx_mock, httpx_transport
):
response_payload = {**MOCK_V1_RESPONSE, "usage": usage}
respx_mock.post("https://api.parallel.ai/v1/search").respond(json=response_payload)
response = await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
mode=mode,
max_results=max_results,
)
assert response._hidden_params["response_cost"] == pytest.approx(expected_cost)
@pytest.mark.asyncio
async def test_search_cost_treats_keyword_queries_as_one_request(
self, bundled_cost_map, respx_mock, httpx_transport
):
response_payload = {
**MOCK_V1_RESPONSE,
"usage": [{"name": "sku_search", "count": 1}],
}
respx_mock.post("https://api.parallel.ai/v1/search").respond(json=response_payload)
response = await litellm.asearch(
query=["AI developments", "machine learning trends"],
search_provider="parallel_ai",
mode="basic",
)
assert response._hidden_params["response_cost"] == pytest.approx(0.005)
@pytest.mark.asyncio
async def test_caller_cannot_supply_provider_usage(self, bundled_cost_map, respx_mock, httpx_transport):
"""`_parallel_ai_usage` prices the request, so a caller must not be able to set it.
The provider reports no usage here, which is the case where a caller-supplied
value would otherwise survive into the cost calculation.
"""
response_payload = {k: v for k, v in MOCK_V1_RESPONSE.items() if k != "usage"}
route = respx_mock.post("https://api.parallel.ai/v1/search").respond(json=response_payload)
response = await litellm.asearch(
query="AI developments",
search_provider="parallel_ai",
mode="basic",
_parallel_ai_usage=[{"name": "sku_search", "count": 0}],
)
assert response._hidden_params["response_cost"] == pytest.approx(0.005)
assert "_parallel_ai_usage" not in json.loads(route.calls[0].request.content)

View file

@ -0,0 +1,191 @@
"""Gateway coverage for Parallel AI Search."""
from __future__ import annotations
from collections.abc import Iterator
from typing import Final
from unittest.mock import AsyncMock
import httpx
import pytest
from fastapi.testclient import TestClient
import litellm
from litellm import Router
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy import proxy_server
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.utils import LlmProviders
PARALLEL_SEARCH_URL: Final = "https://api.parallel.ai/v1/search"
@pytest.fixture
def client() -> TestClient:
return TestClient(proxy_server.app, raise_server_exceptions=False)
@pytest.fixture
def auth_as() -> Iterator[None]:
async def _authorized_request() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="hashed-sk-test",
user_id="parallel-test-user",
)
previous: Final = proxy_server.app.dependency_overrides.get(user_api_key_auth)
proxy_server.app.dependency_overrides[user_api_key_auth] = _authorized_request
try:
yield
finally:
if previous is None:
proxy_server.app.dependency_overrides.pop(user_api_key_auth, None)
else:
proxy_server.app.dependency_overrides[user_api_key_auth] = previous
def _parallel_search_body() -> dict[str, object]:
return {
"search_id": "search_parallel_gateway",
"results": [
{
"url": "https://example.com/parallel",
"title": "Parallel result",
"publish_date": "2026-08-13",
"excerpts": ["First excerpt", "Second excerpt"],
}
],
"usage": [{"name": "sku_search", "count": 1}],
}
def _parallel_router(mode: str = "turbo") -> Router:
return Router(
model_list=[],
search_tools=[
{
"search_tool_name": "parallel-search",
"litellm_params": {
"search_provider": "parallel_ai",
"api_key": "parallel-search-key",
"mode": mode,
},
}
],
num_retries=0,
)
def _mock_async_post(
monkeypatch,
*,
url: str,
response_body: dict[str, object],
) -> AsyncMock:
response = httpx.Response(
status_code=200,
json=response_body,
request=httpx.Request("POST", url),
)
mock_post = AsyncMock(return_value=response)
monkeypatch.setattr(AsyncHTTPHandler, "post", mock_post)
return mock_post
def test_parallel_search_gateway_route(client, auth_as, monkeypatch):
"""The named search route selects its configured Parallel Search tool.
The tool-level `mode` must survive the router hop, so the upstream request
is sent as `turbo` rather than falling back to the adapter default.
"""
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
monkeypatch.setattr(proxy_server, "llm_router", _parallel_router())
mock_post = _mock_async_post(
monkeypatch,
url=PARALLEL_SEARCH_URL,
response_body=_parallel_search_body(),
)
response = client.post(
"/v1/search/parallel-search",
json={"query": "Parallel AI news", "max_results": 3},
)
assert response.status_code == 200, response.text
assert response.json()["results"] == [
{
"title": "Parallel result",
"url": "https://example.com/parallel",
"snippet": "First excerpt ... Second excerpt",
"date": "2026-08-13",
"last_updated": None,
"excerpts": ["First excerpt", "Second excerpt"],
}
]
request_kwargs = mock_post.await_args.kwargs
assert request_kwargs["url"] == PARALLEL_SEARCH_URL
assert request_kwargs["headers"]["x-api-key"] == "parallel-search-key"
assert request_kwargs["json"] == {
"objective": "Parallel AI news",
"search_queries": ["Parallel AI news"],
"mode": "turbo",
"advanced_settings": {"max_results": 3},
}
@pytest.mark.asyncio
async def test_web_search_interception_executes_parallel_search(monkeypatch):
"""An intercepted web-search call uses the configured Parallel Search tool."""
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
monkeypatch.setattr(proxy_server, "llm_router", _parallel_router(mode="fast"))
mock_post = _mock_async_post(
monkeypatch,
url=PARALLEL_SEARCH_URL,
response_body=_parallel_search_body(),
)
logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.OPENAI],
search_tool_name="parallel-search",
)
plan = await logger.async_build_responses_agentic_loop_plan(
tools={
"tool_calls": [
{
"id": "fc_parallel",
"call_id": "fc_parallel",
"type": "function_call",
"name": "litellm_web_search",
"arguments": '{"query":"Parallel AI news"}',
"input": {"query": "Parallel AI news"},
}
]
},
model="gpt-5",
messages=[{"role": "user", "content": "Research Parallel"}],
response=None,
optional_params={"tools": [{"type": "function", "name": "litellm_web_search"}]},
logging_obj=None,
stream=False,
kwargs={"custom_llm_provider": "openai"},
)
assert plan.run_agentic_loop is True
assert plan.request_patch is not None
assert plan.request_patch.messages[-1] == {
"type": "function_call_output",
"call_id": "fc_parallel",
"output": (
"Title: Parallel result\nURL: https://example.com/parallel\nSnippet: First excerpt ... Second excerpt"
),
}
request_kwargs = mock_post.await_args.kwargs
assert request_kwargs["url"] == PARALLEL_SEARCH_URL
assert request_kwargs["headers"]["x-api-key"] == "parallel-search-key"
assert request_kwargs["json"]["mode"] == "fast"

View file

@ -1096,10 +1096,13 @@ def test_natively_signed_parallel_turn_never_carries_a_placeholder(model):
"gemini-3.5-flash",
"gemini-3.6-flash",
"gemini-3.7-flash",
"gemini-3.8-flash",
"vertex_ai/gemini-3.5-flash",
"vertex_ai/gemini-3.7-flash",
"vertex_ai/gemini-3.8-flash",
"gemini/gemini-3.5-flash",
"gemini/gemini-3.7-flash",
"gemini/gemini-3.8-flash",
],
)
def test_placeholder_scoped_to_first_call_across_gemini_3_variants(model):

View file

@ -1185,6 +1185,18 @@ def test_vertex_ai_map_thinking_param_with_budget_tokens_0():
}
def test_vertex_ai_map_thinking_param_without_budget_tokens_for_gemini_3():
v = VertexGeminiConfig()
result = v.map_openai_params(
non_default_params={"thinking": {"type": "enabled"}},
optional_params={},
model="gemini-3.5-flash",
drop_params=False,
)
assert result["thinkingConfig"] == {"includeThoughts": True}
def test_vertex_ai_map_tools():
v = VertexGeminiConfig()
optional_params = {}

View file

@ -0,0 +1,57 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
from litellm.proxy.management_helpers.access_group_key_sync import (
sync_key_access_group_membership,
sync_key_regeneration_access_group_membership,
)
def _routed_prisma_client():
writer_inner = MagicMock(name="writer_prisma")
reader_inner = MagicMock(name="reader_prisma")
writer_inner.query_raw = AsyncMock(return_value=[])
reader_inner.query_raw = AsyncMock(return_value=[])
writer = PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False)
reader = PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False)
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
return SimpleNamespace(db=routing), writer_inner, reader_inner
@pytest.mark.asyncio
async def test_regeneration_repoint_update_runs_on_the_writer():
prisma_client, writer_inner, reader_inner = _routed_prisma_client()
await sync_key_regeneration_access_group_membership(
prisma_client=prisma_client,
previous_key_token="old-token",
new_key_token="new-token",
data=None,
existing_key_row=MagicMock(),
)
writer_inner.query_raw.assert_awaited_once()
assert writer_inner.query_raw.await_args.args[0].startswith('UPDATE "LiteLLM_AccessGroupTable"')
reader_inner.query_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_membership_attach_and_detach_updates_run_on_the_writer():
prisma_client, writer_inner, reader_inner = _routed_prisma_client()
await sync_key_access_group_membership(
prisma_client=prisma_client,
key_token="token",
previous_access_group_ids=["ag-old"],
updated_access_group_ids=["ag-new"],
)
assert writer_inner.query_raw.await_count == 2
assert all(
call.args[0].startswith('UPDATE "LiteLLM_AccessGroupTable"') for call in writer_inner.query_raw.await_args_list
)
reader_inner.query_raw.assert_not_awaited()

View file

@ -60,6 +60,141 @@ def test_config_update_happy_admin(client, auth_as, mock_prisma, monkeypatch):
assert normalize(response.json()) == {"message": "Config updated successfully"}
def test_config_update_persists_optional_pre_call_checks(client, auth_as, mock_prisma, monkeypatch):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
fake_proxy_config = MagicMock()
fake_proxy_config.add_deployment = AsyncMock()
monkeypatch.setattr(ps, "proxy_config", fake_proxy_config)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/update",
json={"router_settings": {"optional_pre_call_checks": ["prompt_caching"]}},
)
assert response.status_code == 200
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
assert persisted["optional_pre_call_checks"] == ["prompt_caching"]
def test_config_update_persists_model_group_affinity_config(client, auth_as, mock_prisma, monkeypatch):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
fake_proxy_config = MagicMock()
fake_proxy_config.add_deployment = AsyncMock()
monkeypatch.setattr(ps, "proxy_config", fake_proxy_config)
model_group_affinity_config = {"gpt-4": ["session_affinity"]}
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/update",
json={"router_settings": {"model_group_affinity_config": model_group_affinity_config}},
)
assert response.status_code == 200
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
assert persisted["model_group_affinity_config"] == model_group_affinity_config
def test_config_update_persists_disable_cooldowns(client, auth_as, mock_prisma, monkeypatch):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
fake_proxy_config = MagicMock()
fake_proxy_config.add_deployment = AsyncMock()
monkeypatch.setattr(ps, "proxy_config", fake_proxy_config)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/update",
json={"router_settings": {"disable_cooldowns": True}},
)
assert response.status_code == 200
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
assert persisted["disable_cooldowns"] is True
def test_config_update_rejects_assistants_config(client, auth_as, mock_prisma, monkeypatch):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/update",
json={"router_settings": {"assistants_config": {"enabled": True}}},
)
assert response.status_code == 400
assert "assistants_config" in response.json()["error"]["message"]
table.upsert.assert_not_called()
def test_config_update_rejects_router_general_settings(client, auth_as, mock_prisma, monkeypatch):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/update",
json={"router_settings": {"router_general_settings": {"async_only_mode": True}}},
)
assert response.status_code == 400
assert "router_general_settings" in response.json()["error"]["message"]
table.upsert.assert_not_called()
def test_config_update_rejects_unknown_router_setting(client, auth_as, mock_prisma, monkeypatch):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/update",
json={"router_settings": {"optional_precall_checks": ["prompt_caching"]}},
)
assert response.status_code == 400
assert "optional_precall_checks" in response.json()["error"]["message"]
table.upsert.assert_not_called()
def test_config_update_unknown_router_setting_non_admin_forbidden(client, auth_as, mock_prisma, monkeypatch):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
_install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
with auth_as(LitellmUserRoles.INTERNAL_USER):
response = client.post(
"/config/update",
json={"router_settings": {"optional_precall_checks": ["prompt_caching"]}},
)
assert response.status_code == 403
assert "admin" in response.json()["error"]["message"].lower()
def test_config_update_non_admin_forbidden(client, auth_as, mock_prisma, monkeypatch):
"""POST /config/update by a non-admin caller is rejected; the error
surfaces as a ProxyException with the admin-only message."""

View file

@ -45,6 +45,7 @@ def patched_models(monkeypatch):
deployment = MagicMock()
deployment.litellm_params.model = "gpt-4"
router.get_deployment_by_model_group_name = MagicMock(return_value=deployment)
router.get_configured_display_name = MagicMock(return_value=None)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
@ -187,6 +188,83 @@ def test_anthropic_format_carries_router_configured_token_limits(client, auth_as
assert (claude["max_input_tokens"], claude["max_tokens"]) == (500000, 4096)
@pytest.mark.parametrize("path", ["/v1/models", "/models"])
def test_anthropic_format_uses_configured_display_name(client, auth_as, patched_models, path):
"""A deployment's ``model_info.display_name`` becomes the Anthropic-native
``display_name`` so Claude Code's picker shows a clean name while the id keeps
routing; models without one keep the id fallback, and the OpenAI-shaped
listing carries no display_name either way."""
def _configured(model_name):
return "Kimi K3" if model_name == "gpt-4" else None
patched_models.get_configured_display_name = MagicMock(side_effect=_configured)
with auth_as():
anthropic_response = client.get(path, headers={"anthropic-version": "2023-06-01"})
openai_response = client.get(path)
assert anthropic_response.status_code == 200
gpt_4, claude = anthropic_response.json()["data"]
assert (gpt_4["id"], gpt_4["display_name"]) == ("gpt-4", "Kimi K3")
assert (claude["id"], claude["display_name"]) == ("claude-sonnet", "claude-sonnet")
assert openai_response.status_code == 200
openai_models = openai_response.json()["data"]
assert [m["id"] for m in openai_models] == ["gpt-4", "claude-sonnet"]
assert all("display_name" not in m for m in openai_models)
@pytest.mark.parametrize("params", [{}, {"scope": "expand"}])
def test_anthropic_display_name_resolved_via_internal_team_key(
client, auth_as, patched_models, monkeypatch, params
):
"""For a team-scoped row the configured display name must be looked up by the
internal routing key while the entry itself is keyed by the public name, so
the clean name lands on the id the client actually sees."""
from litellm.proxy import utils as proxy_utils
from litellm.proxy.auth import model_checks
internal_name = "model_name_team-1_c0ffee"
patched_models.get_model_list = MagicMock(
return_value=[
{
"model_name": internal_name,
"model_info": {
"team_id": "team-1",
"team_public_model_name": "gpt-4-team",
},
}
]
)
patched_models.get_model_names = MagicMock(return_value=[internal_name])
patched_models.get_configured_display_name = MagicMock(
side_effect=lambda model_name: "Team GPT" if model_name == internal_name else None
)
async def _fake_get_available_models_for_user(**kwargs):
return [internal_name]
monkeypatch.setattr(
proxy_utils,
"get_available_models_for_user",
_fake_get_available_models_for_user,
)
monkeypatch.setattr(
model_checks, "get_complete_model_list", lambda **kwargs: [internal_name]
)
with auth_as():
response = client.get(
"/v1/models", params=params, headers={"anthropic-version": "2023-06-01"}
)
assert response.status_code == 200
(entry,) = response.json()["data"]
assert (entry["id"], entry["display_name"]) == ("gpt-4-team", "Team GPT")
@pytest.mark.parametrize("path", ["/v1/models", "/models"])
def test_get_models_invalid_scope_returns_400(client, auth_as, patched_models, path):
"""Pins: ``GET /v1/models``, ``GET /models`` (error path: invalid scope)."""

View file

@ -19,7 +19,10 @@ from litellm.proxy._types import (
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.common_utils.model_listing_utils import TeamModelNameTranslator
from litellm.proxy.common_utils.model_listing_utils import (
TeamModelNameTranslator,
configured_display_names,
)
from litellm.proxy.proxy_server import (
_get_proxy_model_info,
_translate_model_name_for_response,
@ -1391,6 +1394,27 @@ def test_resolve_public_name_respects_legacy_flag():
)
def test_configured_display_names_keyed_by_response_id():
"""The map is keyed by the public response id while the router lookup uses
the internal routing key, and entries without a configured name are omitted."""
router = MagicMock()
router.get_configured_display_name = MagicMock(
side_effect=lambda model_name: "Team Sonnet" if model_name == "model_name_team-abc-123_4a6b8" else None
)
assert configured_display_names(
entries=[
("team-claude-sonnet", "model_name_team-abc-123_4a6b8"),
("gpt-4o", "gpt-4o"),
],
llm_router=router,
) == {"team-claude-sonnet": "Team Sonnet"}
def test_configured_display_names_empty_without_router():
assert configured_display_names(entries=[("gpt-4o", "gpt-4o")], llm_router=None) == {}
@pytest.mark.asyncio
async def test_retrieve_model_by_public_name_returns_200(monkeypatch):
"""Regression: `GET /v1/models/{public_name}` must NOT 404. The listing

View file

@ -111,6 +111,99 @@ def test_together_rerank_honors_api_base(respx_mock: respx.MockRouter):
assert mock_route.calls[0].request.headers["authorization"] == "Bearer fake-together-key"
DASHSCOPE_404_BODY = {
"error": {
"message": "The model `does-not-exist` does not exist or you do not have access to it.",
"type": "invalid_request_error",
"param": None,
"code": "model_not_found",
},
"request_id": "mock-request-id",
}
def test_rerank_error_names_provider_and_keeps_body(respx_mock: respx.MockRouter, monkeypatch):
"""Regression for the rerank error path mapping with the unresolved provider param:
a provider 404 surfaced as 'None - ' instead of naming the provider and its error body."""
monkeypatch.delenv("DASHSCOPE_API_BASE", raising=False)
monkeypatch.delenv("DASHSCOPE_API_BASE_RERANK", raising=False)
mock_route = respx_mock.post("https://dashscope.example/v1/reranks")
mock_route.return_value = httpx.Response(404, json=DASHSCOPE_404_BODY)
with pytest.raises(litellm.NotFoundError) as exc_info:
litellm.rerank(
model="dashscope/does-not-exist",
query=MARKER_QUERY,
documents=[MARKER_DOC],
api_key="fake-dashscope-key",
api_base="https://dashscope.example/v1",
)
assert mock_route.called
assert "DashscopeException" in str(exc_info.value)
assert "does not exist or you do not have access to it" in str(exc_info.value)
assert "None - " not in str(exc_info.value)
@pytest.mark.asyncio
async def test_arerank_error_is_mapped_to_litellm_exception(respx_mock: respx.MockRouter, monkeypatch):
"""Regression for arerank's bare re-raise: provider errors escaped as raw
provider exception classes instead of the mapped litellm exception contract."""
monkeypatch.delenv("DASHSCOPE_API_BASE", raising=False)
monkeypatch.delenv("DASHSCOPE_API_BASE_RERANK", raising=False)
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
mock_route = respx_mock.post("https://dashscope.example/v1/reranks")
mock_route.return_value = httpx.Response(404, json=DASHSCOPE_404_BODY)
with pytest.raises(litellm.NotFoundError) as exc_info:
await litellm.arerank(
model="dashscope/does-not-exist",
query=MARKER_QUERY,
documents=[MARKER_DOC],
api_key="fake-dashscope-key",
api_base="https://dashscope.example/v1",
)
assert mock_route.called
assert "DashscopeException" in str(exc_info.value)
assert "does not exist or you do not have access to it" in str(exc_info.value)
assert "None - " not in str(exc_info.value)
@pytest.mark.asyncio
async def test_arerank_declared_authenticating_provider_skips_resolution(monkeypatch):
"""Regression for the event-loop hazard in arerank's provider pre-resolution:
get_llm_provider runs the blocking OAuth device flow for github_copilot/chatgpt,
so arerank must adopt the declared provider instead of resolving it, while the
except path still maps with that declared provider."""
from litellm.llms.base_llm.chat.transformation import BaseLLMException
resolution_calls = []
def record_resolution(*args, **kwargs):
resolution_calls.append((args, kwargs))
return "gpt-4o", "github_copilot", None, None
def rerank_raises_provider_error(*args, **kwargs):
raise BaseLLMException(status_code=401, message='{"error":"bad key"}')
monkeypatch.setattr(litellm, "get_llm_provider", record_resolution)
monkeypatch.setattr("litellm.rerank_api.main.rerank", rerank_raises_provider_error)
with pytest.raises(litellm.AuthenticationError) as exc_info:
await litellm.arerank(
model="github_copilot/gpt-4o",
query=MARKER_QUERY,
documents=[MARKER_DOC],
)
assert resolution_calls == []
assert "Github_copilotException" in str(exc_info.value)
assert "None - " not in str(exc_info.value)
@pytest.mark.asyncio
async def test_together_rerank_async_honors_env_api_base(respx_mock: respx.MockRouter, monkeypatch):
"""Regression: TOGETHER_AI_API_BASE was honored by chat but ignored by rerank."""

View file

@ -326,3 +326,55 @@ def test_run_post_success_hooks_does_not_report_generation_time_as_overhead():
assert iterator.completed_response._hidden_params["_response_ms"] == 10000.0
assert "litellm_overhead_time_ms" not in iterator.completed_response._hidden_params
def _responses_api_response_with_usage() -> ResponsesAPIResponse:
from litellm.types.llms.openai import ResponseAPIUsage
return ResponsesAPIResponse(
id="resp_lit6427",
created_at=int(datetime(2025, 1, 1).timestamp()),
status="completed",
model="mantle-claude",
object="response",
output=[],
usage=ResponseAPIUsage(input_tokens=20, output_tokens=60, total_tokens=80),
)
def test_stamp_responses_usage_cost_stamps_computed_cost():
from litellm.responses.streaming_iterator import _stamp_responses_usage_cost
response = _responses_api_response_with_usage()
logging_obj = Mock(spec=LiteLLMLoggingObj)
logging_obj._response_cost_calculator.return_value = 0.000704
_stamp_responses_usage_cost(response, logging_obj)
assert getattr(response.usage, "cost", None) == pytest.approx(0.000704)
logging_obj._response_cost_calculator.assert_called_once_with(result=response)
def test_stamp_responses_usage_cost_keeps_provider_reported_cost():
from litellm.responses.streaming_iterator import _stamp_responses_usage_cost
response = _responses_api_response_with_usage()
setattr(response.usage, "cost", 0.5)
logging_obj = Mock(spec=LiteLLMLoggingObj)
_stamp_responses_usage_cost(response, logging_obj)
assert getattr(response.usage, "cost", None) == pytest.approx(0.5)
logging_obj._response_cost_calculator.assert_not_called()
def test_stamp_responses_usage_cost_survives_calculator_failure():
from litellm.responses.streaming_iterator import _stamp_responses_usage_cost
response = _responses_api_response_with_usage()
logging_obj = Mock(spec=LiteLLMLoggingObj)
logging_obj._response_cost_calculator.side_effect = RuntimeError("cost map unavailable")
_stamp_responses_usage_cost(response, logging_obj)
assert getattr(response.usage, "cost", None) is None

View file

@ -26,9 +26,7 @@ import pytest
@pytest.fixture(scope="module")
def model_data():
json_path = os.path.join(
os.path.dirname(__file__), "../../model_prices_and_context_window.json"
)
json_path = os.path.join(os.path.dirname(__file__), "../../model_prices_and_context_window.json")
with open(json_path) as f:
return json.load(f)
@ -51,21 +49,14 @@ def test_usgov_sonnet_4_5_pricing(model_data, model_key):
info = model_data[model_key]
assert info["input_cost_per_token"] == 3.6e-06, (
f"{model_key}: input_cost_per_token should be $3.60/MTok "
f"(got {info['input_cost_per_token']})"
f"{model_key}: input_cost_per_token should be $3.60/MTok (got {info['input_cost_per_token']})"
)
assert (
info["output_cost_per_token"] == 1.8e-05
), f"{model_key}: output_cost_per_token should be $18.00/MTok"
assert (
info["cache_creation_input_token_cost"] == 4.5e-06
), f"{model_key}: 5m cache write should be $4.50/MTok"
assert (
info["cache_creation_input_token_cost_above_1hr"] == 7.2e-06
), f"{model_key}: 1h cache write should be $7.20/MTok"
assert (
info["cache_read_input_token_cost"] == 3.6e-07
), f"{model_key}: cache read should be $0.36/MTok"
assert info["output_cost_per_token"] == 1.8e-05, f"{model_key}: output_cost_per_token should be $18.00/MTok"
assert info["cache_creation_input_token_cost"] == 4.5e-06, f"{model_key}: 5m cache write should be $4.50/MTok"
assert info["cache_creation_input_token_cost_above_1hr"] == 7.2e-06, (
f"{model_key}: 1h cache write should be $7.20/MTok"
)
assert info["cache_read_input_token_cost"] == 3.6e-07, f"{model_key}: cache read should be $0.36/MTok"
def test_usgov_carries_20_percent_premium_over_global(model_data):
@ -84,9 +75,7 @@ def test_usgov_carries_20_percent_premium_over_global(model_data):
"cache_read_input_token_cost",
):
ratio = usgov_info[field] / global_info[field]
assert (
abs(ratio - 1.2) < 1e-9
), f"{field}: us-gov / global ratio is {ratio}, expected 1.2"
assert abs(ratio - 1.2) < 1e-9, f"{field}: us-gov / global ratio is {ratio}, expected 1.2"
# The us-gov.anthropic.* cross-region inference profile is the only us-gov
@ -112,9 +101,7 @@ def test_usgov_cross_region_above_200k_carries_gov_premium(model_data, field, ex
"""
info = model_data[USGOV_CROSS_REGION_KEY]
assert field in info, f"{USGOV_CROSS_REGION_KEY}: missing field {field}"
assert (
info[field] == expected
), f"{USGOV_CROSS_REGION_KEY}: {field} should be {expected} (got {info[field]})"
assert info[field] == expected, f"{USGOV_CROSS_REGION_KEY}: {field} should be {expected} (got {info[field]})"
def test_usgov_cross_region_above_200k_ratio_to_global(model_data):
@ -127,6 +114,176 @@ def test_usgov_cross_region_above_200k_ratio_to_global(model_data):
usgov_info = model_data[USGOV_CROSS_REGION_KEY]
for field in EXPECTED_USGOV_ABOVE_200K:
ratio = usgov_info[field] / global_info[field]
assert (
abs(ratio - 1.2) < 1e-9
), f"{field}: us-gov / global ratio is {ratio}, expected 1.2"
assert abs(ratio - 1.2) < 1e-9, f"{field}: us-gov / global ratio is {ratio}, expected 1.2"
CLAUDE_GOV_EXPECTED = {
"anthropic.claude-sonnet-5": {
"input_cost_per_token": 2.4e-06,
"output_cost_per_token": 1.2e-05,
"cache_creation_input_token_cost": 3e-06,
"cache_creation_input_token_cost_above_1hr": 4.8e-06,
"cache_read_input_token_cost": 2.4e-07,
},
"anthropic.claude-opus-4-8": {
"input_cost_per_token": 6e-06,
"output_cost_per_token": 3e-05,
"cache_creation_input_token_cost": 7.5e-06,
"cache_creation_input_token_cost_above_1hr": 1.2e-05,
"cache_read_input_token_cost": 6e-07,
},
}
USGOV_CLAUDE_KEY_TEMPLATES = {
"bedrock/us-gov-east-1/{base_key}": "bedrock",
"bedrock/us-gov-west-1/{base_key}": "bedrock",
"us-gov.{base_key}": "bedrock_converse",
}
@pytest.mark.parametrize("base_key", CLAUDE_GOV_EXPECTED)
@pytest.mark.parametrize("key_template,expected_provider", USGOV_CLAUDE_KEY_TEMPLATES.items())
def test_usgov_claude_sonnet5_opus48_pricing(model_data, key_template, expected_provider, base_key):
"""Sonnet 5 and Opus 4.8 gov entries, both in-region keys and the us-gov.
geo inference profile the model cards list for GovCloud, must match the
rates AWS publishes on the Bedrock pricing page (1.2x global).
"""
gov_key = key_template.format(base_key=base_key)
assert gov_key in model_data, f"Missing model entry: {gov_key}"
info = model_data[gov_key]
assert info["litellm_provider"] == expected_provider
for field, expected in CLAUDE_GOV_EXPECTED[base_key].items():
assert info[field] == expected, f"{gov_key}: {field} should be {expected} (got {info[field]})"
ratio = info[field] / model_data[base_key][field]
assert abs(ratio - 1.2) < 1e-9, f"{gov_key}: {field} gov/global ratio is {ratio}, expected 1.2"
CONVERSE_GOV_EXPECTED = {
"nvidia.nemotron-nano-3-30b": (7.2e-08, 2.88e-07),
"nvidia.nemotron-nano-12b-v2": (2.4e-07, 7.2e-07),
"nvidia.nemotron-super-3-120b": (1.8e-07, 7.8e-07),
"openai.gpt-oss-20b-1:0": (8.4e-08, 3.6e-07),
"openai.gpt-oss-120b-1:0": (1.8e-07, 7.2e-07),
}
@pytest.mark.parametrize("base_key", CONVERSE_GOV_EXPECTED)
@pytest.mark.parametrize("region", ["us-gov-east-1", "us-gov-west-1"])
def test_usgov_converse_model_pricing(model_data, region, base_key):
"""Nemotron and gpt-oss gov entries must match the AWS Bedrock offer file,
which prices both GovCloud regions identically at 1.2x commercial.
"""
gov_key = f"bedrock/{region}/{base_key}"
assert gov_key in model_data, f"Missing model entry: {gov_key}"
info = model_data[gov_key]
expected_input, expected_output = CONVERSE_GOV_EXPECTED[base_key]
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["litellm_provider"] == "bedrock"
base = model_data[base_key]
assert abs(info["input_cost_per_token"] / base["input_cost_per_token"] - 1.2) < 1e-9
assert abs(info["output_cost_per_token"] / base["output_cost_per_token"] - 1.2) < 1e-9
def test_usgov_west_llama3_8b_output_price_fixed(model_data):
"""The us-gov-west-1 llama3-8b entry carried the 70B output rate ($2.65/MTok);
the AWS Bedrock offer file prices output at $0.60/MTok. AWS lists the model
in us-gov-west-1 only, so there is no east entry to check.
"""
info = model_data["bedrock/us-gov-west-1/meta.llama3-8b-instruct-v1:0"]
assert info["input_cost_per_token"] == 3e-07
assert info["output_cost_per_token"] == 6e-07
MANTLE_GOV_TIERED_EXPECTED = {
"openai.gpt-5.6-luna": {
"input_cost_per_token": 2.64e-07,
"input_cost_per_token_above_272k_tokens": 5.28e-07,
"cache_creation_input_token_cost": 3.3e-07,
"cache_creation_input_token_cost_above_272k_tokens": 6.6e-07,
"cache_read_input_token_cost": 2.64e-08,
"cache_read_input_token_cost_above_272k_tokens": 5.28e-08,
"output_cost_per_token": 1.584e-06,
"output_cost_per_token_above_272k_tokens": 2.376e-06,
},
"openai.gpt-5.6-terra": {
"input_cost_per_token": 2.64e-06,
"input_cost_per_token_above_272k_tokens": 5.28e-06,
"cache_creation_input_token_cost": 3.3e-06,
"cache_creation_input_token_cost_above_272k_tokens": 6.6e-06,
"cache_read_input_token_cost": 2.64e-07,
"cache_read_input_token_cost_above_272k_tokens": 5.28e-07,
"output_cost_per_token": 1.584e-05,
"output_cost_per_token_above_272k_tokens": 2.376e-05,
},
}
@pytest.mark.parametrize("model", MANTLE_GOV_TIERED_EXPECTED)
def test_usgov_west_mantle_terra_luna_pricing(model_data, model):
"""Terra and Luna carry 1.2x commercial across every tier in the
us-gov-west-1 offer file; the us-gov-east-1 offer file has no SKUs for them.
"""
gov_key = f"bedrock_mantle/us-gov-west-1/{model}"
assert gov_key in model_data, f"Missing model entry: {gov_key}"
info = model_data[gov_key]
for field, expected in MANTLE_GOV_TIERED_EXPECTED[model].items():
assert info[field] == expected, f"{gov_key}: {field} should be {expected} (got {info[field]})"
assert info["litellm_provider"] == "bedrock_mantle"
assert f"bedrock_mantle/us-gov-east-1/{model}" not in model_data
@pytest.mark.parametrize("region", ["us-gov-east-1", "us-gov-west-1"])
def test_usgov_mantle_gpt_5_4_pricing_has_no_long_context_tier(model_data, region):
"""gpt-5.4 gov rates come from the offer file, which publishes only the
standard tier in GovCloud: no long-context SKUs exist there, unlike commercial.
"""
gov_key = f"bedrock_mantle/{region}/openai.gpt-5.4"
assert gov_key in model_data, f"Missing model entry: {gov_key}"
info = model_data[gov_key]
assert info["input_cost_per_token"] == 3.3e-06
assert info["cache_read_input_token_cost"] == 3.3e-07
assert info["output_cost_per_token"] == 1.98e-05
assert not any(field.endswith("_above_272k_tokens") for field in info)
def test_usgov_mantle_grok_4_3_west_only(model_data):
"""grok-4.3 is priced in the us-gov-west-1 offer file only; the east offer
file carries grok-4.6 instead.
"""
info = model_data["bedrock_mantle/us-gov-west-1/xai.grok-4.3"]
assert info["input_cost_per_token"] == 1.5e-06
assert info["output_cost_per_token"] == 3e-06
assert info["cache_read_input_token_cost"] == 2.4e-07
assert "bedrock_mantle/us-gov-east-1/xai.grok-4.3" not in model_data
AZURE_GOV_EXPECTED = {
"azure/us-gov/gpt-5.1": {
"input_cost_per_token": 1.71875e-06,
"cache_read_input_token_cost": 1.71875e-07,
"output_cost_per_token": 1.375e-05,
},
"azure/us-gov/o3-mini": {
"input_cost_per_token": 1.513e-06,
"cache_read_input_token_cost": 7.57e-07,
"output_cost_per_token": 6.05e-06,
},
"azure/us-gov/text-embedding-3-large": {"input_cost_per_token": 1.63e-07},
"azure/us-gov/text-embedding-3-small": {"input_cost_per_token": 2.5e-08},
}
@pytest.mark.parametrize("gov_key", AZURE_GOV_EXPECTED)
def test_azure_usgov_pricing(model_data, gov_key):
"""Azure Government meters from the Azure retail prices API
(usgovvirginia/usgovarizona, serviceName 'Foundry Models'). No Government
retirement schedule is published, so these entries carry no deprecation_date.
"""
assert gov_key in model_data, f"Missing model entry: {gov_key}"
info = model_data[gov_key]
for field, expected in AZURE_GOV_EXPECTED[gov_key].items():
assert info[field] == expected, f"{gov_key}: {field} should be {expected} (got {info[field]})"
assert info["litellm_provider"] == "azure"
assert "deprecation_date" not in info

View file

@ -75,6 +75,22 @@ def test_additional_current_models_are_present():
assert entry["output_cost_per_token"] > 0
@pytest.mark.parametrize(
"key, published_price_per_audio_minute",
[
("cloudflare/@cf/openai/whisper", 0.00045),
("cloudflare/@cf/openai/whisper-large-v3-turbo", 0.00051),
],
)
def test_whisper_transcription_pricing_is_stored_per_second(key, published_price_per_audio_minute):
entry = litellm.model_cost[key]
assert entry["litellm_provider"] == "cloudflare"
assert entry["mode"] == "audio_transcription"
assert entry["supported_endpoints"] == ["/v1/audio/transcriptions"]
assert entry["output_cost_per_second"] == 0.0
assert entry["input_cost_per_second"] == pytest.approx(published_price_per_audio_minute / 60)
def test_root_and_backup_have_identical_cloudflare_keys():
if not os.path.exists(ROOT_MAP):
pytest.skip("root cost map only ships in source checkouts")

View file

@ -3150,8 +3150,8 @@ def _stream_builder_logging_obj() -> LiteLLMLogging:
return logging_obj
def test_stream_chunk_builder_reports_streaming_usage_cost_when_enabled(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
def test_stream_chunk_builder_stamps_streaming_usage_cost_by_default(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", False)
chunks: Final = [
_stream_builder_text_chunk("gpt-4o", "Hello "),
_stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"),
@ -3168,11 +3168,45 @@ def test_stream_chunk_builder_reports_streaming_usage_cost_when_enabled(monkeypa
assert response._hidden_params["response_cost"] == pytest.approx(usage_cost)
def test_stream_chunk_builder_defers_cost_to_logging_obj_when_usage_cost_absent(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", False)
def test_stream_chunk_builder_skips_stamp_when_cost_is_unpriceable():
import time as time_module
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
logging_obj: Final = LiteLLMLogging(
model="us.anthropic.claude-opus-5",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="completion",
start_time=time_module.time(),
litellm_call_id="stream-builder-alias-unpriceable",
function_id="1",
)
logging_obj.model_call_details["custom_llm_provider"] = "bedrock"
logging_obj.optional_params = {}
usage_chunk: Final = _stream_builder_text_chunk("bedrock-claude-opus-5", "")
usage_chunk.usage = Usage(prompt_tokens=40, completion_tokens=5, total_tokens=45)
chunks: Final = [
_stream_builder_text_chunk("bedrock-claude-opus-5", "Hello ", finish_reason="stop"),
usage_chunk,
]
response: Final = litellm.stream_chunk_builder(
chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=logging_obj
)
assert response is not None
assert getattr(response.usage, "cost", None) is None
assert response._hidden_params.get("response_cost") is None
def test_stream_chunk_builder_keeps_provider_reported_usage_cost():
usage_chunk: Final = _stream_builder_text_chunk("gpt-4o", "")
usage_chunk.usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15, cost=0.5)
chunks: Final = [
_stream_builder_text_chunk("gpt-4o", "Hello "),
_stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"),
usage_chunk,
]
response: Final = litellm.stream_chunk_builder(
@ -3180,4 +3214,26 @@ def test_stream_chunk_builder_defers_cost_to_logging_obj_when_usage_cost_absent(
)
assert response is not None
assert response._hidden_params.get("response_cost") is None
assert getattr(response.usage, "cost", None) == pytest.approx(0.5)
assert response._hidden_params["response_cost"] == pytest.approx(0.5)
def test_stream_chunk_builder_prices_alias_from_openai_sdk_usage_chunk():
from openai.types.completion_usage import CompletionUsage
usage_chunk: Final = _stream_builder_text_chunk("mantle-claude", "")
usage_chunk.usage = CompletionUsage(prompt_tokens=20, completion_tokens=60, total_tokens=80, cost=0.000704)
assert type(usage_chunk.usage) is CompletionUsage
chunks: Final = [
_stream_builder_text_chunk("mantle-claude", "Hello "),
_stream_builder_text_chunk("mantle-claude", "world.", finish_reason="stop"),
usage_chunk,
]
response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}])
assert response is not None
assert response.usage.prompt_tokens == 20
assert response.usage.completion_tokens == 60
assert getattr(response.usage, "cost", None) == pytest.approx(0.000704)
assert response._hidden_params["response_cost"] == pytest.approx(0.000704)

View file

@ -0,0 +1,156 @@
import json
from functools import lru_cache
from pathlib import Path
import pytest
import litellm
REPO_ROOT = Path(__file__).parents[2]
MAIN_PATH = REPO_ROOT / "model_prices_and_context_window.json"
BACKUP_PATH = REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json"
FLEX_LONG_CONTEXT = {
"gpt-5.4": {
"input_cost_per_token_above_272k_tokens_flex": 2.5e-06,
"output_cost_per_token_above_272k_tokens_flex": 1.125e-05,
"cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07,
},
"gpt-5.4-pro": {
"input_cost_per_token_above_272k_tokens_flex": 3e-05,
"output_cost_per_token_above_272k_tokens_flex": 0.000135,
},
"gpt-5.5": {
"input_cost_per_token_above_272k_tokens_flex": 5e-06,
"output_cost_per_token_above_272k_tokens_flex": 2.25e-05,
"cache_read_input_token_cost_above_272k_tokens_flex": 5e-07,
},
}
PRIORITY_LONG_CONTEXT = {
"gpt-5.6": {
"input_cost_per_token_above_272k_tokens_priority": 1.6e-05,
"output_cost_per_token_above_272k_tokens_priority": 6e-05,
"cache_read_input_token_cost_above_272k_tokens_priority": 1.6e-06,
"cache_creation_input_token_cost_above_272k_tokens_priority": 2e-05,
},
"gpt-5.6-sol": {
"input_cost_per_token_above_272k_tokens_priority": 1.6e-05,
"output_cost_per_token_above_272k_tokens_priority": 6e-05,
"cache_read_input_token_cost_above_272k_tokens_priority": 1.6e-06,
"cache_creation_input_token_cost_above_272k_tokens_priority": 2e-05,
},
"gpt-5.6-terra": {
"input_cost_per_token_above_272k_tokens_priority": 8e-06,
"output_cost_per_token_above_272k_tokens_priority": 3.6e-05,
"cache_read_input_token_cost_above_272k_tokens_priority": 8e-07,
"cache_creation_input_token_cost_above_272k_tokens_priority": 1e-05,
},
"gpt-5.6-luna": {
"input_cost_per_token_above_272k_tokens_priority": 8e-07,
"output_cost_per_token_above_272k_tokens_priority": 3.6e-06,
"cache_read_input_token_cost_above_272k_tokens_priority": 8e-08,
"cache_creation_input_token_cost_above_272k_tokens_priority": 1e-06,
},
}
EXPECTED = {**FLEX_LONG_CONTEXT, **PRIORITY_LONG_CONTEXT}
NO_PUBLISHED_PRIORITY_LONG_CONTEXT = ("gpt-5.4", "gpt-5.5")
@pytest.fixture(autouse=True)
def _local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
@lru_cache(maxsize=2)
def _load(path: Path) -> dict[str, dict[str, object]]:
with open(path) as f:
return json.load(f)
@pytest.mark.parametrize("path", [MAIN_PATH, BACKUP_PATH], ids=["main", "backup"])
@pytest.mark.parametrize("model", sorted(EXPECTED))
def test_service_tier_long_context_rates_are_published(model: str, path: Path) -> None:
"""Each tier must carry its own above-272K rates, in both price files."""
info = _load(path).get(model)
assert info is not None, f"{model} not found in {path.name}"
for key, expected in EXPECTED[model].items():
assert info.get(key) == pytest.approx(expected), f"{model}.{key} is {info.get(key)!r}, expected {expected!r}"
@pytest.mark.parametrize("model", sorted(EXPECTED))
def test_tier_long_context_rate_is_half_or_double_the_standard(model: str) -> None:
"""Flex is half the standard long-context rate; priority is double it."""
info = _load(MAIN_PATH)[model]
tier = "flex" if model in FLEX_LONG_CONTEXT else "priority"
ratio = 0.5 if tier == "flex" else 2.0
for base in ("input_cost_per_token", "output_cost_per_token"):
standard = info[f"{base}_above_272k_tokens"]
tiered = info[f"{base}_above_272k_tokens_{tier}"]
assert tiered == pytest.approx(standard * ratio), (
f"{model}.{base}_above_272k_tokens_{tier} is {tiered!r}, "
f"expected {ratio}x the standard long-context rate {standard!r}"
)
@pytest.mark.parametrize("model", NO_PUBLISHED_PRIORITY_LONG_CONTEXT)
def test_no_priority_long_context_rates_where_openai_publishes_none(model: str) -> None:
"""Guard against back-filling a rate OpenAI does not publish."""
info = _load(MAIN_PATH)[model]
assert "input_cost_per_token_above_272k_tokens_priority" not in info
LONG_CONTEXT_PROMPT_TOKENS = 300_000
COMPLETION_TOKENS = 1_000
TIERED_COST_CASES = [
("gpt-5.4", "flex", 2.5e-06, 1.125e-05),
("gpt-5.4-pro", "flex", 3e-05, 0.000135),
("gpt-5.5", "flex", 5e-06, 2.25e-05),
("gpt-5.6", "priority", 1.6e-05, 6e-05),
("gpt-5.6-sol", "priority", 1.6e-05, 6e-05),
("gpt-5.6-terra", "priority", 8e-06, 3.6e-05),
("gpt-5.6-luna", "priority", 8e-07, 3.6e-06),
]
@pytest.mark.parametrize("model,tier,input_rate,output_rate", TIERED_COST_CASES)
def test_cost_per_token_bills_long_context_at_the_tier_rate(
model: str, tier: str, input_rate: float, output_rate: float
) -> None:
"""A prompt over 272K on flex or priority must bill at that tier's long-context rate."""
input_cost, output_cost = litellm.cost_per_token(
model=model,
prompt_tokens=LONG_CONTEXT_PROMPT_TOKENS,
completion_tokens=COMPLETION_TOKENS,
service_tier=tier,
)
assert input_cost == pytest.approx(LONG_CONTEXT_PROMPT_TOKENS * input_rate)
assert output_cost == pytest.approx(COMPLETION_TOKENS * output_rate)
@pytest.mark.parametrize("model,tier,input_rate,output_rate", TIERED_COST_CASES)
def test_cost_per_token_tier_differs_from_the_standard_long_context_cost(
model: str, tier: str, input_rate: float, output_rate: float
) -> None:
"""Flex halves the standard long-context bill and priority doubles it."""
ratio = 0.5 if tier == "flex" else 2.0
standard = sum(
litellm.cost_per_token(
model=model,
prompt_tokens=LONG_CONTEXT_PROMPT_TOKENS,
completion_tokens=COMPLETION_TOKENS,
)
)
tiered = sum(
litellm.cost_per_token(
model=model,
prompt_tokens=LONG_CONTEXT_PROMPT_TOKENS,
completion_tokens=COMPLETION_TOKENS,
service_tier=tier,
)
)
assert tiered == pytest.approx(standard * ratio)

View file

@ -7271,6 +7271,71 @@ def test_get_configured_token_limits_coerces_numeric_strings():
assert router.get_configured_token_limits("quoted-limits-model") == (32000, 8000)
def test_get_configured_display_name_reads_deployment_model_info():
router = litellm.Router(
model_list=[
{
"model_name": "Kimi K3-claude-compatible",
"litellm_params": {"model": "openai/some-unmapped-model"},
"model_info": {"display_name": "Kimi K3"},
}
]
)
assert router.get_configured_display_name("Kimi K3-claude-compatible") == "Kimi K3"
def test_get_configured_display_name_returns_none_for_unset_or_unknown():
router = litellm.Router(
model_list=[
{
"model_name": "no-display-model",
"litellm_params": {"model": "openai/some-unmapped-model"},
}
]
)
assert router.get_configured_display_name("no-display-model") is None
assert router.get_configured_display_name("not-a-real-model") is None
def test_get_configured_display_name_skips_wildcard_pattern_matching():
router = litellm.Router(
model_list=[
{
"model_name": "bedrock/*",
"litellm_params": {"model": "bedrock/*"},
"model_info": {"display_name": "Bedrock"},
}
]
)
with patch.object(
router.pattern_router, "route", side_effect=AssertionError("pattern route called")
):
assert (
router.get_configured_display_name("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0")
is None
)
def test_get_configured_display_name_treats_malformed_values_as_absent():
malformed = ["", " ", 12345, ["Kimi K3"], {"name": "Kimi K3"}, True]
router = litellm.Router(
model_list=[
{
"model_name": f"bad-display-{i}",
"litellm_params": {"model": "openai/some-unmapped-model"},
"model_info": {"display_name": bad},
}
for i, bad in enumerate(malformed)
]
)
for i in range(len(malformed)):
assert router.get_configured_display_name(f"bad-display-{i}") is None
@pytest.mark.asyncio
async def test_acreate_batch_disable_fallbacks_surfaces_owning_provider_error():
router = litellm.Router(

View file

@ -21,6 +21,7 @@ This file pins both halves of the fix.
import json
from dataclasses import dataclass
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -28,8 +29,19 @@ from pydantic import ValidationError
import litellm
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.router_utils.pre_call_checks.model_rate_limit_check import ModelRateLimitingCheck
from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import PromptCachingDeploymentCheck
from litellm.types.router import RetryPolicy, UpdateRouterConfig
@pytest.fixture(autouse=True)
def isolate_litellm_callbacks():
callbacks_before: Final = litellm.callbacks.copy()
yield
litellm.callbacks = callbacks_before # test-quality-ok: required callback-state restoration fixture
# ---------------------------------------------------------------------------
# UpdateRouterConfig schema membership (LIT-3152 part 1)
# ---------------------------------------------------------------------------
@ -100,6 +112,114 @@ def _build_router() -> litellm.Router:
)
def test_update_settings_adds_optional_pre_call_check_once():
router = _build_router()
router.update_settings(num_retries=7, optional_pre_call_checks=["prompt_caching"])
router.update_settings(optional_pre_call_checks=["prompt_caching"])
prompt_caching_callbacks = [
callback for callback in router.optional_callbacks if isinstance(callback, PromptCachingDeploymentCheck)
]
assert len(prompt_caching_callbacks) == 1
assert router.num_retries == 7
def test_update_settings_clears_omitted_toggleable_pre_call_checks():
router = _build_router()
router.update_settings(optional_pre_call_checks=["prompt_caching"])
router.update_settings(optional_pre_call_checks=[])
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in (router.optional_callbacks or []))
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in litellm.callbacks)
def test_set_optional_pre_call_checks_reconciles_callback_types():
router = _build_router()
router.set_optional_pre_call_checks(["prompt_caching"])
router.set_optional_pre_call_checks([])
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in (router.optional_callbacks or []))
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in litellm.callbacks)
def test_remove_optional_pre_call_check_removes_local_and_global_callbacks():
router = _build_router()
router.set_optional_pre_call_checks(["prompt_caching"])
router._remove_optional_callbacks_of_type(PromptCachingDeploymentCheck)
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in (router.optional_callbacks or []))
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
def test_remove_optional_pre_call_check_keeps_global_callback_for_another_router():
router_a = _build_router()
router_b = _build_router()
router_a.update_settings(optional_pre_call_checks=["prompt_caching"])
router_b.update_settings(optional_pre_call_checks=["prompt_caching"])
router_a.update_settings(optional_pre_call_checks=[])
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in (router_a.optional_callbacks or []))
assert any(type(callback) is PromptCachingDeploymentCheck for callback in (router_b.optional_callbacks or []))
assert any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
router_b.update_settings(optional_pre_call_checks=[])
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in (router_b.optional_callbacks or []))
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
def test_remove_optional_pre_call_check_keeps_global_callback_when_second_router_clears_first():
router_a = _build_router()
router_b = _build_router()
router_a.update_settings(optional_pre_call_checks=["prompt_caching"])
router_b.update_settings(optional_pre_call_checks=["prompt_caching"])
router_b.update_settings(optional_pre_call_checks=[])
assert any(type(callback) is PromptCachingDeploymentCheck for callback in (router_a.optional_callbacks or []))
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in (router_b.optional_callbacks or []))
assert any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
router_a.update_settings(optional_pre_call_checks=[])
assert not any(type(callback) is PromptCachingDeploymentCheck for callback in litellm.callbacks)
def test_update_settings_replaces_toggleable_pre_call_checks():
router = _build_router()
router.update_settings(optional_pre_call_checks=["prompt_caching"])
router.update_settings(optional_pre_call_checks=["enforce_model_rate_limits"])
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in (router.optional_callbacks or []))
assert not any(isinstance(callback, PromptCachingDeploymentCheck) for callback in litellm.callbacks)
assert any(isinstance(callback, ModelRateLimitingCheck) for callback in (router.optional_callbacks or []))
@pytest.mark.asyncio
async def test_update_settings_preserves_router_budget_limiting_when_omitted(monkeypatch):
async def _disable_periodic_sync(*args, **kwargs):
return None
monkeypatch.setattr(
"litellm.router_strategy.budget_limiter.RouterBudgetLimiting.periodic_sync_in_memory_spend_with_redis",
_disable_periodic_sync,
)
router = _build_router()
router.add_optional_pre_call_checks(["router_budget_limiting"])
router.update_settings(optional_pre_call_checks=[])
assert any(isinstance(callback, RouterBudgetLimiting) for callback in (router.optional_callbacks or []))
def test_update_settings_persists_retry_policy_dict():
"""When the proxy's ``_add_router_settings_from_db_config`` calls
``llm_router.update_settings(retry_policy={...})`` after reading the
@ -255,8 +375,12 @@ async def test_config_update_persists_and_reads_back_retry_policy(monkeypatch):
RateLimitErrorRetries=7,
)
)
request = MagicMock()
request.json = AsyncMock(return_value={"router_settings": {"retry_policy": posted.model_dump()}})
await proxy_server.update_config(
config_info=ConfigYAML(router_settings=posted),
request=request,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"),
)

View file

@ -4655,6 +4655,7 @@ GEMINI_4096_CACHE_MIN_MODELS: Final = tuple(
"gemini-3.5-flash",
"gemini-3.6-flash",
"gemini-3.7-flash",
"gemini-3.8-flash",
"gemini-3.1-pro-preview",
"gemini-3.1-pro-preview-customtools",
)

View file

@ -37473,6 +37473,8 @@ export interface components {
} | null;
/** Num Retries */
num_retries?: number | null;
/** Optional Pre Call Checks */
optional_pre_call_checks?: ("prompt_caching" | "router_budget_limiting" | "responses_api_deployment_check" | "deployment_affinity" | "session_affinity" | "forward_client_headers_by_model_group" | "enforce_model_rate_limits" | "encrypted_content_affinity")[] | null;
/** Retry After */
retry_after?: number | null;
retry_policy?: components["schemas"]["RetryPolicy"] | null;

View file

@ -217,3 +217,17 @@ bedrock/us-east-1/zai.glm-5
bedrock/us-west-2/zai.glm-5
bedrock/us-gov-east-1/anthropic.claude-haiku-4-5-20251001-v1:0
bedrock/us-gov-west-1/anthropic.claude-haiku-4-5-20251001-v1:0
bedrock/us-gov-west-1/nvidia.nemotron-nano-3-30b
bedrock/us-gov-west-1/nvidia.nemotron-nano-12b-v2
bedrock/us-gov-west-1/nvidia.nemotron-super-3-120b
bedrock/us-gov-west-1/openai.gpt-oss-20b-1:0
bedrock/us-gov-west-1/openai.gpt-oss-120b-1:0
bedrock/us-gov-west-1/anthropic.claude-sonnet-5
bedrock/us-gov-west-1/anthropic.claude-opus-4-8
bedrock/us-gov-east-1/nvidia.nemotron-nano-3-30b
bedrock/us-gov-east-1/nvidia.nemotron-nano-12b-v2
bedrock/us-gov-east-1/nvidia.nemotron-super-3-120b
bedrock/us-gov-east-1/openai.gpt-oss-20b-1:0
bedrock/us-gov-east-1/openai.gpt-oss-120b-1:0
bedrock/us-gov-east-1/anthropic.claude-sonnet-5
bedrock/us-gov-east-1/anthropic.claude-opus-4-8