mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'upstream/litellm_internal_staging' into deepkeep-as-internal
This commit is contained in:
commit
b892470ba3
88 changed files with 12279 additions and 793 deletions
2
.github/workflows/test-linting.yml
vendored
2
.github/workflows/test-linting.yml
vendored
|
|
@ -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: |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
```
|
||||
|
|
|
|||
14
Makefile
14
Makefile
|
|
@ -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
5974
any-discipline-budget.json
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 ##
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
191
litellm/proxy/dev_config.yaml
Normal file
191
litellm/proxy/dev_config.yaml
Normal 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
|
||||
|
|
@ -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 [],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -162,6 +162,8 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS = (
|
|||
"secret_fields",
|
||||
"_guardrail_pipelines",
|
||||
"_pipeline_managed_guardrails",
|
||||
"client_disconnected",
|
||||
"error_information",
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
52
litellm/proxy/wildcard_config.yaml
Normal file
52
litellm/proxy/wildcard_config.yaml
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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".
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"}]}
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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] == ["*"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
68
tests/test_litellm/test_gpt_5_5_model_metadata.py
Normal file
68
tests/test_litellm/test_gpt_5_5_model_metadata.py
Normal 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"
|
||||
55
tests/test_litellm/test_mistral_medium_3_5_model_metadata.py
Normal file
55
tests/test_litellm/test_mistral_medium_3_5_model_metadata.py
Normal 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"
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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 />}
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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]) => (
|
||||
|
|
|
|||
|
|
@ -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]) => (
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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);
|
||||
}
|
||||
|
|
@ -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[];
|
||||
}
|
||||
|
|
|
|||
29
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
29
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue