Merge remote-tracking branch 'upstream/litellm_internal_staging' into deepkeep-as-internal

This commit is contained in:
Yaniv Israel 2026-06-17 18:30:26 +03:00
commit b892470ba3
88 changed files with 12279 additions and 793 deletions

View file

@ -177,7 +177,7 @@ jobs:
restore-keys: |
any-mypy-cache-${{ runner.os }}-py3.12-
- name: Check Any discipline on changed lines
- name: Check Any discipline (per-file budget on changed files)
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |

View file

@ -36,11 +36,11 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a
Run tests, format your code, and lint your code before each commit
When you fix violations gated by `ruff-strict-budget.json`, `mypy-code-budget.json`, or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered baselines so the ceilings ratchet down instead of leaving stale headroom
When you fix violations gated by `ruff-strict-budget.json`, `mypy-code-budget.json`, `basedpyright-code-budget.json`, or `any-discipline-budget.json`, run `make lint-budget-update` and commit the lowered baselines so the ceilings ratchet down instead of leaving stale headroom
If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and bringing it closer to the max, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in
The Any-discipline gate (`make lint-any`, also a CI job) fails when a line you changed under `litellm/` holds a value typed `Any`, including the `X | Any`. Ideally `# any-ok: <reason>` is never used; treat it as a last resort for a genuine typed/untyped boundary that Pydantic truly can't model
The Any-discipline gate (`make lint-any`, also a CI job) fails when a changed file under `litellm/` carries more `Any`-typed values than its grandfathered ceiling in `any-discipline-budget.json` (each file's captured count plus 50% headroom). It flags values whose inferred type *contains* `Any`, including the `X | Any` unions mypy/basedpyright accept. Editing a legacy file is fine as long as you don't push its `Any` count past the ceiling; a brand-new file must be `Any`-free. Fix a value by giving it a concrete type (if you're given untyped input, validate with Pydantic). Ideally `# any-ok: <reason>` is never used; treat it as a last resort for a genuine typed/untyped boundary that Pydantic truly can't model
If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason

View file

@ -155,7 +155,7 @@ Individual linting commands:
make format-check # Check Black formatting
make lint-ruff # Run Ruff linting
make lint-mypy # Run MyPy type checking
make lint-any # Fail on Any-typed values on changed lines
make lint-any # Gate changed files against their per-file Any budget
make check-circular-imports # Check for circular imports
make check-import-safety # Check import safety
```

View file

@ -6,7 +6,7 @@
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
info lint lint-dev format \
lint-mypy lint-mypy-budget-update lint-basedpyright lint-basedpyright-budget-update \
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-any \
lint-ruff-budget lint-any lint-ruff-budget-update lint-budget-update lint-any-budget-update \
install-dev install-proxy-dev install-test-deps install-hooks \
install-helm-unittest check-circular-imports check-import-safety
@ -30,9 +30,10 @@ help:
@echo " make lint-basedpyright-budget-update - Re-capture the basedpyright per-rule budget (ratchet)"
@echo " make lint-black - Check Black formatting (matches CI)"
@echo " make lint-ruff-budget - Gate the codebase total of each strict ruff rule against its ceiling"
@echo " make lint-any - Gate changed files under litellm/ against their per-file Any budget"
@echo " make lint-ruff-budget-update - Re-capture per-rule baselines in ruff-strict-budget.json (ratchet)"
@echo " make lint-budget-update - Re-capture all three ratchet budgets (ruff + mypy + basedpyright)"
@echo " make lint-any - Fail if changed lines under litellm/ hold an Any-typed value"
@echo " make lint-budget-update - Re-capture all four ratchet budgets (ruff + mypy + basedpyright + any)"
@echo " make lint-any-budget-update - Re-capture the per-file Any budget across the whole tree (ratchet)"
@echo " make check-circular-imports - Check for circular imports"
@echo " make check-import-safety - Check import safety"
@echo " make test - Run all tests"
@ -149,12 +150,15 @@ lint-ruff-budget: install-dev
lint-ruff-budget-update: install-dev
$(UV_RUN) python scripts/ruff_strict_gate.py --update
# Ratchet all three budgets in one shot (ruff strict + mypy + basedpyright)
lint-budget-update: lint-ruff-budget-update lint-mypy-budget-update lint-basedpyright-budget-update
# Ratchet all four budgets in one shot (ruff strict + mypy + basedpyright + any)
lint-budget-update: lint-ruff-budget-update lint-mypy-budget-update lint-basedpyright-budget-update lint-any-budget-update
lint-any: install-dev
$(UV_RUN) python scripts/check_any_discipline.py --changed
lint-any-budget-update: install-dev
$(UV_RUN) python scripts/check_any_discipline.py --update
check-circular-imports: install-dev
cd litellm && $(UV_RUN) python ../tests/documentation_tests/test_circular_imports.py && cd ..

5974
any-discipline-budget.json Normal file

File diff suppressed because it is too large Load diff

View file

@ -5,7 +5,7 @@
},
"reportArgumentType": {
"baseline": 1863,
"slack": 3
"slack": 180
},
"reportAssignmentType": {
"baseline": 220,
@ -113,7 +113,7 @@
},
"reportPrivateUsage": {
"baseline": 1625,
"slack": 10
"slack": 160
},
"reportRedeclaration": {
"baseline": 8,

View file

@ -369,6 +369,8 @@ class RedisCache(BaseCache):
"""
Make sure each key starts with the given namespace
"""
if key is None:
return key # type: ignore[return-value]
if self.namespace is not None and not key.startswith(self.namespace):
key = self.namespace + ":" + key

View file

@ -510,6 +510,8 @@ DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE = os.getenv(
"DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE", "streaming.chunk.yield"
)
LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED = 499
EMAIL_BUDGET_ALERT_TTL = int(
os.getenv("EMAIL_BUDGET_ALERT_TTL", 24 * 60 * 60)
) # 24 hours in seconds

View file

@ -94,6 +94,7 @@ from litellm.types.utils import (
LlmProviders,
LlmProvidersSet,
ModelInfo,
ServiceTier,
StandardBuiltInToolsParams,
TranscriptionUsageDurationObject,
TranscriptionUsageTokensObject,
@ -614,7 +615,9 @@ def cost_per_token(
service_tier=service_tier,
)
elif custom_llm_provider == "anthropic":
return anthropic_cost_per_token(model=model, usage=usage_block)
return anthropic_cost_per_token(
model=model, usage=usage_block, service_tier=service_tier
)
elif custom_llm_provider == "bedrock":
return bedrock_cost_per_token(
model=model, usage=usage_block, service_tier=service_tier
@ -1224,6 +1227,12 @@ def completion_cost(
if service_tier is None and optional_params is not None:
service_tier = optional_params.get("service_tier")
# "auto" is a routing preference, not a billable tier: the provider picks
# the tier and reports the one actually served on the response/usage, so
# defer to that instead of pricing the request-level "auto" as standard
if service_tier is not None and service_tier.lower() == ServiceTier.AUTO.value:
service_tier = None
# Extract service_tier from completion_response if not provided
if service_tier is None and completion_response is not None:
if isinstance(completion_response, BaseModel):

View file

@ -5446,6 +5446,39 @@ class StandardLoggingPayloadSetup:
error_rate_limit_type=rate_limit_type,
)
@staticmethod
def get_error_information_for_logging_payload(
metadata: dict,
original_exception: Exception | None,
error_str: str | None,
) -> tuple[StandardLoggingPayloadErrorInformation, str | None]:
error_information = StandardLoggingPayloadSetup.get_error_information(
original_exception=original_exception,
)
if not metadata.get("client_disconnected"): # any-ok: untyped metadata
return error_information, error_str
client_disconnect_error = metadata.get( # any-ok: untyped metadata
"error_information"
)
if isinstance(client_disconnect_error, dict): # any-ok: untyped metadata
error_information = cast(
StandardLoggingPayloadErrorInformation,
client_disconnect_error, # any-ok: untyped metadata
)
else:
error_information = cast(
StandardLoggingPayloadErrorInformation,
{ # any-ok: untyped metadata
"error_code": "499",
"error_message": "Client disconnected the request",
"error_class": "ClientDisconnected",
},
)
if not error_str:
error_str = "Client disconnected the request"
return error_information, error_str
@staticmethod
def get_response_time(
start_time_float: float,
@ -5773,8 +5806,12 @@ def get_standard_logging_object_payload(
api_base=litellm_params.get("api_base"),
)
error_information = StandardLoggingPayloadSetup.get_error_information(
original_exception=original_exception,
error_information, error_str = (
StandardLoggingPayloadSetup.get_error_information_for_logging_payload(
metadata=metadata, # any-ok: untyped metadata
original_exception=original_exception,
error_str=error_str,
)
)
## get final response object ##

View file

@ -303,40 +303,54 @@ def _get_token_base_cost(
# Apply tiered pricing to cache costs
cache_creation_tiered_key = (
f"cache_creation_input_token_cost_above_{threshold_str}_tokens"
_get_service_tier_cost_key(
f"cache_creation_input_token_cost_above_{threshold_str}_tokens",
service_tier,
)
if service_tier
else f"cache_creation_input_token_cost_above_{threshold_str}_tokens"
)
cache_creation_1hr_tiered_key = (
_get_service_tier_cost_key(
f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens",
service_tier,
)
if service_tier
else f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens"
)
cache_creation_1hr_tiered_key = f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens"
cache_read_tiered_key = (
f"cache_read_input_token_cost_above_{threshold_str}_tokens"
_get_service_tier_cost_key(
f"cache_read_input_token_cost_above_{threshold_str}_tokens",
service_tier,
)
if service_tier
else f"cache_read_input_token_cost_above_{threshold_str}_tokens"
)
if cache_creation_tiered_key in model_info:
cache_creation_cost = cast(
float,
_get_cost_per_unit(
model_info,
cache_creation_tiered_key,
cache_creation_cost,
),
)
cache_creation_cost = cast(
float,
_get_cost_per_unit(
model_info,
cache_creation_tiered_key,
cache_creation_cost,
),
)
if cache_creation_1hr_tiered_key in model_info:
cache_creation_cost_above_1hr = cast(
float,
_get_cost_per_unit(
model_info,
cache_creation_1hr_tiered_key,
cache_creation_cost_above_1hr,
),
)
cache_creation_cost_above_1hr = cast(
float,
_get_cost_per_unit(
model_info,
cache_creation_1hr_tiered_key,
cache_creation_cost_above_1hr,
),
)
if cache_read_tiered_key in model_info:
cache_read_cost = cast(
float,
_get_cost_per_unit(
model_info, cache_read_tiered_key, cache_read_cost
),
)
cache_read_cost = cast(
float,
_get_cost_per_unit(
model_info, cache_read_tiered_key, cache_read_cost
),
)
break
except (IndexError, ValueError):

View file

@ -744,6 +744,17 @@ def _count_content_list(
thinking_text = str(c.get("thinking", ""))
if thinking_text:
num_tokens += count_function(thinking_text)
elif c["type"] == "tool_reference":
# Anthropic tool-search reference block: a lightweight pointer to
# a deferred tool, e.g. {"type": "tool_reference", "tool_name": ...}.
# The full tool definition is counted via the `tools` param, so we
# only count the referenced name here. Without this branch,
# token_counter raises on tool-search traffic; on the streaming
# anthropic_messages path that nulls response_cost and causes the
# proxy to drop the SpendLogs row entirely (silent cost undercount).
tool_name = str(c.get("tool_name") or "")
if tool_name:
num_tokens += count_function(tool_name)
else:
content_type = (
c.get("type", type(c).__name__)
@ -752,7 +763,7 @@ def _count_content_list(
)
raise ValueError(
f"Invalid content item type: {content_type}. "
f"Expected str or dict with 'type' field (text, image_url, tool_use, tool_result, thinking)."
f"Expected str or dict with 'type' field (text, image_url, tool_use, tool_result, thinking, tool_reference)."
)
return num_tokens
except Exception as e:

View file

@ -2201,6 +2201,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
inference_geo: str | None = None
if "inference_geo" in _usage and _usage["inference_geo"] is not None:
inference_geo = _usage["inference_geo"]
service_tier = cast(
str | None,
_usage.get("service_tier"), # any-ok: untyped usage dict
)
iterations: list[Any] | None = _usage.get("iterations")
if iterations:
@ -2312,6 +2316,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
),
inference_geo=inference_geo,
speed=speed,
service_tier=service_tier,
)
return usage

View file

@ -18,7 +18,9 @@ if TYPE_CHECKING:
import litellm
def _compute_cache_only_cost(model_info: "ModelInfo", usage: "Usage") -> float:
def _compute_cache_only_cost(
model_info: "ModelInfo", usage: "Usage", service_tier: str | None = None
) -> float:
"""
Return only the cache-related portion of the prompt cost (cache read + cache write).
@ -36,7 +38,9 @@ def _compute_cache_only_cost(model_info: "ModelInfo", usage: "Usage") -> float:
cache_creation_cost,
cache_creation_cost_above_1hr,
cache_read_cost,
) = _get_token_base_cost(model_info=model_info, usage=usage)
) = _get_token_base_cost(
model_info=model_info, usage=usage, service_tier=service_tier
)
cache_cost = float(prompt_tokens_details["cache_hit_tokens"]) * cache_read_cost
@ -56,19 +60,26 @@ def _compute_cache_only_cost(model_info: "ModelInfo", usage: "Usage") -> float:
return cache_cost
def cost_per_token(model: str, usage: "Usage") -> Tuple[float, float]:
def cost_per_token(
model: str, usage: "Usage", service_tier: str | None = None
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
Input:
- model: str, the model name without provider prefix
- usage: LiteLLM Usage block, containing anthropic caching information
- service_tier: the service tier the request was served at (e.g. "priority"),
read from the Anthropic response usage and used to select tier-specific pricing
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
prompt_cost, completion_cost = generic_cost_per_token(
model=model, usage=usage, custom_llm_provider="anthropic"
model=model,
usage=usage,
custom_llm_provider="anthropic",
service_tier=service_tier,
)
# Apply provider_specific_entry multipliers for geo/speed routing
@ -89,7 +100,9 @@ def cost_per_token(model: str, usage: "Usage") -> Tuple[float, float]:
multiplier *= provider_specific_entry.get("fast", 1.0)
if multiplier != 1.0:
cache_cost = _compute_cache_only_cost(model_info=model_info, usage=usage)
cache_cost = _compute_cache_only_cost(
model_info=model_info, usage=usage, service_tier=service_tier
)
prompt_cost = (prompt_cost - cache_cost) * multiplier + cache_cost
completion_cost *= multiplier
except Exception:

View file

@ -2,6 +2,7 @@
Translate from OpenAI's `/v1/chat/completions` to VLLM's `/v1/chat/completions`
"""
import json
from typing import (
Any,
Coroutine,
@ -22,7 +23,9 @@ from litellm.litellm_core_utils.prompt_templates.factory import _parse_mime_type
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionAssistantToolCall,
ChatCompletionFileObject,
ChatCompletionToolCallFunctionChunk,
ChatCompletionVideoObject,
ChatCompletionVideoUrlObject,
)
@ -101,26 +104,18 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
) -> dict:
_tools = non_default_params.pop("tools", None)
if _tools is not None:
# remove 'additionalProperties' from tools
_tools = _remove_additional_properties(_tools)
# remove 'strict' from tools
_tools = _remove_strict_from_schema(_tools)
if isinstance(_tools, list):
_tools = self._convert_custom_tools_to_function_tools(_tools)
if _tools is not None:
non_default_params["tools"] = _tools
# Handle thinking parameter - convert Anthropic-style to OpenAI-style reasoning_effort
# vLLM is OpenAI-compatible, so it understands reasoning_effort, not thinking
# Reference: https://github.com/BerriAI/litellm/issues/19761
thinking = non_default_params.pop("thinking", None)
if thinking is not None and isinstance(thinking, dict):
if thinking.get("type") == "enabled":
# Only convert if reasoning_effort not already set
if "reasoning_effort" not in non_default_params:
budget_tokens = thinking.get("budget_tokens", 0)
# Map budget_tokens to reasoning_effort level
# Same logic as Anthropic adapter (translate_anthropic_thinking_to_reasoning_effort)
if budget_tokens >= 10000:
non_default_params["reasoning_effort"] = "high"
elif budget_tokens >= 5000:
@ -137,20 +132,13 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
api_base = api_base or get_secret_str("HOSTED_VLLM_API_BASE") # type: ignore
api_base = api_base or get_secret_str("HOSTED_VLLM_API_BASE")
dynamic_api_key = (
api_key or get_secret_str("HOSTED_VLLM_API_KEY") or "fake-api-key"
) # vllm does not require an api key
)
return api_base, dynamic_api_key
def _is_video_file(self, content_item: ChatCompletionFileObject) -> bool:
"""
Check if the file is a video
- format: video/<extension>
- file_data: base64 encoded video data
- file_id: infer mp4 from extension
"""
file = content_item.get("file", {})
format = file.get("format")
file_data = file.get("file_data")
@ -205,29 +193,82 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
"""
Support translating:
- video files from file_id or file_data to video_url
- thinking_blocks on assistant messages to content blocks
- thinking_blocks on assistant messages are removed, and content lists
are converted to strings for vLLM compatibility
"""
for message in messages:
if message["role"] == "assistant":
thinking_blocks = message.pop("thinking_blocks", None) # type: ignore
if thinking_blocks:
new_content: list = [
(
{
"type": block["type"],
"thinking": block.get("thinking", ""),
message.pop("thinking_blocks", None)
existing_content = message.get("content")
if isinstance(existing_content, list):
text_parts = []
tool_calls: list[ChatCompletionAssistantToolCall] = []
content_blocks: list[object] = []
has_structured_content = False
for c in existing_content: # any-ok: untyped content
if (
isinstance(c, dict) # any-ok: untyped content
and c.get("type") == "text" # any-ok: untyped content
):
text_parts.append( # any-ok: untyped content
c.get("text", "") # any-ok: untyped content
)
content_blocks.append(c) # any-ok: untyped content
elif (
isinstance(c, dict) # any-ok: untyped content
and c.get("type") == "tool_use" # any-ok: untyped content
):
tool_input = c.get("input", {}) # any-ok: untyped content
tool_calls.append(
ChatCompletionAssistantToolCall(
id=c.get("id"), # any-ok: untyped content
type="function",
function=ChatCompletionToolCallFunctionChunk(
name=c.get("name"), # any-ok: untyped content
arguments=(
tool_input
if isinstance(
tool_input, # any-ok: untyped content
str, # any-ok: untyped content
)
else json.dumps(
tool_input # any-ok: untyped content
)
),
),
)
)
else:
content_blocks.append(c) # any-ok: untyped content
has_structured_content = True
if tool_calls:
existing_tool_calls = message.get("tool_calls")
if isinstance(existing_tool_calls, list):
existing_tool_call_ids = {
tool_call.get("id") # any-ok: untyped content
for tool_call in existing_tool_calls
if isinstance(
tool_call, dict
) # any-ok: untyped content
and tool_call.get("id")
is not None # any-ok: untyped content
}
if block.get("type") == "thinking"
else {"type": block["type"], "data": block.get("data", "")}
)
for block in thinking_blocks
]
existing_content = message.get("content")
if isinstance(existing_content, str):
new_content.append({"type": "text", "text": existing_content})
elif isinstance(existing_content, list):
new_content.extend(existing_content)
message["content"] = new_content # type: ignore
new_tool_calls = [
tool_call
for tool_call in tool_calls
if tool_call.get("id") not in existing_tool_call_ids
]
if new_tool_calls:
message["tool_calls"] = (
existing_tool_calls + new_tool_calls
)
else:
message["tool_calls"] = tool_calls
content_str = "\n".join(text_parts) # any-ok: untyped content
new_content = (
content_blocks if has_structured_content else content_str
)
message["content"] = new_content # type: ignore[typeddict-item]
elif message["role"] == "user":
message_content = message.get("content")
if message_content and isinstance(message_content, list):
@ -243,6 +284,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig):
message_content[idx] = self._convert_file_to_video_url(
content_item
)
if is_async:
return super()._transform_messages(
messages, model, is_async=cast(Literal[True], True)

View file

@ -1,3 +1,4 @@
import json
from typing import List, Optional, Union
from httpx import Headers, Response
@ -107,9 +108,7 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
"""
data = {"model": model, "file": audio_file, **optional_params}
if "response_format" not in data or (
data["response_format"] == "text" or data["response_format"] == "json"
):
if "response_format" not in data:
data["response_format"] = (
"verbose_json" # ensures 'duration' is received - used for cost calculation
)
@ -133,10 +132,11 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
) -> TranscriptionResponse:
try:
raw_response_json = raw_response.json()
except Exception as e:
raise ValueError(
f"Error transforming response to json: {str(e)}\nResponse: {raw_response.text}"
)
except json.JSONDecodeError:
content_type = raw_response.headers.get("content-type", "").lower()
if "application/json" in content_type:
raise
return TranscriptionResponse(text=raw_response.text)
if any(
key in raw_response_json

View file

@ -50,11 +50,15 @@ class OpenrouterConfig(OpenAIGPTConfig):
def map_openai_params(
self,
non_default_params: dict,
non_default_params: dict[str, object],
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
# OpenRouter expects "xhigh" instead of "max" for reasoning_effort.
if non_default_params.get("reasoning_effort") == "max":
non_default_params = {**non_default_params, "reasoning_effort": "xhigh"}
mapped_openai_params = super().map_openai_params(
non_default_params, optional_params, model, drop_params
)

View file

@ -2528,6 +2528,100 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"azure_ai/gpt-5.5": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
"cache_read_input_token_cost_priority": 1e-06,
"cache_read_input_token_cost_above_272k_tokens_priority": 2e-06,
"input_cost_per_token": 5e-06,
"input_cost_per_token_above_272k_tokens": 1e-05,
"input_cost_per_token_priority": 1e-05,
"input_cost_per_token_above_272k_tokens_priority": 2e-05,
"litellm_provider": "azure_ai",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3e-05,
"output_cost_per_token_above_272k_tokens": 4.5e-05,
"output_cost_per_token_priority": 6e-05,
"output_cost_per_token_above_272k_tokens_priority": 9e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.5",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure_ai/gpt-5.5-2026-04-23": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
"cache_read_input_token_cost_priority": 1e-06,
"cache_read_input_token_cost_above_272k_tokens_priority": 2e-06,
"input_cost_per_token": 5e-06,
"input_cost_per_token_above_272k_tokens": 1e-05,
"input_cost_per_token_priority": 1e-05,
"input_cost_per_token_above_272k_tokens_priority": 2e-05,
"litellm_provider": "azure_ai",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3e-05,
"output_cost_per_token_above_272k_tokens": 4.5e-05,
"output_cost_per_token_priority": 6e-05,
"output_cost_per_token_above_272k_tokens_priority": 9e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.5",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure_ai/gpt-5.4": {
"cache_read_input_token_cost": 2.5e-07,
"cache_read_input_token_cost_above_272k_tokens": 5e-07,
@ -10068,6 +10162,8 @@
},
"claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -10097,6 +10193,8 @@
},
"claude-sonnet-4-5-20250929": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -10127,6 +10225,7 @@
},
"claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "anthropic",
@ -10155,6 +10254,8 @@
},
"claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -25103,6 +25204,21 @@
"supports_tool_choice": true,
"supports_vision": true
},
"mistral/mistral-medium-3-5": {
"input_cost_per_token": 1.5e-06,
"litellm_provider": "mistral",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"mistral/mistral-small": {
"input_cost_per_token": 1e-07,
"litellm_provider": "mistral",
@ -42456,4 +42572,105 @@
"supports_reasoning": true,
"source": "https://serverless.tensormesh.ai/v1/models/openrouter"
}
}
,
"deepseek-v4-flash": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2.8e-09,
"input_cost_per_token": 1.4e-07,
"input_cost_per_token_cache_hit": 2.8e-09,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false
},
"deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 4.35e-07,
"input_cost_per_token_cache_hit": 3.625e-09,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 8.7e-07,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false
},
"deepseek/deepseek-v4-flash": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2.8e-09,
"input_cost_per_token": 1.4e-07,
"input_cost_per_token_cache_hit": 2.8e-09,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false
},
"deepseek/deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 4.35e-07,
"input_cost_per_token_cache_hit": 3.625e-09,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 8.7e-07,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false
}
}

View file

@ -1421,8 +1421,11 @@ class MCPServerManager:
"No allowed MCP Servers found for user api key auth."
)
return list(combined_servers)
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}.")
except Exception: # noqa: BLE001
verbose_logger.exception(
"Failed to get allowed MCP servers; team-level object_permission "
"grants may be dropped. Falling back to global servers only."
)
return allow_all_server_ids
async def resolve_toolset_tool_permissions(

View file

@ -1036,7 +1036,14 @@ if MCP_AVAILABLE:
allowed_mcp_servers: List[MCPServer],
) -> List[MCPServer]:
"""
Get the filtered MCP servers from the MCP server names
Get the filtered MCP servers from the MCP server names.
Fails closed when ``mcp_servers`` is explicitly provided (path- or
header-derived) but none of the names resolve to a server alias or
access group the caller can access. The previous behavior returned
the full ``allowed_mcp_servers`` set, which silently widened scope
when a client targeted ``/mcp/<unknown>/`` and made URL/header
namespacing appear to work when it did not.
"""
filtered_server: dict[str, MCPServer] = {}
@ -1076,6 +1083,17 @@ if MCP_AVAILABLE:
if filtered_server:
return list(filtered_server.values())
if mcp_servers is not None:
# Caller asked for a specific scope but nothing resolved. Fail
# closed so URL/header namespacing cannot silently fall back to
# the caller's full allowed-server set.
verbose_logger.debug(
"MCP scope filter resolved to no servers for requested names %s; "
"returning empty list (fail-closed).",
mcp_servers,
)
return []
return allowed_mcp_servers
def _tool_name_matches(tool_name: str, filter_list: List[str]) -> bool:

View file

@ -10,6 +10,7 @@ from litellm.repositories.object_permission_repository import ObjectPermissionRe
from litellm.router import Router
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
from litellm.types.router import CredentialLiteLLMParams, LiteLLM_Params
from litellm.types.utils import LlmProviders
from litellm.utils import get_valid_models
_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields)
@ -308,10 +309,21 @@ def get_known_models_from_wildcard(
# add model prefix to wildcard models
wildcard_models = [f"{model_prefix}{model}" for model in wildcard_models]
known_providers = {provider.value for provider in LlmProviders}
suffix_appended_wildcard_models = []
for model in wildcard_models:
if not model.startswith(wildcard_provider_prefix):
model = f"{wildcard_provider_prefix}/{model}"
# `get_provider_models` returns provider-prefixed ids (e.g. "ollama/gemma3:1b").
# When the wildcard uses a custom prefix (e.g. "ollama_server1/*" to distinguish
# multiple instances), replace that existing provider prefix instead of stacking
# both, which would otherwise yield an uncallable "ollama_server1/ollama/gemma3:1b".
# Only strip the leading segment when it is a known provider, so ids whose first
# segment is an org rather than a provider (e.g. "meta-llama/Llama-3-8B") keep it.
leading, sep, model_suffix = model.partition("/")
if sep and leading in known_providers:
model = f"{wildcard_provider_prefix}/{model_suffix}"
else:
model = f"{wildcard_provider_prefix}/{model}"
suffix_appended_wildcard_models.append(model)
return suffix_appended_wildcard_models or []

View file

@ -31,6 +31,7 @@ from litellm.constants import (
DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE,
DEFAULT_MAX_RECURSE_DEPTH,
LITELLM_DETAILED_TIMING,
LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED,
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG,
STREAM_SSE_DATA_PREFIX,
)
@ -67,7 +68,12 @@ if TYPE_CHECKING:
else:
ProxyConfig = Any
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.types.utils import ModelResponse, ModelResponseStream, Usage
from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
StandardLoggingPayloadErrorInformation,
Usage,
)
# Datadog streaming spans are a no-op when ddtrace is not enabled, but the
# ``with tracer.trace(...)`` context manager still allocates a NullSpan and
@ -77,6 +83,77 @@ from litellm.types.utils import ModelResponse, ModelResponseStream, Usage
_DD_STREAMING_TRACE_ENABLED = not isinstance(tracer, NullTracer)
_CLIENT_DISCONNECTED_ERROR_INFORMATION: StandardLoggingPayloadErrorInformation = {
"error_code": str(LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED),
"error_message": "Client disconnected the request",
"error_class": "ClientDisconnected",
}
def _apply_client_disconnect_metadata(target_metadata: dict[str, object]) -> None:
target_metadata["client_disconnected"] = True
target_metadata["error_information"] = dict(_CLIENT_DISCONNECTED_ERROR_INFORMATION)
async def _record_streaming_client_disconnect_if_needed(
request: Request | None,
request_data: dict,
client_disconnected: bool = False,
) -> bool:
if not client_disconnected:
if request is None:
return False
try:
disconnected = await request.is_disconnected()
except Exception: # noqa: BLE001
return False
if not disconnected:
return False
logging_obj = request_data.get("litellm_logging_obj") # any-ok: untyped request
if logging_obj is not None: # any-ok: untyped request
litellm_params = (
logging_obj.model_call_details.setdefault( # any-ok: untyped request
"litellm_params", {}
)
)
_apply_client_disconnect_metadata(
litellm_params.setdefault("metadata", {}) # any-ok: untyped request
)
_apply_client_disconnect_metadata(
logging_obj.model_call_details.setdefault( # any-ok: untyped request
"metadata", {}
)
)
_apply_client_disconnect_metadata(
request_data.setdefault("metadata", {}) # any-ok: untyped request
)
litellm_params = request_data.setdefault( # any-ok: untyped request
"litellm_params", {} # any-ok: untyped request
)
_apply_client_disconnect_metadata(
litellm_params.setdefault("metadata", {}) # any-ok: untyped request
)
verbose_proxy_logger.debug(
"Recorded streaming client disconnect with error_code=499 for litellm_call_id=%s",
request_data.get("litellm_call_id"), # any-ok: untyped request
)
return True
async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None:
pending_tasks = [task for task in tasks if not task.done()] # any-ok: untyped task
for task in pending_tasks: # any-ok: untyped task
task.cancel() # any-ok: untyped task
for task in pending_tasks: # any-ok: untyped task
try:
await task # any-ok: untyped request
except (asyncio.CancelledError, Exception): # noqa: BLE001
pass
def _serialize_http_exception_detail(
detail: Any,
) -> Tuple[str, Optional[dict]]:
@ -242,20 +319,6 @@ def _extract_error_from_sse_chunk(event_line: Union[str, bytes]) -> dict:
return default_error
async def _aclose_upstream_response(response: Any) -> None:
"""Release the upstream HTTP connection when a stream ends for any
reason, including client disconnect. Mirrors the finally block of
async_data_generator in proxy_server.py."""
with anyio.CancelScope(shield=True):
if hasattr(response, "aclose"):
try:
await response.aclose()
except BaseException as e:
verbose_proxy_logger.debug(
"error closing upstream response stream: %s", e
)
class _UpstreamClosingStreamingResponse(StreamingResponse):
"""StreamingResponse that always closes its body iterator and the wrapped
upstream generator.
@ -1338,19 +1401,24 @@ class ProxyBaseLLMRequestProcessing:
user_model=user_model,
user_api_key_dict=user_api_key_dict,
)
tasks.append(llm_call)
llm_call_task = asyncio.create_task(llm_call) # any-ok: untyped task
tasks.append(llm_call_task) # any-ok: untyped task
# wait for call to end
llm_responses = asyncio.gather(
*tasks
) # run the moderation check in parallel to the actual llm api call
if general_settings.get("cancel_on_disconnect", False):
responses = await _await_llm_call_cancelling_on_disconnect(
request, llm_responses
)
else:
responses = await llm_responses
try:
if general_settings.get( # any-ok: untyped request
"cancel_on_disconnect", False
):
responses = await _await_llm_call_cancelling_on_disconnect( # any-ok: untyped request
request, llm_responses # any-ok: untyped task
)
else:
responses = await llm_responses # any-ok: untyped request
finally:
await _cancel_pending_gather_tasks(tasks) # any-ok: untyped task
response = responses[1]
@ -1526,6 +1594,7 @@ class ProxyBaseLLMRequestProcessing:
user_api_key_dict=user_api_key_dict,
request_data=self.data,
proxy_logging_obj=proxy_logging_obj,
request=request,
)
)
return await create_response(
@ -1539,6 +1608,7 @@ class ProxyBaseLLMRequestProcessing:
response=response,
user_api_key_dict=user_api_key_dict,
request_data=self.data,
request=request,
)
if route_type == "aresponses":
# Streaming /v1/responses returns here without
@ -2295,6 +2365,13 @@ class ProxyBaseLLMRequestProcessing:
self._apply_router_cooldown_retry_after(headers, e)
if isinstance(e, ProxyException):
e.headers = {
**e.headers,
**{k: v if isinstance(v, str) else str(v) for k, v in headers.items()},
}
raise e
if isinstance(e, HTTPException):
raw_detail = getattr(e, "detail", str(e))
message, structured_fields = _serialize_http_exception_detail(raw_detail)
@ -2383,6 +2460,41 @@ class ProxyBaseLLMRequestProcessing:
else:
return chunk
@staticmethod
async def _finalize_streaming_generator_cleanup(
request: Request | None,
request_data: dict,
response: Any,
stream_completed: bool = False,
client_disconnected: bool = False,
) -> None:
with anyio.CancelScope(shield=True):
should_record_client_disconnect = client_disconnected or (
not stream_completed
)
recorded_client_disconnect = False
if should_record_client_disconnect:
recorded_client_disconnect = (
await _record_streaming_client_disconnect_if_needed(
request,
request_data, # any-ok: untyped request
client_disconnected, # any-ok: untyped request
)
)
if recorded_client_disconnect:
ProxyLogging._fire_deferred_stream_logging(
request_data # any-ok: untyped request
)
if hasattr(response, "aclose"): # any-ok: untyped request
try:
await response.aclose() # any-ok: untyped request
except BaseException as e: # noqa: BLE001
verbose_proxy_logger.debug(
"async_streaming_data_generator: error closing response stream: %s",
e,
)
@staticmethod
async def async_streaming_data_generator(
response: Any,
@ -2392,6 +2504,7 @@ class ProxyBaseLLMRequestProcessing:
*,
serialize_chunk: StreamChunkSerializer,
serialize_error: StreamErrorSerializer,
request: Request | None = None,
) -> AsyncGenerator[str, None]:
"""
Shared streaming data generator: runs proxy iterator hook, per-chunk hook,
@ -2416,6 +2529,8 @@ class ProxyBaseLLMRequestProcessing:
and not cost_injection_enabled
)
debug_enabled = verbose_proxy_logger.isEnabledFor(logging.DEBUG)
stream_completed = False
client_disconnected = False
try:
str_so_far = ""
async for (
@ -2463,6 +2578,7 @@ class ProxyBaseLLMRequestProcessing:
)
)
yield serialize_chunk(chunk)
stream_completed = True
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit
# are BaseException and bypass the success/failure logging
@ -2470,9 +2586,11 @@ class ProxyBaseLLMRequestProcessing:
# release it here. This is the outermost generator Starlette closes
# on disconnect, so the nested iterator hook (which only sees
# GeneratorExit on GC) cannot own the refund.
proxy_logging_obj._release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
if not stream_completed:
proxy_logging_obj._release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
client_disconnected = True
raise
except Exception as e:
verbose_proxy_logger.exception(
@ -2501,9 +2619,16 @@ class ProxyBaseLLMRequestProcessing:
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
)
stream_completed = True
yield serialize_error(proxy_exception)
finally:
await _aclose_upstream_response(response)
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
request=request,
request_data=request_data, # any-ok: untyped request
response=response, # any-ok: untyped request
stream_completed=stream_completed,
client_disconnected=client_disconnected,
)
@staticmethod
def async_sse_data_generator(
@ -2511,6 +2636,7 @@ class ProxyBaseLLMRequestProcessing:
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
proxy_logging_obj: ProxyLogging,
request: Request | None = None,
) -> AsyncGenerator[str, None]:
"""
Anthropic /messages and Google /generateContent streaming data generator require SSE events.
@ -2529,6 +2655,7 @@ class ProxyBaseLLMRequestProcessing:
serialize_error=lambda proxy_exc: (
f"{STREAM_SSE_DATA_PREFIX}{json.dumps({'error': proxy_exc.to_dict()})}\n\n"
),
request=request,
)
@staticmethod

View file

@ -0,0 +1,191 @@
model_list:
# ---------- Anthropic native ----------
- model_name: anthropic-haiku-4-5
litellm_params:
model: anthropic/claude-haiku-4-5
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: anthropic-sonnet-4-5
litellm_params:
model: anthropic/claude-sonnet-4-5
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: anthropic-opus-4-5
litellm_params:
model: anthropic/claude-opus-4-5
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: anthropic-sonnet-4-6
litellm_params:
model: anthropic/claude-sonnet-4-6
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: anthropic-opus-4-6
litellm_params:
model: anthropic/claude-opus-4-6
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: anthropic-opus-4-7
litellm_params:
model: anthropic/claude-opus-4-7
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: anthropic-opus-4-8
litellm_params:
model: anthropic/claude-opus-4-8
api_key: os.environ/ANTHROPIC_API_KEY
# ---------- Bedrock Invoke ----------
- model_name: bedrock-invoke-haiku-4-5
litellm_params:
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
aws_region_name: us-east-1
- model_name: bedrock-invoke-sonnet-4-5
litellm_params:
model: bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0
aws_region_name: us-east-1
- model_name: bedrock-invoke-opus-4-5
litellm_params:
model: bedrock/us.anthropic.claude-opus-4-5-20251101-v1:0
aws_region_name: us-east-1
- model_name: bedrock-invoke-sonnet-4-6
litellm_params:
model: bedrock/us.anthropic.claude-sonnet-4-6
aws_region_name: us-east-1
- model_name: bedrock-invoke-opus-4-6
litellm_params:
model: bedrock/us.anthropic.claude-opus-4-6-v1
aws_region_name: us-east-1
- model_name: bedrock-invoke-opus-4-7
litellm_params:
model: bedrock/global.anthropic.claude-opus-4-7
aws_region_name: us-east-1
- model_name: bedrock-invoke-opus-4-8
litellm_params:
model: bedrock/global.anthropic.claude-opus-4-8
aws_region_name: us-east-1
# ---------- Bedrock Converse ----------
- model_name: bedrock-converse-haiku-4-5
litellm_params:
model: bedrock/converse/us.anthropic.claude-haiku-4-5-20251001-v1:0
aws_region_name: us-east-1
- model_name: bedrock-converse-sonnet-4-5
litellm_params:
model: bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0
aws_region_name: us-east-1
- model_name: bedrock-converse-opus-4-5
litellm_params:
model: bedrock/converse/us.anthropic.claude-opus-4-5-20251101-v1:0
aws_region_name: us-east-1
- model_name: bedrock-converse-sonnet-4-6
litellm_params:
model: bedrock/converse/us.anthropic.claude-sonnet-4-6
aws_region_name: us-east-1
- model_name: bedrock-converse-opus-4-6
litellm_params:
model: bedrock/converse/us.anthropic.claude-opus-4-6-v1
aws_region_name: us-east-1
- model_name: bedrock-converse-opus-4-7
litellm_params:
model: bedrock/converse/global.anthropic.claude-opus-4-7
aws_region_name: us-east-1
- model_name: bedrock-converse-opus-4-8
litellm_params:
model: bedrock/converse/global.anthropic.claude-opus-4-8
aws_region_name: us-east-1
# ---------- Vertex AI (Anthropic on Vertex) ----------
- model_name: vertex-haiku-4-5
litellm_params:
model: vertex_ai/claude-haiku-4-5@20251001
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
- model_name: vertex-sonnet-4-5
litellm_params:
model: vertex_ai/claude-sonnet-4-5@20250929
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
- model_name: vertex-opus-4-5
litellm_params:
model: vertex_ai/claude-opus-4-5@20251101
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
- model_name: vertex-sonnet-4-6
litellm_params:
model: vertex_ai/claude-sonnet-4-6
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
- model_name: vertex-opus-4-6
litellm_params:
model: vertex_ai/claude-opus-4-6
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
- model_name: vertex-opus-4-7
litellm_params:
model: vertex_ai/claude-opus-4-7
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
- model_name: vertex-opus-4-8
litellm_params:
model: vertex_ai/claude-opus-4-8
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
# ---------- Gemini Enterprise Agent Platform ----------
- model_name: gemini-claude-code
litellm_params:
model: vertex_ai/gemini-2.5-pro
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
vertex_credentials: os.environ/GEMINI_CLAUDE_CODE_VERTEX_CREDENTIALS
extra_body:
labels:
workload: claude-code
source: litellm
environment: internal
reconciliation_group: claude-code-gemini
# ---------- Azure AI Foundry (Anthropic on Azure) ----------
- model_name: azure-haiku-4-5
litellm_params:
model: azure_ai/claude-haiku-4-5
api_base: os.environ/AZURE_AI_API_BASE
api_key: os.environ/AZURE_AI_API_KEY
- model_name: azure-sonnet-4-5
litellm_params:
model: azure_ai/claude-sonnet-4-5
api_base: os.environ/AZURE_AI_API_BASE
api_key: os.environ/AZURE_AI_API_KEY
- model_name: azure-opus-4-5
litellm_params:
model: azure_ai/claude-opus-4-5
api_base: os.environ/AZURE_AI_API_BASE
api_key: os.environ/AZURE_AI_API_KEY
- model_name: azure-sonnet-4-6
litellm_params:
model: azure_ai/claude-sonnet-4-6
api_base: os.environ/AZURE_AI_API_BASE
api_key: os.environ/AZURE_AI_API_KEY
- model_name: azure-opus-4-6
litellm_params:
model: azure_ai/claude-opus-4-6
api_base: os.environ/AZURE_AI_API_BASE
api_key: os.environ/AZURE_AI_API_KEY
- model_name: azure-opus-4-7
litellm_params:
model: azure_ai/claude-opus-4-7
api_base: os.environ/AZURE_AI_API_BASE
api_key: os.environ/AZURE_AI_API_KEY
- model_name: azure-opus-4-8
litellm_params:
model: azure_ai/claude-opus-4-8
api_base: os.environ/AZURE_AI_API_BASE
api_key: os.environ/AZURE_AI_API_KEY
# ---------- OpenAI ----------
- model_name: gpt-5.5
litellm_params:
model: openai/gpt-5.5
api_key: os.environ/OPENAI_API_KEY
general_settings:
master_key: sk-1234
litellm_settings:
drop_params: True
telemetry: False

View file

@ -25,6 +25,16 @@ async def get_ui_config():
or general_settings.get("auto_redirect_ui_login_to_sso", False) is True
)
admin_ui_disabled = os.getenv("DISABLE_ADMIN_UI", "false").lower() == "true"
hide_default_credentials_hint = bool( # any-ok: untyped settings
os.getenv( # any-ok: untyped settings
"LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false"
).lower()
== "true"
or general_settings.get( # any-ok: untyped settings
"hide_default_credentials_hint", False
)
is True
)
sso_configured = _has_user_setup_sso()
@ -38,6 +48,7 @@ async def get_ui_config():
auto_redirect_to_sso=sso_configured and auto_redirect_ui_login_to_sso,
admin_ui_disabled=admin_ui_disabled,
sso_configured=sso_configured,
hide_default_credentials_hint=hide_default_credentials_hint, # any-ok: untyped settings
is_control_plane=is_control_plane,
workers=proxy_config.worker_registry if is_control_plane else [],
)

View file

@ -107,6 +107,7 @@ async def google_stream_generate_content(
data["stream"] = True
# google-genai SDK (?alt=sse) must not receive OpenAI's data: [DONE] terminator.
data["_litellm_skip_openai_stream_done"] = True
data["_litellm_raw_sse_stream"] = True # any-ok: untyped request
processor = ProxyBaseLLMRequestProcessing(data=data)
try:

View file

@ -9,7 +9,6 @@ import json
import os
from typing import TYPE_CHECKING, Any, AsyncGenerator, Optional, Type, Union
from fastapi import HTTPException
from pydantic import BaseModel
from websockets.asyncio.client import ClientConnection, connect
@ -21,7 +20,7 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import (
apply_redacted_messages_back,
build_inspection_messages,
@ -129,6 +128,16 @@ class AimGuardrail(CustomGuardrail):
verbose_proxy_logger.error(f"Aim: {action_type} action")
return data
@staticmethod
def _rejection(message: str, *, openai_code: str | None = None) -> ProxyException:
return ProxyException(
message=message,
type="invalid_request_error",
param=None,
code=400,
openai_code=openai_code,
)
def _handle_block_action(self, analysis_result: Any, required_action: Any) -> None:
detection_message = required_action.get("detection_message", None)
verbose_proxy_logger.info(
@ -136,7 +145,7 @@ class AimGuardrail(CustomGuardrail):
policies=list(analysis_result["policy_drill_down"].keys()),
),
)
raise HTTPException(status_code=400, detail=detection_message)
raise self._rejection(detection_message, openai_code="content_policy_violation")
def _anonymize_request(self, res: Any, data: dict) -> dict:
verbose_proxy_logger.info("Aim: anonymize action")
@ -148,14 +157,11 @@ class AimGuardrail(CustomGuardrail):
# parts from a multimodal request — degrade to block so the
# multimodal payload is never silently rewritten.
if has_non_string_content(data):
raise HTTPException(
status_code=400,
detail=(
"Aim: anonymize action requested for multimodal input "
"but mask-in-place would drop non-text parts. Send the "
"request with plain string content to use anonymize, "
"or rely on block-mode policies."
),
raise self._rejection(
"Aim: anonymize action requested for multimodal input "
"but mask-in-place would drop non-text parts. Send the "
"request with plain string content to use anonymize, "
"or rely on block-mode policies."
)
redacted_messages = [
{
@ -287,9 +293,9 @@ class AimGuardrail(CustomGuardrail):
if aim_output_guardrail_result and aim_output_guardrail_result.get(
"detection_message"
):
raise HTTPException(
status_code=400,
detail=aim_output_guardrail_result.get("detection_message"),
raise self._rejection(
aim_output_guardrail_result.get("detection_message"),
openai_code="content_policy_violation",
)
if aim_output_guardrail_result and aim_output_guardrail_result.get(
"redacted_output"

View file

@ -741,6 +741,18 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
For multiple messages in /chat/completions, we'll need to call them in parallel.
"""
# Respect the configured event hook. In `logging_only` mode (and any config that
# excludes pre_call) the live request must not be masked - masking is applied to a
# copy at logging time via `async_logging_hook`. Without this gate the request sent
# to the model would carry anonymization tokens and the response would echo them.
if (
self.should_run_guardrail(
data=data, # any-ok: untyped request
event_type=GuardrailEventHooks.pre_call, # any-ok: untyped request
)
is not True
):
return data # any-ok: untyped request
try:
content_safety = data.get("content_safety", None)

View file

@ -162,6 +162,8 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS = (
"secret_fields",
"_guardrail_pipelines",
"_pipeline_managed_guardrails",
"client_disconnected",
"error_information",
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
)

View file

@ -1,5 +1,5 @@
import asyncio
from datetime import datetime, timedelta
from datetime import datetime
from types import SimpleNamespace
from typing import Any, Callable, Dict, List, Optional, Set, Tuple, Union
@ -390,38 +390,24 @@ def _adjust_dates_for_timezone(
timezone_offset_minutes: Optional[int],
) -> Tuple[str, str]:
"""
Adjust date range to account for timezone differences.
Pass-through for the local date range; the timezone offset is intentionally ignored here.
The database stores dates in UTC. When a user in a different timezone
selects a local date range, we need to expand the UTC query range to
capture all records that fall within their local date range.
The aggregation table (e.g. LiteLLM_DailyUserSpend) stores spend in whole-UTC-day
buckets keyed on date as YYYY-MM-DD. Any conversion from a local date range to a
UTC date range using only date arithmetic must round to whole UTC days, allowing up
to 24h of slop at each boundary. The previous implementation expanded the SQL range
by an extra full UTC day on whichever side the offset pointed, which pulled in 24h
of unrelated bucket data per boundary and produced approximately 100% over-counting
on single-day queries (e.g. IST May 29 returning UTC May 28 + UTC May 29 in full).
Sums of single-day queries then exceeded the equivalent multi-day aggregate, which
is mathematically impossible.
Args:
start_date: Start date in YYYY-MM-DD format (user's local date)
end_date: End date in YYYY-MM-DD format (user's local date)
timezone_offset_minutes: Minutes behind UTC (positive = west of UTC)
This matches JavaScript's Date.getTimezoneOffset() convention.
For example: PST = +480 (8 hours * 60 = 480 minutes behind UTC)
Returns:
Tuple of (adjusted_start_date, adjusted_end_date) in YYYY-MM-DD format
Treating the local date as the UTC date trades a small one-time boundary slop for
correct, monotonic, additive results across single-day and multi-day queries. A
later fix can introduce hour-level buckets or pro-rata weighting on adjacent UTC
days; both require data the current schema does not store.
"""
if timezone_offset_minutes is None or timezone_offset_minutes == 0:
return start_date, end_date
start = datetime.strptime(start_date, "%Y-%m-%d")
end = datetime.strptime(end_date, "%Y-%m-%d")
if timezone_offset_minutes > 0:
# West of UTC (Americas): local evening extends into next UTC day
# e.g., Feb 4 23:59 PST = Feb 5 07:59 UTC
end = end + timedelta(days=1)
else:
# East of UTC (Asia/Europe): local morning starts in previous UTC day
# e.g., Feb 4 00:00 IST = Feb 3 18:30 UTC
start = start - timedelta(days=1)
return start.strftime("%Y-%m-%d"), end.strftime("%Y-%m-%d")
return start_date, end_date
def _build_where_conditions(

View file

@ -18,6 +18,7 @@ import os
import re
import secrets
import traceback
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, cast
@ -59,6 +60,9 @@ from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_k
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks
from litellm.proxy.hooks.model_max_budget_limiter import (
VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX,
)
from litellm.proxy.management_endpoints.common_utils import (
_check_passthrough_routes_caller_permission,
_is_user_org_admin_for_team,
@ -3225,6 +3229,69 @@ async def delete_key_fn(
raise handle_exception_on_proxy(e)
async def _get_model_max_budget_current_spend(
api_key_hash: str,
model: str,
budget_config: BudgetConfig,
user_api_key_cache: UserApiKeyCache,
) -> float:
virtual_key_model_spend_cache_key = (
f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:"
f"{api_key_hash}:{model}:{budget_config.budget_duration}"
)
current_spend: float | None = (
await user_api_key_cache.async_get_cache( # any-ok: untyped dump
key=virtual_key_model_spend_cache_key,
)
)
if current_spend is None:
model_without_prefix = model.split("/")[-1] if "/" in model else model
virtual_key_model_spend_cache_key = (
f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:"
f"{api_key_hash}:{model_without_prefix}:{budget_config.budget_duration}"
)
current_spend = (
await user_api_key_cache.async_get_cache( # any-ok: untyped dump
key=virtual_key_model_spend_cache_key,
)
)
try:
return float(current_spend or 0.0) # any-ok: untyped dump
except (TypeError, ValueError):
return 0.0
async def _build_model_max_budget_usage(
api_key_hash: str,
model_max_budget: Mapping[str, Mapping[str, object]],
user_api_key_cache: UserApiKeyCache | None,
) -> dict[str, dict[str, object]]:
if user_api_key_cache is None or not model_max_budget:
return {}
result: dict[str, dict[str, object]] = {}
for model, budget_info in model_max_budget.items():
try:
budget_config = BudgetConfig.model_validate(budget_info)
if budget_config.budget_duration is None:
continue
duration_in_seconds(budget_config.budget_duration)
except Exception: # noqa: BLE001
continue
spend = await _get_model_max_budget_current_spend(
api_key_hash=api_key_hash,
model=model,
budget_config=budget_config,
user_api_key_cache=user_api_key_cache,
)
result[model] = {
"current_spend": round(spend, 4),
"budget_limit": budget_config.max_budget,
"time_period": budget_config.budget_duration,
}
return result
@router.post(
"/v2/key/info",
tags=["key management"],
@ -3252,7 +3319,7 @@ async def info_key_fn_v2(
-d {"keys": ["sk-1", "sk-2", "sk-3"]}
```
"""
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
try:
if prisma_client is None:
@ -3298,7 +3365,29 @@ async def info_key_fn_v2(
k_dict = k.model_dump()
except Exception:
k_dict = k.dict()
k_dict.pop("token", None)
k_token_hash = k_dict.pop("token", None) # any-ok: untyped dump
model_max_budget = (
k_dict.get("model_max_budget") or {} # any-ok: untyped dump
)
budget_table = (
k_dict.get("litellm_budget_table") or {} # any-ok: untyped dump
)
if not model_max_budget and isinstance( # any-ok: untyped dump
budget_table, dict # any-ok: untyped dump
):
model_max_budget = (
budget_table.get("model_max_budget") or {} # any-ok: untyped dump
)
if model_max_budget and k_token_hash: # any-ok: untyped dump
k_dict["model_max_budget_usage"] = ( # any-ok: untyped dump
await _build_model_max_budget_usage( # any-ok: untyped dump
api_key_hash=k_token_hash, # any-ok: untyped dump
model_max_budget=model_max_budget, # any-ok: untyped dump
user_api_key_cache=user_api_key_cache,
)
)
filtered_key_info.append(k_dict)
return {"key": data.keys, "info": filtered_key_info}
@ -3336,7 +3425,7 @@ async def info_key_fn(
-H "Authorization: Bearer sk-test-example-key-123"
```
"""
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
try:
if prisma_client is None:
@ -3381,7 +3470,28 @@ async def info_key_fn(
except Exception:
# if using pydantic v1
key_info = key_info.dict()
key_info.pop("token")
key_token_hash = key_info.pop("token") # any-ok: untyped dump
model_max_budget = (
key_info.get("model_max_budget") or {} # any-ok: untyped dump
)
budget_table = (
key_info.get("litellm_budget_table") or {} # any-ok: untyped dump
)
if not model_max_budget and isinstance( # any-ok: untyped dump
budget_table, dict # any-ok: untyped dump
):
model_max_budget = (
budget_table.get("model_max_budget") or {} # any-ok: untyped dump
)
if model_max_budget and key_token_hash: # any-ok: untyped dump
key_info["model_max_budget_usage"] = ( # any-ok: untyped dump
await _build_model_max_budget_usage( # any-ok: untyped dump
api_key_hash=key_token_hash, # any-ok: untyped dump
model_max_budget=model_max_budget, # any-ok: untyped dump
user_api_key_cache=user_api_key_cache,
)
)
# Attach object_permission if object_permission_id is set
key_info = await attach_object_permission_to_dict(key_info, prisma_client)

View file

@ -888,16 +888,26 @@ async def proxy_startup_event(app: FastAPI):
asyncio.create_task(_run_pw_migration())
## use_redis_transaction_buffer: fall back to a standalone Redis (REDIS_* env)
## when the proxy cache backend is not Redis ##
transaction_buffer_redis_cache = redis_usage_cache
if transaction_buffer_redis_cache is None:
transaction_buffer_redis_cache = (
ProxyStartupEvent._get_transaction_buffer_redis_cache(
general_settings=general_settings # any-ok: untyped stream
)
)
ProxyStartupEvent._initialize_startup_logging(
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
redis_usage_cache=redis_usage_cache,
redis_usage_cache=transaction_buffer_redis_cache,
)
## Validate use_redis_transaction_buffer requires Redis cache ##
ProxyStartupEvent._validate_redis_transaction_buffer_config(
general_settings=general_settings,
redis_usage_cache=redis_usage_cache,
redis_usage_cache=transaction_buffer_redis_cache,
)
## SEMANTIC TOOL FILTER ##
@ -7022,10 +7032,33 @@ def _format_streaming_sse_chunk(chunk: Union[str, bytes]) -> Union[str, bytes]:
return f"data: {chunk}\n\n"
_SSE_FRAME_DELIMITERS = ("\r\n\r\n", "\n\n", "\r\r")
_MAX_RAW_SSE_BUFFER_CHARS = 8 * 1024 * 1024
def _pop_complete_sse_frame(buffer: str) -> tuple[str | None, str]:
delimiter_positions = [
(position, delimiter)
for delimiter in _SSE_FRAME_DELIMITERS
if (position := buffer.find(delimiter)) != -1
]
if not delimiter_positions:
return None, buffer
position, delimiter = min(delimiter_positions, key=lambda item: item[0])
frame_end = position + len(delimiter)
return buffer[:frame_end], buffer[frame_end:]
async def async_data_generator(
response, user_api_key_dict: UserAPIKeyAuth, request_data: dict
response,
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
request: Request | None = None,
):
verbose_proxy_logger.debug("inside generator")
stream_completed = False
client_disconnected = False
try:
error_message: Optional[str] = None
requested_model_from_client = _get_client_requested_model_for_streaming(
@ -7047,6 +7080,10 @@ async def async_data_generator(
# happened to ship a streaming-iterator override (the default).
needs_iterator_wrap = proxy_logging_obj.needs_iterator_wrap()
needs_per_chunk_hook = proxy_logging_obj.needs_per_chunk_streaming_hook()
is_raw_sse_stream = bool(
request_data.get("_litellm_raw_sse_stream") # any-ok: untyped stream
)
raw_sse_buffer = ""
if needs_iterator_wrap:
stream_iterator = proxy_logging_obj.async_post_call_streaming_iterator_hook(
@ -7077,14 +7114,38 @@ async def async_data_generator(
if isinstance(chunk, BaseModel):
chunk = _serialize_streaming_chunk(chunk)
elif isinstance(chunk, bytes):
# Some upstream streaming iterators (e.g. AsyncGoogleGenAIGenerateContentStreamingIterator
# for /v1beta/.../streamGenerateContent) yield raw SSE bytes from Gemini.
# Decode to str so the f-string below does not emit a Python b'...' literal,
# and pass already-formatted SSE through unchanged to avoid double "data:" prefix.
chunk = chunk.decode("utf-8", errors="replace")
if chunk.startswith(("data:", "event:", ":")):
yield chunk if chunk.endswith("\n\n") else chunk + "\n\n"
if is_raw_sse_stream:
raw_sse_buffer += chunk
while True:
frame, raw_sse_buffer = _pop_complete_sse_frame(raw_sse_buffer)
if frame is None:
break
yield frame # any-ok: untyped stream
if len(raw_sse_buffer) > _MAX_RAW_SSE_BUFFER_CHARS:
raise ValueError(
"Raw SSE stream exceeded maximum buffered size without a frame delimiter"
)
continue
if chunk.startswith(("data:", "event:", ":")):
yield ( # any-ok: untyped stream
chunk
if chunk.endswith(_SSE_FRAME_DELIMITERS)
else chunk + "\n\n"
)
continue
elif isinstance(chunk, str) and is_raw_sse_stream: # any-ok: untyped stream
raw_sse_buffer += chunk
while True:
frame, raw_sse_buffer = _pop_complete_sse_frame(raw_sse_buffer)
if frame is None:
break
yield frame # any-ok: untyped stream
if len(raw_sse_buffer) > _MAX_RAW_SSE_BUFFER_CHARS:
raise ValueError(
"Raw SSE stream exceeded maximum buffered size without a frame delimiter"
)
continue
elif isinstance(chunk, str) and chunk.startswith("data: "):
error_message = chunk
break
@ -7094,12 +7155,20 @@ async def async_data_generator(
except Exception as e:
yield f"data: {str(e)}\n\n"
stream_completed = True
if not needs_iterator_wrap:
# The iterator-wrap path fires deferred logging itself; fire it
# here for the no-wrap fast path so non-callback deployments
# still flush their post-stream logging.
ProxyLogging._fire_deferred_stream_logging(request_data)
if raw_sse_buffer:
yield ( # any-ok: untyped stream
raw_sse_buffer
if raw_sse_buffer.endswith(_SSE_FRAME_DELIMITERS)
else raw_sse_buffer + "\n\n"
)
if error_message is not None:
yield error_message
# OpenAI-compatible streams terminate with data: [DONE]; Google GenAI (?alt=sse) does not.
@ -7113,9 +7182,11 @@ async def async_data_generator(
# it here. This is the outermost generator Starlette closes on
# disconnect, so it fires reliably regardless of needs_iterator_wrap
# (a nested iterator hook would only see GeneratorExit on GC).
proxy_logging_obj._release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
if not stream_completed:
proxy_logging_obj._release_max_parallel_requests_on_disconnect(
user_api_key_dict
)
client_disconnected = True
raise
except Exception as e:
verbose_proxy_logger.exception(
@ -7149,30 +7220,33 @@ async def async_data_generator(
code=getattr(e, "status_code", 500),
)
error_returned = json.dumps({"error": proxy_exception.to_dict()})
stream_completed = True
yield f"data: {error_returned}\n\n"
finally:
# Close the response stream to release the underlying HTTP connection
# back to the connection pool. This prevents pool exhaustion when
# clients disconnect mid-stream.
# Shield from cancellation so the close awaits can complete.
with anyio.CancelScope(shield=True):
if hasattr(response, "aclose"):
try:
await response.aclose()
except BaseException as e:
verbose_proxy_logger.debug(
"async_data_generator: error closing response stream: %s",
e,
)
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
request=request,
request_data=request_data, # any-ok: untyped stream
response=response, # any-ok: untyped stream
stream_completed=stream_completed,
client_disconnected=client_disconnected,
)
def select_data_generator(
response, user_api_key_dict: UserAPIKeyAuth, request_data: dict
response,
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
request: Request | None = None,
):
return async_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
request=request,
)
@ -7250,15 +7324,53 @@ class ProxyStartupEvent:
if _use_redis_transaction_buffer and redis_usage_cache is None:
raise ValueError(
"`use_redis_transaction_buffer` is enabled in general_settings "
"but no Redis cache is configured. This will cause spend updates "
"but no Redis is configured. This will cause spend updates "
"to not be tracked. Add a Redis cache in litellm_settings:\n\n"
"litellm_settings:\n"
" cache: true\n"
" cache_params:\n"
" type: redis\n"
" url: os.environ/REDIS_URL\n"
" url: os.environ/REDIS_URL\n\n"
"or set REDIS_* environment variables (e.g. REDIS_HOST, "
"REDIS_PORT, REDIS_PASSWORD, or REDIS_URL) to use a standalone "
"Redis for the transaction buffer."
)
@staticmethod
def _get_transaction_buffer_redis_cache(
general_settings: dict,
) -> RedisCache | None:
"""
Builds a standalone Redis cache from REDIS_* environment variables so
use_redis_transaction_buffer can run when the proxy cache backend is not
Redis (e.g. disk, s3).
Returns None when the buffer is disabled, or when no Redis host or url
is set in the environment.
"""
from litellm._redis import _redis_kwargs_from_environment
from litellm.secret_managers.main import str_to_bool
_use_redis_transaction_buffer: bool | str | None = (
general_settings.get( # any-ok: untyped stream
"use_redis_transaction_buffer", False
)
)
if isinstance(_use_redis_transaction_buffer, str):
_use_redis_transaction_buffer = str_to_bool(_use_redis_transaction_buffer)
if not _use_redis_transaction_buffer:
return None
redis_env_kwargs = _redis_kwargs_from_environment() # any-ok: untyped stream
if (
"host" not in redis_env_kwargs # any-ok: untyped stream
and "url" not in redis_env_kwargs # any-ok: untyped stream
):
return None
return RedisCache(**redis_env_kwargs) # any-ok: untyped stream
@classmethod
async def _initialize_semantic_tool_filter(
cls,
@ -8609,6 +8721,7 @@ async def chat_completion(
response=_streaming_response,
user_api_key_dict=user_api_key_dict,
request_data=_data,
request=request,
)
return StreamingResponse(
@ -8643,6 +8756,7 @@ async def chat_completion(
response=_streaming_response,
user_api_key_dict=user_api_key_dict,
request_data=_data,
request=request,
)
return StreamingResponse(
@ -8791,6 +8905,7 @@ async def completion(
response=_streaming_response,
user_api_key_dict=user_api_key_dict,
request_data=_data,
request=request,
)
return StreamingResponse(
@ -8837,6 +8952,7 @@ async def completion(
response=_streaming_response,
user_api_key_dict=user_api_key_dict,
request_data=data,
request=request,
)
return StreamingResponse(
@ -13309,6 +13425,7 @@ async def async_queue_request(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=data,
request=request,
),
media_type="text/event-stream",
)

View file

@ -1243,6 +1243,16 @@
"provider_display_name": "Google AI Studio",
"litellm_provider": "gemini",
"credential_fields": [
{
"key": "api_base",
"label": "API Base",
"placeholder": "https://generativelanguage.googleapis.com/v1beta",
"tooltip": "Leave blank to let LiteLLM pick the right Gemini API version automatically (v1alpha for Gemini 3+ models, v1beta otherwise). Override only when fronting Gemini through a custom gateway; if you do, include the version prefix (e.g. /v1beta) but not the trailing slash. LiteLLM appends '/models/{model}:generateContent'.",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "api_key",
"label": "API Key",

View file

@ -2088,7 +2088,7 @@ class ProxyLogging:
litellm_call_id=request_data.get("litellm_call_id", ""), status="fail"
)
if AlertType.llm_exceptions in self.alert_types and not isinstance(
original_exception, HTTPException
original_exception, (HTTPException, ProxyException)
):
"""
Just alert on LLM API exceptions. Do not alert on user errors
@ -2192,6 +2192,7 @@ class ProxyLogging:
e.g should only return True for:
- Authentication Errors from user_api_key_auth
- HTTP HTTPException (rate limit errors)
- ProxyException (guardrail blocks, budget / rate-limit errors)
"""
#########################################################
@ -2208,7 +2209,7 @@ class ProxyLogging:
):
return False
return isinstance(original_exception, HTTPException) or (
return isinstance(original_exception, (HTTPException, ProxyException)) or (
error_type == ProxyErrorTypes.auth_error
)

View file

@ -0,0 +1,52 @@
model_list:
# ---------- Anthropic native ----------
- model_name: "anthropic/*"
litellm_params:
model: "anthropic/*"
api_key: os.environ/ANTHROPIC_API_KEY
# ---------- Bedrock ----------
- model_name: "bedrock/*"
litellm_params:
model: "bedrock/*"
aws_region_name: us-east-1
# ---------- Vertex AI ----------
- model_name: "vertex_ai/*"
litellm_params:
model: "vertex_ai/*"
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: global
# ---------- Azure AI Foundry ----------
- model_name: "azure_ai/*"
litellm_params:
model: "azure_ai/*"
api_base: os.environ/AZURE_AI_API_BASE
api_key: os.environ/AZURE_AI_API_KEY
# ---------- Azure OpenAI ----------
- model_name: "azure/*"
litellm_params:
model: "azure/*"
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
# ---------- Gemini ----------
- model_name: "gemini/*"
litellm_params:
model: "gemini/*"
api_key: os.environ/GEMINI_API_KEY
# ---------- OpenAI ----------
- model_name: "openai/*"
litellm_params:
model: "openai/*"
api_key: os.environ/OPENAI_API_KEY
general_settings:
master_key: sk-1234
litellm_settings:
drop_params: True
telemetry: False

View file

@ -244,9 +244,17 @@ def _check_non_standard_fallback_format(fallbacks: Optional[List[Any]]) -> bool:
if all(isinstance(item, str) for item in fallbacks):
return True
elif all(isinstance(item, dict) for item in fallbacks):
for key in LiteLLMParamsTypedDict.__annotations__.keys():
if key in fallbacks[0].keys():
return True
for item in fallbacks: # any-ok: untyped config
for (
key
) in (
LiteLLMParamsTypedDict.__annotations__.keys() # any-ok: untyped config
):
if key in item: # any-ok: untyped config
# If the value is a list, it's likely a standard fallback model group mapping
# (e.g. {"model": ["backup"]}) rather than a parameter override.
if not isinstance(item[key], list): # any-ok: untyped config
return True
return False

View file

@ -29,6 +29,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._types import KeyManagementSystem
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.secret_managers.main import KeyManagementSettings
from .base_secret_manager import BaseSecretManager
@ -43,6 +44,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
aws_profile_name: Optional[str] = None,
aws_web_identity_token: Optional[str] = None,
aws_sts_endpoint: Optional[str] = None,
replica_regions: list[str] | None = None,
**kwargs,
):
BaseSecretManager.__init__(self, **kwargs)
@ -56,6 +58,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
self.aws_profile_name = aws_profile_name
self.aws_web_identity_token = aws_web_identity_token
self.aws_sts_endpoint = aws_sts_endpoint
self.replica_regions: list[str] = replica_regions or []
@classmethod
def validate_environment(cls):
@ -75,7 +78,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
def load_aws_secret_manager(
cls,
use_aws_secret_manager: Optional[bool],
key_management_settings: Optional[Any] = None,
key_management_settings: KeyManagementSettings | None = None,
):
"""
Initialize AWSSecretsManagerV2 with settings from key_management_settings
@ -110,6 +113,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
"aws_sts_endpoint": getattr(
key_management_settings, "aws_sts_endpoint", None
),
"replica_regions": key_management_settings.replica_regions,
}
# Remove None values
aws_kwargs = {k: v for k, v in aws_kwargs.items() if v is not None}
@ -316,6 +320,90 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager):
params={"timeout": timeout},
)
try:
response = await async_client.post( # any-ok: untyped httpx
url=endpoint_url,
headers=headers, # any-ok: untyped httpx
data=body.decode("utf-8"), # any-ok: untyped httpx
)
response.raise_for_status() # any-ok: untyped httpx
create_response = response.json() # any-ok: untyped httpx
except httpx.HTTPStatusError as err:
raise ValueError(f"HTTP error occurred: {err.response.text}")
except httpx.TimeoutException:
raise ValueError("Timeout error occurred")
if self.replica_regions:
try:
await self.async_replicate_secret(
secret_name=secret_name,
replica_regions=self.replica_regions,
optional_params=optional_params, # any-ok: untyped httpx
timeout=timeout,
)
verbose_logger.debug(
"Replicated secret '%s' to regions: %s",
secret_name,
self.replica_regions,
)
except Exception as replication_err: # noqa: BLE001
verbose_logger.warning(
"Failed to replicate secret '%s' to regions %s: %s — key was created successfully.",
secret_name,
self.replica_regions,
str(replication_err),
)
return create_response # any-ok: untyped httpx
async def async_replicate_secret(
self,
secret_name: str,
replica_regions: list[str],
optional_params: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
) -> dict[str, object]:
"""
Replicate a secret to additional AWS regions using ReplicateSecretToRegions.
Called after a successful CreateSecret when replica_regions is configured.
Replication is best-effort — callers should not depend on this for correctness.
Args:
secret_name: Name or ARN of the secret to replicate
replica_regions: List of target AWS region names, e.g. ["us-west-2"]
optional_params: Additional AWS parameters
timeout: Request timeout
Returns:
dict: AWS response, or {} if replica_regions is empty
"""
if not replica_regions:
return {}
verbose_logger.info(
"ReplicateSecretToRegions called for secret '%s' in regions %s",
secret_name,
replica_regions,
)
data: dict[str, object] = {
"SecretId": secret_name,
"AddReplicaRegions": [{"Region": r} for r in replica_regions],
}
endpoint_url, headers, body = self._prepare_request( # any-ok: untyped httpx
action="ReplicateSecretToRegions",
secret_name=secret_name,
optional_params=optional_params,
request_data=data,
)
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.SecretManager,
params={"timeout": timeout}, # any-ok: untyped httpx
)
try:
response = await async_client.post(
url=endpoint_url, headers=headers, data=body.decode("utf-8")

View file

@ -11,5 +11,6 @@ class UiDiscoveryEndpoints(BaseModel):
auto_redirect_to_sso: bool
admin_ui_disabled: bool
sso_configured: bool
hide_default_credentials_hint: bool = False
is_control_plane: bool = False
workers: List[WorkerRegistryEntry] = []

View file

@ -72,3 +72,12 @@ class KeyManagementSettings(LiteLLMPydanticObjectBase):
aws_sts_endpoint: Optional[str] = None
"""Custom STS endpoint URL (useful for VPC endpoints or testing)"""
replica_regions: Optional[List[str]] = None
"""
Optional list of additional AWS regions to replicate secrets to after CreateSecret.
Uses the AWS Secrets Manager ReplicateSecretToRegions API. Replication is
best-effort — failure to replicate does not fail key creation.
Example: ["us-west-2", "eu-west-1"]
Only applies when key_management_system is "aws_secret_manager".
"""

View file

@ -196,7 +196,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
float
] # OpenAI priority service tier pricing
cache_read_input_token_cost_above_200k_tokens: Optional[float]
cache_read_input_token_cost_above_200k_tokens_priority: Optional[float]
cache_read_input_token_cost_above_272k_tokens: Optional[float]
cache_read_input_token_cost_above_272k_tokens_priority: Optional[float]
cache_read_input_token_cost_above_512k_tokens: Optional[float]
input_cost_per_character: Optional[float] # only for vertex ai models
input_cost_per_audio_token: Optional[float]
@ -204,9 +206,11 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
input_cost_per_token_above_200k_tokens: Optional[
float
] # only for vertex ai gemini-2.5-pro models
input_cost_per_token_above_200k_tokens_priority: Optional[float]
input_cost_per_token_above_272k_tokens: Optional[
float
] # GPT-5.4/5.4-pro: prompts >272K priced at 2x input
input_cost_per_token_above_272k_tokens_priority: Optional[float]
input_cost_per_token_above_512k_tokens: Optional[
float
] # MiniMax-M3: prompts >512K priced at 2x input
@ -240,9 +244,11 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
output_cost_per_token_above_200k_tokens: Optional[
float
] # only for vertex ai gemini-2.5-pro models
output_cost_per_token_above_200k_tokens_priority: Optional[float]
output_cost_per_token_above_272k_tokens: Optional[
float
] # GPT-5.4/5.4-pro: prompts >272K priced at 1.5x output
output_cost_per_token_above_272k_tokens_priority: Optional[float]
output_cost_per_token_above_512k_tokens: Optional[
float
] # MiniMax-M3: prompts >512K priced at 2x output
@ -3093,6 +3099,8 @@ class CustomPricingLiteLLMParams(BaseModel):
cache_read_input_token_cost_flex: Optional[float] = None
cache_read_input_token_cost_priority: Optional[float] = None
cache_read_input_token_cost_above_200k_tokens: Optional[float] = None
cache_read_input_token_cost_above_200k_tokens_priority: Optional[float] = None
cache_read_input_token_cost_above_272k_tokens_priority: Optional[float] = None
cache_read_input_audio_token_cost: Optional[float] = None
input_cost_per_character: Optional[float] = None
input_cost_per_character_above_128k_tokens: Optional[float] = None
@ -3100,6 +3108,8 @@ class CustomPricingLiteLLMParams(BaseModel):
input_cost_per_token_cache_hit: Optional[float] = None
input_cost_per_token_above_128k_tokens: Optional[float] = None
input_cost_per_token_above_200k_tokens: Optional[float] = None
input_cost_per_token_above_200k_tokens_priority: Optional[float] = None
input_cost_per_token_above_272k_tokens_priority: Optional[float] = None
input_cost_per_query: Optional[float] = None
input_cost_per_image: Optional[float] = None
input_cost_per_image_above_128k_tokens: Optional[float] = None
@ -3117,6 +3127,8 @@ class CustomPricingLiteLLMParams(BaseModel):
output_cost_per_audio_token: Optional[float] = None
output_cost_per_token_above_128k_tokens: Optional[float] = None
output_cost_per_token_above_200k_tokens: Optional[float] = None
output_cost_per_token_above_200k_tokens_priority: Optional[float] = None
output_cost_per_token_above_272k_tokens_priority: Optional[float] = None
output_cost_per_character_above_128k_tokens: Optional[float] = None
output_cost_per_image: Optional[float] = None
output_cost_per_image_token: Optional[float] = None
@ -3657,6 +3669,7 @@ class SpecialEnums(Enum):
class ServiceTier(Enum):
"""Enum for service tier types used in cost calculations."""
AUTO = "auto"
FLEX = "flex"
PRIORITY = "priority"

View file

@ -6043,9 +6043,15 @@ def _get_model_info_helper(
cache_read_input_token_cost_above_200k_tokens=_model_info.get(
"cache_read_input_token_cost_above_200k_tokens", None
),
cache_read_input_token_cost_above_200k_tokens_priority=_model_info.get( # any-ok: untyped cost map
"cache_read_input_token_cost_above_200k_tokens_priority", None
),
cache_read_input_token_cost_above_272k_tokens=_model_info.get(
"cache_read_input_token_cost_above_272k_tokens", None
),
cache_read_input_token_cost_above_272k_tokens_priority=_model_info.get( # any-ok: untyped cost map
"cache_read_input_token_cost_above_272k_tokens_priority", None
),
cache_read_input_token_cost_above_512k_tokens=_model_info.get(
"cache_read_input_token_cost_above_512k_tokens", None
),
@ -6067,9 +6073,15 @@ def _get_model_info_helper(
input_cost_per_token_above_200k_tokens=_model_info.get(
"input_cost_per_token_above_200k_tokens", None
),
input_cost_per_token_above_200k_tokens_priority=_model_info.get( # any-ok: untyped cost map
"input_cost_per_token_above_200k_tokens_priority", None
),
input_cost_per_token_above_272k_tokens=_model_info.get(
"input_cost_per_token_above_272k_tokens", None
),
input_cost_per_token_above_272k_tokens_priority=_model_info.get( # any-ok: untyped cost map
"input_cost_per_token_above_272k_tokens_priority", None
),
input_cost_per_token_above_512k_tokens=_model_info.get(
"input_cost_per_token_above_512k_tokens", None
),
@ -6125,9 +6137,15 @@ def _get_model_info_helper(
output_cost_per_token_above_200k_tokens=_model_info.get(
"output_cost_per_token_above_200k_tokens", None
),
output_cost_per_token_above_200k_tokens_priority=_model_info.get( # any-ok: untyped cost map
"output_cost_per_token_above_200k_tokens_priority", None
),
output_cost_per_token_above_272k_tokens=_model_info.get(
"output_cost_per_token_above_272k_tokens", None
),
output_cost_per_token_above_272k_tokens_priority=_model_info.get( # any-ok: untyped cost map
"output_cost_per_token_above_272k_tokens_priority", None
),
output_cost_per_token_above_512k_tokens=_model_info.get(
"output_cost_per_token_above_512k_tokens", None
),

View file

@ -2528,6 +2528,100 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"azure_ai/gpt-5.5": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
"cache_read_input_token_cost_priority": 1e-06,
"cache_read_input_token_cost_above_272k_tokens_priority": 2e-06,
"input_cost_per_token": 5e-06,
"input_cost_per_token_above_272k_tokens": 1e-05,
"input_cost_per_token_priority": 1e-05,
"input_cost_per_token_above_272k_tokens_priority": 2e-05,
"litellm_provider": "azure_ai",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3e-05,
"output_cost_per_token_above_272k_tokens": 4.5e-05,
"output_cost_per_token_priority": 6e-05,
"output_cost_per_token_above_272k_tokens_priority": 9e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.5",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure_ai/gpt-5.5-2026-04-23": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
"cache_read_input_token_cost_priority": 1e-06,
"cache_read_input_token_cost_above_272k_tokens_priority": 2e-06,
"input_cost_per_token": 5e-06,
"input_cost_per_token_above_272k_tokens": 1e-05,
"input_cost_per_token_priority": 1e-05,
"input_cost_per_token_above_272k_tokens_priority": 2e-05,
"litellm_provider": "azure_ai",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 3e-05,
"output_cost_per_token_above_272k_tokens": 4.5e-05,
"output_cost_per_token_priority": 6e-05,
"output_cost_per_token_above_272k_tokens_priority": 9e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.5",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"azure_ai/gpt-5.4": {
"cache_read_input_token_cost": 2.5e-07,
"cache_read_input_token_cost_above_272k_tokens": 5e-07,
@ -10068,6 +10162,8 @@
},
"claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -10097,6 +10193,8 @@
},
"claude-sonnet-4-5-20250929": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -10127,6 +10225,7 @@
},
"claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "anthropic",
@ -10155,6 +10254,8 @@
},
"claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -25103,6 +25204,21 @@
"supports_tool_choice": true,
"supports_vision": true
},
"mistral/mistral-medium-3-5": {
"input_cost_per_token": 1.5e-06,
"litellm_provider": "mistral",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04",
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"mistral/mistral-small": {
"input_cost_per_token": 1e-07,
"litellm_provider": "mistral",
@ -42830,4 +42946,105 @@
"supports_reasoning": true,
"source": "https://serverless.tensormesh.ai/v1/models/openrouter"
}
}
,
"deepseek-v4-flash": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2.8e-09,
"input_cost_per_token": 1.4e-07,
"input_cost_per_token_cache_hit": 2.8e-09,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false
},
"deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 4.35e-07,
"input_cost_per_token_cache_hit": 3.625e-09,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 8.7e-07,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false
},
"deepseek/deepseek-v4-flash": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2.8e-09,
"input_cost_per_token": 1.4e-07,
"input_cost_per_token_cache_hit": 2.8e-09,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false
},
"deepseek/deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 4.35e-07,
"input_cost_per_token_cache_hit": 3.625e-09,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 8.7e-07,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": false
}
}

View file

@ -1,8 +1,8 @@
#!/usr/bin/env python3
"""Non-gating ratchet guard: budget ceilings may only fall, never rise.
Every `*-budget.json` file (ruff-strict, type-discipline, mypy-code, basedpyright-code) is a
one-way ratchet: each rule's ceiling is `baseline + slack`, and the whole point is
Every `*-budget.json` file (ruff-strict, type-discipline, mypy-code, basedpyright-code,
any-discipline) is a one-way ratchet: each rule's ceiling is `baseline + slack`, and the whole point is
to drive that number DOWN over time. This check compares every budget file against
its own content at the merge-base with the target branch and fails (exits 1, red) if:
@ -12,6 +12,12 @@ its own content at the merge-base with the target branch and fails (exits 1, red
New rules and lowered/equal ceilings are fine.
The any-discipline budget is keyed by file rather than rule: its gate treats an
absent file as ceiling 0 (the file must be Any-free), so an entry vanishing means
that file was cleaned to zero -- a tightening, and exactly the cleanup this
ratchet exists to encourage. Such a budget is therefore exempt from the
dropped-entry rule (a raised ceiling is still caught).
This is deliberately NOT a gating check. It should turn the run red so that a
loosening is impossible to miss in review, but it must stay OUT of the
branch-protection required-checks list: a justified bump (e.g. banning a new API,
@ -40,8 +46,15 @@ DEFAULT_BUDGETS: tuple[str, ...] = (
"type-discipline-budget.json",
"mypy-code-budget.json",
"basedpyright-code-budget.json",
"any-discipline-budget.json",
)
# File-keyed budgets whose gate treats an absent entry as ceiling 0 (the file
# must stay clean). Dropping an entry there is a tightening, not the "untracked,
# now unbounded" loosening a vanished rule is for the rule-keyed budgets, so a
# dropped entry must not read as a regression.
ZERO_FLOOR_BUDGETS: frozenset[str] = frozenset({"any-discipline-budget.json"})
class Regression(NamedTuple):
budget: str
@ -99,10 +112,12 @@ def regressions_for(rel: str, base: dict | None, head: dict | None) -> list[Regr
base_caps = _caps(base)
head_caps = _caps(head)
drop_floors_to_zero = rel in ZERO_FLOOR_BUDGETS
out: list[Regression] = []
for rule, base_cap in sorted(base_caps.items()):
if rule not in head_caps:
out.append(Regression(rel, rule, f"rule dropped (ceiling {base_cap} -> removed)"))
if not drop_floors_to_zero:
out.append(Regression(rel, rule, f"rule dropped (ceiling {base_cap} -> removed)"))
elif head_caps[rule] > base_cap:
out.append(Regression(rel, rule, f"ceiling raised {base_cap} -> {head_caps[rule]}"))
return out

View file

@ -1,5 +1,5 @@
#!/usr/bin/env python3
"""Any-discipline gate: fail when a *changed* file holds a value typed `Any`.
"""Any-discipline gate: fail when a changed file exceeds its `Any` budget.
Where ruff, `mypy --strict`, and even basedpyright's `reportAny` stop short, this
catches the case that actually bites: a *union* hiding an `Any`. For example
@ -7,18 +7,26 @@ catches the case that actually bites: a *union* hiding an `Any`. For example
-> `list[Any]`/`dict[..., Any]`. Any value whose inferred type *contains* `Any`
(recursively, through unions / generics / tuples) is reported.
Scope: changed-only, changed-lines
----------------------------------
litellm already contains a large amount of pre-existing `Any` (a single legacy
file can have >100 findings), and a whole-tree scan would have to re-export types
for litellm's entire import closure on every run (~2 min, ~3 GB). So this gate is
*changed-only* and reports a finding only on a line that the diff against
`--base` actually adds or edits (untracked files count as wholly new). A brand
new file is therefore checked in full, while editing a legacy file only requires
*your* lines to be clean -- you can't introduce an `X | Any`, but you aren't
forced to clean the file's existing debt. This mirrors how `ruff_strict_gate.py`
blames a change only for the violations it introduces; cold legacy code is left
to the ratchet gates (mypy/basedpyright/ruff budgets).
Scope: changed files, per-file budget
-------------------------------------
litellm carries a large amount of pre-existing `Any` (a single legacy file can
have >100 findings). Rather than force every touched line clean (the original
changed-lines rule, which tripped on merely *editing* a legacy `X | Any` line),
this gate grandfathers each file: `any-discipline-budget.json` records every
file's current count of Any-typed values, and a file fails only when its count
exceeds `baseline + slack`, where `slack` is 50% headroom (rounded up). New or
unbudgeted files have baseline 0, so they stay airtight.
Only *changed* files (vs the merge-base with `--base`) are re-type-checked -- an
unchanged file's count can't move from edits this branch didn't make -- so the
per-PR cost equals re-checking just those files, exactly like the original
changed-lines gate. The whole-tree scan needed to (re)capture the budget
(~2 min, ~3 GB) runs only under `--update`.
The budget is a one-way ratchet (the same `{baseline, slack}` shape as the
ruff / mypy / basedpyright budgets) guarded by `scripts/budget_ratchet_check.py`:
a file's ceiling may fall but never rise. Drive a file's count down and rerun
`--update` (`make lint-any-budget-update`) to lock in the lower ceiling.
How it works
------------
@ -36,8 +44,9 @@ Rules
-----
Codes share the `LIT***` namespace with `scripts/check_type_discipline.py` (PR
#30500), which owns LIT001/002/003/004/006/007/008. This gate claims the rest:
LIT009 A value expression's inferred type is, or contains, `Any`.
Suppress with `# any-ok: <reason>` on the offending line.
LIT009 A value expression's inferred type is, or contains, `Any`. Budgeted
per file (a file fails when its count exceeds `baseline + slack`).
Suppress an individual line with `# any-ok: <reason>`.
LIT005 An `# any-ok` suppression without a reason (the shared
suppression-needs-a-reason code, same as `# cast-ok` / `# guard-ok`).
LIT000 Setup failure: mypy could not build, or a target file could not be read.
@ -48,13 +57,16 @@ whose signature mentions `Any` is not flagged -- only the value its call produce
Usage
-----
# gate mode (CI / pre-push): check changed lines under litellm/
# gate mode (CI / pre-push): per-file Any budget on changed files
uv run --no-sync python scripts/check_any_discipline.py --changed --base origin/litellm_internal_staging
# whole-file spot-check (no line filter), paths relative to repo root
# re-capture the per-file budget across the whole tree (ratchet)
uv run --no-sync python scripts/check_any_discipline.py --update
# whole-file spot-check (no budget, no line filter), paths relative to repo root
uv run --no-sync python scripts/check_any_discipline.py litellm/budget_manager.py
Exit code 1 if any Any-tainted value is found, 2 on a setup/usage error.
Exit code 1 if a file is over budget (or a hard rule trips), 2 on a setup error.
"""
from __future__ import annotations
@ -66,7 +78,7 @@ import re
import subprocess
import sys
import tokenize
from collections.abc import Iterable, Sequence
from collections.abc import Callable, Iterable, Sequence
from pathlib import Path
from typing import NamedTuple
@ -76,7 +88,7 @@ try:
from mypy.find_sources import create_source_list
from mypy.fscache import FileSystemCache
from mypy.modulefinder import BuildSource
from mypy.nodes import AssignmentStmt, Expression, NameExpr, Node
from mypy.nodes import AssignmentStmt, Expression, NameExpr, Node, TempNode
from mypy.options import Options
from mypy.types import (
AnyType,
@ -104,6 +116,7 @@ MYPY_INI = LITELLM_DIR / "mypy.ini"
CACHE_DIR = REPO_ROOT / ".mypy_cache_any"
PY_TAG = f"{sys.version_info.major}.{sys.version_info.minor}"
DEFAULT_BASE = "origin/litellm_internal_staging"
BUDGET_PATH = REPO_ROOT / "any-discipline-budget.json"
MIN_REASON_LEN = 3
ANY_OK_RE = re.compile(r"#\s*any-ok(?::\s*(?P<reason>.*))?")
@ -134,6 +147,19 @@ _HARMLESS_ANY = frozenset(
# against ExtendedTraverserVisitor across the full grammar (see commit notes).
_NON_SYNTACTIC_ATTRS = frozenset({"node", "info"})
# Awaitable / coroutine / generator instances carry synthetic `Any` in their
# send (and, for coroutines, yield) protocol slots: `async def f() -> float`
# produces `Coroutine[Any, Any, float]`, so the bare call expression `f()` would
# be flagged even though the awaited value is a clean `float`. Only the args that
# hold a value the caller observes (the awaited result, the yielded item) are
# meaningful; a real `Any` there -- e.g. a coroutine that returns `Any` -- is
# still caught because that index is still checked.
_SYNTHETIC_SEND_YIELD_VALUE_ARGS: dict[str, tuple[int, ...]] = {
"typing.Coroutine": (2,),
"typing.Generator": (0, 2),
"typing.AsyncGenerator": (0,),
}
class Violation(NamedTuple):
path: Path
@ -151,34 +177,50 @@ class Violation(NamedTuple):
# --------------------------------------------------------------------------- #
_MAX_CONTAINS_ANY_DEPTH = 64
# Recursive type aliases (e.g. a JSON-like `T = Union[..., list[T], dict[str, T]]`)
# make `get_proper_type` yield a fresh object at every unfold, so an id()-based
# cycle guard never trips and a naive recursion overflows the stack. We walk
# iteratively and cap the depth: a real `Any` lives at shallow depth in the
# alias's definition, so a deep alias that has not produced one by `_MAX_DEPTH`
# never will. (The changed-lines gate never hit this; a whole-tree scan does.)
_MAX_DEPTH = 100
def contains_any(t: Type, _seen: set[int] | None = None, _depth: int = 0) -> bool:
def contains_any(t: Type) -> bool:
"""True if a *value* of type ``t`` carries `Any` anywhere meaningful."""
if _depth > _MAX_CONTAINS_ANY_DEPTH:
# Bail out on deeply-nested / potentially-circular types rather than
# overflowing the Python call stack. A type this deep is unlikely to
# carry a *meaningful* Any that the developer could actually fix.
return False
seen = _seen if _seen is not None else set()
p = get_proper_type(t)
if id(p) in seen:
return False
seen.add(id(p))
seen: set[int] = set()
stack: list[tuple[Type, int]] = [(t, 0)]
while stack:
cur, depth = stack.pop()
if depth > _MAX_DEPTH:
continue
p = get_proper_type(cur)
if id(p) in seen:
continue
seen.add(id(p))
# A function/method *reference* whose signature mentions Any is not itself an
# unsafe value -- only its eventual call result is. Don't recurse into it.
if isinstance(p, (CallableType, Overloaded)):
return False
if isinstance(p, AnyType):
return p.type_of_any not in _HARMLESS_ANY
if isinstance(p, UnionType):
return any(contains_any(item, seen, _depth + 1) for item in p.items)
if isinstance(p, Instance):
return any(contains_any(arg, seen, _depth + 1) for arg in p.args)
if isinstance(p, TupleType):
return any(contains_any(item, seen, _depth + 1) for item in p.items)
# A function/method *reference* whose signature mentions Any is not itself
# an unsafe value -- only its eventual call result is. Don't recurse in.
if isinstance(p, (CallableType, Overloaded)):
continue
if isinstance(p, AnyType):
if p.type_of_any not in _HARMLESS_ANY:
return True
continue
if isinstance(p, UnionType):
stack.extend((item, depth + 1) for item in p.items)
elif isinstance(p, Instance):
value_arg_indices = _SYNTHETIC_SEND_YIELD_VALUE_ARGS.get(p.type.fullname)
if value_arg_indices is None:
stack.extend((arg, depth + 1) for arg in p.args)
else:
stack.extend(
(p.args[index], depth + 1)
for index in value_arg_indices
if index < len(p.args)
)
elif isinstance(p, TupleType):
stack.extend((item, depth + 1) for item in p.items)
return False
@ -232,7 +274,11 @@ def find_any_in_tree(tree: Node, idmap: dict[int, Type]) -> list[tuple[int, int,
exprs, skip_lvalues = _walk_file(tree)
findings: list[tuple[int, int, str]] = []
for expr in exprs:
if id(expr) in skip_lvalues:
# A TempNode is mypy's synthetic placeholder for a position with no real
# expression -- e.g. the rvalue of an annotation-only `field: T` in a
# TypedDict / class body, whose `special_form` `Any` is not a value the
# author wrote. It never corresponds to a runtime value, so skip it.
if id(expr) in skip_lvalues or isinstance(expr, TempNode):
continue
t = idmap.get(id(expr))
if t is not None and contains_any(t):
@ -257,11 +303,17 @@ def _reason_ok(reason: str | None) -> bool:
return reason is not None and len(reason.strip()) >= MIN_REASON_LEN
def scan_any_ok(path: Path, source: str) -> tuple[frozenset[int], tuple[Violation, ...]]:
def scan_any_ok(
path: Path, source: str
) -> tuple[frozenset[int], tuple[Violation, ...]]:
"""Return (lines with a valid any-ok suppression, LIT005 violations)."""
try:
tokens = tokenize.generate_tokens(iter(source.splitlines(keepends=True)).__next__)
comments = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT)
tokens = tokenize.generate_tokens(
iter(source.splitlines(keepends=True)).__next__
)
comments = tuple(
(t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT
)
except tokenize.TokenError:
return frozenset(), ()
@ -372,7 +424,9 @@ def check_files(rel_paths: Sequence[str]) -> tuple[Violation, ...]:
try:
source = abs_path.read_text(encoding="utf-8")
except (OSError, UnicodeDecodeError) as exc:
out.append(Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}"))
out.append(
Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}")
)
continue
ok_lines, ok_violations = scan_any_ok(report_path, source)
@ -491,59 +545,233 @@ def _in_scope(v: Violation, line_map: dict[str, LineScope] | None) -> bool:
if line_map is None or v.code == "LIT000":
return True
lines = line_map.get(v.path.as_posix())
return lines is ALL_LINES or (lines is not None and v.line in lines)
return lines is ALL_LINES or (isinstance(lines, set) and v.line in lines)
# --------------------------------------------------------------------------- #
# Per-file Any budget (one-way ratchet, 50% headroom; ratchet-checked)
# --------------------------------------------------------------------------- #
def _slack_for(baseline: int) -> int:
"""50% headroom, rounded up so even a 1-Any file gets a little room."""
return (baseline + 1) // 2
def _ceiling(spec: dict[str, int]) -> int:
"""A file's ceiling: ``baseline + slack`` (0 for an absent/empty entry)."""
return int(spec.get("baseline", 0)) + int(spec.get("slack", 0))
def load_budget() -> dict[str, dict[str, int]]:
"""Read ``any-discipline-budget.json`` ({path: {baseline, slack}}); {} if absent."""
if not BUDGET_PATH.exists():
return {}
try:
data = json.loads(BUDGET_PATH.read_text())
except (OSError, ValueError):
return {}
return data if isinstance(data, dict) else {}
def save_budget(counts: dict[str, int]) -> None:
"""Write a fresh budget from per-file counts, with 50% headroom each.
Files with zero Any are omitted: an absent entry means baseline 0, so a
file's first Any always trips the gate until it is deliberately baselined."""
budget = {
path: {"baseline": n, "slack": _slack_for(n)}
for path, n in counts.items()
if n > 0
}
BUDGET_PATH.write_text(json.dumps(budget, indent=2, sort_keys=True) + "\n")
def lit009_counts(violations: Iterable[Violation]) -> dict[str, int]:
"""Count LIT009 (Any-typed value) findings per repo-relative file path."""
counts: dict[str, int] = {}
for v in violations:
if v.code == "LIT009":
key = v.path.as_posix()
counts[key] = counts.get(key, 0) + 1
return counts
def all_litellm_py_files() -> list[str] | None:
"""Every tracked ``.py`` under litellm/, as litellm-package-relative paths;
None if git is unavailable / not a repo (mirrors ``changed_line_map``)."""
try:
tracked = _git("ls-files", "--", "litellm")
except (subprocess.CalledProcessError, FileNotFoundError):
return None
return _to_litellm_relative(
REPO_ROOT / name for name in tracked if name.endswith(".py")
)
def update_budget(
list_files: Callable[[], list[str] | None] = all_litellm_py_files,
) -> int:
"""Whole-tree scan: recapture every file's Any count into the budget."""
rel_paths = list_files()
if rel_paths is None:
print(
"check_any_discipline: not a git repository; cannot capture the budget",
file=sys.stderr,
)
return 2
if not rel_paths:
print("check_any_discipline: no litellm/*.py files found", file=sys.stderr)
return 2
violations = check_files(rel_paths)
build_errors = [v for v in violations if v.code == "LIT000"]
if build_errors:
for v in build_errors:
print(v.render(), file=sys.stderr)
print(
"FAIL: mypy could not build the tree; budget left unchanged.",
file=sys.stderr,
)
return 2
counts = lit009_counts(violations)
save_budget(counts)
print(
f"Wrote {BUDGET_PATH.name}: "
f"{sum(1 for n in counts.values() if n > 0)} file(s), "
f"{sum(counts.values())} Any-typed value(s) baselined (50% headroom each)."
)
return 0
def _report_over_budget(
path: str,
count: int,
spec: dict[str, int] | None,
lit009: list[Violation],
line_map: dict[str, LineScope],
) -> None:
"""Print one over-budget file plus the Any findings on its changed lines."""
ceiling = _ceiling(spec or {})
if spec:
why = f"baseline {spec['baseline']} + 50% slack {spec['slack']} = ceiling {ceiling}"
else:
why = "no budget entry -> baseline 0 (a new/unbudgeted file must be Any-free)"
print(f"{path}: {count} Any-typed value(s) total, over budget ({why})")
# Surface the findings on changed lines first: the ones this branch most
# likely just added, and the cheapest path back under the ceiling.
scope = line_map.get(path)
for v in sorted(lit009):
if scope is ALL_LINES or (isinstance(scope, set) and v.line in scope):
print(f" changed-line Any {v.line}:{v.col} {v.message}")
def run_gate(base: str) -> int:
"""Gate changed files under litellm/ against the committed per-file budget."""
line_map = changed_line_map(base)
if line_map is None:
print(
"check_any_discipline: not a git repository; nothing to check",
file=sys.stderr,
)
return 0
rel_paths = _to_litellm_relative((REPO_ROOT / name).resolve() for name in line_map)
if not rel_paths:
print("OK: no changed Python files under litellm/ to check")
return 0
violations = check_files(rel_paths)
budget = load_budget()
# Hard rules, independent of the budget: a build/read failure (always), and a
# reasonless `# any-ok` on a line this branch touched.
hard = sorted(
v
for v in violations
if v.code == "LIT000" or (v.code == "LIT005" and _in_scope(v, line_map))
)
# Per-file Any budget: a changed file fails when its total Any count exceeds
# its ceiling. Unchanged files keep their committed baseline (never re-scanned).
counts = lit009_counts(violations)
lit009_by_file: dict[str, list[Violation]] = {}
for v in violations:
if v.code == "LIT009":
lit009_by_file.setdefault(v.path.as_posix(), []).append(v)
over_budget = [
(path, count)
for path, count in sorted(counts.items())
if count > _ceiling(budget.get(path, {}))
]
if not hard and not over_budget:
print(
f"OK: {len(rel_paths)} changed file(s) under litellm/ are within their Any budget"
)
return 0
for v in hard:
print(v.render())
for path, count in over_budget:
_report_over_budget(
path, count, budget.get(path), lit009_by_file.get(path, []), line_map
)
print(
f"\nFAIL: {len(hard)} hard violation(s), {len(over_budget)} file(s) over their Any budget.\n"
"Give the new values concrete types (validate untyped input with Pydantic) to get back\n"
"under the file's ceiling, or annotate a genuine boundary line `# any-ok: <reason>`.\n"
"Re-baseline with `make lint-any-budget-update` only to lock in a reduction.",
file=sys.stderr,
)
return 1
def spot_check(rel_paths: Sequence[str]) -> int:
"""Explicit-paths mode: report every finding in the files (no budget)."""
violations = sorted(check_files(rel_paths))
for v in violations:
print(v.render())
if violations:
print(f"\nFAIL: {len(violations)} Any-discipline finding(s).", file=sys.stderr)
return 1
print(f"OK: {len(rel_paths)} file(s) have no Any-typed values")
return 0
def main(argv: Sequence[str]) -> int:
parser = argparse.ArgumentParser(description="Any-discipline gate (changed-only, changed-lines).")
parser = argparse.ArgumentParser(
description="Any-discipline gate (changed files, per-file Any budget)."
)
parser.add_argument(
"paths",
nargs="*",
help="explicit files (repo-root relative); whole-file, no line filter",
help="explicit files (repo-root relative); whole-file spot-check, no budget",
)
parser.add_argument(
"--changed",
action="store_true",
help="check changed lines under litellm/ vs --base",
help="gate changed files under litellm/ vs --base against the per-file budget",
)
parser.add_argument(
"--update",
action="store_true",
help="recapture the whole-tree per-file budget (any-discipline-budget.json)",
)
parser.add_argument("--base", default=os.environ.get("ANY_GATE_BASE", DEFAULT_BASE))
args = parser.parse_args(list(argv))
line_map: dict[str, LineScope] | None = None
if args.update:
return update_budget()
if args.changed:
line_map = changed_line_map(args.base)
if line_map is None:
print(
"check_any_discipline: not a git repository; nothing to check",
file=sys.stderr,
)
return 0
rel_paths = _to_litellm_relative((REPO_ROOT / name).resolve() for name in line_map)
elif args.paths:
return run_gate(args.base)
if args.paths:
rel_paths = _to_litellm_relative((REPO_ROOT / p).resolve() for p in args.paths)
else:
parser.error("pass --changed or explicit file paths")
return 2
if not rel_paths:
print("OK: no changed Python lines under litellm/ to check")
return 0
violations = tuple(v for v in check_files(rel_paths) if _in_scope(v, line_map))
for v in sorted(violations):
print(v.render())
if violations:
n = len(violations)
print(
f"\nFAIL: {n} Any-discipline violation(s) on changed lines.\n"
"Give the value a concrete type, or annotate the line `# any-ok: <reason>`.",
file=sys.stderr,
)
return 1
print(f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines")
return 0
if not rel_paths:
print("check_any_discipline: no litellm/*.py paths given", file=sys.stderr)
return 2
return spot_check(rel_paths)
parser.error("pass --changed, --update, or explicit file paths")
return 2
if __name__ == "__main__":

View file

@ -6,10 +6,10 @@ import sys
from unittest.mock import AsyncMock, patch, call
import pytest
from fastapi.exceptions import HTTPException
from httpx import Request, Response
from litellm import DualCache
from litellm.proxy._types import ProxyException
from litellm.proxy.guardrails.guardrail_hooks.aim.aim import (
AimGuardrail,
AimGuardrailMissingSecrets,
@ -101,7 +101,7 @@ async def test_block_callback(mode: str):
],
}
with pytest.raises(HTTPException, match="Jailbreak detected"):
with pytest.raises(ProxyException, match="Jailbreak detected") as exc_info:
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=Response(
@ -135,6 +135,137 @@ async def test_block_callback(mode: str):
call_type="completion",
)
exc = exc_info.value
assert exc.code == "400"
assert exc.type == "invalid_request_error"
assert exc.param is None
assert exc.openai_code == "content_policy_violation"
@pytest.mark.asyncio
async def test_output_block_raises_proxy_exception():
"""An output-side block is a content-policy violation, like the input block:
it must surface a conformant ProxyException, not a bare HTTPException whose
type/param serialize as the literal string "None". Regression for LIT-3751."""
init_guardrails_v2(
all_guardrails=[
{
"guardrail_name": "gibberish-guard",
"litellm_params": {
"guardrail": "aim",
"mode": "post_call",
"api_key": "hs-aim-key",
},
},
],
config_file_path="",
)
aim_guardrails = [
callback for callback in litellm.callbacks if isinstance(callback, AimGuardrail)
]
assert len(aim_guardrails) == 1
aim_guardrail = aim_guardrails[0]
block_on_output = Response(
json={
"analysis_result": {"policy_drill_down": {"PII": {}}},
"required_action": {
"action_type": "block_action",
"detection_message": "Output blocked: leaked secret",
"policy_name": "blocking policy",
},
},
status_code=200,
request=Request(method="POST", url="http://aim"),
)
response = ModelResponse(
choices=[
{
"finish_reason": "stop",
"index": 0,
"message": {"content": "here is the secret", "role": "assistant"},
}
]
)
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=block_on_output,
):
with pytest.raises(ProxyException, match="Output blocked") as exc_info:
await aim_guardrail.async_post_call_success_hook(
data={"messages": [{"role": "user", "content": "tell me a secret"}]},
response=response,
user_api_key_dict=UserAPIKeyAuth(),
)
exc = exc_info.value
assert exc.code == "400"
assert exc.type == "invalid_request_error"
assert exc.param is None
assert exc.openai_code == "content_policy_violation"
@pytest.mark.asyncio
async def test_anonymize_multimodal_rejection_raises_proxy_exception():
"""Anonymize on multimodal input degrades to a 400 because mask-in-place would
drop non-text parts. That is a usage error, not a content-policy violation, so
it must raise a conformant ProxyException WITHOUT the content_policy_violation
code. Regression for LIT-3751."""
init_guardrails_v2(
all_guardrails=[
{
"guardrail_name": "gibberish-guard",
"litellm_params": {
"guardrail": "aim",
"mode": "pre_call",
"api_key": "hs-aim-key",
},
},
],
config_file_path="",
)
aim_guardrails = [
callback for callback in litellm.callbacks if isinstance(callback, AimGuardrail)
]
assert len(aim_guardrails) == 1
aim_guardrail = aim_guardrails[0]
data = {
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": "Hi my name is Brian"},
{
"type": "image_url",
"image_url": {"url": "data:image/png;base64,iVBORw0KGgo="},
},
],
},
],
}
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=response_with_detections,
):
with pytest.raises(
ProxyException, match="anonymize action requested for multimodal"
) as exc_info:
await aim_guardrail.async_pre_call_hook(
data=data,
cache=DualCache(),
user_api_key_dict=UserAPIKeyAuth(),
call_type="completion",
)
exc = exc_info.value
assert exc.code == "400"
assert exc.type == "invalid_request_error"
assert exc.param is None
assert exc.openai_code != "content_policy_violation"
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ["pre_call", "during_call"])

View file

@ -0,0 +1,49 @@
"""
Test that check_and_fix_namespace handles None key gracefully.
Regression test for https://github.com/BerriAI/litellm/issues/30424
"""
from unittest.mock import MagicMock
from litellm.caching.redis_cache import RedisCache
def test_check_and_fix_namespace_with_none_key():
"""When key is None, check_and_fix_namespace should return None without raising."""
cache = MagicMock(spec=RedisCache)
cache.namespace = "litellm"
# Call the real method
result = RedisCache.check_and_fix_namespace(cache, key=None)
assert result is None
def test_check_and_fix_namespace_with_none_key_no_namespace():
"""When key is None and namespace is None, should return None without raising."""
cache = MagicMock(spec=RedisCache)
cache.namespace = None
result = RedisCache.check_and_fix_namespace(cache, key=None)
assert result is None
def test_check_and_fix_namespace_with_valid_key():
"""Normal behavior: prefix key with namespace if not already prefixed."""
cache = MagicMock(spec=RedisCache)
cache.namespace = "litellm"
result = RedisCache.check_and_fix_namespace(cache, key="my_key")
assert result == "litellm:my_key"
def test_check_and_fix_namespace_with_already_prefixed_key():
"""If key already starts with namespace, don't double-prefix."""
cache = MagicMock(spec=RedisCache)
cache.namespace = "litellm"
result = RedisCache.check_and_fix_namespace(cache, key="litellm:my_key")
assert result == "litellm:my_key"
def test_check_and_fix_namespace_no_namespace():
"""When namespace is None, return key as-is."""
cache = MagicMock(spec=RedisCache)
cache.namespace = None
result = RedisCache.check_and_fix_namespace(cache, key="my_key")
assert result == "my_key"

View file

@ -1573,3 +1573,72 @@ def test_data_residency_composes_with_service_tier(_local_model_cost_map):
assert priority_base_total > 0
assert priority_eu_total == pytest.approx(priority_base_total * 1.10, rel=1e-9)
def test_priority_service_tier_above_threshold_uses_priority_tier_rates_for_cached_tokens(
_local_model_cost_map,
):
"""Regression: for a model that publishes both service_tier and above_threshold rate
variants, a priority request over the threshold must bill cached tokens at
cache_read_input_token_cost_above_200k_tokens_priority (and analogously for
input/output above-threshold), not the standard above-threshold rate."""
usage = Usage(
prompt_tokens=250_000,
completion_tokens=1_000,
total_tokens=251_000,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=200_000, text_tokens=50_000
),
completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=1_000),
)
prompt_cost, completion_cost = generic_cost_per_token(
model="gemini-3-pro-preview",
usage=usage,
custom_llm_provider="gemini",
service_tier="priority",
)
# gemini-3-pro-preview priority + above_200k rates from the pricing JSON:
# input 7.2e-6, output 3.24e-5, cache_read 7.2e-7
expected_prompt = 50_000 * 7.2e-6 + 200_000 * 7.2e-7
expected_completion = 1_000 * 3.24e-5
assert prompt_cost == pytest.approx(expected_prompt, rel=1e-9)
assert completion_cost == pytest.approx(expected_completion, rel=1e-9)
def test_priority_service_tier_above_threshold_falls_back_to_standard_for_cache_creation(
_local_model_cost_map,
):
"""Regression: priority requests against models that publish standard above-threshold
cache_creation rates but no priority variant must fall back to the standard
above-threshold rate, not the priority-base rate. vertex_ai/claude-sonnet-4-5
has cache_creation_input_token_cost_above_200k_tokens but no _priority sibling."""
usage = Usage(
prompt_tokens=350_000,
completion_tokens=1_000,
total_tokens=351_000,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=200_000,
cache_creation_tokens=100_000,
text_tokens=50_000,
),
completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=1_000),
)
prompt_cost, completion_cost = generic_cost_per_token(
model="vertex_ai/claude-sonnet-4-5",
usage=usage,
custom_llm_provider="vertex_ai",
service_tier="priority",
)
# vertex_ai/claude-sonnet-4-5 above_200k (no _priority variants):
# input 6e-6, output 2.25e-5, cache_read 6e-7, cache_creation 7.5e-6
# text 50_000 * 6e-6 = 0.30
# cache_read 200_000 * 6e-7 = 0.12
# cache_creation 100_000 * 7.5e-6 = 0.75
expected_prompt = 50_000 * 6e-6 + 200_000 * 6e-7 + 100_000 * 7.5e-6
expected_completion = 1_000 * 2.25e-5
assert prompt_cost == pytest.approx(expected_prompt, rel=1e-9)
assert completion_cost == pytest.approx(expected_completion, rel=1e-9)

View file

@ -3115,6 +3115,71 @@ class TestFirstApiCallStartTimeSetOnce:
assert user_meta == {}
def test_get_error_information_for_logging_payload_ignores_spoofed_disconnect_without_flag():
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
baseline = StandardLoggingPayloadSetup.get_error_information(
original_exception=ValueError("provider failure"),
)
error_information, error_str = (
StandardLoggingPayloadSetup.get_error_information_for_logging_payload(
metadata={
"error_information": {
"error_code": "499",
"error_message": "Client disconnected the request",
"error_class": "ClientDisconnected",
}
},
original_exception=ValueError("provider failure"),
error_str="provider failure",
)
)
assert error_information == baseline
assert error_str == "provider failure"
def test_get_error_information_for_logging_payload_client_disconnect():
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
custom_error = {
"error_code": "499",
"error_message": "Client disconnected the request",
"error_class": "ClientDisconnected",
}
error_information, error_str = (
StandardLoggingPayloadSetup.get_error_information_for_logging_payload(
metadata={"client_disconnected": True, "error_information": custom_error},
original_exception=None,
error_str=None,
)
)
assert error_information == custom_error
assert error_str == "Client disconnected the request"
error_information, error_str = (
StandardLoggingPayloadSetup.get_error_information_for_logging_payload(
metadata={"client_disconnected": True},
original_exception=None,
error_str="existing error",
)
)
assert error_information["error_code"] == "499"
assert error_str == "existing error"
baseline = StandardLoggingPayloadSetup.get_error_information(
original_exception=None,
)
error_information, error_str = (
StandardLoggingPayloadSetup.get_error_information_for_logging_payload(
metadata={},
original_exception=None,
error_str=None,
)
)
assert error_information == baseline
assert error_str is None
def test_get_error_information_proxy_exception_preserves_message():
"""ProxyException keeps its text in ``.message`` (str() was empty pre-fix),
so error_information must still surface the message and code."""

View file

@ -523,7 +523,6 @@ from unittest.mock import MagicMock, patch
from litellm.utils import _select_tokenizer_helper, claude_json_str, encoding
# Clear the cache at module load to ensure clean state
_select_tokenizer_helper.cache_clear()
@ -1010,3 +1009,64 @@ def test_token_counter_with_thinking_content():
assert (
tokens_no_thinking < 15
), f"Expected minimal token count for empty thinking block, got {tokens_no_thinking}"
def test_token_counter_with_tool_reference_block():
"""
Regression test: a message containing an Anthropic tool-search
`tool_reference` content block must NOT raise.
Before the fix, token_counter raised
`Invalid content item type: tool_reference`. On the streaming
anthropic_messages proxy path this nulled response_cost and caused the
SpendLogs row to be dropped, silently undercounting cost. token_counter
must instead count the referenced tool name and return a positive count.
"""
messages = [
{
"role": "assistant",
"content": [
{"type": "text", "text": "Let me look up the right tool."},
{"type": "tool_reference", "tool_name": "search_knowledge_base"},
],
}
]
# Must not raise, and must produce a positive token count.
tokens = token_counter_new(
model="anthropic/claude-sonnet-4-5-20250929", messages=messages
)
assert tokens > 0, f"Expected positive token count, got {tokens}"
# A tool_reference with no/empty tool_name must also be handled gracefully.
messages_empty = [
{
"role": "assistant",
"content": [{"type": "tool_reference", "tool_name": ""}],
}
]
tokens_empty = token_counter_new(
model="anthropic/claude-sonnet-4-5-20250929", messages=messages_empty
)
assert tokens_empty >= 0
def test_count_content_list_rejects_unknown_type():
"""
An unrecognized content block type must raise, and the error message must
enumerate the supported types (including `tool_reference`). This pins the
catch-all contract so a future block type isn't silently dropped.
"""
from litellm.litellm_core_utils.token_counter import _count_content_list
with pytest.raises(ValueError) as exc_info:
_count_content_list(
count_function=len,
content_list=[{"type": "totally_unknown_block"}],
use_default_image_token_count=False,
default_token_count=None,
)
message = str(exc_info.value)
assert "Invalid content item type: totally_unknown_block" in message
assert "tool_reference" in message

View file

@ -0,0 +1,131 @@
"""
Integration / regression tests for Anthropic tool-search (`tool_reference`)
content blocks on the cost-calculation and streaming-assembly paths used by
Claude Code.
Claude Code's tool-search feature emits assistant content blocks of the form
``{"type": "tool_reference", "tool_name": ...}`` -- a lightweight pointer to a
deferred tool. Before the fix, `token_counter` did not recognise this block
type and raised ``Invalid content item type: tool_reference``.
Why this matters (the bug these tests guard against):
* On the cost path, that exception propagates out of ``completion_cost`` ->
``response_cost_calculator``. The proxy logging layer catches it and nulls
``response_cost``; the spend-tracking callback then skips the request, so
the entire SpendLogs row is dropped. The request succeeds for the caller
but the spend is silently never recorded -- a cost undercount on ALL
tool-search traffic.
* On the streaming-assembly path, ``stream_chunk_builder`` recomputes the
prompt tokens from the request messages when the provider stream does not
carry usage. The same exception there was swallowed and prompt tokens
silently collapsed to 0 -- a quieter undercount of the same traffic.
These tests exercise the real public entry points (not the private
``_count_content_list`` helper) so the whole chain is covered end to end.
"""
import os
import sys
sys.path.insert(0, os.path.abspath("../../.."))
import litellm
from litellm import stream_chunk_builder
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
ANTHROPIC_MODEL = "anthropic/claude-sonnet-4-5-20250929"
# Mirrors a Claude Code tool-search turn: a normal text block followed by a
# `tool_reference` pointer to a deferred tool.
TOOL_SEARCH_MESSAGES = [
{
"role": "assistant",
"content": [
{"type": "text", "text": "Let me look up the right tool."},
{"type": "tool_reference", "tool_name": "search_knowledge_base"},
],
}
]
def test_completion_cost_with_tool_reference_records_spend():
"""
``completion_cost`` must return a real, positive cost for messages that
contain a tool-search ``tool_reference`` block.
This is the exact chain that fails on the streaming anthropic_messages
proxy path: before the fix ``completion_cost`` raised, the logging layer
caught the exception and set ``response_cost = None``, and the spend
callback then dropped the SpendLogs row. A positive cost here means the
row is recorded instead of silently dropped.
"""
cost = litellm.completion_cost(model=ANTHROPIC_MODEL, messages=TOOL_SEARCH_MESSAGES)
assert cost is not None, "response_cost is None -> SpendLogs row would be dropped"
assert cost > 0, f"Expected a positive cost for tool-search traffic, got {cost}"
def test_completion_cost_with_empty_tool_name_records_spend():
"""A ``tool_reference`` with an empty/missing ``tool_name`` must also cost
out cleanly rather than raising and nulling the spend."""
messages = [
{
"role": "assistant",
"content": [{"type": "tool_reference", "tool_name": ""}],
}
]
cost = litellm.completion_cost(model=ANTHROPIC_MODEL, messages=messages)
assert cost is not None
assert cost >= 0
def test_stream_chunk_builder_counts_prompt_tokens_for_tool_reference():
"""
On the streaming-assembly path used by Claude Code, when the provider
stream carries no prompt-token usage, ``stream_chunk_builder`` recomputes
prompt tokens from the request messages via ``token_counter``.
With a ``tool_reference`` block in those messages the count must be
positive. Before the fix the underlying ``token_counter`` call raised and
the assembler swallowed it, collapsing ``prompt_tokens`` to 0 -- a silent
undercount of every tool-search request.
"""
model = "claude-sonnet-4-5-20250929"
# Chunks deliberately carry no usage, forcing the prompt-token fallback.
chunks = [
ModelResponseStream(
id="chatcmpl-tool-search",
created=1700000000,
model=model,
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content="Searching...", role="assistant"),
)
],
),
ModelResponseStream(
id="chatcmpl-tool-search",
created=1700000000,
model=model,
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop", index=0, delta=Delta(content="")
),
],
),
]
response = stream_chunk_builder(chunks, messages=TOOL_SEARCH_MESSAGES)
assert response is not None
assert (
response.usage.prompt_tokens > 0
), "prompt_tokens collapsed to 0 -> tool-search traffic silently undercounted"

View file

@ -3702,6 +3702,39 @@ def test_fast_mode_with_inference_geo():
assert abs(completion_cost - base_completion * expected_multiplier) < 1e-10
def test_calculate_usage_captures_service_tier():
"""
Anthropic returns the assigned service tier on the response usage object
(e.g. ``"priority"``). It must be surfaced on the Usage object so it is
visible in logs and used to select tier-specific pricing.
"""
config = AnthropicConfig()
usage_object = {
"input_tokens": 410,
"cache_creation_input_tokens": 0,
"cache_read_input_tokens": 0,
"output_tokens": 585,
"service_tier": "priority",
}
usage = config.calculate_usage(usage_object=usage_object, reasoning_content=None)
assert usage.service_tier == "priority"
def test_calculate_usage_service_tier_defaults_to_none():
"""A response without a service tier must not invent one."""
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={"input_tokens": 10, "output_tokens": 5},
reasoning_content=None,
)
assert usage.service_tier is None
def test_fast_mode_parameter_in_supported_params():
"""
Test that 'speed' is in the list of supported OpenAI params.

View file

@ -169,8 +169,8 @@ def test_hosted_vllm_supports_thinking():
def test_hosted_vllm_thinking_blocks_prepended_to_assistant_content():
"""
Test that thinking_blocks on assistant messages are converted to content
blocks prepended before the existing content.
Test that thinking_blocks on assistant messages are removed and content
stays a string for vLLM compatibility.
"""
config = HostedVLLMChatConfig()
messages = [
@ -203,21 +203,15 @@ def test_hosted_vllm_thinking_blocks_prepended_to_assistant_content():
)
assistant_msg = transformed["messages"][1]
assert assistant_msg["role"] == "assistant"
assert isinstance(assistant_msg["content"], list)
assert assistant_msg["content"][0] == {
"type": "thinking",
"thinking": "Let me reason about this...",
}
assert assistant_msg["content"][1] == {
"type": "text",
"text": "Here is my answer.",
}
assert isinstance(assistant_msg["content"], str)
assert assistant_msg["content"] == "Here is my answer."
assert "thinking_blocks" not in assistant_msg
def test_hosted_vllm_thinking_blocks_with_list_content():
"""
Test thinking_blocks prepended when assistant content is already a list.
Test thinking_blocks are removed and assistant content list is converted
to a string.
"""
config = HostedVLLMChatConfig()
messages = [
@ -246,19 +240,125 @@ def test_hosted_vllm_thinking_blocks_with_list_content():
headers={},
)
assistant_msg = transformed["messages"][0]
assert len(assistant_msg["content"]) == 3
assert assistant_msg["content"][0] == {
"type": "thinking",
"thinking": "Step 1 reasoning",
}
assert assistant_msg["content"][1] == {
"type": "thinking",
"thinking": "Step 2 reasoning",
}
assert assistant_msg["content"][2] == {"type": "text", "text": "Response text"}
assert isinstance(assistant_msg["content"], str)
assert assistant_msg["content"] == "Response text"
assert "thinking_blocks" not in assistant_msg
def test_hosted_vllm_assistant_structured_content_is_preserved():
config = HostedVLLMChatConfig()
image_block = {
"type": "image_url",
"image_url": {"url": "https://example.com/image.png"},
}
messages = [
{
"role": "assistant",
"content": [{"type": "text", "text": "Here is the image"}, image_block],
},
]
transformed = config.transform_request(
model="hosted_vllm/llama-3.1-70b-instruct",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)
assistant_msg = transformed["messages"][0]
assert assistant_msg["content"] == [
{"type": "text", "text": "Here is the image"},
image_block,
]
def test_hosted_vllm_assistant_tool_use_content_becomes_tool_calls():
config = HostedVLLMChatConfig()
messages = [
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_1",
"name": "get_weather",
"input": {"city": "Boston"},
}
],
},
]
transformed = config.transform_request(
model="hosted_vllm/llama-3.1-70b-instruct",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)
assistant_msg = transformed["messages"][0]
assert assistant_msg["content"] == ""
assert assistant_msg["tool_calls"] == [
{
"id": "toolu_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": json.dumps({"city": "Boston"}),
},
}
]
def test_hosted_vllm_assistant_tool_use_does_not_duplicate_existing_tool_calls():
config = HostedVLLMChatConfig()
messages = [
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_1",
"name": "get_weather",
"input": {"city": "Boston"},
}
],
"tool_calls": [
{
"id": "toolu_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": json.dumps({"city": "Boston"}),
},
}
],
},
]
transformed = config.transform_request(
model="hosted_vllm/llama-3.1-70b-instruct",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)
assistant_msg = transformed["messages"][0]
assert assistant_msg["content"] == ""
assert assistant_msg["tool_calls"] == [
{
"id": "toolu_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": json.dumps({"city": "Boston"}),
},
}
]
def test_hosted_vllm_custom_tools_are_converted_to_function_tools():
config = HostedVLLMChatConfig()
optional_params = config.map_openai_params(

View file

@ -0,0 +1,106 @@
"""
Tests for OpenAIWhisperAudioTranscriptionConfig.transform_audio_transcription_request
and transform_audio_transcription_response.
"""
import io
import json
from unittest.mock import MagicMock
import pytest
from litellm.llms.openai.transcriptions.whisper_transformation import (
OpenAIWhisperAudioTranscriptionConfig,
)
class TestWhisperTransformRequestResponseFormat:
def _transform(self, optional_params: dict) -> dict:
config = OpenAIWhisperAudioTranscriptionConfig()
audio_file = io.BytesIO(b"fake audio")
audio_file.name = "test.wav"
result = config.transform_audio_transcription_request(
model="whisper-1",
audio_file=audio_file,
optional_params=optional_params,
litellm_params={},
)
return result.data
def test_defaults_to_verbose_json_when_unset(self):
"""When response_format is not specified, default to verbose_json for cost calculation."""
data = self._transform({})
assert data["response_format"] == "verbose_json"
def test_respects_explicit_json(self):
"""When response_format='json' is set, do not override to verbose_json."""
data = self._transform({"response_format": "json"})
assert data["response_format"] == "json"
def test_respects_explicit_text(self):
"""When response_format='text' is set, do not override to verbose_json."""
data = self._transform({"response_format": "text"})
assert data["response_format"] == "text"
def test_preserves_verbose_json_when_set(self):
"""verbose_json explicitly set by the caller stays as-is."""
data = self._transform({"response_format": "verbose_json"})
assert data["response_format"] == "verbose_json"
class TestWhisperTransformResponse:
def _make_response(self, *, text: str, content_type: str, is_json: bool):
mock = MagicMock()
mock.headers = {"content-type": content_type}
if is_json:
mock.json.return_value = {"text": text}
else:
mock.json.side_effect = json.JSONDecodeError("", "", 0)
mock.text = text
return mock
def test_parses_json_response(self):
"""JSON body (verbose_json or json format) is parsed into TranscriptionResponse."""
config = OpenAIWhisperAudioTranscriptionConfig()
result = config.transform_audio_transcription_response(
self._make_response(
text="Hello world", content_type="application/json", is_json=True
)
)
assert result.text == "Hello world"
def test_parses_plain_text_response(self):
"""Plain-text body (response_format=text) is returned as TranscriptionResponse without error."""
config = OpenAIWhisperAudioTranscriptionConfig()
result = config.transform_audio_transcription_response(
self._make_response(
text="Four score and seven years ago",
content_type="text/plain",
is_json=False,
)
)
assert result.text == "Four score and seven years ago"
def test_malformed_json_body_with_json_content_type_raises(self):
"""A non-JSON body labelled application/json is a genuine upstream error, not a transcription."""
config = OpenAIWhisperAudioTranscriptionConfig()
with pytest.raises(json.JSONDecodeError):
config.transform_audio_transcription_response(
self._make_response(
text="<html>502 Bad Gateway</html>",
content_type="application/json",
is_json=False,
)
)
def test_json_content_type_match_is_case_insensitive(self):
"""Media types are case-insensitive (RFC 7231), so a mixed-case application/json still re-raises."""
config = OpenAIWhisperAudioTranscriptionConfig()
with pytest.raises(json.JSONDecodeError):
config.transform_audio_transcription_response(
self._make_response(
text="<html>502 Bad Gateway</html>",
content_type="Application/JSON; charset=utf-8",
is_json=False,
)
)

View file

@ -553,3 +553,68 @@ def test_openrouter_non_reasoning_models_do_not_add_reasoning_effort():
)
assert "reasoning_effort" not in supported_params
def test_openrouter_reasoning_effort_max_maps_to_xhigh():
"""
OpenRouter expects 'xhigh' instead of 'max' for reasoning_effort.
"""
config = OpenrouterConfig()
result = config.map_openai_params(
non_default_params={"reasoning_effort": "max"},
optional_params={},
model="openrouter/deepseek/deepseek-r1",
drop_params=False,
)
assert result["reasoning_effort"] == "xhigh"
def test_openrouter_reasoning_effort_max_does_not_mutate_caller_dict():
"""
map_openai_params must not mutate the caller-supplied non_default_params dict.
"""
config = OpenrouterConfig()
original_params = {"reasoning_effort": "max"}
config.map_openai_params(
non_default_params=original_params,
optional_params={},
model="openrouter/deepseek/deepseek-r1",
drop_params=False,
)
assert original_params["reasoning_effort"] == "max"
def test_openrouter_reasoning_effort_xhigh_passes_through():
"""
reasoning_effort='xhigh' should be forwarded unchanged.
"""
config = OpenrouterConfig()
result = config.map_openai_params(
non_default_params={"reasoning_effort": "xhigh"},
optional_params={},
model="openrouter/deepseek/deepseek-r1",
drop_params=False,
)
assert result["reasoning_effort"] == "xhigh"
def test_openrouter_reasoning_effort_high_passes_through():
"""
Non-max reasoning_effort values should be forwarded unchanged.
"""
config = OpenrouterConfig()
result = config.map_openai_params(
non_default_params={"reasoning_effort": "high"},
optional_params={},
model="openrouter/deepseek/deepseek-r1",
drop_params=False,
)
assert result["reasoning_effort"] == "high"

View file

@ -4996,3 +4996,146 @@ def test_mid_stream_429_error_raises_during_iteration():
# Verify: 429 error is properly raised
assert exc_info.value.status_code == 429
assert "RESOURCE_EXHAUSTED" in str(exc_info.value.message)
class TestModelResponseIteratorCleanup:
def _make_logging_obj(self):
from unittest.mock import Mock
obj = Mock()
obj.optional_params = {}
return obj
def test_aclose_closes_iterator_and_response(self):
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
mock_iterator = MagicMock()
mock_iterator.aclose = AsyncMock()
iterator = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=False,
logging_obj=self._make_logging_obj(),
response=mock_response,
)
iterator.async_response_iterator = mock_iterator
asyncio.run(iterator.aclose())
mock_iterator.aclose.assert_awaited_once()
mock_response.aclose.assert_awaited_once()
def test_close_closes_iterator_and_response(self):
from unittest.mock import MagicMock
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_response = MagicMock()
mock_iterator = MagicMock()
iterator = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=True,
logging_obj=self._make_logging_obj(),
response=mock_response,
)
iterator.response_iterator = mock_iterator
iterator.close()
mock_iterator.close.assert_called_once()
mock_response.close.assert_called_once()
def test_aclose_without_response_does_not_raise(self):
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_iterator = MagicMock()
mock_iterator.aclose = AsyncMock()
iterator = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=False,
logging_obj=self._make_logging_obj(),
)
iterator.async_response_iterator = mock_iterator
asyncio.run(iterator.aclose())
mock_iterator.aclose.assert_awaited_once()
def test_aclose_tolerates_iterator_error(self):
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
mock_iterator = MagicMock()
mock_iterator.aclose = AsyncMock(side_effect=RuntimeError("transport error"))
iterator = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=False,
logging_obj=self._make_logging_obj(),
response=mock_response,
)
iterator.async_response_iterator = mock_iterator
asyncio.run(iterator.aclose())
mock_response.aclose.assert_awaited_once()
def test_custom_stream_wrapper_aclose_triggers_model_response_iterator_aclose(self):
"""CustomStreamWrapper.aclose() must propagate to ModelResponseIterator.aclose()."""
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
mock_iterator = MagicMock()
mock_iterator.aclose = AsyncMock()
model_response_iter = ModelResponseIterator(
streaming_response=MagicMock(),
sync_stream=False,
logging_obj=self._make_logging_obj(),
response=mock_response,
)
model_response_iter.async_response_iterator = mock_iterator
wrapper = CustomStreamWrapper(
completion_stream=model_response_iter,
model="gemini-2.0-flash",
custom_llm_provider="vertex_ai",
logging_obj=MagicMock(),
)
asyncio.run(wrapper.aclose())
mock_iterator.aclose.assert_awaited_once()
mock_response.aclose.assert_awaited_once()

View file

@ -5179,6 +5179,12 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow():
side_effect=lambda update: MCPServer(
server_id=legacy_server.server_id,
name=legacy_server.name,
# Carry alias/server_name forward so get_server_prefix resolves to
# "legacy_m2m" (not the server_id) when the request scope filter
# matches by alias. Without these, the filter relied on the now-
# removed silent fail-open fallback.
alias=legacy_server.alias,
server_name=legacy_server.server_name,
transport=MCPTransport.http,
auth_type=legacy_server.auth_type,
oauth2_flow=update.get("oauth2_flow", legacy_server.oauth2_flow),
@ -6083,3 +6089,207 @@ async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_serv
assert exc_info.value.status_code == 403
assert exc_info.value.detail["error"] == "tool_server_mismatch"
# ---------------------------------------------------------------------------
# Regression tests for _get_allowed_mcp_servers_from_mcp_server_names
#
# Prior to the fail-closed fix, an unresolved scope filter (path- or
# header-derived) silently returned the caller's full allowed-server set,
# which made URL/header namespacing appear to work when it did not.
# ---------------------------------------------------------------------------
def _make_mcp_server_for_scope_filter(server_id: str, alias: str) -> MCPServer:
return MCPServer(
server_id=server_id,
name=alias,
alias=alias,
server_name=alias,
url=f"https://{alias}.test/mcp",
transport=MCPTransport.http,
mcp_info={"server_name": alias},
)
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_from_mcp_server_names_unknown_name_fails_closed():
"""
Bug fix: requesting an unknown server name (e.g. ``/mcp/<typo>/``) must
NOT silently fall back to the caller's full allowed-server set.
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers_from_mcp_server_names,
)
except ImportError:
pytest.skip("MCP server not available")
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
):
result = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=["does-not-exist"],
allowed_mcp_servers=allowed,
)
assert result == []
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_from_mcp_server_names_none_returns_all():
"""
Regression: ``mcp_servers=None`` (no scope filter requested) must still
return the full allowed-server set. This is the legitimate "no scoping"
path that the fail-closed fix must not break.
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers_from_mcp_server_names,
)
except ImportError:
pytest.skip("MCP server not available")
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
result = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=None,
allowed_mcp_servers=allowed,
)
assert {s.server_id for s in result} == {"id-a", "id-b"}
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_from_mcp_server_names_known_alias_returns_match():
"""
Regression: a known server alias must still resolve to exactly that
server. Guards against the fix accidentally narrowing the happy path.
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers_from_mcp_server_names,
)
except ImportError:
pytest.skip("MCP server not available")
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
):
result = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=["alpha"],
allowed_mcp_servers=allowed,
)
assert [s.server_id for s in result] == ["id-a"]
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown():
"""
Mixed scope (one valid + one unknown) returns only the resolved server,
not the full allowed set. Confirms the fail-closed branch only fires
when NOTHING resolves.
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers_from_mcp_server_names,
)
except ImportError:
pytest.skip("MCP server not available")
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
):
result = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=["alpha", "does-not-exist"],
allowed_mcp_servers=allowed,
)
assert [s.server_id for s in result] == ["id-a"]
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_from_mcp_server_names_access_group_resolves():
"""
Regression: when a requested name is not a server alias but IS an access
group, it must still resolve to the underlying servers (not be treated
as unresolved).
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers_from_mcp_server_names,
)
except ImportError:
pytest.skip("MCP server not available")
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=["id-b"],
):
result = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=["group-name"],
allowed_mcp_servers=allowed,
)
assert [s.server_id for s in result] == ["id-b"]
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_from_mcp_server_names_empty_list_fails_closed():
"""
Edge case: ``mcp_servers=[]`` (explicit empty scope) is still an
explicit filter request. Fail closed rather than returning everything.
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers_from_mcp_server_names,
)
except ImportError:
pytest.skip("MCP server not available")
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
result = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=[],
allowed_mcp_servers=allowed,
)
assert result == []

View file

@ -388,6 +388,60 @@ def test_wildcard_credential_hydration_preserves_deployment_params(
}
def test_wildcard_custom_prefix_does_not_stack_provider_prefix(monkeypatch):
"""Regression test for #30358.
A wildcard with a custom prefix (e.g. ``ollama_server1/*`` to distinguish multiple Ollama
instances) must not stack the provider's own prefix onto the expanded model ids. The expanded
ids should be ``ollama_server1/gemma3:1b`` rather than ``ollama_server1/ollama/gemma3:1b``.
"""
from litellm.proxy.auth import model_checks
from litellm.proxy.auth.model_checks import get_known_models_from_wildcard
from litellm.types.router import LiteLLM_Params
monkeypatch.setattr(
model_checks,
"get_provider_models",
lambda provider, litellm_params=None: ["ollama/gemma3:1b", "ollama/llama3:8b"],
)
result = get_known_models_from_wildcard(
wildcard_model="ollama_server1/*",
litellm_params=LiteLLM_Params(
model="ollama_chat/*", custom_llm_provider="ollama_chat"
),
)
assert result == ["ollama_server1/gemma3:1b", "ollama_server1/llama3:8b"]
def test_wildcard_custom_prefix_keeps_org_segment_for_non_provider_first_segment(
monkeypatch,
):
"""Only a known provider prefix should be stripped before re-prefixing.
If ``get_provider_models`` returns ids whose first segment is an org rather than a litellm
provider (e.g. ``meta-llama/Llama-3-8B``), stripping the first slash segment would drop the
org and produce an uncallable id. The org segment must be preserved.
"""
from litellm.proxy.auth import model_checks
from litellm.proxy.auth.model_checks import get_known_models_from_wildcard
from litellm.types.router import LiteLLM_Params
monkeypatch.setattr(
model_checks,
"get_provider_models",
lambda provider, litellm_params=None: ["meta-llama/Llama-3-8B"],
)
result = get_known_models_from_wildcard(
wildcard_model="my_hf/*",
litellm_params=LiteLLM_Params(model="huggingface/*", custom_llm_provider="huggingface"),
)
assert result == ["my_hf/meta-llama/Llama-3-8B"]
def test_wildcard_credential_hydration_preserves_missing_credential_name(
monkeypatch,
):

View file

@ -1,7 +1,7 @@
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -11,7 +11,6 @@ sys.path.insert(
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer
from litellm.proxy.proxy_server import ProxyStartupEvent
from litellm.types.caching import RedisPipelineRpushOperation
@pytest.fixture
@ -305,3 +304,73 @@ def test_validate_redis_transaction_buffer_passes_when_disabled():
general_settings={},
redis_usage_cache=None,
)
def test_get_transaction_buffer_redis_cache_builds_from_env(monkeypatch):
"""
When use_redis_transaction_buffer=true, a standalone RedisCache is built from
REDIS_* environment variables so the buffer works without a Redis cache backend.
"""
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
with patch("litellm.proxy.proxy_server.RedisCache") as mock_redis_cache:
result = ProxyStartupEvent._get_transaction_buffer_redis_cache(
general_settings={"use_redis_transaction_buffer": True},
)
mock_redis_cache.assert_called_once()
assert mock_redis_cache.call_args.kwargs["host"] == "localhost"
assert result is mock_redis_cache.return_value
def test_get_transaction_buffer_redis_cache_none_when_disabled():
"""When use_redis_transaction_buffer is not enabled, no standalone cache is built."""
result = ProxyStartupEvent._get_transaction_buffer_redis_cache(
general_settings={},
)
assert result is None
def test_get_transaction_buffer_redis_cache_none_without_redis_env():
"""
When use_redis_transaction_buffer=true but no REDIS_* env vars are set,
no standalone cache is built (startup validation then raises the config error).
"""
with patch("litellm._redis._redis_kwargs_from_environment", return_value={}):
result = ProxyStartupEvent._get_transaction_buffer_redis_cache(
general_settings={"use_redis_transaction_buffer": True},
)
assert result is None
def test_get_transaction_buffer_redis_cache_none_without_host_or_url():
"""
A REDIS_* var that is not a connection target (e.g. REDIS_SOCKET_TIMEOUT) must not
trigger a build. Without a host or url, get_redis_client raises, so return None and
let startup validation surface the config error instead of crashing.
"""
with patch(
"litellm._redis._redis_kwargs_from_environment",
return_value={"socket_timeout": 5.0},
):
result = ProxyStartupEvent._get_transaction_buffer_redis_cache(
general_settings={"use_redis_transaction_buffer": True},
)
assert result is None
def test_get_transaction_buffer_redis_cache_parses_string_flag(monkeypatch):
"""
use_redis_transaction_buffer accepts a string value (e.g. from env/YAML); "true"
is parsed to a bool before the standalone cache is built.
"""
monkeypatch.setenv("REDIS_HOST", "localhost")
with patch("litellm.proxy.proxy_server.RedisCache") as mock_redis_cache:
result = ProxyStartupEvent._get_transaction_buffer_redis_cache(
general_settings={"use_redis_transaction_buffer": "true"},
)
mock_redis_cache.assert_called_once()
assert result is mock_redis_cache.return_value

View file

@ -352,6 +352,79 @@ def test_ui_discovery_endpoints_is_control_plane_true_when_workers_configured():
assert data["workers"][0]["url"] == "https://worker-1:4001"
def test_ui_discovery_endpoints_hide_default_credentials_hint_default_false():
"""Default credentials hint is shown by default (flag false)."""
app = FastAPI()
app.include_router(router)
client = TestClient(app)
with (
patch("litellm.proxy.utils.get_server_root_path", return_value="/"),
patch("litellm.proxy.utils.get_proxy_base_url", return_value=None),
patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False),
patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}, clear=False),
):
os.environ.pop("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", None)
response = client.get("/.well-known/litellm-ui-config")
assert response.status_code == 200
data = response.json()
assert data["hide_default_credentials_hint"] is False
def test_ui_discovery_endpoints_hide_default_credentials_hint_via_env_var():
"""LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT=true hides the login-page credentials card."""
app = FastAPI()
app.include_router(router)
client = TestClient(app)
with (
patch("litellm.proxy.utils.get_server_root_path", return_value="/"),
patch("litellm.proxy.utils.get_proxy_base_url", return_value=None),
patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False),
patch.dict(
os.environ,
{
"LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT": "true",
"DISABLE_ADMIN_UI": "false",
},
clear=False,
),
):
response = client.get("/.well-known/litellm-ui-config")
assert response.status_code == 200
data = response.json()
assert data["hide_default_credentials_hint"] is True
def test_ui_discovery_endpoints_hide_default_credentials_hint_via_general_settings():
"""general_settings.hide_default_credentials_hint=true also hides the card."""
app = FastAPI()
app.include_router(router)
client = TestClient(app)
with (
patch("litellm.proxy.utils.get_server_root_path", return_value="/"),
patch("litellm.proxy.utils.get_proxy_base_url", return_value=None),
patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False),
patch(
"litellm.proxy.proxy_server.general_settings",
{"hide_default_credentials_hint": True},
),
patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}, clear=False),
):
os.environ.pop("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", None)
response = client.get("/.well-known/litellm-ui-config")
assert response.status_code == 200
data = response.json()
assert data["hide_default_credentials_hint"] is True
def test_ui_discovery_endpoints_is_control_plane_false_when_no_workers():
app = FastAPI()
app.include_router(router)

View file

@ -2,6 +2,7 @@
"""
Test to verify the Google GenAI proxy API endpoints
"""
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -88,6 +89,8 @@ def test_google_stream_generate_content_endpoint():
# stream=True must be forced into the data the processor receives.
init_kwargs = mock_init.call_args.kwargs
assert init_kwargs["data"]["stream"] is True
assert init_kwargs["data"]["_litellm_raw_sse_stream"] is True
assert init_kwargs["data"]["_litellm_skip_openai_stream_done"] is True
assert init_kwargs["data"]["model"] == "test-model"
assert init_kwargs["data"]["contents"] == [
{"role": "user", "parts": [{"text": "Hello"}]}

View file

@ -584,6 +584,50 @@ async def test_logging_hook_multiple_content_items(presidio_guardrail):
print("✓ Logging hook multiple content items test passed")
@pytest.mark.asyncio
async def test_logging_only_does_not_mask_pre_call_request(
mock_user_api_key, mock_cache
):
"""
A guardrail configured with `logging_only` must only mask PII for logs/traces,
never for the request sent to the model. `async_pre_call_hook` should leave the
request untouched so the model receives (and replies based on) the real input.
Regression test for the case where the pre-call hook masked the live request,
causing the model's response to contain anonymization tokens (e.g. <PERSON>)
instead of the real output.
"""
presidio_guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
logging_only=True,
pii_entities_config={PiiEntityType.PHONE_NUMBER: PiiAction.MASK},
)
async def mock_check_pii(text, output_parse_pii, presidio_config, request_data):
return text.replace("555-123-4567", "[PHONE]")
presidio_guardrail.check_pii = mock_check_pii
original_text = "My phone is 555-123-4567"
test_data = {
"messages": [{"role": "user", "content": original_text}],
"model": "gpt-4",
}
result = await presidio_guardrail.async_pre_call_hook(
user_api_key_dict=mock_user_api_key,
cache=mock_cache,
data=test_data,
call_type="completion",
)
# The live request must be unchanged: PII reaches the model intact.
assert result["messages"][0]["content"] == original_text
assert "[PHONE]" not in result["messages"][0]["content"]
print("✓ logging_only leaves the pre-call request unmasked")
@pytest.mark.asyncio
async def test_presidio_sets_guardrail_information_in_request_data():
"""

View file

@ -9,6 +9,8 @@ sys.path.insert(
) # Adds the parent directory to the system path
from litellm.proxy.management_endpoints.common_daily_activity import (
_adjust_dates_for_timezone,
_build_aggregated_sql_query,
_is_user_agent_tag,
get_api_key_metadata,
get_daily_activity,
@ -632,6 +634,126 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
assert key_data.metrics.spend == 10.0
class TestAdjustDatesForTimezone:
"""
Regression tests for the timezone double-counting bug.
Background: the previous implementation expanded the SQL date range by a full
UTC day on whichever side a non-UTC timezone offset pointed. Because spend is
bucketed in whole UTC days in the aggregation table, that expansion caused
single-day queries from non-UTC timezones to include a second full UTC day's
worth of data, producing approximately 2x over-counting. The sum of single-day
spends across a window then exceeded the equivalent multi-day aggregate, which
is mathematically impossible.
These tests pin the function to a pass-through and assert the additivity
invariant that any future implementation must preserve.
"""
@pytest.mark.parametrize(
"offset_minutes",
[
None,
0,
-330, # IST UTC+5:30
-540, # JST UTC+9
-60, # CET UTC+1
240, # AST UTC-4
300, # EST UTC-5
480, # PST UTC-8
],
)
def test_returns_input_dates_unchanged_for_any_offset(self, offset_minutes):
start, end = _adjust_dates_for_timezone(
"2026-05-29", "2026-05-29", offset_minutes
)
assert start == "2026-05-29"
assert end == "2026-05-29"
def test_single_day_query_does_not_widen_to_two_utc_days(self):
"""
Pins the boundary that caused the original 2x bug: a single IST day must
not be translated into a SQL filter covering two UTC days.
"""
start, end = _adjust_dates_for_timezone("2026-05-29", "2026-05-29", -330)
assert start == end == "2026-05-29", (
"Single-day IST query expanded to a multi-day UTC range; this is "
"the regression that produced approximately 2x over-counting."
)
def test_multi_day_range_endpoints_are_preserved(self):
start, end = _adjust_dates_for_timezone("2026-05-29", "2026-06-02", -330)
assert (start, end) == ("2026-05-29", "2026-06-02")
@pytest.mark.parametrize("offset_minutes", [-330, 480])
def test_single_day_sums_match_multi_day_window(self, offset_minutes):
"""
Additivity invariant: querying each day in a window separately and summing
the resulting SQL ranges must cover exactly the same range as querying the
whole window at once. The bug broke this; without it, single-day sums
exceeded the multi-day total by ~50% over a 5-day IST window.
"""
days = ["2026-05-29", "2026-05-30", "2026-05-31", "2026-06-01", "2026-06-02"]
single_day_ranges = [
_adjust_dates_for_timezone(d, d, offset_minutes) for d in days
]
multi_day_range = _adjust_dates_for_timezone(days[0], days[-1], offset_minutes)
per_day_starts = [r[0] for r in single_day_ranges]
per_day_ends = [r[1] for r in single_day_ranges]
assert min(per_day_starts) == multi_day_range[0]
assert max(per_day_ends) == multi_day_range[1]
assert per_day_starts == days
assert per_day_ends == days
class TestBuildAggregatedSqlQuery:
"""
Asserts the SQL emitted by the aggregated query path stays anchored to the
user-supplied date range. The original bug shipped a function that returned
expanded dates from _adjust_dates_for_timezone, so the regression surface is
not just the helper but the SQL it feeds into.
"""
@pytest.mark.parametrize("offset_minutes", [None, 0, -330, 480])
def test_sql_date_bounds_are_user_supplied_dates(self, offset_minutes):
sql, params = _build_aggregated_sql_query(
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id="user-1",
start_date="2026-05-29",
end_date="2026-05-29",
model=None,
api_key=None,
timezone_offset_minutes=offset_minutes,
)
assert params[0] == "2026-05-29"
assert params[1] == "2026-05-29"
assert "date >= $1" in sql
assert "date <= $2" in sql
def test_optional_filters_appear_in_params_in_order(self):
sql, params = _build_aggregated_sql_query(
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id="user-1",
start_date="2026-05-29",
end_date="2026-06-02",
model="bedrock/global.anthropic.claude-opus-4-8",
api_key="sk-test",
timezone_offset_minutes=-330,
)
assert params == [
"2026-05-29",
"2026-06-02",
"user-1",
"bedrock/global.anthropic.claude-opus-4-8",
"sk-test",
]
assert "model = $4" in sql
assert "api_key = $5" in sql
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_empty_result_set():
"""Regression test for the empty-range 500.

View file

@ -11862,7 +11862,6 @@ async def test_ghsa_q775_default_team_id_does_not_grant_session_token_exemption(
assert "cannot exceed" in msg.lower()
@pytest.mark.asyncio
async def test_prepare_key_update_data_budget_duration_null_clears_fields():
"""
@ -11941,3 +11940,511 @@ async def test_prepare_key_update_data_budget_duration_valid_sets_reset():
assert result["budget_reset_at"] is not None
@pytest.mark.asyncio
async def test_info_key_fn_includes_model_max_budget_usage(monkeypatch):
"""
/key/info should include model_max_budget_usage showing current-period spend
for each model that has a per-model budget configured.
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn
test_key_token = "hashed_token_budget_test"
model_max_budget = {
"gpt-4o": {"budget_limit": 0.50, "time_period": "1d"},
}
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.23)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
mock_key_info = MagicMock(spec=LiteLLM_VerificationToken)
mock_key_info.token = test_key_token
mock_key_info.object_permission_id = None
mock_key_info.user_id = "user-x"
mock_key_info.team_id = None
mock_key_info.litellm_budget_table = None
mock_key_info.model_dump.return_value = {
"token": test_key_token,
"model_max_budget": model_max_budget,
"user_id": "user-x",
"team_id": None,
"object_permission_id": None,
"litellm_budget_table": None,
}
mock_key_info.dict.return_value = mock_key_info.model_dump.return_value
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=mock_key_info
)
mock_prisma_client.db.query_raw = AsyncMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-test-budget-key",
)
result = await info_key_fn(
key="sk-test-budget-key",
user_api_key_dict=user_api_key_dict,
)
assert "model_max_budget_usage" in result["info"]
usage = result["info"]["model_max_budget_usage"]
assert usage["gpt-4o"]["current_spend"] == 0.23
assert usage["gpt-4o"]["budget_limit"] == 0.50
assert usage["gpt-4o"]["time_period"] == "1d"
mock_prisma_client.db.query_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_info_key_fn_no_model_max_budget_skips_usage(monkeypatch):
"""Keys with no model_max_budget should not include model_max_budget_usage."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn
test_key_token = "hashed_token_no_budget"
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock()
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
mock_key_info = MagicMock(spec=LiteLLM_VerificationToken)
mock_key_info.token = test_key_token
mock_key_info.object_permission_id = None
mock_key_info.user_id = "user-y"
mock_key_info.team_id = None
mock_key_info.litellm_budget_table = None
mock_key_info.model_dump.return_value = {
"token": test_key_token,
"model_max_budget": {},
"user_id": "user-y",
"team_id": None,
"object_permission_id": None,
"litellm_budget_table": None,
}
mock_key_info.dict.return_value = mock_key_info.model_dump.return_value
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=mock_key_info
)
mock_prisma_client.db.query_raw = AsyncMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-test-no-budget",
)
result = await info_key_fn(
key="sk-test-no-budget",
user_api_key_dict=user_api_key_dict,
)
assert "model_max_budget_usage" not in result["info"]
mock_prisma_client.db.query_raw.assert_not_awaited()
mock_user_api_key_cache.async_get_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_info_key_fn_v2_includes_model_max_budget_usage(monkeypatch):
"""/v2/key/info should include model_max_budget_usage for keys with per-model budgets."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import KeyRequest, LiteLLM_VerificationToken
from litellm.proxy.management_endpoints.key_management_endpoints import (
info_key_fn_v2,
)
test_key_token = "hashed_token_v2_test"
model_max_budget = {"gpt-4o": {"budget_limit": 1.00, "time_period": "7d"}}
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.55)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
mock_key = MagicMock(spec=LiteLLM_VerificationToken)
mock_key.token = test_key_token
mock_key.user_id = "user-v2"
mock_key.team_id = None
mock_key.model_dump.return_value = {
"token": test_key_token,
"model_max_budget": model_max_budget,
"user_id": "user-v2",
"team_id": None,
"litellm_budget_table": None,
}
mock_key.dict.return_value = mock_key.model_dump.return_value
mock_prisma_client.get_data = AsyncMock(return_value=[mock_key])
mock_prisma_client.db.query_raw = AsyncMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
)
result = await info_key_fn_v2(
data=KeyRequest(keys=[test_key_token]),
user_api_key_dict=user_api_key_dict,
)
assert len(result["info"]) == 1
key_info = result["info"][0]
assert "model_max_budget_usage" in key_info
usage = key_info["model_max_budget_usage"]
assert usage["gpt-4o"]["current_spend"] == 0.55
assert usage["gpt-4o"]["budget_limit"] == 1.00
assert usage["gpt-4o"]["time_period"] == "7d"
mock_prisma_client.db.query_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_info_key_fn_budget_table_fallback(monkeypatch):
"""When model_max_budget is empty on the key but set in litellm_budget_table,
/key/info should still populate model_max_budget_usage.
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn
test_key_token = "hashed_token_budget_table_test"
budget_table_model_max_budget = {
"bedrock/anthropic.claude-opus-4": {"max_budget": 5, "budget_duration": "30d"},
}
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=1.20)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
mock_key_info = MagicMock(spec=LiteLLM_VerificationToken)
mock_key_info.token = test_key_token
mock_key_info.object_permission_id = None
mock_key_info.user_id = "user-bt"
mock_key_info.team_id = None
mock_key_info.litellm_budget_table = None
mock_key_info.model_dump.return_value = {
"token": test_key_token,
"model_max_budget": {},
"user_id": "user-bt",
"team_id": None,
"object_permission_id": None,
"litellm_budget_table": {
"budget_id": "bt-123",
"budget_duration": "30d",
"budget_reset_at": "2026-07-01T00:00:00+00:00",
"model_max_budget": budget_table_model_max_budget,
},
}
mock_key_info.dict.return_value = mock_key_info.model_dump.return_value
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=mock_key_info
)
mock_prisma_client.db.query_raw = AsyncMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-test-bt-key",
)
result = await info_key_fn(
key="sk-test-bt-key",
user_api_key_dict=user_api_key_dict,
)
assert "model_max_budget_usage" in result["info"]
usage = result["info"]["model_max_budget_usage"]
assert usage["bedrock/anthropic.claude-opus-4"]["current_spend"] == 1.20
assert usage["bedrock/anthropic.claude-opus-4"]["budget_limit"] == 5
assert usage["bedrock/anthropic.claude-opus-4"]["time_period"] == "30d"
mock_prisma_client.db.query_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_info_key_fn_v2_budget_table_fallback(monkeypatch):
"""When model_max_budget is empty on the key but set in litellm_budget_table,
/v2/key/info should still populate model_max_budget_usage."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import KeyRequest, LiteLLM_VerificationToken
from litellm.proxy.management_endpoints.key_management_endpoints import (
info_key_fn_v2,
)
test_key_token = "hashed_token_v2_bt_test"
budget_table_model_max_budget = {
"bedrock/anthropic.claude-opus-4": {"max_budget": 5, "budget_duration": "30d"},
}
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=2.50)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
mock_key = MagicMock(spec=LiteLLM_VerificationToken)
mock_key.token = test_key_token
mock_key.user_id = "user-v2-bt"
mock_key.team_id = None
mock_key.model_dump.return_value = {
"token": test_key_token,
"model_max_budget": {},
"user_id": "user-v2-bt",
"team_id": None,
"litellm_budget_table": {
"budget_id": "bt-456",
"budget_duration": "30d",
"budget_reset_at": "2026-07-01T00:00:00+00:00",
"model_max_budget": budget_table_model_max_budget,
},
}
mock_key.dict.return_value = mock_key.model_dump.return_value
mock_prisma_client.get_data = AsyncMock(return_value=[mock_key])
mock_prisma_client.db.query_raw = AsyncMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin-v2-bt",
)
result = await info_key_fn_v2(
data=KeyRequest(keys=[test_key_token]),
user_api_key_dict=user_api_key_dict,
)
assert len(result["info"]) == 1
key_info = result["info"][0]
assert "model_max_budget_usage" in key_info
usage = key_info["model_max_budget_usage"]
assert usage["bedrock/anthropic.claude-opus-4"]["current_spend"] == 2.50
mock_prisma_client.db.query_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_info_key_fn_provider_prefix_spend_fallback(monkeypatch):
"""Cached spend for 'gpt-4o' matches budget key 'openai/gpt-4o' via suffix match."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn
test_key_token = "hashed_token_prefix_test"
model_max_budget = {
"openai/gpt-4o": {"budget_limit": 2.00, "time_period": "7d"},
}
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(side_effect=[None, 0.75])
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
mock_key_info = MagicMock(spec=LiteLLM_VerificationToken)
mock_key_info.token = test_key_token
mock_key_info.object_permission_id = None
mock_key_info.user_id = "user-prefix"
mock_key_info.team_id = None
mock_key_info.litellm_budget_table = None
mock_key_info.model_dump.return_value = {
"token": test_key_token,
"model_max_budget": model_max_budget,
"user_id": "user-prefix",
"team_id": None,
"object_permission_id": None,
"litellm_budget_table": None,
}
mock_key_info.dict.return_value = mock_key_info.model_dump.return_value
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=mock_key_info
)
mock_prisma_client.db.query_raw = AsyncMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-prefix-test",
)
result = await info_key_fn(
key="sk-prefix-test",
user_api_key_dict=user_api_key_dict,
)
assert "model_max_budget_usage" in result["info"]
usage = result["info"]["model_max_budget_usage"]
assert usage["openai/gpt-4o"]["current_spend"] == 0.75
assert mock_user_api_key_cache.async_get_cache.await_count == 2
@pytest.mark.asyncio
async def test_build_model_max_budget_usage_no_cache_returns_empty():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_model_max_budget_usage,
)
result = await _build_model_max_budget_usage(
api_key_hash="some-hash",
model_max_budget={"gpt-4o": {"budget_limit": 1.0, "time_period": "1d"}},
user_api_key_cache=None,
)
assert result == {}
@pytest.mark.asyncio
async def test_build_model_max_budget_usage_reads_current_cache_window():
from unittest.mock import AsyncMock
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_model_max_budget_usage,
)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.30)
result = await _build_model_max_budget_usage(
api_key_hash="some-hash",
model_max_budget={"gpt-4o": {"budget_limit": 1.0, "time_period": "30d"}},
user_api_key_cache=mock_user_api_key_cache,
)
assert result["gpt-4o"]["current_spend"] == 0.30
mock_user_api_key_cache.async_get_cache.assert_awaited_once_with(
key="virtual_key_spend:some-hash:gpt-4o:30d"
)
@pytest.mark.asyncio
async def test_build_model_max_budget_usage_no_duration_in_budget_returns_empty():
from unittest.mock import AsyncMock
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_model_max_budget_usage,
)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock()
result = await _build_model_max_budget_usage(
api_key_hash="some-hash",
model_max_budget={"gpt-4o": {"budget_limit": 1.0}},
user_api_key_cache=mock_user_api_key_cache,
)
assert result == {}
mock_user_api_key_cache.async_get_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_build_model_max_budget_usage_skips_model_without_duration():
from unittest.mock import AsyncMock
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_model_max_budget_usage,
)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.10)
result = await _build_model_max_budget_usage(
api_key_hash="some-hash",
model_max_budget={
"gpt-4o": {"budget_limit": 1.0, "time_period": "1d"},
"gpt-3.5-turbo": {"budget_limit": 0.5},
},
user_api_key_cache=mock_user_api_key_cache,
)
assert "gpt-4o" in result
assert "gpt-3.5-turbo" not in result
assert mock_user_api_key_cache.async_get_cache.await_count == 1
@pytest.mark.asyncio
async def test_build_model_max_budget_usage_unparseable_duration_skipped():
from unittest.mock import AsyncMock
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_model_max_budget_usage,
)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock()
result = await _build_model_max_budget_usage(
api_key_hash="some-hash",
model_max_budget={
"gpt-4o": {"budget_limit": 1.0, "budget_duration": "not-valid"}
},
user_api_key_cache=mock_user_api_key_cache,
)
assert result == {}
mock_user_api_key_cache.async_get_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_build_model_max_budget_usage_invalid_budget_config_skipped():
from unittest.mock import AsyncMock
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_model_max_budget_usage,
)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.20)
result = await _build_model_max_budget_usage(
api_key_hash="some-hash",
model_max_budget={
"gpt-4o": {"max_budget": "not-a-number", "budget_duration": "1d"},
"gpt-3.5-turbo": {"budget_limit": 0.5, "time_period": "7d"},
},
user_api_key_cache=mock_user_api_key_cache,
)
assert "gpt-4o" not in result
assert "gpt-3.5-turbo" in result
assert mock_user_api_key_cache.async_get_cache.await_count == 1
@pytest.mark.asyncio
async def test_build_model_max_budget_usage_provider_prefix_cache_fallback():
from unittest.mock import AsyncMock
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_model_max_budget_usage,
)
mock_user_api_key_cache = AsyncMock()
mock_user_api_key_cache.async_get_cache = AsyncMock(side_effect=[None, 0.55])
result = await _build_model_max_budget_usage(
api_key_hash="test-hash",
model_max_budget={"openai/gpt-4o": {"budget_limit": 2.0, "time_period": "7d"}},
user_api_key_cache=mock_user_api_key_cache,
)
assert result["openai/gpt-4o"]["current_spend"] == 0.55
assert mock_user_api_key_cache.async_get_cache.await_count == 2

View file

@ -201,6 +201,58 @@ def test_anthropic_provider_fields_support_byok():
), "api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)."
def test_google_ai_studio_provider_fields_expose_api_base():
"""The Google AI Studio (gemini) credential form must let admins set a custom
api_base so they can point at a Gemini-compatible gateway (e.g. a self-hosted
proxy at /v1beta) without env var access.
The runtime gemini provider already supports custom api_base via
`vertex_llm_base._check_custom_proxy`; the UI just needs to expose the field.
"""
app_instance = FastAPI()
app_instance.include_router(router)
test_client = TestClient(app_instance)
response = test_client.get("/public/providers/fields")
assert response.status_code == 200
providers = response.json()
google_ai = next(
(p for p in providers if p["provider"] == "Google_AI_Studio"), None
)
assert google_ai is not None, "Google_AI_Studio provider entry not found"
assert google_ai["litellm_provider"] == "gemini"
fields_by_key = {f["key"]: f for f in google_ai["credential_fields"]}
assert "api_key" in fields_by_key
assert "api_base" in fields_by_key, (
"Google_AI_Studio provider form must expose api_base so admins can "
"point at a Gemini-compatible gateway without env var access."
)
api_base_field = fields_by_key["api_base"]
assert api_base_field["required"] is False
assert api_base_field["field_type"] == "text"
# default_value MUST be null (not the canonical URL): saving it as the
# default would persist v1beta into every credential record and bypass
# `_get_gemini_url`'s automatic v1alpha routing for Gemini 3+ models. The
# placeholder shows the canonical URL so users still get the visual hint.
# (See greptileai threads on PR #30419.)
assert api_base_field["default_value"] is None
assert (
api_base_field["placeholder"]
== "https://generativelanguage.googleapis.com/v1beta"
)
# UI forms render fields in credential_fields order; api_base should come
# first so an admin sees the URL override before the key field (matches
# OpenAI and Anthropic conventions).
field_order = [f["key"] for f in google_ai["credential_fields"]]
assert field_order.index("api_base") < field_order.index(
"api_key"
), "api_base must appear before api_key in credential_fields."
def test_public_model_hub_with_healthy_model():
"""Test that health information is populated for a healthy model"""
app = FastAPI()

View file

@ -359,6 +359,7 @@ ignored_keys = [
"metadata.additional_usage_values.cache_read_input_tokens",
"metadata.additional_usage_values.inference_geo",
"metadata.additional_usage_values.speed",
"metadata.additional_usage_values.service_tier",
"metadata.additional_usage_values.iterations",
"metadata.litellm_overhead_time_ms",
"metadata.cost_breakdown",

View file

@ -2415,6 +2415,41 @@ class TestHandleLLMApiExceptionDictDetail:
proxy_exc = await self._invoke(exc)
assert proxy_exc.code == "500"
async def test_already_normalized_proxy_exception_is_honored(self):
"""A ProxyException raised mid-request (e.g. a guardrail block) is already
the OpenAI wire format. The funnel must re-raise it untouched instead of
re-deriving the status from a (nonexistent) status_code attribute and
defaulting to 500. Regression for LIT-3751."""
from litellm.proxy._types import ProxyException
exc = ProxyException(
message='"Leroy Jenkins" detected as name',
type="invalid_request_error",
param=None,
code=400,
openai_code="content_policy_violation",
)
proxy_exc = await self._invoke(exc)
assert proxy_exc is exc
assert proxy_exc.code == "400"
assert proxy_exc.type == "invalid_request_error"
assert proxy_exc.param is None
assert proxy_exc.openai_code == "content_policy_violation"
assert proxy_exc.message == '"Leroy Jenkins" detected as name'
# The body the OpenAI-SDK client actually receives. The HTTP status line
# comes from int(exc.code) == 400; the wire ``code`` stays the status
# string. ``openai_code`` ("content_policy_violation") is intentionally
# NOT serialized here - to_dict() emits only ``code`` - so this asserts
# the real contract rather than the write-only attribute.
assert int(proxy_exc.code) == 400
assert proxy_exc.to_dict() == {
"message": '"Leroy Jenkins" detected as name',
"type": "invalid_request_error",
"param": None,
"code": "400",
}
class TestStreamCloseOnDisconnect:
"""
@ -2737,6 +2772,438 @@ class TestAsyncStreamingDataGeneratorFastPath:
ProxyLogging._callback_capabilities_cache.clear()
class TestDisconnectGatherCleanup:
def _disconnect_request(self) -> Request:
messages = [
{"type": "http.request", "body": b"", "more_body": False},
{"type": "http.disconnect"},
]
async def receive():
if messages:
return messages.pop(0)
await asyncio.Event().wait()
return Request(scope={"type": "http", "headers": []}, receive=receive)
@pytest.mark.asyncio
async def test_base_process_llm_request_raises_499_on_client_disconnect(
self, monkeypatch
):
"""With cancel_on_disconnect enabled, base_process_llm_request returns 499."""
import asyncio
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
async def slow_llm():
await asyncio.sleep(9999)
async def fake_route_request(**_kwargs):
return slow_llm()
mock_logging_obj = MagicMock()
mock_logging_obj.litellm_call_id = "test-call-id"
mock_logging_obj._defer_async_logging = False
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.during_call_hook = AsyncMock(return_value=None)
mock_proxy_logging._callback_capabilities_cache = {}
monkeypatch.setattr(cpr, "route_request", fake_route_request)
processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
with pytest.raises(HTTPException) as exc_info:
await processing_obj.base_process_llm_request(
request=self._disconnect_request(),
fastapi_response=MagicMock(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
proxy_logging_obj=mock_proxy_logging,
general_settings={"cancel_on_disconnect": True},
proxy_config=MagicMock(spec=ProxyConfig),
route_type="acompletion",
version=None,
)
assert exc_info.value.status_code == 499
assert "disconnected" in exc_info.value.detail.lower()
@pytest.mark.asyncio
async def test_base_process_llm_request_reraises_cancelled_error_without_client_disconnect(
self, monkeypatch
):
import asyncio
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
async def fake_gather(*_tasks, **_kwargs):
raise asyncio.CancelledError()
mock_logging_obj = MagicMock()
mock_logging_obj.litellm_call_id = "test-call-id"
mock_logging_obj._defer_async_logging = False
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.during_call_hook = AsyncMock(return_value=None)
mock_proxy_logging._callback_capabilities_cache = {}
monkeypatch.setattr(cpr.asyncio, "gather", fake_gather)
processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
monkeypatch.setattr(
cpr,
"route_request",
AsyncMock(return_value=asyncio.sleep(9999)),
)
mock_request = MagicMock(spec=Request)
mock_request.headers = {}
with pytest.raises(asyncio.CancelledError):
await processing_obj.base_process_llm_request(
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
proxy_logging_obj=mock_proxy_logging,
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
route_type="acompletion",
version=None,
)
@pytest.mark.asyncio
async def test_disconnect_cancels_during_call_hook_task(self, monkeypatch):
import asyncio
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
hook_cancelled = False
async def slow_during_call_hook(**_kwargs):
try:
await asyncio.sleep(9999)
except asyncio.CancelledError:
nonlocal hook_cancelled
hook_cancelled = True
raise
async def slow_llm():
await asyncio.sleep(9999)
async def fake_route_request(**_kwargs):
return slow_llm()
mock_logging_obj = MagicMock()
mock_logging_obj.litellm_call_id = "test-call-id"
mock_logging_obj._defer_async_logging = False
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.during_call_hook = slow_during_call_hook
mock_proxy_logging._callback_capabilities_cache = {}
monkeypatch.setattr(cpr, "route_request", fake_route_request)
processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
with pytest.raises(HTTPException):
await processing_obj.base_process_llm_request(
request=self._disconnect_request(),
fastapi_response=MagicMock(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
proxy_logging_obj=mock_proxy_logging,
general_settings={"cancel_on_disconnect": True},
proxy_config=MagicMock(spec=ProxyConfig),
route_type="acompletion",
version=None,
)
assert hook_cancelled is True
@pytest.mark.asyncio
async def test_cancel_pending_gather_tasks_skips_already_done_tasks(self):
import asyncio
from litellm.proxy.common_request_processing import _cancel_pending_gather_tasks
async def failing_task():
raise ValueError("llm api error")
task = asyncio.create_task(failing_task())
with pytest.raises(ValueError, match="llm api error"):
await task
await _cancel_pending_gather_tasks([task])
@pytest.mark.asyncio
async def test_cancel_pending_gather_tasks_swallows_guardrail_converted_cancel(
self,
):
import asyncio
from litellm.proxy.common_request_processing import _cancel_pending_gather_tasks
async def hook_converts_cancel_to_runtime_error():
try:
await asyncio.sleep(9999)
except asyncio.CancelledError:
raise RuntimeError("guardrail converted cancel")
task = asyncio.create_task(hook_converts_cancel_to_runtime_error())
await asyncio.sleep(0)
await _cancel_pending_gather_tasks([task])
assert task.done()
@pytest.mark.asyncio
async def test_base_process_llm_request_preserves_llm_error_after_gather(
self, monkeypatch
):
import asyncio
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
async def failing_llm():
raise ValueError("llm api error")
async def successful_hook(**_kwargs):
return None
async def fake_route_request(**_kwargs):
return failing_llm()
mock_logging_obj = MagicMock()
mock_logging_obj.litellm_call_id = "test-call-id"
mock_logging_obj._defer_async_logging = False
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.during_call_hook = successful_hook
mock_proxy_logging._callback_capabilities_cache = {}
monkeypatch.setattr(cpr, "route_request", fake_route_request)
processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"})
monkeypatch.setattr(
processing_obj,
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=False)
mock_request.headers = {}
with pytest.raises(ValueError, match="llm api error"):
await processing_obj.base_process_llm_request(
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
proxy_logging_obj=mock_proxy_logging,
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
route_type="acompletion",
version=None,
)
class TestStreamingClientDisconnectLogging:
@pytest.mark.asyncio
async def test_record_streaming_client_disconnect_sets_error_information(self):
from litellm.proxy.common_request_processing import (
_record_streaming_client_disconnect_if_needed,
)
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"litellm_params": {}, "metadata": {}}
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=True)
request_data = {
"litellm_call_id": "test-call-id",
"litellm_logging_obj": mock_logging_obj,
"metadata": {},
"litellm_params": {"metadata": {}},
}
recorded = await _record_streaming_client_disconnect_if_needed(
mock_request, request_data
)
assert recorded is True
assert request_data["metadata"]["client_disconnected"] is True
assert (
request_data["metadata"]["error_information"]["error_code"] == "499"
)
assert (
mock_logging_obj.model_call_details["litellm_params"]["metadata"][
"error_information"
]["error_code"]
== "499"
)
@pytest.mark.asyncio
async def test_record_streaming_client_disconnect_no_op_when_connected(self):
from litellm.proxy.common_request_processing import (
_record_streaming_client_disconnect_if_needed,
)
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=False)
request_data = {"metadata": {}}
recorded = await _record_streaming_client_disconnect_if_needed(
mock_request, request_data
)
assert recorded is False
assert "client_disconnected" not in request_data["metadata"]
@pytest.mark.asyncio
async def test_finalize_streaming_generator_cleanup_fires_deferred_logging(
self, monkeypatch
):
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
fire_spy = MagicMock()
monkeypatch.setattr(
"litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging",
fire_spy,
)
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=True)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
request_data = {
"metadata": {},
"litellm_params": {"metadata": {}},
"litellm_logging_obj": MagicMock(model_call_details={"metadata": {}, "litellm_params": {}}),
}
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
request=mock_request,
request_data=request_data,
response=mock_response,
)
fire_spy.assert_called_once_with(request_data)
mock_response.aclose.assert_awaited_once()
assert request_data["metadata"]["error_information"]["error_code"] == "499"
@pytest.mark.asyncio
async def test_finalize_streaming_generator_cleanup_skips_disconnect_after_completion(
self, monkeypatch
):
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
fire_spy = MagicMock()
monkeypatch.setattr(
"litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging",
fire_spy,
)
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=True)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
request_data = {"metadata": {}, "litellm_params": {"metadata": {}}}
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
request=mock_request,
request_data=request_data,
response=mock_response,
stream_completed=True,
)
fire_spy.assert_not_called()
mock_request.is_disconnected.assert_not_awaited()
mock_response.aclose.assert_awaited_once()
assert "client_disconnected" not in request_data["metadata"]
@pytest.mark.asyncio
async def test_async_streaming_data_generator_records_499_on_early_aclose(
self, monkeypatch
):
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
monkeypatch.setattr(
"litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging",
MagicMock(),
)
async def mock_streaming_iterator(*_args, **_kwargs):
yield {"choices": [{"delta": {"content": "hi"}}]}
yield {"choices": [{"delta": {"content": " there"}}]}
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.async_post_call_streaming_iterator_hook = (
mock_streaming_iterator
)
ProxyLogging._callback_capabilities_cache.clear()
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=True)
mock_response = MagicMock()
mock_response.aclose = AsyncMock()
request_data = {
"model": "gemini-2.0-flash",
"metadata": {},
"litellm_params": {"metadata": {}},
"litellm_logging_obj": MagicMock(
model_call_details={"metadata": {}, "litellm_params": {}}
),
}
gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
response=mock_response,
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
request_data=request_data,
proxy_logging_obj=mock_proxy_logging,
serialize_chunk=lambda chunk: f"data: {chunk}\n\n",
serialize_error=lambda proxy_exc: f"data: {proxy_exc.to_dict()}\n\n",
request=mock_request,
)
await gen.__anext__()
await gen.aclose()
assert request_data["metadata"]["client_disconnected"] is True
assert request_data["metadata"]["error_information"]["error_code"] == "499"
ProxyLogging._callback_capabilities_cache.clear()
class TestCancelOnDisconnect:
"""
Coverage for the opt-in `general_settings.cancel_on_disconnect` flag:

View file

@ -188,6 +188,34 @@ async def test_add_litellm_data_to_request_strips_root_pricing_fields():
assert "output_cost_per_token" not in updated
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_client_disconnect_metadata():
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "hi"}],
"metadata": {
"client_disconnected": True,
"error_information": {
"error_code": "499",
"error_message": "Client disconnected the request",
"error_class": "ClientDisconnected",
},
},
}
updated = await add_litellm_data_to_request(
data=data,
request=_make_request_mock(),
user_api_key_dict=_user_api_key_auth(),
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert "client_disconnected" not in updated.get("metadata", {})
assert "error_information" not in updated.get("metadata", {})
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_metadata_model_info():
data = {

View file

@ -5246,10 +5246,10 @@ async def test_async_data_generator_uses_direct_stream_fast_path_without_callbac
@pytest.mark.asyncio
async def test_async_data_generator_passes_through_google_native_sse_bytes():
async def test_async_data_generator_preserves_non_raw_sse_like_bytes():
"""
Google-native streamGenerateContent yields raw SSE bytes; they must not be
re-wrapped as data: b'data: {...}'.
Already formatted SSE bytes from non-raw streams keep the legacy passthrough
behavior, including appending a missing event terminator.
"""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
@ -5305,6 +5305,241 @@ async def test_async_data_generator_passes_through_google_native_sse_bytes():
assert yielded_text[-1] == "data: [DONE]\n\n"
@pytest.mark.asyncio
async def test_async_data_generator_buffers_split_google_native_sse_json_frame():
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
from litellm.proxy.utils import ProxyLogging
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_request_data = {
"model": "gemini-3.5-flash",
"_litellm_skip_openai_stream_done": True,
"_litellm_raw_sse_stream": True,
}
payload = (
'data: {"candidates": [{"content": {"role": "model", "parts": '
'[{"text": "", "thoughtSignature": "abc123def456"}]}}]}\n\n'
)
raw_chunks = [
payload[:2].encode("utf-8"),
payload[
2 : payload.index("thoughtSignature") + len('thoughtSignature": "abc')
].encode("utf-8"),
payload[
payload.index("thoughtSignature") + len('thoughtSignature": "abc') :
].encode("utf-8"),
]
class MockStream:
def __aiter__(self):
return self._stream()
async def _stream(self):
for chunk in raw_chunks:
yield chunk
async def aclose(self):
pass
mock_response = MockStream()
mock_response.aclose = AsyncMock()
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
yielded_data = []
async for data in async_data_generator(
mock_response, mock_user_api_key_dict, mock_request_data
):
yielded_data.append(data)
yielded_text = [
chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
for chunk in yielded_data
]
assert yielded_text == [payload]
for chunk in yielded_text:
assert chunk.endswith("\n\n")
assert json.loads(chunk.removeprefix("data: ").strip())
@pytest.mark.asyncio
async def test_async_data_generator_flushes_raw_sse_stream_without_trailing_delimiter():
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
from litellm.proxy.utils import ProxyLogging
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_request_data = {
"model": "gemini-3.5-flash",
"_litellm_skip_openai_stream_done": True,
"_litellm_raw_sse_stream": True,
}
class MockStream:
def __aiter__(self):
return self._stream()
async def _stream(self):
yield b'data: {"candidates": [{"content": "unterminated"}]'
async def aclose(self):
pass
mock_response = MockStream()
mock_response.aclose = AsyncMock()
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj),
patch.object(ProxyLogging, "_fire_deferred_stream_logging"),
):
yielded_data = []
async for data in async_data_generator(
mock_response, mock_user_api_key_dict, mock_request_data
):
yielded_data.append(data)
yielded_text = [
chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
for chunk in yielded_data
]
assert len(yielded_text) == 1
assert yielded_text[0] == 'data: {"candidates": [{"content": "unterminated"}]\n\n'
assert "[DONE]" not in yielded_text[0]
mock_proxy_logging_obj.post_call_failure_hook.assert_not_awaited()
@pytest.mark.asyncio
async def test_async_data_generator_errors_when_raw_sse_frame_exceeds_buffer_limit():
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
from litellm.proxy.utils import ProxyLogging
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_request_data = {
"model": "gemini-3.5-flash",
"_litellm_skip_openai_stream_done": True,
"_litellm_raw_sse_stream": True,
}
class MockStream:
def __aiter__(self):
return self._stream()
async def _stream(self):
yield b"data: "
yield b'{"candidates": [{"content": "unterminated"}]'
async def aclose(self):
pass
mock_response = MockStream()
mock_response.aclose = AsyncMock()
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj),
patch("litellm.proxy.proxy_server._MAX_RAW_SSE_BUFFER_CHARS", 8),
patch.object(ProxyLogging, "_fire_deferred_stream_logging"),
):
yielded_data = []
async for data in async_data_generator(
mock_response, mock_user_api_key_dict, mock_request_data
):
yielded_data.append(data)
yielded_text = [
chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
for chunk in yielded_data
]
assert len(yielded_text) == 1
assert "maximum buffered size" in yielded_text[0]
assert "[DONE]" not in yielded_text[0]
mock_proxy_logging_obj.post_call_failure_hook.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("as_bytes", [True, False])
async def test_async_data_generator_checks_raw_sse_buffer_limit_after_complete_frames(
as_bytes,
):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
from litellm.proxy.utils import ProxyLogging
complete_frame = 'data: {"candidates": [{"content": "ok"}]}\n\n'
partial_frame = "data: "
raw_chunk = complete_frame + partial_frame
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_request_data = {
"model": "gemini-3.5-flash",
"_litellm_skip_openai_stream_done": True,
"_litellm_raw_sse_stream": True,
}
class MockStream:
def __aiter__(self):
return self._stream()
async def _stream(self):
yield raw_chunk.encode("utf-8") if as_bytes else raw_chunk
async def aclose(self):
pass
mock_response = MockStream()
mock_response.aclose = AsyncMock()
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj),
patch("litellm.proxy.proxy_server._MAX_RAW_SSE_BUFFER_CHARS", 8),
patch.object(ProxyLogging, "_fire_deferred_stream_logging"),
):
yielded_data = []
async for data in async_data_generator(
mock_response, mock_user_api_key_dict, mock_request_data
):
yielded_data.append(data)
yielded_text = [
chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
for chunk in yielded_data
]
assert yielded_text[0] == complete_frame
assert yielded_text[1] == partial_frame + "\n\n"
assert "[DONE]" not in "".join(yielded_text)
mock_proxy_logging_obj.post_call_failure_hook.assert_not_awaited()
@pytest.mark.asyncio
async def test_async_data_generator_google_genai_stream_omits_openai_done():
"""
@ -5359,6 +5594,53 @@ async def test_async_data_generator_google_genai_stream_omits_openai_done():
assert "[DONE]" not in "".join(yielded_text)
@pytest.mark.asyncio
async def test_async_data_generator_does_not_mark_completed_stream_as_disconnect():
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
from litellm.proxy.utils import ProxyLogging
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_request_data = {"model": "gpt-4o", "metadata": {}}
class MockStream:
def __aiter__(self):
return self._stream()
async def _stream(self):
yield {"choices": [{"delta": {"content": "done"}}]}
async def aclose(self):
pass
mock_request = MagicMock()
mock_request.is_disconnected = AsyncMock(return_value=True)
mock_response = MockStream()
mock_response.aclose = AsyncMock()
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
yielded_data = []
async for data in async_data_generator(
mock_response,
mock_user_api_key_dict,
mock_request_data,
request=mock_request,
):
yielded_data.append(data)
assert yielded_data[-1] == "data: [DONE]\n\n"
mock_request.is_disconnected.assert_not_awaited()
assert "client_disconnected" not in mock_request_data["metadata"]
@pytest.mark.asyncio
async def test_async_data_generator_google_genai_stream_forwards_error_without_done():
"""Stream errors must still reach the client when OpenAI [DONE] is skipped."""

View file

@ -427,6 +427,120 @@ class TestPostCallFailureHookLiftsFirstApiCallStartTime:
assert "litellm_logging_obj" not in request_data
class TestPostCallFailureHookLLMExceptionAlerting:
"""The llm_exceptions alert is for infra / LLM-API failures, not user
errors (https://github.com/BerriAI/litellm/issues/3395). Already-normalized
client errors must be excluded so a guardrail content-policy block never
pages on-call. ProxyException is such an error; before LIT-3751 only
HTTPException was excluded, so AIM blocks paged as if the LLM API failed."""
async def _alerted(self, exc) -> bool:
import asyncio
from unittest.mock import AsyncMock
from litellm.proxy._types import AlertType, UserAPIKeyAuth
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.alert_types = [AlertType.llm_exceptions]
alerting_handler = AsyncMock()
with (
patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()),
patch.object(proxy_logging_obj, "alerting_handler", new=alerting_handler),
):
await proxy_logging_obj.post_call_failure_hook(
request_data={},
original_exception=exc,
user_api_key_dict=UserAPIKeyAuth(),
)
await asyncio.sleep(0) # let the fire-and-forget alert task run
return alerting_handler.called
@pytest.mark.asyncio
async def test_proxy_exception_does_not_alert(self):
from litellm.proxy._types import ProxyException
exc = ProxyException(
message="content blocked",
type="invalid_request_error",
param=None,
code=400,
openai_code="content_policy_violation",
)
assert await self._alerted(exc) is False
@pytest.mark.asyncio
async def test_http_exception_does_not_alert(self):
assert (
await self._alerted(HTTPException(status_code=400, detail="blocked"))
is False
)
@pytest.mark.asyncio
async def test_genuine_llm_api_error_still_alerts(self):
assert await self._alerted(Exception("upstream 503")) is True
class TestPostCallFailureHookProxyExceptionLogging:
"""A guardrail block raises a ProxyException; on an LLM route it must still
drive proxy-only failure logging (_handle_logging_proxy_only_error) so the
blocked request is recorded, exactly as the old HTTPException did. Before
LIT-3751 the classifier only matched HTTPException, so switching AIM to
ProxyException silently dropped the rejected prompt from failure logs."""
async def _logged(self, exc, *, request_route) -> bool:
from unittest.mock import AsyncMock
from litellm.proxy._types import UserAPIKeyAuth
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.alert_types = []
handle_mock = AsyncMock()
with (
patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()),
patch.object(
proxy_logging_obj,
"_handle_logging_proxy_only_error",
new=handle_mock,
),
):
await proxy_logging_obj.post_call_failure_hook(
request_data={},
original_exception=exc,
user_api_key_dict=UserAPIKeyAuth(
api_key="sk-test", request_route=request_route
),
)
return handle_mock.await_count > 0
def _block(self):
from litellm.proxy._types import ProxyException
return ProxyException(
message="content blocked",
type="invalid_request_error",
param=None,
code=400,
openai_code="content_policy_violation",
)
@pytest.mark.asyncio
async def test_proxy_exception_on_llm_route_is_logged(self):
assert (
await self._logged(self._block(), request_route="/v1/chat/completions")
is True
)
@pytest.mark.asyncio
async def test_generic_exception_on_llm_route_is_not_logged(self):
# A raw provider/unknown exception is logged by the LLM call path, not here.
assert (
await self._logged(
Exception("upstream 503"), request_route="/v1/chat/completions"
)
is False
)
class TestShouldUseSmtpSsl:
def test_port_465_uses_ssl(self, monkeypatch):
from litellm.proxy.utils import _should_use_smtp_ssl

View file

@ -0,0 +1,399 @@
"""
Unit tests for AWSSecretsManagerV2 cross-region replication via ReplicateSecretToRegions.
All tests are mocked — no real AWS credentials required.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2
# ---------------------------------------------------------------------------
# Shared fixtures
# ---------------------------------------------------------------------------
_CREATE_RESPONSE = {
"ARN": "arn:aws:secretsmanager:us-east-1:123456789012:secret:litellm/test-key",
"Name": "litellm/test-key",
"VersionId": "mock-version-id",
}
_REPLICATE_RESPONSE = {
"ARN": "arn:aws:secretsmanager:us-east-1:123456789012:secret:litellm/test-key",
"ReplicationStatus": [
{"Region": "us-west-2", "Status": "InProgress"},
],
}
def _mock_http_client(json_response: dict) -> MagicMock:
mock_response = MagicMock()
mock_response.raise_for_status = MagicMock()
mock_response.json.return_value = json_response
mock_async_client = AsyncMock()
mock_async_client.post.return_value = mock_response
return mock_async_client
# ---------------------------------------------------------------------------
# Tests: async_write_secret + replication
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_write_secret_replicates_when_configured():
"""async_replicate_secret is called after a successful CreateSecret when replica_regions is set."""
manager = AWSSecretsManagerV2(replica_regions=["us-west-2"])
with patch.object(
AWSSecretsManagerV2,
"_prepare_request",
return_value=(
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b'{"Name":"litellm/test-key"}',
),
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client(_CREATE_RESPONSE),
):
with patch.object(
AWSSecretsManagerV2,
"async_replicate_secret",
new_callable=AsyncMock,
return_value=_REPLICATE_RESPONSE,
) as mock_replicate:
result = await manager.async_write_secret(
secret_name="litellm/test-key",
secret_value="sk-test-value",
)
assert result == _CREATE_RESPONSE
mock_replicate.assert_called_once_with(
secret_name="litellm/test-key",
replica_regions=["us-west-2"],
optional_params=None,
timeout=None,
)
@pytest.mark.asyncio
async def test_write_secret_no_replication_when_not_configured():
"""async_replicate_secret is NOT called when replica_regions is None."""
manager = AWSSecretsManagerV2(replica_regions=None)
with patch.object(
AWSSecretsManagerV2,
"_prepare_request",
return_value=(
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b'{"Name":"litellm/test-key"}',
),
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client(_CREATE_RESPONSE),
):
with patch.object(
AWSSecretsManagerV2,
"async_replicate_secret",
new_callable=AsyncMock,
) as mock_replicate:
result = await manager.async_write_secret(
secret_name="litellm/test-key",
secret_value="sk-test-value",
)
assert result == _CREATE_RESPONSE
mock_replicate.assert_not_called()
@pytest.mark.asyncio
async def test_replication_failure_does_not_fail_write():
"""If async_replicate_secret raises, async_write_secret still returns the CreateSecret response."""
manager = AWSSecretsManagerV2(replica_regions=["us-west-2"])
with patch.object(
AWSSecretsManagerV2,
"_prepare_request",
return_value=(
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b'{"Name":"litellm/test-key"}',
),
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client(_CREATE_RESPONSE),
):
with patch.object(
AWSSecretsManagerV2,
"async_replicate_secret",
new_callable=AsyncMock,
side_effect=ValueError("AccessDenied: not authorized"),
):
result = await manager.async_write_secret(
secret_name="litellm/test-key",
secret_value="sk-test-value",
)
assert result == _CREATE_RESPONSE
# ---------------------------------------------------------------------------
# Tests: async_replicate_secret directly
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_async_replicate_secret_empty_regions_returns_empty():
"""async_replicate_secret returns {} immediately for an empty list — no HTTP call."""
manager = AWSSecretsManagerV2()
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client"
) as mock_get_client:
result = await manager.async_replicate_secret(
secret_name="litellm/test-key",
replica_regions=[],
)
assert result == {}
mock_get_client.assert_not_called()
@pytest.mark.asyncio
async def test_async_replicate_secret_correct_payload():
"""async_replicate_secret sends the correct AddReplicaRegions payload."""
manager = AWSSecretsManagerV2()
captured: dict = {}
def capture_prepare(action, secret_name, optional_params=None, request_data=None):
captured.update(request_data or {})
captured["_action"] = action
return (
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b"{}",
)
with patch.object(
AWSSecretsManagerV2, "_prepare_request", side_effect=capture_prepare
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client(_REPLICATE_RESPONSE),
):
result = await manager.async_replicate_secret(
secret_name="litellm/test-key",
replica_regions=["us-west-2", "eu-west-1"],
)
assert result == _REPLICATE_RESPONSE
assert captured["_action"] == "ReplicateSecretToRegions"
assert captured["SecretId"] == "litellm/test-key"
assert captured["AddReplicaRegions"] == [
{"Region": "us-west-2"},
{"Region": "eu-west-1"},
]
@pytest.mark.asyncio
async def test_replication_fires_on_create(caplog):
"""async_replicate_secret emits an INFO log line mentioning ReplicateSecretToRegions."""
import logging
manager = AWSSecretsManagerV2()
with patch.object(
AWSSecretsManagerV2,
"_prepare_request",
return_value=(
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b"{}",
),
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client(_REPLICATE_RESPONSE),
):
with caplog.at_level(logging.INFO, logger="LiteLLM"):
await manager.async_replicate_secret(
secret_name="litellm/test-key",
replica_regions=["us-west-2"],
)
assert "ReplicateSecretToRegions" in caplog.text
# ---------------------------------------------------------------------------
# Tests: load_aws_secret_manager forwards replica_regions
# ---------------------------------------------------------------------------
def test_load_aws_secret_manager_passes_replica_regions():
"""load_aws_secret_manager must forward replica_regions from key_management_settings."""
import litellm
original = litellm.secret_manager_client
settings = MagicMock()
settings.aws_region_name = "us-east-1"
settings.aws_role_name = None
settings.aws_session_name = None
settings.aws_external_id = None
settings.aws_profile_name = None
settings.aws_web_identity_token = None
settings.aws_sts_endpoint = None
settings.replica_regions = ["us-west-2", "eu-west-1"]
try:
AWSSecretsManagerV2.load_aws_secret_manager(
use_aws_secret_manager=True,
key_management_settings=settings,
)
assert isinstance(litellm.secret_manager_client, AWSSecretsManagerV2)
assert litellm.secret_manager_client.replica_regions == [
"us-west-2",
"eu-west-1",
]
finally:
litellm.secret_manager_client = original
def _http_status_error(status_code: int, body: str) -> httpx.HTTPStatusError:
request = httpx.Request("POST", "https://secretsmanager.us-east-1.amazonaws.com")
response = httpx.Response(status_code=status_code, text=body, request=request)
return httpx.HTTPStatusError(message=body, request=request, response=response)
def _mock_http_client_raising(exc: Exception) -> MagicMock:
mock_response = MagicMock()
mock_response.raise_for_status.side_effect = exc
mock_async_client = AsyncMock()
mock_async_client.post.return_value = mock_response
return mock_async_client
# ---------------------------------------------------------------------------
# Tests: error paths in async_write_secret
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_write_secret_http_error_raises():
"""async_write_secret raises ValueError when CreateSecret returns a non-2xx HTTP status."""
manager = AWSSecretsManagerV2()
with patch.object(
AWSSecretsManagerV2,
"_prepare_request",
return_value=(
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b'{"Name":"litellm/test-key"}',
),
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client_raising(
_http_status_error(400, "ResourceExistsException")
),
):
with pytest.raises(ValueError, match="HTTP error occurred"):
await manager.async_write_secret(
secret_name="litellm/test-key",
secret_value="sk-test-value",
)
@pytest.mark.asyncio
async def test_write_secret_timeout_raises():
"""async_write_secret raises ValueError when the CreateSecret call times out."""
manager = AWSSecretsManagerV2()
with patch.object(
AWSSecretsManagerV2,
"_prepare_request",
return_value=(
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b'{"Name":"litellm/test-key"}',
),
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client_raising(
httpx.ReadTimeout("timed out", request=None)
),
):
with pytest.raises(ValueError, match="Timeout error occurred"):
await manager.async_write_secret(
secret_name="litellm/test-key",
secret_value="sk-test-value",
)
# ---------------------------------------------------------------------------
# Tests: error paths in async_replicate_secret
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_replicate_secret_http_error_raises():
"""async_replicate_secret raises ValueError when ReplicateSecretToRegions returns a non-2xx status."""
manager = AWSSecretsManagerV2()
with patch.object(
AWSSecretsManagerV2,
"_prepare_request",
return_value=(
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b"{}",
),
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client_raising(
_http_status_error(403, "AccessDeniedException")
),
):
with pytest.raises(ValueError, match="HTTP error occurred"):
await manager.async_replicate_secret(
secret_name="litellm/test-key",
replica_regions=["us-west-2"],
)
@pytest.mark.asyncio
async def test_replicate_secret_timeout_raises():
"""async_replicate_secret raises ValueError when the ReplicateSecretToRegions call times out."""
manager = AWSSecretsManagerV2()
with patch.object(
AWSSecretsManagerV2,
"_prepare_request",
return_value=(
"https://secretsmanager.us-east-1.amazonaws.com",
{"Content-Type": "application/x-amz-json-1.1"},
b"{}",
),
):
with patch(
"litellm.secret_managers.aws_secret_manager_v2.get_async_httpx_client",
return_value=_mock_http_client_raising(
httpx.ReadTimeout("timed out", request=None)
),
):
with pytest.raises(ValueError, match="Timeout error occurred"):
await manager.async_replicate_secret(
secret_name="litellm/test-key",
replica_regions=["us-west-2"],
)

View file

@ -0,0 +1,89 @@
"""
Validate that the native (first-party) Anthropic Claude Sonnet 4.5 / 4.6 entries
carry the 1-hour prompt-cache write tier (`cache_creation_input_token_cost_above_1hr`)
in `model_prices_and_context_window.json`.
Anthropic's first-party API charges a separate 1-hour cache write rate (2x base
input) alongside the 5-minute write (1.25x base input) and cache read (0.1x base
input). The 1h/5m ratio is therefore 1.6. Without the 1-hour field, cost tracking
on 1-hour-TTL prompt caching falls back to the 5-minute rate and undercounts spend.
The native (non-bedrock) `claude-sonnet-4-5*` / `claude-sonnet-4-6` entries were
missing this field, while every sibling (`vertex_ai/`, `azure_ai/`, the
`*.anthropic.*` Bedrock profiles) and the older `claude-sonnet-4-20250514` already
carried it. This test guards against regression.
Values (per token):
Sonnet base input 3e-06 -> 5m 3.75e-06, 1h 6e-06
Sonnet 4.5 long-context (>200K) base 6e-06 -> 5m 7.5e-06, 1h 1.2e-05
"""
import json
import os
import pytest
@pytest.fixture(scope="module")
def model_data():
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)
# (model_key, expected 1hr write per token, expected 1hr long-context tier or None)
EXPECTED = [
("claude-sonnet-4-5", 6e-06, 1.2e-05),
("claude-sonnet-4-5-20250929", 6e-06, 1.2e-05),
("claude-sonnet-4-5-20250929-v1:0", 6e-06, 1.2e-05),
("claude-sonnet-4-6", 6e-06, None),
]
@pytest.mark.parametrize("model_key, expected_1hr, expected_1hr_lc", EXPECTED)
def test_anthropic_sonnet_1hr_cache_write_pricing(
model_data, model_key, expected_1hr, expected_1hr_lc
):
assert model_key in model_data, f"Missing model entry: {model_key}"
info = model_data[model_key]
# Regular 1hr cache write rate must be present and exact.
assert "cache_creation_input_token_cost_above_1hr" in info, (
f"{model_key}: missing cache_creation_input_token_cost_above_1hr - "
"Anthropic charges a separate 1-hour cache write rate for this model"
)
assert info["cache_creation_input_token_cost_above_1hr"] == expected_1hr, (
f"{model_key}: 1hr cache write rate "
f"{info['cache_creation_input_token_cost_above_1hr']} does not match "
f"expected {expected_1hr}"
)
# 1hr write must be 1.6x the 5-minute write (Anthropic 2x-base / 1.25x-base).
ratio = (
info["cache_creation_input_token_cost_above_1hr"]
/ info["cache_creation_input_token_cost"]
)
assert (
abs(ratio - 1.6) < 1e-9
), f"{model_key}: 1hr/5min ratio is {ratio}, expected 1.6"
# Long-context (>200K) 1hr tier, where the model publishes a >200K tier.
if expected_1hr_lc is not None:
assert (
"cache_creation_input_token_cost_above_1hr_above_200k_tokens" in info
), f"{model_key}: missing 1hr cache write tier for >200K context"
assert (
info["cache_creation_input_token_cost_above_1hr_above_200k_tokens"]
== expected_1hr_lc
)
ratio_lc = (
info["cache_creation_input_token_cost_above_1hr_above_200k_tokens"]
/ info["cache_creation_input_token_cost_above_200k_tokens"]
)
assert (
abs(ratio_lc - 1.6) < 1e-9
), f"{model_key}: long-context 1hr/5min ratio is {ratio_lc}, expected 1.6"
else:
assert "cache_creation_input_token_cost_above_1hr_above_200k_tokens" not in info

View file

@ -51,6 +51,23 @@ def test_new_rule_in_head_is_clean():
assert ratchet.regressions_for("b.json", {}, {"LIT009": _spec_of(5, 0)}) == []
def test_dropped_file_in_the_any_budget_is_not_a_regression():
# any-discipline is file-keyed: an absent file means ceiling 0, so cleaning a
# file to zero (which drops its entry on --update) is a tightening, never the
# loosening a dropped rule is for the rule-keyed budgets.
base = {"litellm/x.py": _spec_of(10, 5)}
assert ratchet.regressions_for("any-discipline-budget.json", base, {}) == []
def test_raised_ceiling_in_the_any_budget_is_still_a_regression():
base = {"litellm/x.py": _spec_of(10, 5)} # ceiling 15
regs = ratchet.regressions_for(
"any-discipline-budget.json", base, {"litellm/x.py": _spec_of(20, 10)} # ceiling 30
)
assert [r.rule for r in regs] == ["litellm/x.py"]
assert "15 -> 30" in regs[0].detail
def test_deleted_budget_file_is_a_regression():
regs = ratchet.regressions_for("b.json", {"LIT006": _spec_of(1, 0)}, None)
assert [r.rule for r in regs] == ["*"]

View file

@ -39,3 +39,52 @@ def test_no_line_map_means_no_line_filtering():
def test_build_error_is_always_in_scope():
assert mod._in_scope(_v(code="LIT000", line=1), {"litellm/x.py": {2}}) is True
# --- per-file Any budget ------------------------------------------------------
def test_slack_is_50_percent_rounded_up():
assert mod._slack_for(0) == 0
assert mod._slack_for(1) == 1 # ceil(0.5): even a 1-Any file gets a little room
assert mod._slack_for(3) == 2 # ceil(1.5)
assert mod._slack_for(20) == 10
assert mod._slack_for(5145) == 2573
def test_ceiling_is_baseline_plus_slack():
assert mod._ceiling({"baseline": 20, "slack": 10}) == 30
assert mod._ceiling({}) == 0 # an absent/empty entry means a zero ceiling
def test_lit009_counts_groups_by_file_and_ignores_other_codes():
violations = [
_v(path="litellm/a.py", line=1, code="LIT009"),
_v(path="litellm/a.py", line=2, code="LIT009"),
_v(path="litellm/a.py", line=3, code="LIT005"), # suppression hygiene, not an Any
_v(path="litellm/b.py", line=1, code="LIT009"),
_v(path="litellm/c.py", line=0, code="LIT000"), # build error, not an Any
]
assert mod.lit009_counts(violations) == {"litellm/a.py": 2, "litellm/b.py": 1}
def test_save_budget_omits_zero_count_files_and_round_trips(monkeypatch, tmp_path):
monkeypatch.setattr(mod, "BUDGET_PATH", tmp_path / "any-discipline-budget.json")
mod.save_budget({"litellm/a.py": 20, "litellm/b.py": 0, "litellm/c.py": 1})
loaded = mod.load_budget()
assert loaded == {
"litellm/a.py": {"baseline": 20, "slack": 10},
"litellm/c.py": {"baseline": 1, "slack": 1},
}
assert "litellm/b.py" not in loaded # zero-Any files are never baselined
def test_load_budget_missing_file_is_empty(monkeypatch, tmp_path):
monkeypatch.setattr(mod, "BUDGET_PATH", tmp_path / "nope.json")
assert mod.load_budget() == {}
def test_update_budget_reports_setup_error_when_git_is_unavailable():
# all_litellm_py_files returns None when git can't list files; --update must
# surface a clean setup error (exit 2), not crash with a raw traceback.
assert mod.update_budget(list_files=lambda: None) == 2

View file

@ -2127,6 +2127,169 @@ def test_completion_cost_service_tier_for_bedrock():
assert priority_cost > default_cost > flex_cost > 0
def test_completion_cost_service_tier_for_anthropic():
"""
Anthropic priority-tier requests must be priced at the priority rate.
Regression for LIT-3771: the Anthropic cost route dropped ``service_tier``,
so priority requests (whose tier is reported on the response usage) were
always billed at the standard rate. The tier is captured by the
transformation and must flow through to ``generic_cost_per_token``.
"""
from litellm import completion_cost
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "claude-test-service-tier-cost-model"
litellm.register_model(
model_cost={
model: {
"input_cost_per_token": 3e-6,
"output_cost_per_token": 15e-6,
"input_cost_per_token_priority": 6e-6,
"output_cost_per_token_priority": 30e-6,
"litellm_provider": "anthropic",
"max_tokens": 8192,
}
}
)
def _cost_for_tier(service_tier):
usage = AnthropicConfig().calculate_usage(
usage_object={
"input_tokens": 1000,
"output_tokens": 500,
"service_tier": service_tier,
},
reasoning_content=None,
)
response = ModelResponse(usage=usage, model=model)
return completion_cost(
completion_response=response,
model=model,
custom_llm_provider="anthropic",
)
standard_cost = _cost_for_tier("standard")
priority_cost = _cost_for_tier("priority")
expected_standard = 1000 * 3e-6 + 500 * 15e-6
assert standard_cost == pytest.approx(expected_standard)
# priority rates are exactly 2x standard for both input and output
assert priority_cost == pytest.approx(2 * standard_cost)
def test_completion_cost_anthropic_auto_tier_uses_served_priority_rate():
"""
Proxy billing path regression for LIT-3771.
Priority is opted into with ``service_tier="auto"``; Anthropic then serves
"priority" and reports it on the response usage. The proxy forwards the
request-level "auto" into ``completion_cost`` (via ``_response_cost_calculator``),
and that preference must not shadow the served tier, otherwise priority
requests are silently billed at the standard rate.
"""
from litellm import completion_cost
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "claude-test-auto-tier-cost-model"
litellm.register_model(
model_cost={
model: {
"input_cost_per_token": 3e-6,
"output_cost_per_token": 15e-6,
"input_cost_per_token_priority": 6e-6,
"output_cost_per_token_priority": 30e-6,
"litellm_provider": "anthropic",
"max_tokens": 8192,
}
}
)
usage = AnthropicConfig().calculate_usage(
usage_object={
"input_tokens": 1000,
"output_tokens": 500,
"service_tier": "priority",
},
reasoning_content=None,
)
response = ModelResponse(usage=usage, model=model)
cost = completion_cost(
completion_response=response,
model=model,
custom_llm_provider="anthropic",
service_tier="auto",
optional_params={"service_tier": "auto"},
)
expected_priority = 1000 * 6e-6 + 500 * 30e-6
assert cost == pytest.approx(expected_priority)
def test_anthropic_cost_per_token_prices_cache_at_served_tier_with_multiplier():
"""
Regression for the cache/tier interaction in the Anthropic geo/speed path.
When a request is served at "priority" and also carries a geo/speed
multiplier (here ``speed="fast"``), the cache portion is held out of the
multiplier so it is not scaled. That held-out cache cost must use the
served tier's cache rate; pricing it at the standard rate while the cache
embedded in ``prompt_cost`` is priced at the priority rate leaves a
``(cache_priority - cache_standard)(multiplier - 1)`` billing error.
"""
from litellm.llms.anthropic.cost_calculation import (
cost_per_token as anthropic_cost_per_token,
)
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "claude-test-priority-cache-fast-model"
litellm.register_model(
model_cost={
model: {
"input_cost_per_token": 3e-6,
"output_cost_per_token": 15e-6,
"cache_read_input_token_cost": 0.3e-6,
"input_cost_per_token_priority": 6e-6,
"output_cost_per_token_priority": 30e-6,
"cache_read_input_token_cost_priority": 0.6e-6,
"litellm_provider": "anthropic",
"max_tokens": 8192,
"provider_specific_entry": {"fast": 2.0},
}
}
)
usage = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=200),
)
usage.speed = "fast"
prompt_cost, completion_cost = anthropic_cost_per_token(
model=model, usage=usage, service_tier="priority"
)
# non-cache input priced at the priority rate and scaled by the fast
# multiplier; the 200 cache-hit tokens priced at the priority cache rate
# and held out of the multiplier
expected_prompt = (1000 - 200) * 6e-6 * 2 + 200 * 0.6e-6
expected_completion = 500 * 30e-6 * 2
assert prompt_cost == pytest.approx(expected_prompt)
assert completion_cost == pytest.approx(expected_completion)
def test_gemini_cache_tokens_details_no_negative_values():
"""
Test for Issue #18750: Negative text_tokens with Gemini caching

View file

@ -0,0 +1,68 @@
import json
from pathlib import Path
import pytest
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
@pytest.mark.parametrize("model", ["azure_ai/gpt-5.5", "azure_ai/gpt-5.5-2026-04-23"])
def test_azure_ai_gpt_5_5_model_info(model):
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
info = model_cost.get(model)
assert (
info is not None
), f"{model} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "azure_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 5e-06
assert info["output_cost_per_token"] == 3e-05
assert info["cache_read_input_token_cost"] == 5e-07
assert info["input_cost_per_token_above_272k_tokens"] == 1e-05
assert info["output_cost_per_token_above_272k_tokens"] == 4.5e-05
assert info["cache_read_input_token_cost_above_272k_tokens"] == 1e-06
assert info["input_cost_per_token_priority"] == 1e-05
assert info["output_cost_per_token_priority"] == 6e-05
assert info["max_input_tokens"] == 1050000
assert info["max_output_tokens"] == 128000
assert info["max_tokens"] == 128000
assert info["supports_function_calling"] is True
assert info["supports_prompt_caching"] is True
assert info["supports_reasoning"] is True
assert info["supports_response_schema"] is True
assert info["supports_tool_choice"] is True
assert info["supports_vision"] is True
assert info["supports_web_search"] is True
# gpt-5.5 dropped minimal reasoning effort support (true on gpt-5.4)
assert info["supports_minimal_reasoning_effort"] is False
routed_model, provider, _, _ = get_llm_provider(model=model)
assert routed_model == model.split("/", 1)[1]
# azure_ai/* models resolve under the azure provider in get_llm_provider
assert provider == "azure"
def test_azure_ai_gpt_5_5_backup_matches_main():
"""Ensure the bundled model cost map stays in sync with the canonical file."""
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"
with open(main_path) as f:
main_cost = json.load(f)
with open(backup_path) as f:
backup_cost = json.load(f)
for model in ("azure_ai/gpt-5.5", "azure_ai/gpt-5.5-2026-04-23"):
assert backup_cost.get(model) == main_cost.get(
model
), f"{model} differs between main and backup model cost maps"

View file

@ -0,0 +1,55 @@
import json
from pathlib import Path
import pytest
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
@pytest.mark.parametrize("model", ["mistral/mistral-medium-3-5"])
def test_mistral_medium_3_5_model_info(model):
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
info = model_cost.get(model)
assert (
info is not None
), f"{model} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "mistral"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 1.5e-06
assert info["output_cost_per_token"] == 7.5e-06
assert info["max_input_tokens"] == 262144
assert info["max_output_tokens"] == 262144
assert info["max_tokens"] == 262144
assert info["supports_function_calling"] is True
assert info["supports_response_schema"] is True
assert info["supports_tool_choice"] is True
assert info["supports_vision"] is True
assert info["supports_assistant_prefill"] is True
routed_model, provider, _, _ = get_llm_provider(model=model)
assert routed_model == model.split("/", 1)[1]
assert provider == "mistral"
def test_mistral_medium_3_5_backup_matches_main():
"""Ensure the bundled model cost map stays in sync with the canonical file."""
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"
with open(main_path) as f:
main_cost = json.load(f)
with open(backup_path) as f:
backup_cost = json.load(f)
for model in ("mistral/mistral-medium-3-5",):
assert backup_cost.get(model) == main_cost.get(
model
), f"{model} differs between main and backup model cost maps"

View file

@ -365,3 +365,37 @@ async def test_router_order_fallback_with_wildcard_model_group():
messages=[{"role": "user", "content": "hi"}],
)
assert response._hidden_params["model_id"] == "2"
def test_check_non_standard_fallback_format():
from litellm.router_utils.fallback_event_handlers import (
_check_non_standard_fallback_format,
)
# Standard formats
assert (
_check_non_standard_fallback_format([{"gpt-3.5-turbo": ["claude-3-haiku"]}])
== False
)
assert _check_non_standard_fallback_format([{"model": ["qwen-backup"]}]) == False
assert (
_check_non_standard_fallback_format(
[{"model": ["qwen-backup"], "region": ["us-east-1"]}]
)
== False
)
# Non-standard formats
assert _check_non_standard_fallback_format([{"model": "qwen-backup"}]) == True
assert (
_check_non_standard_fallback_format(
[{"model": "qwen-backup", "messages": [{"role": "user", "content": "hi"}]}]
)
== True
)
assert (
_check_non_standard_fallback_format(
[{"model": ["qwen-backup"], "api_key": "some-key"}]
)
== True
)

View file

@ -890,6 +890,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"/v1/audio/speech",
"/v1/ocr",
"/vertex_ai/live",
"/v1/realtime/transcription_sessions",
],
},
},
@ -4153,6 +4154,96 @@ class TestValidateAndFixThinkingParam:
assert "budget_tokens" not in thinking
def test_deepseek_v4_models_in_cost_map():
"""
Test that deepseek-v4-flash and deepseek-v4-pro entries are correctly
configured in model_prices_and_context_window.json.
Prices sourced from https://api-docs.deepseek.com/quick_start/pricing:
- deepseek-v4-flash: $0.14/M input, $0.28/M output
- deepseek-v4-pro: $0.435/M input, $0.87/M output (75% discounted active price)
Closes https://github.com/BerriAI/litellm/issues/26709
"""
import json
from pathlib import Path
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
# --- bare model names ---
for key, expected_input, expected_output, expected_cache in [
("deepseek-v4-flash", 1.4e-07, 2.8e-07, 2.8e-09),
("deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
assert info is not None, f"{key} missing from model_prices_and_context_window.json"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
assert info["max_input_tokens"] == 1_000_000
assert info["supports_function_calling"] is True
assert info["supports_tool_choice"] is True
# --- provider-prefixed names ---
for key, expected_input, expected_output, expected_cache in [
("deepseek/deepseek-v4-flash", 1.4e-07, 2.8e-07, 2.8e-09),
("deepseek/deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
assert info is not None, f"{key} missing from model_prices_and_context_window.json"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
assert info["supports_function_calling"] is True
assert info["supports_tool_choice"] is True
def test_deepseek_v4_models_in_backup_cost_map():
"""
Test that deepseek-v4-flash and deepseek-v4-pro entries are correctly
configured in litellm/model_prices_and_context_window_backup.json.
"""
import json
from pathlib import Path
json_path = Path(__file__).parents[2] / "litellm" / "model_prices_and_context_window_backup.json"
with open(json_path) as f:
model_cost = json.load(f)
# --- bare model names ---
for key, expected_input, expected_output, expected_cache in [
("deepseek-v4-flash", 1.4e-07, 2.8e-07, 2.8e-09),
("deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
assert info is not None, f"{key} missing from backup JSON"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
assert info["max_input_tokens"] == 1_000_000
# --- provider-prefixed names ---
for key, expected_input, expected_output, expected_cache in [
("deepseek/deepseek-v4-flash", 1.4e-07, 2.8e-07, 2.8e-09),
("deepseek/deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
assert info is not None, f"{key} missing from backup JSON"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
class TestBedrockBaseModelLabelKeepsTools:
"""Regression for #29618: a Bedrock deployment whose ``base_model`` is a friendly
label must not silently drop ``tools``/``tool_choice`` under ``drop_params``."""
@ -4217,3 +4308,4 @@ def test_aws_bedrock_project_id_excluded_from_bedrock_optional_params():
assert "aws_bedrock_project_id" not in result
assert result["aws_region_name"] == "us-east-1"

View file

@ -188,28 +188,30 @@ function LoginPageContent() {
<Text type="secondary">Access your LiteLLM Admin UI.</Text>
</div>
<Alert
message="Default Credentials"
description={
<>
<Paragraph className="text-sm">
By default, Username is <code className="bg-gray-100 px-1 py-0.5 rounded text-xs">admin</code> and
Password is your set LiteLLM Proxy
<code className="bg-gray-100 px-1 py-0.5 rounded text-xs">MASTER_KEY</code>.
</Paragraph>
<Paragraph className="text-sm">
Need to set UI credentials or SSO?{" "}
<a href="https://docs.litellm.ai/docs/proxy/ui" target="_blank" rel="noopener noreferrer">
Check the documentation
</a>
.
</Paragraph>
</>
}
type="info"
icon={<InfoCircleOutlined />}
showIcon
/>
{!uiConfig?.hide_default_credentials_hint && (
<Alert
message="Default Credentials"
description={
<>
<Paragraph className="text-sm">
By default, Username is <code className="bg-gray-100 px-1 py-0.5 rounded text-xs">admin</code> and
Password is your set LiteLLM Proxy
<code className="bg-gray-100 px-1 py-0.5 rounded text-xs">MASTER_KEY</code>.
</Paragraph>
<Paragraph className="text-sm">
Need to set UI credentials or SSO?{" "}
<a href="https://docs.litellm.ai/docs/proxy/ui" target="_blank" rel="noopener noreferrer">
Check the documentation
</a>
.
</Paragraph>
</>
}
type="info"
icon={<InfoCircleOutlined />}
showIcon
/>
)}
{error && <Alert message={error} type="error" showIcon />}

View file

@ -55,4 +55,39 @@ describe("LoggingCallbacksTable", () => {
);
expect(getByText("custom_callback_x")).toBeInTheDocument();
});
// Regression: `/get_callbacks` returns the same `name` twice when a
// callback is registered for both success and failure (e.g. `generic_api`
// → POST to spend-log on both 200 and 4xx/5xx). The UI used to ignore
// the `type` field and render every row as "Success", masking the
// failure registration. Reading `record.type` fixes the badge AND
// composing the rowKey with type avoids React's duplicate-key warning.
it("renders distinct Success and Failure badges for same-name dual registration", () => {
const baseVars = {
SLACK_WEBHOOK_URL: null,
LANGFUSE_PUBLIC_KEY: null,
LANGFUSE_SECRET_KEY: null,
LANGFUSE_HOST: null,
OPENMETER_API_KEY: null,
};
const { getAllByText, getByText } = render(
<LoggingCallbacksTable
callbacks={[
{ name: "generic_api", type: "success", variables: baseVars },
{ name: "generic_api", type: "failure", variables: baseVars },
]}
availableCallbacks={{
generic_api: {
litellm_callback_name: "generic_api",
litellm_callback_params: [],
ui_callback_name: "Custom Callback API",
},
}}
/>,
);
// Both rows show the same display name, but distinct mode badges.
expect(getAllByText("Custom Callback API")).toHaveLength(2);
expect(getByText("Success")).toBeInTheDocument();
expect(getByText("Failure")).toBeInTheDocument();
});
});

View file

@ -48,7 +48,6 @@ export const LoggingCallbacksTable: React.FC<LoggingCallbacksProps> = ({
key: "name",
render: (_: string, record: CallbackRow) => {
const id = record.name;
console.log("availableCallbacks", availableCallbacks);
const displayName = availableCallbacks[id]?.ui_callback_name || id;
return <div className="font-medium text-gray-800">{displayName}</div>;
},
@ -57,7 +56,10 @@ export const LoggingCallbacksTable: React.FC<LoggingCallbacksProps> = ({
title: <span className="font-medium text-gray-700">Mode</span>,
key: "mode",
render: (_: unknown, record: CallbackRow) => {
const mode = record.mode || "success";
// Backend sends `type` (success | failure); legacy in-memory rows
// from add-callback flow set `mode`. Read both so newly-added rows
// and server-fetched rows both render correctly.
const mode = record.type || record.mode || "success";
const label = CALLBACK_MODES.find((m) => m.value === mode)?.label || mode;
const badgeClass =
mode === "success"
@ -109,7 +111,10 @@ export const LoggingCallbacksTable: React.FC<LoggingCallbacksProps> = ({
<Table
columns={columns}
dataSource={callbacks as CallbackRow[]}
rowKey={(record) => record.name}
// `generic_api` can appear as both a success and a failure
// callback simultaneously — keying by `name` alone produced
// duplicate React keys. Compose with type to keep keys unique.
rowKey={(record) => `${record.name}-${record.type || record.mode || "success"}`}
pagination={false}
rowClassName={() => "hover:bg-gray-50"}
/>

View file

@ -1,5 +1,12 @@
export interface AlertingObject {
name: string;
// Backend distinguishes success vs failure callback registrations
// (`/get_callbacks` returns `type: "success" | "failure"`). Same callback
// (e.g. `generic_api`) can appear twice — once per event class — and
// those entries fire on disjoint events, not double-fire on one event.
// UI must read this to render the correct badge; missing it caused
// every row to render as "Success".
type?: "success" | "failure" | "success_and_failure";
variables: AlertingVariables;
}

View file

@ -4,6 +4,7 @@ import type { UploadProps } from "antd/es/upload";
import React, { useState } from "react";
import ProviderSpecificFields from "../add_model/provider_specific_fields";
import { Providers, providerLogoMap } from "../provider_info_helpers";
import { resetCredentialFormOnProviderChange } from "./credential_form_helpers";
const { Link } = Typography;
interface AddCredentialsModalProps {
@ -59,8 +60,7 @@ const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCance
<AntdSelect
showSearch
onChange={(value) => {
setSelectedProvider(value as Providers);
form.setFieldValue("custom_llm_provider", value);
resetCredentialFormOnProviderChange(form, value as Providers, setSelectedProvider);
}}
>
{Object.entries(Providers).map(([providerEnum, providerDisplayName]) => (

View file

@ -5,6 +5,7 @@ import { useEffect, useState } from "react";
import ProviderSpecificFields from "../add_model/provider_specific_fields";
import { CredentialItem } from "../networking";
import { Providers, providerLogoMap } from "../provider_info_helpers";
import { resetCredentialFormOnProviderChange } from "./credential_form_helpers";
const { Link } = Typography;
interface EditCredentialsModalProps {
@ -92,8 +93,7 @@ export default function EditCredentialsModal({
<AntdSelect
showSearch
onChange={(value) => {
setSelectedProvider(value as Providers);
form.setFieldValue("custom_llm_provider", value);
resetCredentialFormOnProviderChange(form, value as Providers, setSelectedProvider);
}}
>
{Object.entries(Providers).map(([providerEnum, providerDisplayName]) => (

View file

@ -0,0 +1,84 @@
import type { FormInstance } from "antd";
import { describe, expect, it, vi } from "vitest";
import { Providers } from "../provider_info_helpers";
import { resetCredentialFormOnProviderChange } from "./credential_form_helpers";
/**
* Build a minimal FormInstance stub that records calls. We don't depend
* on the full Antd API surface — only the three methods the helper uses.
*/
function makeFormStub(initialFields: Record<string, unknown> = {}) {
const fields: Record<string, unknown> = { ...initialFields };
const stub = {
getFieldValue: vi.fn((key: string) => fields[key]),
setFieldValue: vi.fn((key: string, value: unknown) => {
fields[key] = value;
}),
resetFields: vi.fn(() => {
Object.keys(fields).forEach((k) => delete fields[k]);
}),
};
return { stub: stub as unknown as FormInstance, fields, calls: stub };
}
describe("resetCredentialFormOnProviderChange", () => {
it("clears all fields when switching providers", () => {
// Simulate the OpenAI->Google AI Studio leak: api_base picked up
// OpenAI's default value and the user typed a custom URL.
const { stub, fields, calls } = makeFormStub({
credential_name: "my-prod-key",
custom_llm_provider: "OpenAI",
api_base: "https://api.openai.com/v1",
api_key: "sk-stale-openai-key",
organization: "org-leak",
});
const setSelectedProvider = vi.fn();
resetCredentialFormOnProviderChange(stub, Providers.Google_AI_Studio, setSelectedProvider);
expect(calls.resetFields).toHaveBeenCalledTimes(1);
// Provider-specific fields must be gone so the next render starts
// from the new provider's default_value, not OpenAI's leftover.
expect(fields.api_base).toBeUndefined();
expect(fields.api_key).toBeUndefined();
expect(fields.organization).toBeUndefined();
});
it("preserves credential_name across the switch", () => {
// credential_name is user-supplied metadata, not provider-specific.
// The admin shouldn't have to retype it just because they re-picked
// the provider.
const { stub, fields } = makeFormStub({
credential_name: "my-prod-key",
custom_llm_provider: "OpenAI",
api_base: "https://api.openai.com/v1",
});
resetCredentialFormOnProviderChange(stub, Providers.Google_AI_Studio, vi.fn());
expect(fields.credential_name).toBe("my-prod-key");
});
it("updates custom_llm_provider and selectedProvider state to the new value", () => {
const { stub, fields } = makeFormStub({ credential_name: "x" });
const setSelectedProvider = vi.fn();
resetCredentialFormOnProviderChange(stub, Providers.Google_AI_Studio, setSelectedProvider);
expect(fields.custom_llm_provider).toBe(Providers.Google_AI_Studio);
expect(setSelectedProvider).toHaveBeenCalledExactlyOnceWith(Providers.Google_AI_Studio);
});
it("does not call setFieldValue('credential_name', undefined) when the name was unset", () => {
// Edge case: brand-new modal with no name typed yet. We shouldn't
// explicitly write `undefined` back into the form (Antd treats that
// as a touched empty field, triggering the "required" validation
// prematurely).
const { stub, calls } = makeFormStub({});
resetCredentialFormOnProviderChange(stub, Providers.Anthropic, vi.fn());
const credentialNameCalls = calls.setFieldValue.mock.calls.filter(([key]) => key === "credential_name");
expect(credentialNameCalls).toHaveLength(0);
});
});

View file

@ -0,0 +1,33 @@
import type { FormInstance } from "antd";
import { Providers } from "../provider_info_helpers";
/**
* Reset the credential form when the user switches providers.
*
* Why: provider-specific fields (api_base, api_key, organization, ...)
* share a single Antd Form state across providers. Without this reset,
* the previous provider's values stick around — most visibly, OpenAI's
* default `api_base` (https://api.openai.com/v1) carries over when the
* user switches to Google AI Studio, overriding that provider's own
* default_value.
*
* Strategy: blow away the whole form, then restore the provider-agnostic
* fields (credential name + the new provider id) so the newly rendered
* `ProviderSpecificFields` can apply its own defaults from a clean slate.
*
* The credential name is preserved because it's a user-supplied label
* that shouldn't reset just because the admin re-selected a provider.
*/
export function resetCredentialFormOnProviderChange(
form: FormInstance,
newProvider: Providers,
setSelectedProvider: (p: Providers) => void,
): void {
const preservedName = form.getFieldValue("credential_name");
form.resetFields();
if (preservedName !== undefined) {
form.setFieldValue("credential_name", preservedName);
}
setSelectedProvider(newProvider);
form.setFieldValue("custom_llm_provider", newProvider);
}

View file

@ -285,6 +285,7 @@ export interface LiteLLMWellKnownUiConfig {
auto_redirect_to_sso: boolean;
admin_ui_disabled: boolean;
sso_configured: boolean;
hide_default_credentials_hint?: boolean;
is_control_plane?: boolean;
workers?: WorkerInfo[];
}

View file

@ -25016,6 +25016,10 @@ export interface components {
cache_read_input_token_cost?: number | null;
/** Cache Read Input Token Cost Above 200K Tokens */
cache_read_input_token_cost_above_200k_tokens?: number | null;
/** Cache Read Input Token Cost Above 200K Tokens Priority */
cache_read_input_token_cost_above_200k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Above 272K Tokens Priority */
cache_read_input_token_cost_above_272k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Flex */
cache_read_input_token_cost_flex?: number | null;
/** Cache Read Input Token Cost Priority */
@ -25064,6 +25068,10 @@ export interface components {
input_cost_per_token_above_128k_tokens?: number | null;
/** Input Cost Per Token Above 200K Tokens */
input_cost_per_token_above_200k_tokens?: number | null;
/** Input Cost Per Token Above 200K Tokens Priority */
input_cost_per_token_above_200k_tokens_priority?: number | null;
/** Input Cost Per Token Above 272K Tokens Priority */
input_cost_per_token_above_272k_tokens_priority?: number | null;
/** Input Cost Per Token Batches */
input_cost_per_token_batches?: number | null;
/** Input Cost Per Token Cache Hit */
@ -25137,6 +25145,10 @@ export interface components {
output_cost_per_token_above_128k_tokens?: number | null;
/** Output Cost Per Token Above 200K Tokens */
output_cost_per_token_above_200k_tokens?: number | null;
/** Output Cost Per Token Above 200K Tokens Priority */
output_cost_per_token_above_200k_tokens_priority?: number | null;
/** Output Cost Per Token Above 272K Tokens Priority */
output_cost_per_token_above_272k_tokens_priority?: number | null;
/** Output Cost Per Token Batches */
output_cost_per_token_batches?: number | null;
/** Output Cost Per Token Flex */
@ -31108,6 +31120,11 @@ export interface components {
admin_ui_disabled: boolean;
/** Auto Redirect To Sso */
auto_redirect_to_sso: boolean;
/**
* Hide Default Credentials Hint
* @default false
*/
hide_default_credentials_hint: boolean;
/**
* Is Control Plane
* @default false
@ -32657,6 +32674,10 @@ export interface components {
cache_read_input_token_cost?: number | null;
/** Cache Read Input Token Cost Above 200K Tokens */
cache_read_input_token_cost_above_200k_tokens?: number | null;
/** Cache Read Input Token Cost Above 200K Tokens Priority */
cache_read_input_token_cost_above_200k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Above 272K Tokens Priority */
cache_read_input_token_cost_above_272k_tokens_priority?: number | null;
/** Cache Read Input Token Cost Flex */
cache_read_input_token_cost_flex?: number | null;
/** Cache Read Input Token Cost Priority */
@ -32705,6 +32726,10 @@ export interface components {
input_cost_per_token_above_128k_tokens?: number | null;
/** Input Cost Per Token Above 200K Tokens */
input_cost_per_token_above_200k_tokens?: number | null;
/** Input Cost Per Token Above 200K Tokens Priority */
input_cost_per_token_above_200k_tokens_priority?: number | null;
/** Input Cost Per Token Above 272K Tokens Priority */
input_cost_per_token_above_272k_tokens_priority?: number | null;
/** Input Cost Per Token Batches */
input_cost_per_token_batches?: number | null;
/** Input Cost Per Token Cache Hit */
@ -32778,6 +32803,10 @@ export interface components {
output_cost_per_token_above_128k_tokens?: number | null;
/** Output Cost Per Token Above 200K Tokens */
output_cost_per_token_above_200k_tokens?: number | null;
/** Output Cost Per Token Above 200K Tokens Priority */
output_cost_per_token_above_200k_tokens_priority?: number | null;
/** Output Cost Per Token Above 272K Tokens Priority */
output_cost_per_token_above_272k_tokens_priority?: number | null;
/** Output Cost Per Token Batches */
output_cost_per_token_batches?: number | null;
/** Output Cost Per Token Flex */