mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
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:
commit
70fbf7f4a3
55 changed files with 4433 additions and 359 deletions
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 }
|
||||
|
|
|
|||
|
|
@ -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: {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
90
litellm/llms/parallel_ai/search/cost_calculator.py
Normal file
90
litellm/llms/parallel_ai/search/cost_calculator.py
Normal 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
|
||||
|
|
@ -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}))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]}],
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue