Merge pull request #33491 from BerriAI/litellm_internal_staging
Some checks are pending
CodeQL / Analyze (actions) (push) Waiting to run
CodeQL / Analyze (javascript-typescript) (push) Waiting to run
CodeQL / Analyze (python) (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run

chore(ci): promote internal staging to main
This commit is contained in:
yuneng-jiang 2026-07-15 20:34:53 -07:00 • committed by GitHub
commit 229159c790
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
94 changed files with 8179 additions and 1526 deletions

View file

@ -36,7 +36,7 @@ RUN uv venv --python python && \
"opentelemetry-api==1.28.0" \
"opentelemetry-sdk==1.28.0" \
"opentelemetry-exporter-otlp==1.28.0" \
"ddtrace==2.19.0" \
"ddtrace==4.11.0" \
"sentry-sdk==2.21.0" \
"mangum==0.17.0" \
"azure-ai-contentsafety==1.0.0" \

View file

@ -7,7 +7,8 @@
# Thank you users! We ❤️ you! - Krrish & Ishaan
## This provides an LLM Guard Integration for content moderation on the proxy
from typing import Literal, Optional
import asyncio
from typing import Optional
import aiohttp
from fastapi import HTTPException
@ -18,7 +19,6 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import CallTypesLiteral
from litellm.utils import get_formatted_prompt
class _ENTERPRISE_LLMGuard(CustomLogger):
@ -46,45 +46,44 @@ class _ENTERPRISE_LLMGuard(CustomLogger):
except Exception:
pass
async def moderation_check(self, text: str):
async def moderation_check(self, text: str) -> str:
"""
Runs the LLM Guard moderation check on ``text``.
Raises an HTTPException when the content violates the safety policy;
otherwise returns the sanitized prompt from LLM Guard, falling back to
the original text when the API does not provide one.
[TODO] make this more performant for high-throughput scenario
"""
try:
async with aiohttp.ClientSession() as session:
if self.mock_redacted_text is not None:
redacted_text = self.mock_redacted_text
else:
# Make the first request to /analyze
analyze_url = f"{self.llm_guard_api_base}analyze/prompt"
verbose_proxy_logger.debug("Making request to: %s", analyze_url)
analyze_payload = {"prompt": text}
redacted_text = None
if self.mock_redacted_text is not None:
redacted_text = self.mock_redacted_text
else:
analyze_url = f"{self.llm_guard_api_base}analyze/prompt"
verbose_proxy_logger.debug("Making request to: %s", analyze_url)
async with aiohttp.ClientSession() as session:
async with session.post(
analyze_url, json=analyze_payload
analyze_url, json={"prompt": text}
) as response:
redacted_text = await response.json()
verbose_proxy_logger.debug(
f"LLM Guard: Received response - {redacted_text}"
verbose_proxy_logger.debug(
f"LLM Guard: Received response - {redacted_text}"
)
if redacted_text is None:
raise HTTPException(
status_code=500,
detail={
"error": f"Invalid content moderation response: {redacted_text}"
},
)
if redacted_text is not None:
if (
redacted_text.get("is_valid", None) is not None
and redacted_text["is_valid"] is False
):
raise HTTPException(
status_code=400,
detail={"error": "Violated content safety policy"},
)
else:
pass
else:
raise HTTPException(
status_code=500,
detail={
"error": f"Invalid content moderation response: {redacted_text}"
},
)
if redacted_text.get("is_valid", None) is False:
raise HTTPException(
status_code=400,
detail={"error": "Violated content safety policy"},
)
sanitized_prompt = redacted_text.get("sanitized_prompt")
return sanitized_prompt if isinstance(sanitized_prompt, str) else text
except Exception as e:
verbose_proxy_logger.exception(
"litellm.enterprise.enterprise_hooks.llm_guard::moderation_check - Exception occurred - {}".format(
@ -138,23 +137,75 @@ class _ENTERPRISE_LLMGuard(CustomLogger):
return
self.print_verbose("Makes LLM Guard Check")
try:
assert call_type in [
"completion",
"embeddings",
"image_generation",
"moderation",
"audio_transcription",
]
except Exception:
if call_type not in [
"completion",
"embeddings",
"image_generation",
"moderation",
"audio_transcription",
]:
self.print_verbose(
f"Call Type - {call_type}, not in accepted list - ['completion','embeddings','image_generation','moderation','audio_transcription']"
)
return data
formatted_prompt = get_formatted_prompt(data=data, call_type=call_type) # type: ignore
self.print_verbose(f"LLM Guard, formatted_prompt: {formatted_prompt}")
return await self.moderation_check(text=formatted_prompt)
return await self._moderate_request(data=data)
async def _moderate_request(self, data: dict) -> dict:
"""
Sanitizes the request in place using the prompt returned by LLM Guard so
the provider-bound request carries the redacted content, then returns it.
"""
messages = data.get("messages")
if messages is not None:
data["messages"] = list(
await asyncio.gather(
*(self._moderate_message(message) for message in messages)
)
)
return data
input_ = data.get("input")
if input_ is not None:
data["input"] = await self._moderate_input(input_)
return data
prompt = data.get("prompt")
if isinstance(prompt, str):
data["prompt"] = await self.moderation_check(text=prompt)
return data
async def _moderate_message(self, message: dict) -> dict:
content = message.get("content")
if isinstance(content, str):
return {**message, "content": await self.moderation_check(text=content)}
if isinstance(content, list):
return {
**message,
"content": list(
await asyncio.gather(
*(self._moderate_content_part(part) for part in content)
)
),
}
return message
async def _moderate_content_part(self, part: dict) -> dict:
if part.get("type") == "text" and isinstance(part.get("text"), str):
return {**part, "text": await self.moderation_check(text=part["text"])}
return part
async def _moderate_input(self, input_: object) -> object:
if isinstance(input_, str):
return await self.moderation_check(text=input_)
if isinstance(input_, list):
return [
await self.moderation_check(text=item)
if isinstance(item, str)
else item
for item in input_
]
return input_
async def async_post_call_streaming_hook(
self, user_api_key_dict: UserAPIKeyAuth, response: str

View file

@ -85,6 +85,22 @@ class CachingHandlerResponse(BaseModel):
in_memory_cache_obj = InMemoryCache()
def _drop_logging_obj_from_kwargs(request_kwargs: dict[str, object]) -> dict[str, object]:
"""
The caching handler is stored on the Logging object
(``logging_obj._llm_caching_handler``), so keeping ``litellm_logging_obj``
inside ``request_kwargs`` closes a reference cycle
(Logging -> LLMCachingHandler -> kwargs -> Logging) that keeps the full
request payload (messages included) alive until a generational GC pass
instead of being freed by refcount when the request ends. Nothing in the
caching layer reads the logging object from these kwargs; cache-key
generation ignores litellm-internal params.
"""
if "litellm_logging_obj" not in request_kwargs:
return request_kwargs
return {k: v for k, v in request_kwargs.items() if k != "litellm_logging_obj"}
def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
cached_id = cached_result.get("id")
if isinstance(cached_id, str) and cached_id.startswith("chatcmpl"):
@ -118,7 +134,7 @@ class LLMCachingHandler:
self.async_streaming_chunks: List[ModelResponse] = []
self.sync_streaming_chunks: List[ModelResponse] = []
self.request_kwargs = request_kwargs
self.request_kwargs = _drop_logging_obj_from_kwargs(request_kwargs)
self.preset_cache_key: Optional[str] = None
self.original_function = original_function
self.start_time = start_time
@ -297,7 +313,7 @@ class LLMCachingHandler:
new_kwargs.pop("metadata", None)
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
self.request_kwargs = new_kwargs
self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs)
print_verbose("Checking Sync Cache")
cached_result = litellm.cache.get_cache(**new_kwargs)
if cached_result is not None:
@ -693,7 +709,7 @@ class LLMCachingHandler:
new_kwargs.pop("metadata", None)
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
self.request_kwargs = new_kwargs
self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs)
cached_result: Optional[Any] = None
if call_type == CallTypes.aembedding.value:
if isinstance(new_kwargs["input"], str):

View file

@ -1496,6 +1496,7 @@ MAX_TEAM_LIST_LIMIT = int(os.getenv("MAX_TEAM_LIST_LIMIT", 20))
MAX_POLICY_ESTIMATE_IMPACT_ROWS = int(os.getenv("MAX_POLICY_ESTIMATE_IMPACT_ROWS", 1000))
DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD", 0.7))
LENGTH_OF_LITELLM_GENERATED_KEY = int(os.getenv("LENGTH_OF_LITELLM_GENERATED_KEY", 16))
MINIMUM_CUSTOM_KEY_LENGTH = int(os.getenv("MINIMUM_CUSTOM_KEY_LENGTH", 16))
SECRET_MANAGER_REFRESH_INTERVAL = int(os.getenv("SECRET_MANAGER_REFRESH_INTERVAL", 86400))
LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [
"default_internal_user_params",

View file

@ -10,7 +10,7 @@ from typing import TYPE_CHECKING, Any, Optional, Union
from litellm.secret_managers.main import get_secret_bool
if TYPE_CHECKING:
from ddtrace.tracer import Tracer as DD_TRACER
from ddtrace.trace import Tracer as DD_TRACER
else:
DD_TRACER = Any

View file

@ -925,7 +925,6 @@ class Logging(LiteLLMLoggingBaseClass):
def pre_call(self, input, api_key, model=None, additional_args={}):
# Log the exact input to the LLM API
litellm.error_logs["PRE_CALL"] = locals()
try:
self._pre_call(
input=input,
@ -1135,7 +1134,6 @@ class Logging(LiteLLMLoggingBaseClass):
def post_call(self, original_response, input=None, api_key=None, additional_args={}):
# Log the exact result from the LLM API, for streaming - log the type of response received
litellm.error_logs["POST_CALL"] = locals()
if isinstance(original_response, dict):
original_response = json.dumps(original_response, default=str)
try:
@ -3074,7 +3072,7 @@ class Logging(LiteLLMLoggingBaseClass):
def get_combined_callback_list(self, dynamic_success_callbacks: Optional[List], global_callbacks: List) -> List:
if dynamic_success_callbacks is None:
return list(global_callbacks)
return list(set(dynamic_success_callbacks + global_callbacks))
return list(dict.fromkeys(dynamic_success_callbacks + global_callbacks))
def _remove_internal_litellm_callbacks(self, callbacks: List) -> List:
"""

View file

@ -77,6 +77,22 @@ def _redact_streaming_response(streaming_response):
streaming_response.reasoning = None
def _redact_tool_calls(tool_calls) -> None:
"""Redact tool call arguments (assistant tool calls carry prompt-derived data)."""
if not tool_calls:
return
for tool_call in tool_calls:
function = getattr(tool_call, "function", None)
if function is not None and hasattr(function, "arguments"):
function.arguments = "redacted-by-litellm"
def _redact_function_call(function_call) -> None:
"""Redact legacy assistant function_call arguments."""
if function_call is not None and hasattr(function_call, "arguments"):
function_call.arguments = "redacted-by-litellm"
def _redact_choice_content(choice):
"""Helper to redact content in a choice (message or delta)."""
if isinstance(choice, litellm.Choices):
@ -85,12 +101,16 @@ def _redact_choice_content(choice):
choice.message.reasoning_content = "redacted-by-litellm"
if hasattr(choice.message, "thinking_blocks"):
choice.message.thinking_blocks = None
_redact_tool_calls(getattr(choice.message, "tool_calls", None))
_redact_function_call(getattr(choice.message, "function_call", None))
elif isinstance(choice, litellm.utils.StreamingChoices):
choice.delta.content = "redacted-by-litellm"
if hasattr(choice.delta, "reasoning_content"):
choice.delta.reasoning_content = "redacted-by-litellm"
if hasattr(choice.delta, "thinking_blocks"):
choice.delta.thinking_blocks = None
_redact_tool_calls(getattr(choice.delta, "tool_calls", None))
_redact_function_call(getattr(choice.delta, "function_call", None))
def _redact_responses_api_output(output_items):
@ -111,6 +131,9 @@ def _redact_responses_api_output(output_items):
if hasattr(summary_item, "text"):
summary_item.text = "redacted-by-litellm"
if hasattr(output_item, "type") and output_item.type == "function_call" and hasattr(output_item, "arguments"):
output_item.arguments = "redacted-by-litellm"
def _redact_responses_api_output_dict(output_items, redacted_str: str):
"""Helper to redact ResponsesAPIResponse output items in dict form."""
@ -131,6 +154,9 @@ def _redact_responses_api_output_dict(output_items, redacted_str: str):
if isinstance(summary_item, dict) and "text" in summary_item:
summary_item["text"] = redacted_str
if output_item.get("type") == "function_call" and "arguments" in output_item:
output_item["arguments"] = redacted_str
def _redact_standard_logging_object(model_call_details: dict):
"""Redact messages and response inside standard_logging_object if present."""
@ -162,6 +188,19 @@ def _redact_standard_logging_object(model_call_details: dict):
standard_logging_object["response"] = {"text": redacted_str}
def _redact_tool_calls_dict(message: dict, redacted_str: str) -> None:
"""Redact tool call / function_call arguments in a dict-form message or delta."""
tool_calls = message.get("tool_calls")
if isinstance(tool_calls, list):
for tool_call in tool_calls:
if isinstance(tool_call, dict) and isinstance(tool_call.get("function"), dict):
tool_call["function"]["arguments"] = redacted_str
function_call = message.get("function_call")
if isinstance(function_call, dict) and "arguments" in function_call:
function_call["arguments"] = redacted_str
def _redact_model_response_dict_choices(choices, redacted_str: str):
for choice in choices:
if isinstance(choice, dict):
@ -173,6 +212,7 @@ def _redact_model_response_dict_choices(choices, redacted_str: str):
choice["message"]["thinking_blocks"] = None
if "audio" in choice["message"]:
choice["message"]["audio"] = None
_redact_tool_calls_dict(choice["message"], redacted_str)
elif "delta" in choice and isinstance(choice["delta"], dict):
choice["delta"]["content"] = redacted_str
if "reasoning_content" in choice["delta"]:
@ -181,6 +221,7 @@ def _redact_model_response_dict_choices(choices, redacted_str: str):
choice["delta"]["thinking_blocks"] = None
if "audio" in choice["delta"]:
choice["delta"]["audio"] = None
_redact_tool_calls_dict(choice["delta"], redacted_str)
else:
_redact_choice_content(choice)

View file

@ -9,6 +9,8 @@ secrets from strings without depending on the logging-configuration module.
import re
from typing import List
from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH
_REDACTED = "REDACTED"
@ -30,7 +32,7 @@ def _build_secret_patterns() -> "re.Pattern[str]":
# Basic auth headers
r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}",
# OpenAI / Anthropic sk- prefixed keys
r"sk-[A-Za-z0-9\-_]{20,}",
rf"sk-[A-Za-z0-9\-_]{{{MINIMUM_CUSTOM_KEY_LENGTH - len('sk-')},}}",
# Generic api_key / api-key / apikey (handles 'key': 'value' dict repr)
r"(?:api[_-]?key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]{8,}",
# x-api-key / api-key header values (handles 'key': 'value' dict repr)

View file

@ -1879,6 +1879,7 @@ class BaseLLMHTTPHandler:
litellm_params: GenericLiteLLMParams,
api_key: Optional[str],
model: str,
timeout: Optional[Union[float, httpx.Timeout]] = None,
) -> httpx.Response:
max_attempts = max(provider_config.max_retry_on_anthropic_messages_http_error, 1)
litellm_params_dict = dict(litellm_params)
@ -1891,6 +1892,7 @@ class BaseLLMHTTPHandler:
data=signed_json_body or json.dumps(request_body),
stream=stream or False,
logging_obj=logging_obj,
timeout=timeout,
)
response.raise_for_status()
return response
@ -1925,6 +1927,32 @@ class BaseLLMHTTPHandler:
raise RuntimeError("unreachable: anthropic messages HTTP retry loop exited without return")
@staticmethod
def _resolve_anthropic_messages_timeout(
litellm_params: GenericLiteLLMParams,
stream: bool,
custom_llm_provider: str,
) -> Optional[Union[float, httpx.Timeout]]:
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.request_timeout_resolver import (
get_configured_request_timeout,
)
from litellm.utils import supports_httpx_timeout
stream_timeout = litellm_params.get("stream_timeout") if stream else None
model_timeout = stream_timeout if stream_timeout is not None else litellm_params.get("timeout")
request_timeout = litellm_params.get("request_timeout")
global_timeout = get_configured_request_timeout()
if model_timeout is None and request_timeout is None and global_timeout is None:
return None
return CompletionTimeout.resolve(
model_timeout,
{"request_timeout": request_timeout},
custom_llm_provider,
global_timeout=global_timeout,
supports_httpx_timeout=supports_httpx_timeout,
)
async def async_anthropic_messages_handler(
self,
model: str,
@ -2075,6 +2103,11 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params,
api_key=api_key,
model=model,
timeout=self._resolve_anthropic_messages_timeout(
litellm_params=litellm_params,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
)
# used for logging + cost tracking

View file

@ -247,6 +247,8 @@ class OpenAIResponsesHandler(BaseTranslation):
"""
Merge remapped guardrailed tools with original tools that were not sent
to the guardrail (e.g. web_search, web_search_preview), preserving order.
Tools a guardrail appended (``remapped`` longer than ``original_tools``)
have no original slot and are kept so an injected tool is not dropped.
"""
if not original_tools:
return remapped
@ -262,6 +264,8 @@ class OpenAIResponsesHandler(BaseTranslation):
if j < len(remapped):
result.append(remapped[j])
j += 1
# Keep guardrail-appended tools that matched no original slot above.
result.extend(remapped[j:])
return result
def _apply_guardrailed_tools_to_data(

View file

@ -44457,6 +44457,7 @@
"supports_vision": true
},
"bedrock_mantle/xai.grok-4.3": {
"use_openai_responses_path": true,
"input_cost_per_token": 1.25e-06,
"output_cost_per_token": 2.5e-06,
"cache_read_input_token_cost": 2e-07,

View file

@ -186,6 +186,61 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
)
def _blank_to_none(value: str | None) -> str | None:
"""Collapse an absent, empty, or whitespace-only string to ``None``.
OAuth endpoint fields are consumed by truthiness-based merges (``row or discovered``) and by the
corroboration gate. A whitespace-only value is truthy to ``or`` but is not a usable endpoint, so
without this the merge would keep the blank value for redirects while the gate treats it as
unpinned and backfills the other fields, yielding a broken half-discovered config. Normalizing
the pinned fields once, at each build entry point, gives every downstream consumer a single
notion of "blank" so those code paths cannot disagree.
"""
if not isinstance(value, str):
return None
return value.strip() or None
def _normalized_authorize_endpoint(url: str) -> str:
"""Compare authorize endpoints on scheme, host, and path only. The default port is elided and
the host is lowercased so ``https://IDP.example.com:443/authorize/`` and
``https://idp.example.com/authorize`` are the same identity; query and trailing slash are not."""
parsed = urlparse(url)
scheme = parsed.scheme.lower()
host = (parsed.hostname or "").lower()
default_port = {"https": 443, "http": 80}.get(scheme)
try:
port = parsed.port
except ValueError:
port = None
authority = host if port is None or port == default_port else f"{host}:{port}"
return f"{scheme}://{authority}{parsed.path.rstrip('/')}"
def _endpoints_corroborate_authorization_url(
source_authorization_url: str | None,
trusted_authorization_url: str | None,
) -> bool:
"""Whether a source's ``token_url``/``registration_url`` may be paired with a trusted authorize
endpoint. This is the single trust rule for adopting OAuth endpoints from any non-manual source.
Discovery is rooted at the MCP resource (RFC 9728), so a compromised upstream can advertise an
attacker-run authorization server. When ``authorization_url`` is admin-pinned, pairing it with a
``token_url`` from a different source is the RFC 9700 authorization-server mix-up: the user signs
in at the trusted authorize endpoint while the gateway redeems the code, with the stored client
secret and PKCE verifier, at the attacker's token endpoint. Endpoints are trustworthy together
only when they share an authorization server, so a source's endpoints are adopted only when the
same source advertised an ``authorization_endpoint`` matching the pinned value. With no pinned
value (``trusted_authorization_url is None``) there is nothing to protect: the authorize endpoint
comes from the same source as the token endpoint, so they corroborate each other by construction.
"""
if not (trusted_authorization_url and trusted_authorization_url.strip()):
return True
return bool(source_authorization_url) and _normalized_authorize_endpoint(
source_authorization_url
) == _normalized_authorize_endpoint(trusted_authorization_url)
def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_server: MCPServer | None) -> None:
"""Keep the last known good OAuth endpoints when a rebuild's re-discovery comes back empty.
@ -193,26 +248,82 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv
during re-discovery downgrades a working server (``authorization_url`` set) to a broken one
(``None``, /authorize 400s) with no configuration change. Mirrors the ``short_prefix``
carry-forward. Skipped when the server's ``url`` or ``auth_type`` changed, since the previous
endpoints may then belong to a different upstream. ``registration_url`` IS carried here even
though ``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only
restores the same in-memory value the previous build already ran with, while persisting it
would flip ``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for
dcr_bridge servers that never had one configured.
endpoints may then belong to a different upstream. ``registration_url`` IS carried even though
``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only restores
the same in-memory value the previous build already ran with, while persisting it would flip
``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for dcr_bridge
servers that never had one configured.
Carry-forward is a non-manual endpoint source, so the same trust rule as discovery applies: the
previous ``token_url``/``registration_url``/``scopes`` are carried only when the previous
``authorization_url`` corroborates the authorize endpoint this build will use, i.e. when the
incoming build has no pinned authorize endpoint (``None`` -> we adopt the previous one too, a
consistent group) or pins the same one. An admin re-pointing ``authorization_url`` to a different
server must not keep serving the old server's token endpoint or granted scopes.
"""
if previous_server is None:
return
if previous_server.url != new_server.url or previous_server.auth_type != new_server.auth_type:
return
may_carry = _endpoints_corroborate_authorization_url(
previous_server.authorization_url, new_server.authorization_url
)
if new_server.authorization_url is None and previous_server.authorization_url:
new_server.authorization_url = previous_server.authorization_url
if new_server.token_url is None and previous_server.token_url:
if may_carry and new_server.token_url is None and previous_server.token_url:
new_server.token_url = previous_server.token_url
if new_server.registration_url is None and previous_server.registration_url:
if may_carry and new_server.registration_url is None and previous_server.registration_url:
new_server.registration_url = previous_server.registration_url
if not new_server.scopes and previous_server.scopes:
if may_carry and not new_server.scopes and previous_server.scopes:
new_server.scopes = previous_server.scopes
def _restrict_discovery_to_corroborated_authorization_server(
metadata: MCPOAuthMetadata | None,
manual_authorization_url: str | None,
server_identifier: str,
is_dcr_bridge: bool,
) -> MCPOAuthMetadata | None:
"""Reject discovered token/registration endpoints a manually pinned authorize endpoint cannot
vouch for (the RFC 9700 authorization-server mix-up).
Discovery is rooted at the MCP resource, so a compromised upstream can advertise an attacker
``token_endpoint``: with ``authorization_url`` admin-pinned but ``token_url`` blank, the merge
would pair the trusted authorize endpoint with that attacker token endpoint, and the gateway would
post the authorization code and client secret there. So the discovered ``token_url`` and
``registration_url`` are kept only if the document corroborates the pin (its
``authorization_endpoint`` matches). ``scopes`` are deliberately NOT gated here: per the MCP
authorization spec Scope Selection Strategy and RFC 9700 §2.3, the scopes a client requests are
resource-driven (the WWW-Authenticate challenge or the RFC 9728 protected-resource
``scopes_supported``), and scope inflation by a compromised resource is bounded by the
authorization server and user consent (RFC 6749 §3.3), not by the client second-guessing the
request. With no pin there is no trust anchor to protect, so discovery is returned as-is.
"""
if metadata is None or not (manual_authorization_url and manual_authorization_url.strip()):
return metadata
if _endpoints_corroborate_authorization_url(metadata.authorization_url, manual_authorization_url):
return metadata
if not metadata.token_url and not metadata.registration_url:
return metadata
bridge_note = (
" The discovered registration_url is rejected with it, so this dcr_bridge server stays on the"
" short-circuit registration arm."
if is_dcr_bridge and metadata.registration_url
else ""
)
verbose_logger.warning(
"MCP OAuth discovery for server %s advertised authorization_endpoint %s, which does not match the "
"manually configured authorization_url %s; rejecting the discovered token_url/registration_url so "
"authorization codes and client credentials only follow the configured authorization server. "
"Configure Token URL manually if the mismatch is intentional.%s",
server_identifier,
_normalized_authorize_endpoint(metadata.authorization_url) if metadata.authorization_url else "<absent>",
_normalized_authorize_endpoint(manual_authorization_url),
bridge_note,
)
return metadata.model_copy(update={"token_url": None, "registration_url": None})
def invalidate_user_env_vars_cache(user_id: str, server_id: str) -> None:
"""Drop a cached entry after the user stores or clears their env var values
so the next request reads the fresh value instead of a stale one."""
@ -1026,12 +1137,15 @@ class MCPServerManager:
)
auth_type = server_config.get("auth_type", None)
manual_authorization_url = _blank_to_none(server_config.get("authorization_url"))
manual_token_url = _blank_to_none(server_config.get("token_url"))
manual_registration_url = _blank_to_none(server_config.get("registration_url"))
if server_url and (
auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
or self._obo_needs_endpoint_discovery(
auth_type,
server_config.get("token_exchange_endpoint"),
server_config.get("token_url"),
manual_token_url,
)
):
mcp_oauth_metadata = await self._descovery_metadata(
@ -1041,20 +1155,29 @@ class MCPServerManager:
else:
mcp_oauth_metadata = None
gated_oauth_metadata = (
_restrict_discovery_to_corroborated_authorization_server(
mcp_oauth_metadata,
manual_authorization_url,
server_name or server_id,
bool(server_config.get("dcr_bridge")),
)
if auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
else mcp_oauth_metadata
)
# Filter blank scopes (e.g. YAML ``scopes: [""]``) the same way the DB-build path does, so
# an all-blank list normalizes to None rather than a ``("",)`` tuple that skips the
# entra_obo fail-closed scope precondition and POSTs an empty scope to the IdP.
resolved_scopes = self._extract_scopes(server_config.get("scopes")) or (
mcp_oauth_metadata.scopes if mcp_oauth_metadata else None
gated_oauth_metadata.scopes if gated_oauth_metadata else None
)
resolved_authorization_url = server_config.get("authorization_url") or (
mcp_oauth_metadata.authorization_url if mcp_oauth_metadata else None
resolved_authorization_url = manual_authorization_url or (
gated_oauth_metadata.authorization_url if gated_oauth_metadata else None
)
resolved_token_url = server_config.get("token_url") or (
mcp_oauth_metadata.token_url if mcp_oauth_metadata else None
)
resolved_registration_url = server_config.get("registration_url") or (
mcp_oauth_metadata.registration_url if mcp_oauth_metadata else None
resolved_token_url = manual_token_url or (gated_oauth_metadata.token_url if gated_oauth_metadata else None)
resolved_registration_url = manual_registration_url or (
gated_oauth_metadata.registration_url if gated_oauth_metadata else None
)
config_oauth2_flow = server_config.get("oauth2_flow", None)
@ -1447,13 +1570,17 @@ class MCPServerManager:
auth_type = cast(MCPAuthType, mcp_server.auth_type)
server_url = mcp_server.url
manual_authorization_url = _blank_to_none(mcp_server.authorization_url)
manual_token_url = _blank_to_none(mcp_server.token_url)
manual_registration_url = _blank_to_none(mcp_server.registration_url)
has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes)
needs_discovery = bool(server_url) and (
(auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and not mcp_server.authorization_url)
(auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and not has_all_upstream_oauth_fields)
or self._obo_needs_endpoint_discovery(
auth_type,
mcp_server.token_exchange_endpoint
or (credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
mcp_server.token_url,
manual_token_url,
)
)
mcp_oauth_metadata = (
@ -1467,12 +1594,22 @@ class MCPServerManager:
if needs_discovery and mcp_oauth_metadata is None:
verbose_logger.warning(
"MCP OAuth discovery yielded no metadata for server %s (%s); "
"OAuth endpoints stay unresolved until a rebuild succeeds",
"OAuth endpoints/scopes stay unresolved until a rebuild succeeds",
mcp_server.server_id,
server_url,
)
gated_oauth_metadata = (
_restrict_discovery_to_corroborated_authorization_server(
mcp_oauth_metadata,
manual_authorization_url,
mcp_server.server_id,
bool(getattr(mcp_server, "dcr_bridge", None)),
)
if auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
else mcp_oauth_metadata
)
resolved_scopes = scopes or (mcp_oauth_metadata.scopes if mcp_oauth_metadata else None)
resolved_scopes = scopes or (gated_oauth_metadata.scopes if gated_oauth_metadata else None)
new_server = MCPServer(
server_id=mcp_server.server_id,
@ -1492,9 +1629,9 @@ class MCPServerManager:
client_secret=client_secret_value or getattr(mcp_server, "client_secret", None),
oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)),
scopes=resolved_scopes,
authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None),
token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None),
registration_url=mcp_server.registration_url or getattr(mcp_oauth_metadata, "registration_url", None),
authorization_url=manual_authorization_url or getattr(gated_oauth_metadata, "authorization_url", None),
token_url=manual_token_url or getattr(gated_oauth_metadata, "token_url", None),
registration_url=manual_registration_url or getattr(gated_oauth_metadata, "registration_url", None),
token_endpoint_auth_method=(
credentials_dict.get("token_endpoint_auth_method") if credentials_dict else None
),
@ -1545,16 +1682,16 @@ class MCPServerManager:
await self._persist_discovered_obo_token_url(
server_id=mcp_server.server_id,
auth_type=auth_type,
existing_token_url=mcp_server.token_url,
existing_token_url=manual_token_url,
discovered_token_url=new_server.token_url,
)
await self._persist_discovered_oauth_endpoints(
server_id=mcp_server.server_id,
auth_type=auth_type,
existing_authorization_url=mcp_server.authorization_url,
existing_token_url=mcp_server.token_url,
existing_authorization_url=manual_authorization_url,
existing_token_url=manual_token_url,
existing_scopes=scopes,
metadata=mcp_oauth_metadata,
metadata=gated_oauth_metadata,
)
return new_server

View file

@ -7613,6 +7613,18 @@
],
"title": "Messages"
},
"metadata": {
"anyOf": [
{
"additionalProperties": true,
"type": "object"
},
{
"type": "null"
}
],
"title": "Metadata"
},
"text": {
"title": "Text",
"type": "string"

View file

@ -10,7 +10,7 @@ from fastapi import HTTPException, Request, status
import litellm
from litellm import Router, provider_list
from litellm._logging import verbose_proxy_logger
from litellm.constants import STANDARD_CUSTOMER_ID_HEADERS
from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH, STANDARD_CUSTOMER_ID_HEADERS
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.proxy._types import *
@ -1533,4 +1533,6 @@ def get_model_from_request(
def abbreviate_api_key(api_key: str) -> str:
if len(api_key) < MINIMUM_CUSTOM_KEY_LENGTH:
return "sk-..."
return f"sk-...{api_key[-4:]}"

View file

@ -2238,7 +2238,10 @@ async def apply_guardrail(
if litellm_logging_obj is not None:
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
request_data: dict = {"messages": request.messages} if request.messages else {}
request_data: dict = {
**({"messages": request.messages} if request.messages is not None else {}),
**({"metadata": request.metadata} if request.metadata is not None else {}),
}
_input_type = _resolve_guardrail_input_type(active_guardrail, request.input_type)
guardrailed_inputs = await active_guardrail.apply_guardrail(
inputs={"texts": [request.text]},

View file

@ -0,0 +1,75 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from litellm.types.guardrails import (
GuardrailEventHooks,
Mode,
SupportedGuardrailIntegrations,
)
from .compresr import CompresrGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def _coerce_event_hook(
mode: str | list[str] | Mode,
) -> GuardrailEventHooks | list[GuardrailEventHooks] | Mode:
if isinstance(mode, Mode):
return mode
if isinstance(mode, list):
return [GuardrailEventHooks(item) for item in mode]
return GuardrailEventHooks(mode)
def _get_optional_value(litellm_params: LitellmParams, optional_params: object | None, attribute_name: str) -> object:
if optional_params is not None:
value = getattr(optional_params, attribute_name, None)
if value is not None:
return value
return getattr(litellm_params, attribute_name, None)
def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) -> CompresrGuardrail:
import litellm
optional_params = getattr(litellm_params, "optional_params", None)
_callback = CompresrGuardrail(
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
model=litellm_params.model,
target_compression_ratio=_get_optional_value(litellm_params, optional_params, "target_compression_ratio"),
coarse=_get_optional_value(litellm_params, optional_params, "coarse"),
min_chars_to_compress=_get_optional_value(litellm_params, optional_params, "min_chars_to_compress"),
compress_tool_outputs=_get_optional_value(litellm_params, optional_params, "compress_tool_outputs"),
compress_system=_get_optional_value(litellm_params, optional_params, "compress_system"),
compress_history=_get_optional_value(litellm_params, optional_params, "compress_history"),
compress_last_user=_get_optional_value(litellm_params, optional_params, "compress_last_user"),
enable_retrieval=_get_optional_value(litellm_params, optional_params, "enable_retrieval"),
max_bytes_per_call=_get_optional_value(litellm_params, optional_params, "max_bytes_per_call"),
allow_bypass_header=_get_optional_value(litellm_params, optional_params, "allow_bypass_header"),
dynamic=_get_optional_value(litellm_params, optional_params, "dynamic"),
dynamic_min_ratio=_get_optional_value(litellm_params, optional_params, "dynamic_min_ratio"),
dynamic_max_ratio=_get_optional_value(litellm_params, optional_params, "dynamic_max_ratio"),
compression_params=_get_optional_value(litellm_params, optional_params, "compression_params"),
guardrail_name=guardrail["guardrail_name"],
event_hook=_coerce_event_hook(litellm_params.mode),
default_on=litellm_params.default_on or False,
unreachable_fallback=litellm_params.unreachable_fallback,
)
litellm.logging_callback_manager.add_litellm_callback( # pyright: ignore[reportUnknownMemberType] # callback manager is untyped
_callback
)
return _callback
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.COMPRESR.value: initialize_guardrail,
}
guardrail_class_registry = {
SupportedGuardrailIntegrations.COMPRESR.value: CompresrGuardrail,
}

File diff suppressed because it is too large Load diff

View file

@ -32,6 +32,7 @@ from litellm._uuid import uuid
from litellm.constants import (
LENGTH_OF_LITELLM_GENERATED_KEY,
LITELLM_PROXY_ADMIN_NAME,
MINIMUM_CUSTOM_KEY_LENGTH,
UI_SESSION_TOKEN_TEAM_ID,
)
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
@ -1022,6 +1023,14 @@ async def _common_key_generation_helper(
detail={"error": f"Invalid key format. LiteLLM Virtual Key must start with 'sk-'. Received: {_masked}"},
)
if data.key is not None and len(data.key) < MINIMUM_CUSTOM_KEY_LENGTH:
raise HTTPException(
status_code=400,
detail={
"error": f"Invalid key format. LiteLLM Virtual Key must be at least {MINIMUM_CUSTOM_KEY_LENGTH} characters long."
},
)
# check org key limits - done here to handle inheriting org id from team
if data.organization_id is not None:
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
@ -1474,7 +1483,7 @@ async def generate_key_fn(
Parameters:
- duration: Optional[str] - Specify the length of time the token is valid for. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
- key_alias: Optional[str] - User defined key alias
- key: Optional[str] - User defined key value. If not set, a 16-digit unique sk-key is created for you.
- key: Optional[str] - User defined key value. Must start with 'sk-' and be at least 16 characters long. If not set, a 16-digit unique sk-key is created for you.
- team_id: Optional[str] - The team id of the key
- user_id: Optional[str] - The user id of the key
- agent_id: Optional[str] - The agent id associated with the key.
@ -1688,7 +1697,7 @@ async def generate_service_account_key_fn(
Parameters:
- duration: Optional[str] - Specify the length of time the token is valid for. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
- key_alias: Optional[str] - User defined key alias
- key: Optional[str] - User defined key value. If not set, a 16-digit unique sk-key is created for you.
- key: Optional[str] - User defined key value. Must start with 'sk-' and be at least 16 characters long. If not set, a 16-digit unique sk-key is created for you.
- team_id: Optional[str] - The team id of the key
- user_id: Optional[str] - [NON-FUNCTIONAL] THIS WILL BE IGNORED. The user id of the key
- budget_id: Optional[str] - The budget id associated with the key. Created by calling `/budget/new`.
@ -4356,7 +4365,6 @@ async def get_new_token(data: Optional[RegenerateKeyRequest]) -> str:
if data and data.new_key is not None:
# Reject custom key values if disabled by admin
await _check_custom_key_allowed(data.new_key)
new_token = data.new_key
if not data.new_key.startswith("sk-"):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -4364,6 +4372,12 @@ async def get_new_token(data: Optional[RegenerateKeyRequest]) -> str:
"error": "New key must start with 'sk-'. This is to distinguish a key hash (used by litellm for logging / internal logic) from the actual key."
},
)
if len(data.new_key) < MINIMUM_CUSTOM_KEY_LENGTH:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": f"New key must be at least {MINIMUM_CUSTOM_KEY_LENGTH} characters long."},
)
new_token = data.new_key
else:
new_token = f"sk-{secrets.token_urlsafe(LENGTH_OF_LITELLM_GENERATED_KEY)}"
return new_token
@ -4470,7 +4484,7 @@ async def _execute_virtual_key_regeneration(
new_token = await get_new_token(data=data)
new_token_hash = hash_token(new_token)
new_token_key_name = f"sk-...{new_token[-4:]}"
new_token_key_name = abbreviate_api_key(api_key=new_token)
update_data = {"token": new_token_hash, "key_name": new_token_key_name}
non_default_values = {}
@ -4550,7 +4564,7 @@ async def regenerate_key_fn(
- data: Optional[RegenerateKeyRequest] - Request body containing optional parameters to update
- key: Optional[str] - The key to regenerate.
- new_master_key: Optional[str] - The new master key to use, if key is the master key.
- new_key: Optional[str] - The new key to use, if key is not the master key. If both set, new_master_key will be used.
- new_key: Optional[str] - The new key to use, if key is not the master key. Must start with 'sk-' and be at least 16 characters long. If both set, new_master_key will be used.
- key_alias: Optional[str] - User-friendly key alias
- user_id: Optional[str] - User ID associated with key
- team_id: Optional[str] - Team ID associated with key

View file

@ -6,6 +6,7 @@ from fastapi import HTTPException, status
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.router_utils.common_utils import _is_proxy_admin_request
# Router-internal mock_testing_* flag names — kept in sync with
# ``litellm.types.router.MockRouterTestingParams`` by the test
@ -363,6 +364,7 @@ async def route_request(
team_id = get_team_id_from_data(data)
router_model_names = llm_router.model_names if llm_router is not None else []
is_proxy_admin_without_team = team_id is None and _is_proxy_admin_request(data)
# Preprocess Google GenAI generate content requests
if route_type in ["agenerate_content", "agenerate_content_stream"]:
@ -517,6 +519,13 @@ async def route_request(
data["model"] = team_model_name
return getattr(llm_router, f"{route_type}")(**data)
elif (
is_proxy_admin_without_team
and data["model"] not in router_model_names
and data["model"] in llm_router.team_public_model_names
):
return getattr(llm_router, f"{route_type}")(**data)
elif data["model"] in router_model_names or llm_router.has_model_id(data["model"]):
return getattr(llm_router, f"{route_type}")(**data)

View file

@ -833,7 +833,7 @@ class ProxyLogging:
def get_combined_callback_list(self, dynamic_success_callbacks: Optional[List], global_callbacks: List) -> List:
if dynamic_success_callbacks is None:
return list(global_callbacks)
return list(set(dynamic_success_callbacks + global_callbacks))
return list(dict.fromkeys(dynamic_success_callbacks + global_callbacks))
def _parse_pre_mcp_call_hook_response(
self,

View file

@ -26,6 +26,7 @@ from typing import (
AsyncGenerator,
Callable,
Dict,
FrozenSet,
Generator,
List,
Literal,
@ -108,6 +109,7 @@ from litellm.router_utils.clientside_credential_handler import (
is_clientside_credential,
)
from litellm.router_utils.common_utils import (
_is_proxy_admin_request,
filter_team_based_models,
filter_web_search_deployments,
)
@ -494,6 +496,7 @@ class Router:
self.model_name_to_deployment_indices: Dict[str, List[int]] = {}
# Maps (team_id, team_public_model_name) -> list of indices in model_list
self.team_model_to_deployment_indices: Dict[Tuple[str, str], List[int]] = {}
self.team_public_model_names: FrozenSet[str] = frozenset()
# Initialize cache attributes that ``_invalidate_model_group_info_cache``
# touches *before* the first ``set_model_list`` below (which calls
@ -2983,7 +2986,7 @@ class Router:
# here before it's wiped below, instead of relying on that attempt's
# (possibly still-pending) failure event to do it.
refund_stale_reservation_before_retry(self.cache, kwargs)
set_io_token_rate_limit_request_kwargs(kwargs)
set_io_token_rate_limit_request_kwargs(kwargs, store_in_context=deployment_has_io_token_limits(deployment))
## DEPLOYMENT-LEVEL TAGS
deployment_tags = deployment.get("litellm_params", {}).get("tags")
@ -7800,6 +7803,7 @@ class Router:
self.model_id_to_deployment_index_map = {} # Reset the index
self.model_name_to_deployment_indices = {} # Reset the model_name index
self.team_model_to_deployment_indices = {} # Reset the team_model index
self.team_public_model_names = frozenset()
# Reset per-strategy router registries so hot-reload doesn't leave
# stale routers pointing at the old model_list.
self.quality_routers = {}
@ -8151,6 +8155,9 @@ class Router:
self.team_model_to_deployment_indices[key] = updated_indices
else:
del self.team_model_to_deployment_indices[key]
self.team_public_model_names = frozenset(
public_model_name for _, public_model_name in self.team_model_to_deployment_indices
)
def _update_team_model_index(self, model: dict, idx: int) -> None:
"""
@ -8164,6 +8171,7 @@ class Router:
team_public_model_name = (model.get("model_info") or {}).get("team_public_model_name")
if team_id and team_public_model_name:
key = (team_id, team_public_model_name)
self.team_public_model_names = self.team_public_model_names | frozenset({team_public_model_name})
if key not in self.team_model_to_deployment_indices:
self.team_model_to_deployment_indices[key] = []
if idx not in self.team_model_to_deployment_indices[key]:
@ -9118,6 +9126,7 @@ class Router:
"""
self.model_name_to_deployment_indices.clear()
self.team_model_to_deployment_indices.clear()
self.team_public_model_names = frozenset()
for idx, model in enumerate(model_list):
model_name = model.get("model_name")
@ -10026,7 +10035,10 @@ class Router:
return [m for m in self.model_list if m["litellm_params"]["model"] == model]
def _try_early_resolve_deployments_for_model_not_in_names(
self, model: str, request_team_id: Optional[str]
self,
model: str,
request_team_id: Optional[str],
include_team_models: bool = False,
) -> Optional[Tuple[str, Union[List, Dict]]]:
"""
When ``model`` is not in ``self.model_names``, try team routes, pattern routes,
@ -10041,6 +10053,30 @@ class Router:
team_deployments = self._get_all_deployments(model_name=model, team_id=request_team_id)
if team_deployments:
return model, team_deployments
elif include_team_models:
team_deployments = [
self.model_list[index]
for (_, public_model_name), indices in self.team_model_to_deployment_indices.items()
if public_model_name == model
for index in indices
]
team_ids = {
team_id
for deployment in team_deployments
for team_id in [(deployment.get("model_info") or {}).get("team_id")]
if team_id is not None
}
if len(team_ids) > 1:
raise litellm.BadRequestError(
message=(
f"Model name '{model}' matches deployments from multiple teams. "
"Specify the deployment ID directly to disambiguate."
),
model=model,
llm_provider="",
)
if team_deployments:
return model, team_deployments
pattern_deployments = self.pattern_router.get_deployments_by_pattern(
model=model,
@ -10105,7 +10141,11 @@ class Router:
if _model_from_alias is not None:
model = _model_from_alias
early = self._try_early_resolve_deployments_for_model_not_in_names(model=model, request_team_id=request_team_id)
early = self._try_early_resolve_deployments_for_model_not_in_names(
model=model,
request_team_id=request_team_id,
include_team_models=_is_proxy_admin_request(request_kwargs),
)
if early is not None:
return early

View file

@ -98,9 +98,9 @@ def _sanitize_user_api_key_auth(auth: Any) -> Any:
return auth
def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any] | None:
def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any]:
if not metadata:
return metadata
return {}
return {
k: _sanitize_user_api_key_auth(v) if k == "user_api_key_auth" else v
for k, v in metadata.items()
@ -763,8 +763,8 @@ class ComplexityRouter(CustomLogger):
# embedding call. Forwarding it would let the embedding's cost callback finalize the
# reservation, so the routed completion's own callback then skips incrementing the
# key/team budget. Key/team attribution fields are preserved for spend logging.
metadata = _classifier_call_metadata(request_kwargs.get("metadata")) or {}
litellm_metadata = _classifier_call_metadata(request_kwargs.get("litellm_metadata")) or {}
metadata = _classifier_call_metadata(request_kwargs.get("metadata"))
litellm_metadata = _classifier_call_metadata(request_kwargs.get("litellm_metadata"))
query_vector = (
await encoder.aencode_queries([user_message], metadata=metadata, litellm_metadata=litellm_metadata)
)[0]

View file

@ -1,14 +1,27 @@
import hashlib
import json
from collections.abc import Mapping
from typing import TYPE_CHECKING, Dict, List, Optional, Union
if TYPE_CHECKING:
from litellm.types.llms.openai import OpenAIFileObject
from litellm.exceptions import BadRequestError
from litellm.types.router import CredentialLiteLLMParams
from litellm._logging import verbose_logger
def _is_proxy_admin_request(request_kwargs: Optional[Mapping[str, object]]) -> bool:
if request_kwargs is None:
return False
metadata_value = request_kwargs.get("metadata")
litellm_metadata_value = request_kwargs.get("litellm_metadata")
metadata = metadata_value if isinstance(metadata_value, Mapping) else {}
litellm_metadata = litellm_metadata_value if isinstance(litellm_metadata_value, Mapping) else {}
user_api_key_auth = metadata.get("user_api_key_auth") or litellm_metadata.get("user_api_key_auth")
return getattr(user_api_key_auth, "user_role", None) == "proxy_admin"
def get_litellm_params_sensitive_credential_hash(litellm_params: dict) -> str:
"""
Hash of the credential params, used for mapping the file id to the right model
@ -59,6 +72,40 @@ def filter_team_based_models(
metadata = request_kwargs.get("metadata") or {}
litellm_metadata = request_kwargs.get("litellm_metadata") or {}
request_team_id = metadata.get("user_api_key_team_id") or litellm_metadata.get("user_api_key_team_id")
if request_team_id is None and _is_proxy_admin_request(request_kwargs) and isinstance(healthy_deployments, list):
requested_model = (
request_kwargs.get("model") or metadata.get("model_group") or litellm_metadata.get("model_group")
)
candidate_deployments = tuple(
(deployment.get("model_name"), deployment.get("model_info") or {}) for deployment in healthy_deployments
)
team_ids = frozenset(
team_id
for _, model_info in candidate_deployments
for team_id in [model_info.get("team_id")]
if team_id is not None
)
matches_requested_model = (
isinstance(requested_model, str)
and bool(candidate_deployments)
and all(
model_info.get("team_id") is not None
and (model_name == requested_model or model_info.get("team_public_model_name") == requested_model)
for model_name, model_info in candidate_deployments
)
)
if matches_requested_model and len(team_ids) > 1:
raise BadRequestError(
message=(
f"Model name '{requested_model}' matches deployments from multiple teams. "
"Specify the deployment ID directly to disambiguate."
),
model=requested_model,
llm_provider="",
)
if matches_requested_model:
return healthy_deployments
ids_to_remove = set()
if isinstance(healthy_deployments, dict):
return healthy_deployments

View file

@ -43,14 +43,21 @@ ITPM_CACHE_KEY = "_litellm_itpm_cache_key"
OTPM_CACHE_KEY = "_litellm_otpm_cache_key"
def set_io_token_rate_limit_request_kwargs(kwargs: Optional[dict[str, Any]]) -> None:
def set_io_token_rate_limit_request_kwargs(kwargs: Optional[dict[str, Any]], store_in_context: bool = True) -> None:
# The reservation sentinels are server-only, but `metadata` is caller
# controlled on proxy requests. Strip any client-supplied copies here (this
# runs before the router stashes its own reservation) so a forged
# reservation can't drive the post-call reconcile/refund against an
# arbitrary counter and bypass the configured limits.
_clear_reservation_from_kwargs(kwargs)
_io_token_rate_limit_request_kwargs.set(kwargs)
# The context slot pins the entire request kwargs (messages included) for
# the lifetime of the surrounding context, which outlives the request when
# the context is captured by pooled resources (e.g. a redis connection
# created mid-request). Only ITPM/OTPM-limited deployments read it, so for
# every other deployment overwrite the slot with None instead of the
# kwargs; overwriting (rather than skipping) also releases a previous
# request's kwargs when a context is reused.
_io_token_rate_limit_request_kwargs.set(kwargs if store_in_context else None)
def get_io_token_rate_limit_request_kwargs() -> Optional[dict[str, Any]]:

View file

@ -56,6 +56,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import (
from litellm.types.proxy.guardrails.guardrail_hooks.headroom import (
HeadroomGuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.compresr import (
CompresrGuardrailConfigModel,
)
"""
Pydantic object defining how to set guardrails on litellm proxy
@ -123,6 +126,7 @@ class SupportedGuardrailIntegrations(Enum):
VIGIL_GUARD = "vigil_guard"
REPELLOAI = "repelloai"
HEADROOM = "headroom"
COMPRESR = "compresr"
class Role(Enum):
@ -806,7 +810,7 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
default="fail_closed",
description=(
"Behavior when a guardrail endpoint is unreachable due to network errors. "
"Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', and 'headroom'. "
"Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. "
"'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed."
),
)
@ -899,6 +903,7 @@ class LitellmParams(
BedrockGuardrailConfigModel,
LakeraV2GuardrailConfigModel,
HeadroomGuardrailConfigModel,
CompresrGuardrailConfigModel,
RepelloAIGuardrailConfigModel,
LassoGuardrailConfigModel,
PillarGuardrailConfigModel,
@ -1034,6 +1039,7 @@ class ApplyGuardrailRequest(BaseModel):
entities: Optional[List[PiiEntityType]] = None
input_type: str = "request"
messages: Optional[List[Dict[str, Any]]] = None
metadata: Dict[str, Any] | None = None
class ApplyGuardrailResponse(BaseModel):

View file

@ -17,6 +17,12 @@ MCPInfo = Dict[str, Any]
class MCPOAuthMetadata(BaseModel):
scopes: Optional[List[str]] = None
"""Resource-driven scopes for the authorization request: the RFC 9728 protected-resource
``scopes_supported``, or the ``scope`` from the WWW-Authenticate 401 challenge when the resource
supplied one, else the authorization server's ``scopes_supported``. This is the scope value a
client requests per the MCP authorization spec Scope Selection Strategy; scope minimization and
inflation control are the authorization server's and user's job at consent (RFC 6749 §3.3), not
the client's."""
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None

View file

@ -0,0 +1,135 @@
from typing import Any, Dict, Literal
from pydantic import BaseModel, Field
from .base import GuardrailConfigModel
class CompresrGuardrailOptionalParams(BaseModel):
"""Optional tuning knobs for the Compresr guardrail."""
target_compression_ratio: float | None = Field(
default=None,
description=(
"Compression strength. 0-1 is the fraction of tokens to remove "
"(0.5 = remove ~50%, the default); a value >1 is an Nx reduction "
"factor (e.g. 4 = ~4x smaller)."
),
)
coarse: bool | None = Field(
default=None,
description=("Paragraph-level compression (default, faster) instead of token-level (finer-grained)."),
)
min_chars_to_compress: int | None = Field(
default=None,
description=("Skip messages whose text is shorter than this many characters. Defaults to 500."),
)
compress_tool_outputs: bool | None = Field(
default=None,
description=("Compress tool/function result messages (search hits, RAG chunks, API dumps). Defaults to True."),
)
compress_system: bool | None = Field(
default=None,
description="Also compress system messages. Defaults to False.",
)
compress_history: bool | None = Field(
default=None,
description="Also compress prior (non-last) user messages. Defaults to False.",
)
compress_last_user: bool | None = Field(
default=None,
description=(
"Also compress the last user message. The query sent to Compresr "
"is always the original verbatim text. Defaults to False."
),
)
enable_retrieval: bool | None = Field(
default=None,
description=(
"Make compression recoverable: inject a `compresr_retrieve` tool "
"so the model can fetch the original content behind a compression "
"marker via the agentic loop. Defaults to True. Set to False (or "
"run the proxy with --workers 1) for multi-worker deployments: "
"the recovery store is per-process, so pre-call and retrieval hooks "
"on different workers cannot see each other's originals."
),
)
max_bytes_per_call: int | None = Field(
default=None,
description=(
"Cap on aggregate bytes of stored originals per litellm_call_id. "
"When a call exceeds this, oldest entries are evicted so the "
"in-process store cannot grow without bound. Defaults to 10 MiB."
),
)
allow_bypass_header: bool | None = Field(
default=None,
description=(
"Honor the `x-compresr-bypass: true` request header to skip "
"compression for a single call. Off by default because the "
"header is caller-settable; enable only on trusted deployments."
),
)
dynamic: bool | None = Field(
default=None,
description=(
"latte_v2 only. Let the server choose the compression amount per input "
"(Kneedle elbow) instead of using target_compression_ratio. Defaults to True."
),
)
dynamic_min_ratio: float | None = Field(
default=None,
description=(
"latte_v2 only. Floor on the adaptive ratio when `dynamic` is on. "
"Unset lets the server default apply (~1.5)."
),
)
dynamic_max_ratio: float | None = Field(
default=None,
description=(
"latte_v2 only. Ceiling on the adaptive ratio when `dynamic` is on. "
"Unset lets the server default apply (~10.0)."
),
)
compression_params: Dict[str, Any] | None = Field(
default=None,
description=(
"Passthrough of extra parameters forwarded verbatim in the Compresr "
"compress payload (e.g. `heuristic_chunking`, or any newer knob), so "
"a new Compresr feature works without a guardrail update. The named "
"fields above take precedence on collision."
),
)
class CompresrGuardrailConfigModel(GuardrailConfigModel[CompresrGuardrailOptionalParams]):
api_key: str | None = Field(
default=None,
description=("Compresr API key. Falls back to the COMPRESR_API_KEY env var."),
)
api_base: str | None = Field(
default=None,
description=(
"Base URL of the Compresr API. Falls back to the COMPRESR_API_BASE "
"env var, then https://api.compresr.ai. Point at your internal "
"service URL for on-prem deployments."
),
)
model: str | None = Field(
default=None,
description=(
"Compresr compression model (not the LLM). Defaults to 'latte_v2', the query-aware compression model."
),
)
unreachable_fallback: Literal["fail_closed", "fail_open"] = Field(
default="fail_closed",
description=(
"Behavior when the Compresr compression service is unreachable or errors. "
"'fail_closed' raises an error (default). 'fail_open' logs a critical error and "
"forwards the request uncompressed instead of blocking it."
),
)
@staticmethod
def ui_friendly_name() -> str:
return "Compresr (context compression)"

View file

@ -44578,6 +44578,7 @@
"supports_vision": true
},
"bedrock_mantle/xai.grok-4.3": {
"use_openai_responses_path": true,
"input_cost_per_token": 1.25e-06,
"output_cost_per_token": 2.5e-06,
"cache_read_input_token_cost": 2e-07,

View file

@ -3,7 +3,7 @@ name = "litellm"
version = "1.94.0"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.14"
requires-python = ">=3.10, <3.15"
license = "MIT"
license-files = ["LICENSE"]
authors = [
@ -129,7 +129,7 @@ proxy-runtime = [
"opentelemetry-sdk==1.28.0",
"opentelemetry-exporter-otlp==1.28.0",
"opentelemetry-instrumentation-fastapi==0.49b0",
"ddtrace>=2.19.0,<3.0",
"ddtrace>=4.8.2,<5.0",
"sentry-sdk>=2.21.0,<3.0",
"mangum>=0.17.0,<1.0",
"azure-ai-contentsafety>=1.0.0,<2.0",

View file

@ -16,8 +16,8 @@ from pathlib import Path
import pytest
import yaml
REPO_ROOT = Path(__file__).resolve().parents[1]
MANIFEST_PATH = REPO_ROOT / "manifest.yaml"
SUITE_ROOT = Path(__file__).resolve().parents[1]
MANIFEST_PATH = SUITE_ROOT / "manifest.yaml"
# The PRD's "Features in v0" section, in row order.
EXPECTED_FEATURE_IDS = [
@ -90,14 +90,14 @@ def test_manifest_every_feature_has_human_readable_name(manifest):
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
def test_feature_directory_exists(feature_id):
feature_dir = REPO_ROOT / feature_id
feature_dir = SUITE_ROOT / feature_id
assert feature_dir.is_dir(), f"missing feature directory: {feature_dir}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
@pytest.mark.parametrize("provider", EXPECTED_PROVIDERS)
def test_per_provider_test_file_exists(feature_id, provider):
test_file = REPO_ROOT / feature_id / f"test_{provider}.py"
test_file = SUITE_ROOT / feature_id / f"test_{provider}.py"
assert test_file.is_file(), f"missing per-provider test file: {test_file}"
@ -106,7 +106,7 @@ def test_feature_directory_has_init_file(feature_id):
"""Each feature directory needs an __init__.py so pytest collects
the per-provider test files as a package — matches the layout
established by `basic_messaging_non_streaming/`."""
init_file = REPO_ROOT / feature_id / "__init__.py"
init_file = SUITE_ROOT / feature_id / "__init__.py"
assert init_file.is_file(), f"missing __init__.py: {init_file}"
@ -117,7 +117,7 @@ def test_feature_directory_has_init_file(feature_id):
# a broken post-v0 directory still fails CI.
@pytest.mark.parametrize("feature_id", ALL_FEATURE_IDS)
def test_every_manifest_feature_has_directory(feature_id):
feature_dir = REPO_ROOT / feature_id
feature_dir = SUITE_ROOT / feature_id
assert feature_dir.is_dir(), (
f"manifest declares {feature_id!r} but {feature_dir} is missing — "
"feature_id MUST match its on-disk directory (see manifest.yaml header)."
@ -126,7 +126,7 @@ def test_every_manifest_feature_has_directory(feature_id):
@pytest.mark.parametrize("feature_id", ALL_FEATURE_IDS)
def test_every_manifest_feature_has_init_file(feature_id):
init_file = REPO_ROOT / feature_id / "__init__.py"
init_file = SUITE_ROOT / feature_id / "__init__.py"
assert init_file.is_file(), f"missing __init__.py: {init_file}"
@ -137,7 +137,7 @@ def test_every_manifest_feature_has_per_provider_test_file(feature_id, provider)
backed by a per-provider test file. Without this check, a missing
file silently becomes a `not_tested` cell in the published matrix
rather than a CI failure surfacing the layout drift."""
test_file = REPO_ROOT / feature_id / f"test_{provider}.py"
test_file = SUITE_ROOT / feature_id / f"test_{provider}.py"
assert test_file.is_file(), f"missing per-provider test file: {test_file}"
@ -151,7 +151,7 @@ def test_per_provider_test_file_imports_and_parametrizes_three_models(
use plain aliases or per-provider-suffixed aliases (e.g.
`claude-opus-4-7-bedrock-invoke`), so we check for the tier
substrings rather than exact alias names."""
text = (REPO_ROOT / feature_id / f"test_{provider}.py").read_text()
text = (SUITE_ROOT / feature_id / f"test_{provider}.py").read_text()
for tier in ("haiku-4-5", "sonnet-4-6", "opus-4-7"):
assert (
tier in text
@ -171,7 +171,7 @@ def test_azure_test_file_drives_the_proxy(feature_id):
that wraps them — both shapes drive the proxy, and we don't want
this layout pin to block legitimate de-duplication of test bodies.
"""
text = (REPO_ROOT / feature_id / "test_azure.py").read_text()
text = (SUITE_ROOT / feature_id / "test_azure.py").read_text()
assert "run_claude" in text or "run_basic_messaging_cell" in text, (
f"{feature_id}/test_azure.py must drive the claude CLI via run_claude() "
"or a shared helper that wraps it; the not_applicable stub was removed "

View file

@ -19,6 +19,7 @@ from claude_code.cli_driver import (
ClaudeCLIError,
DriverResult,
failure_diagnostic,
is_rate_limit_shaped,
run_claude,
run_claude_models_parallel,
)
@ -793,3 +794,199 @@ def test_failure_diagnostic_uses_last_result_event_status():
diag = failure_diagnostic(result)
assert "api_status=429" in diag
assert "500" not in diag
_RATE_LIMITED_STDOUT = (
json.dumps(
{
"type": "assistant",
"message": {
"content": [
{"type": "text", "text": "API Error: 429 Too Many Requests"}
]
},
}
)
+ "\n"
+ json.dumps({"type": "result", "api_error_status": 429})
+ "\n"
)
_OK_STDOUT = (
json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "pong"}]},
}
)
+ "\n"
)
class _FlakyRunner:
"""Fake runner that rate-limits each model N times before succeeding.
Keeps a per-model call count so tests can assert exactly how many
attempts the retry loop made — the load-bearing detail a canned
single-response runner can't express.
"""
def __init__(self, failures_before_success: dict):
self.failures_before_success = dict(failures_before_success)
self.calls: dict = {}
def __call__(self, cmd, env, capture_output, text, timeout, check, input=None):
model = cmd[cmd.index("--model") + 1]
self.calls[model] = self.calls.get(model, 0) + 1
if self.calls[model] <= self.failures_before_success.get(model, 0):
return _Completed(returncode=1, stdout=_RATE_LIMITED_STDOUT)
return _Completed(returncode=0, stdout=_OK_STDOUT)
@pytest.mark.parametrize(
"outcome,expected",
[
(ClaudeCLIError("claude CLI timed out after 120.0s"), True),
(ClaudeCLIError("claude CLI not found at 'claude'"), False),
(
DriverResult(
text="",
events=[{"type": "result", "api_error_status": 429}],
exit_code=1,
),
True,
),
(DriverResult(text="Too Many Requests", exit_code=1), True),
(DriverResult(text="", stderr="throttled by upstream", exit_code=1), True),
(DriverResult(text="rate limit exceeded", exit_code=0), False),
(DriverResult(text="", stderr="auth failed", exit_code=2), False),
],
)
def test_is_rate_limit_shaped_classification(outcome, expected):
"""The retry trigger must match 429/throttle/timeout markers on
failures only — a passing result mentioning '429' in its reply text
must never be classified as retryable."""
assert is_rate_limit_shaped(outcome) is expected
def test_run_claude_models_parallel_retries_rate_limited_model_until_success():
"""A model that 429s once must be retried after the backoff sleep and
end up green, while an untroubled sibling model runs exactly once."""
runner = _FlakyRunner({"flaky": 1})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["flaky", "steady"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert isinstance(outcomes["flaky"], DriverResult)
assert outcomes["flaky"].exit_code == 0
assert outcomes["flaky"].text == "pong"
assert runner.calls == {"flaky": 2, "steady": 1}
assert sleeps == [0.5]
def test_run_claude_models_parallel_does_not_retry_non_rate_limit_failures():
"""A deterministic failure (bad auth) must fail fast: no sleeps, one
attempt — retrying it would just triple the matrix wall time."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
return _Completed(returncode=2, stdout="", stderr="auth failed")
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["a"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert outcomes["a"].exit_code == 2
assert sleeps == []
def test_run_claude_models_parallel_returns_last_failure_when_retries_exhausted():
"""A persistently rate-limited model exhausts its budget (initial
attempt + N retries, each preceded by one backoff sleep) and still
surfaces the 429 diagnostic instead of masking it."""
runner = _FlakyRunner({"stuck": 99})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["stuck"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.25,
sleep=sleeps.append,
)
assert runner.calls == {"stuck": 3}
assert sleeps == [0.25, 0.25]
assert outcomes["stuck"].exit_code == 1
assert "429" in failure_diagnostic(outcomes["stuck"])
def test_run_claude_models_parallel_retries_timeout_shaped_cli_errors():
"""CLI timeouts are how saturated upstreams usually present (the CLI
retries 429s internally until the harness kills it), so a timeout
must be retried like an explicit 429."""
calls: List[int] = []
def runner(cmd, env, capture_output, text, timeout, check, input=None):
calls.append(1)
if len(calls) == 1:
raise subprocess.TimeoutExpired(cmd="claude", timeout=1)
return _Completed(returncode=0, stdout=_OK_STDOUT)
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["a"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=1,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert isinstance(outcomes["a"], DriverResult)
assert outcomes["a"].text == "pong"
assert len(calls) == 2
assert sleeps == [0.5]
def test_run_claude_models_parallel_zero_retries_disables_backoff():
"""`rate_limit_retries=0` must restore the old single-attempt
behavior exactly: one call, no sleeps, failure returned as-is."""
runner = _FlakyRunner({"stuck": 99})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["stuck"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=0,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert runner.calls == {"stuck": 1}
assert sleeps == []
assert outcomes["stuck"].exit_code == 1

View file

@ -0,0 +1,195 @@
"""Unit tests for the shared `run_passthrough_cell` helper.
These tests inject a fake `run_models` callable and an explicit `env`
mapping (both are first-class parameters, no monkeypatching), so they
exercise the helper's branching -- env-missing guard, base-URL
assembly, extra-env forwarding, per-model pass/fail -- without
spawning the real CLI.
The env-builder tests pin the provider-mode contract itself: the
CLAUDE_CODE_USE_* / CLAUDE_CODE_SKIP_*_AUTH flags and the passthrough
route each mode must target. Those values are the feature -- e.g.
dropping the `/v1` from the vertex base URL produces a request Google
404s on -- so a mutation to any of them must fail here before it burns
a live matrix run.
"""
from __future__ import annotations
from typing import Any, Dict, List, Mapping, Optional
import pytest
from claude_code._passthrough import (
ANTHROPIC_PASSTHROUGH_BASE_PATH,
CLIENT_SIDE_AWS_REGION,
VERTEX_PLACEHOLDER_PROJECT,
VERTEX_PLACEHOLDER_REGION,
bedrock_extra_env,
foundry_extra_env,
run_passthrough_cell,
vertex_extra_env,
)
from claude_code.cli_driver import ClaudeCLIError, DriverResult
PROXY_ENV = {
"LITELLM_PROXY_BASE_URL": "http://localhost:4000",
"LITELLM_PROXY_API_KEY": "sk-test",
}
class _FakeResult:
def __init__(self) -> None:
self.rows: List[Dict[str, Any]] = []
self.single: Optional[Dict[str, Any]] = None
def set(self, payload: Mapping[str, Any]) -> None:
self.single = dict(payload)
def add(self, payload: Mapping[str, Any]) -> None:
self.rows.append(dict(payload))
def _fake_run_models(outcomes_by_model, captured: Dict[str, Any]):
def fake(*, models, prompt, base_url, api_key, extra_env=None, **_kwargs):
captured["models"] = list(models)
captured["prompt"] = prompt
captured["base_url"] = base_url
captured["api_key"] = api_key
captured["extra_env"] = dict(extra_env) if extra_env is not None else None
return {model: outcomes_by_model[model] for model in models}
return fake
def test_env_missing_guard_reports_fail_and_aborts():
fake_result = _FakeResult()
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
env={},
)
assert fake_result.single is not None
assert fake_result.single["status"] == "fail"
assert "LITELLM_PROXY_BASE_URL" in fake_result.single["error"]
def test_anthropic_base_path_appended_to_normalized_proxy_url():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcome = DriverResult(text="pong")
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
passthrough_base_path=ANTHROPIC_PASSTHROUGH_BASE_PATH,
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
env={**PROXY_ENV, "LITELLM_PROXY_BASE_URL": "http://localhost:4000/"},
)
assert captured["base_url"] == "http://localhost:4000/anthropic"
assert captured["extra_env"] is None
assert fake_result.rows == [{"status": "pass"}]
def test_extra_env_builder_receives_normalized_base_and_is_forwarded():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcome = DriverResult(text="pong")
seen_bases: List[str] = []
def build(proxy_base: str) -> Dict[str, str]:
seen_bases.append(proxy_base)
return {"SOME_FLAG": "1"}
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
build_extra_env=build,
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
env={**PROXY_ENV, "LITELLM_PROXY_BASE_URL": "http://localhost:4000/"},
)
assert seen_bases == ["http://localhost:4000"]
assert captured["extra_env"] == {"SOME_FLAG": "1"}
assert captured["base_url"] == "http://localhost:4000"
def test_per_model_failures_reported_individually():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcomes = {
"claude-haiku-4-5": DriverResult(text="pong"),
"claude-sonnet-4-6": ClaudeCLIError("claude CLI timed out after 120s"),
"claude-opus-4-7": DriverResult(text="", exit_code=1),
}
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=list(outcomes.keys()),
prompt="ping",
run_models=_fake_run_models(outcomes, captured),
env=PROXY_ENV,
)
statuses = [row["status"] for row in fake_result.rows]
assert statuses == ["pass", "fail", "fail"]
assert "timed out" in fake_result.rows[1]["error"]
assert "claude CLI failed" in fake_result.rows[2]["error"]
def test_empty_assistant_text_is_a_fail():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcomes = {"claude-haiku-4-5": DriverResult(text=" ")}
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
run_models=_fake_run_models(outcomes, captured),
env=PROXY_ENV,
)
assert fake_result.rows == [
{
"status": "fail",
"error": "[claude-haiku-4-5] claude returned empty assistant text",
}
]
def test_bedrock_extra_env_targets_proxy_bedrock_route():
env = bedrock_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_BEDROCK": "1",
"CLAUDE_CODE_SKIP_BEDROCK_AUTH": "1",
"ANTHROPIC_BEDROCK_BASE_URL": "http://localhost:4000/bedrock",
"AWS_REGION": CLIENT_SIDE_AWS_REGION,
}
def test_vertex_extra_env_keeps_the_api_version_in_the_base_url():
env = vertex_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_VERTEX": "1",
"CLAUDE_CODE_SKIP_VERTEX_AUTH": "1",
"ANTHROPIC_VERTEX_BASE_URL": "http://localhost:4000/vertex_ai/v1",
"ANTHROPIC_VERTEX_PROJECT_ID": VERTEX_PLACEHOLDER_PROJECT,
"CLOUD_ML_REGION": VERTEX_PLACEHOLDER_REGION,
}
def test_foundry_extra_env_targets_proxy_azure_route():
env = foundry_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_FOUNDRY": "1",
"CLAUDE_CODE_SKIP_FOUNDRY_AUTH": "1",
"ANTHROPIC_FOUNDRY_BASE_URL": "http://localhost:4000/azure",
}

View file

@ -0,0 +1,196 @@
"""Shared body for the `passthrough` × <provider> compat cells.
Every other matrix row drives the proxy's `/v1/messages` translation
layer: Claude Code speaks the first-party Anthropic wire and LiteLLM
transforms the request per provider. This row instead exercises
LiteLLM's *native passthrough* routes -- the "LLM gateway"
configuration documented at https://code.claude.com/docs/en/gateway --
where Claude Code speaks each cloud's own wire format and the proxy
forwards it, attaching provider credentials on the way out:
anthropic ANTHROPIC_BASE_URL={proxy}/anthropic. The CLI's
first-party wire, forwarded verbatim to
api.anthropic.com, so the model ids are real
Anthropic ids rather than proxy aliases.
bedrock_invoke CLAUDE_CODE_USE_BEDROCK=1 +
ANTHROPIC_BEDROCK_BASE_URL={proxy}/bedrock. The
CLI POSTs /model/{model}/invoke-with-response-stream;
the proxy recognizes a router alias in the model
segment, rewrites it to the deployment's upstream
model id, and SigV4-signs with its own AWS creds.
vertex_ai CLAUDE_CODE_USE_VERTEX=1 +
ANTHROPIC_VERTEX_BASE_URL={proxy}/vertex_ai/v1.
The CLI POSTs
.../models/{model}:streamRawPredict; the proxy
resolves a router alias in the model segment and
takes project, location, and credentials from the
deployment (which is why the deployment must set
`use_in_pass_through: true` -- see
test_config.yaml).
azure CLAUDE_CODE_USE_FOUNDRY=1 +
ANTHROPIC_FOUNDRY_BASE_URL={proxy}/azure. Foundry
mode sends the model in the JSON body, not the
URL, so the proxy's /azure route cannot resolve a
router alias and falls back to the env-configured
AZURE_API_BASE / AZURE_API_KEY target.
bedrock_converse not applicable -- Claude Code's bedrock mode is
InvokeModel-only; no Converse-wire client exists.
Auth is the same in every mode: the CLI's provider-native signing is
disabled via CLAUDE_CODE_SKIP_<PROVIDER>_AUTH, and the LiteLLM virtual
key travels as `Authorization: Bearer` (ANTHROPIC_AUTH_TOKEN), exactly
like the translation rows. The proxy holds the real provider
credentials.
The per-mode env vars and URL shapes above were captured from a real
`claude` CLI (2.1.210) run against a request-logging sink, not from
docs; if a CLI release changes them, the cells fail with the CLI's own
diagnostic rather than silently testing the wrong wire.
`run_models` and `env` are injection seams for
`_driver_unit_tests/test_passthrough.py`; production callers leave
them unset.
"""
from __future__ import annotations
import os
from typing import Any, Callable, Dict, Mapping, Optional, Sequence
import pytest
from claude_code.cli_driver import (
ClaudeCLIError,
failure_diagnostic,
run_claude_models_parallel,
)
PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL"
PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY"
ANTHROPIC_PASSTHROUGH_BASE_PATH = "/anthropic"
CLIENT_SIDE_AWS_REGION = "us-east-1"
"""Satisfies the CLI's embedded AWS SDK, which refuses to construct a
client without a region. The value never influences routing: the proxy
signs the upstream request with its own credentials and region."""
VERTEX_PLACEHOLDER_PROJECT = "proxy-resolved-project"
VERTEX_PLACEHOLDER_REGION = "us-east5"
"""The CLI refuses to build a Vertex URL without a project id and
region, but the proxy replaces both path segments with the resolved
deployment's `vertex_project` / `vertex_location` before forwarding,
so deliberately-fake values prove the resolution actually happened."""
def bedrock_extra_env(proxy_base_url: str) -> Dict[str, str]:
return {
"CLAUDE_CODE_USE_BEDROCK": "1",
"CLAUDE_CODE_SKIP_BEDROCK_AUTH": "1",
"ANTHROPIC_BEDROCK_BASE_URL": f"{proxy_base_url}/bedrock",
"AWS_REGION": CLIENT_SIDE_AWS_REGION,
}
def vertex_extra_env(proxy_base_url: str) -> Dict[str, str]:
"""Vertex-mode CLI env pointed at the proxy's /vertex_ai route.
The `/v1` suffix on ANTHROPIC_VERTEX_BASE_URL is load-bearing: the
CLI's Vertex SDK ships its API version inside its *default* base
URL (`https://{region}-aiplatform.googleapis.com/v1`), so
overriding the base drops the version from the request path unless
the override carries it. LiteLLM's /vertex_ai route reuses the
incoming path verbatim when it contains `/projects/.../locations/...`,
so a version-less path would reach Google as
`aiplatform.googleapis.com/projects/...` and 404.
"""
return {
"CLAUDE_CODE_USE_VERTEX": "1",
"CLAUDE_CODE_SKIP_VERTEX_AUTH": "1",
"ANTHROPIC_VERTEX_BASE_URL": f"{proxy_base_url}/vertex_ai/v1",
"ANTHROPIC_VERTEX_PROJECT_ID": VERTEX_PLACEHOLDER_PROJECT,
"CLOUD_ML_REGION": VERTEX_PLACEHOLDER_REGION,
}
def foundry_extra_env(proxy_base_url: str) -> Dict[str, str]:
return {
"CLAUDE_CODE_USE_FOUNDRY": "1",
"CLAUDE_CODE_SKIP_FOUNDRY_AUTH": "1",
"ANTHROPIC_FOUNDRY_BASE_URL": f"{proxy_base_url}/azure",
}
def run_passthrough_cell(
*,
compat_result,
models: Sequence[str],
prompt: str,
passthrough_base_path: str = "",
build_extra_env: Optional[Callable[[str], Mapping[str, str]]] = None,
run_models: Callable[..., Mapping[str, Any]] = run_claude_models_parallel,
env: Optional[Mapping[str, str]] = None,
) -> None:
"""Run the shared `passthrough` × <provider> cell body.
`passthrough_base_path` is appended to the proxy base URL and
becomes the CLI's ANTHROPIC_BASE_URL (only the anthropic column
uses it; the cloud columns ignore ANTHROPIC_BASE_URL entirely once
their CLAUDE_CODE_USE_* flag is set). `build_extra_env` receives
the trailing-slash-normalized proxy base URL and returns the
provider-mode env for the CLI subprocess.
"""
environ = env if env is not None else os.environ
base_url = environ.get(PROXY_BASE_URL_ENV)
api_key = environ.get(PROXY_API_KEY_ENV)
if not base_url or not api_key:
compat_result.set(
{
"status": "fail",
"error": (
f"missing required env: set {PROXY_BASE_URL_ENV} and "
f"{PROXY_API_KEY_ENV} to point at a running LiteLLM proxy"
),
}
)
pytest.fail(
f"{PROXY_BASE_URL_ENV} / {PROXY_API_KEY_ENV} not configured",
pytrace=False,
)
proxy_base = base_url.rstrip("/")
extra_env = dict(build_extra_env(proxy_base)) if build_extra_env else None
outcomes = run_models(
models=models,
prompt=prompt,
base_url=proxy_base + passthrough_base_path,
api_key=api_key,
extra_env=extra_env,
)
failures = []
for model in models:
outcome = outcomes[model]
if isinstance(outcome, ClaudeCLIError):
error = f"[{model}] {outcome}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if outcome.exit_code != 0:
error = f"[{model}] claude CLI failed: {failure_diagnostic(outcome)}"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
if not outcome.text.strip():
error = f"[{model}] claude returned empty assistant text"
compat_result.add({"status": "fail", "error": error})
failures.append(error)
continue
compat_result.add({"status": "pass"})
if failures:
pytest.fail("; ".join(failures), pytrace=False)

View file

@ -15,6 +15,7 @@ from __future__ import annotations
import json
import os
import re
import shutil
import subprocess
import sys
@ -40,6 +41,29 @@ DEFAULT_TIMEOUT_SECONDS = float(
os.environ.get("LITELLM_COMPAT_CLI_TIMEOUT_SECONDS") or 120
)
RATE_LIMIT_SHAPED_RE = re.compile(
r"(?:\b429\b|rate[\s_-]?limit|too\s+many\s+requests|throttl(?:ed|ing)|"
r"claude\s+CLI\s+timed\s+out)",
re.IGNORECASE,
)
"""Heuristic shared with the conftest rate-limit summary: 429s and
throttle markers anywhere in the failure text, plus CLI timeouts --
the CLI retries 429s internally until the harness timeout kills it,
so a saturated upstream usually surfaces as a timeout rather than a
clean 429."""
DEFAULT_RATE_LIMIT_RETRIES = int(
os.environ.get("LITELLM_COMPAT_RATE_LIMIT_RETRIES") or 2
)
DEFAULT_RATE_LIMIT_BACKOFF_SECONDS = float(
os.environ.get("LITELLM_COMPAT_RATE_LIMIT_BACKOFF_SECONDS") or 65
)
"""Rate-limit-shaped failures are retried after a backoff long enough
for a per-minute quota window (the dominant 429 source across
Anthropic / Bedrock / Vertex) to reset. Both knobs are env-tunable so
a matrix run can trade wall time for resilience without code edits;
retries=0 disables the behavior entirely."""
# Env vars the `claude` Node CLI legitimately needs to function:
# locating its own binary + node, basic locale/terminal plumbing.
# Deliberately excludes every credential-bearing var that the
@ -265,6 +289,22 @@ def run_claude(
ModelResult = Union[DriverResult, ClaudeCLIError]
def is_rate_limit_shaped(outcome: ModelResult) -> bool:
"""Classify an outcome as a retryable rate-limit-shaped failure.
A `ClaudeCLIError` matches on its message (which is where the
driver's own timeout diagnostic lands); a failing `DriverResult`
matches on its full `failure_diagnostic` so 429s buried in the
CLI's stdout text or `api_error_status` are both caught. Passing
results are never rate-limit-shaped.
"""
if isinstance(outcome, ClaudeCLIError):
return bool(RATE_LIMIT_SHAPED_RE.search(str(outcome)))
if outcome.exit_code == 0:
return False
return bool(RATE_LIMIT_SHAPED_RE.search(failure_diagnostic(outcome)))
def run_claude_models_parallel(
*,
models: Sequence[str],
@ -277,6 +317,9 @@ def run_claude_models_parallel(
cli_path: str = CLAUDE_CLI_DEFAULT,
timeout: float = DEFAULT_TIMEOUT_SECONDS,
runner: Optional[Callable[..., Any]] = None,
rate_limit_retries: Optional[int] = None,
rate_limit_backoff_seconds: Optional[float] = None,
sleep: Callable[[float], None] = time.sleep,
) -> Dict[str, ModelResult]:
"""Invoke `run_claude` for every `models[i]` concurrently and collect outcomes.
@ -290,6 +333,14 @@ def run_claude_models_parallel(
keep the synchronous CLI driver unchanged so unit tests can keep
injecting a fake `runner`.
Rate-limit-shaped failures (see `is_rate_limit_shaped`) are retried
per model up to `rate_limit_retries` times, sleeping
`rate_limit_backoff_seconds` before each retry so per-minute quota
windows can reset; both default to the `LITELLM_COMPAT_RATE_LIMIT_*`
env knobs. Each retry goes back through `run_claude`, so it
re-acquires a token from the provider rate limiter like any other
invocation. `sleep` is an injection seam for unit tests.
Returns a dict keyed by model id. Each value is either the
`DriverResult` produced by `run_claude` or the `ClaudeCLIError`
that aborted that model's run — callers decide how to map either
@ -300,14 +351,20 @@ def run_claude_models_parallel(
if not models:
raise ValueError("models must be a non-empty sequence")
def _one(model: str) -> Tuple[str, ModelResult, float]:
# Per-model wall clock: this is what the matrix run actually pays for.
# We record it whether the run succeeded or raised so the breakdown
# log below covers both code paths and surfaces "which model is the
# long pole?" without requiring per-test instrumentation.
started = time.monotonic()
retries = (
DEFAULT_RATE_LIMIT_RETRIES
if rate_limit_retries is None
else max(0, rate_limit_retries)
)
backoff = (
DEFAULT_RATE_LIMIT_BACKOFF_SECONDS
if rate_limit_backoff_seconds is None
else max(0.0, rate_limit_backoff_seconds)
)
def _run_once(model: str) -> ModelResult:
try:
result = run_claude(
return run_claude(
prompt=prompt,
model=model,
base_url=base_url,
@ -319,14 +376,8 @@ def run_claude_models_parallel(
timeout=timeout,
runner=runner,
)
elapsed = time.monotonic() - started
# Stamp the duration onto the DriverResult so callers (tests,
# diagnostics) can attribute slow cells without re-timing.
result.duration_ms = int(elapsed * 1000)
return model, result, elapsed
except ClaudeCLIError as exc:
elapsed = time.monotonic() - started
return model, exc, elapsed
return exc
except Exception as exc:
# Honor the documented "errors as values" contract for any
# exception type — not just ClaudeCLIError. The rate
@ -334,13 +385,38 @@ def run_claude_models_parallel(
# raise ValueError on edge-case model strings, and a future
# bug elsewhere in the call stack must not abort the entire
# parallel batch and lose the other models' outcomes.
elapsed = time.monotonic() - started
wrapped = ClaudeCLIError(
f"unexpected error running model {model!r}: "
f"{type(exc).__name__}: {exc}"
)
wrapped.__cause__ = exc
return model, wrapped, elapsed
return wrapped
def _one(model: str) -> Tuple[str, ModelResult, float]:
# Per-model wall clock: this is what the matrix run actually pays
# for, retries and backoff sleeps included. We record it whether
# the run succeeded or raised so the breakdown log below covers
# both code paths and surfaces "which model is the long pole?"
# without requiring per-test instrumentation.
started = time.monotonic()
outcome = _run_once(model)
for attempt in range(retries):
if not is_rate_limit_shaped(outcome):
break
print(
f"[retry] {model}: rate-limit-shaped failure; sleeping "
f"{backoff:.0f}s before attempt {attempt + 2}/{retries + 1}",
file=sys.stderr,
flush=True,
)
sleep(backoff)
outcome = _run_once(model)
elapsed = time.monotonic() - started
if isinstance(outcome, DriverResult):
# Stamp the duration onto the DriverResult so callers (tests,
# diagnostics) can attribute slow cells without re-timing.
outcome.duration_ms = int(elapsed * 1000)
return model, outcome, elapsed
outcomes: Dict[str, ModelResult] = {}
durations: Dict[str, float] = {}

View file

@ -33,7 +33,6 @@ from __future__ import annotations
import functools
import json
import os
import re
import sys
from collections import Counter, defaultdict
from dataclasses import dataclass, field
@ -43,6 +42,8 @@ from typing import Any, Dict, FrozenSet, List, Optional, Tuple
import pytest
import yaml
from claude_code.cli_driver import RATE_LIMIT_SHAPED_RE
VALID_STATUSES = {"pass", "fail", "not_applicable", "not_tested"}
RESULTS_ARTIFACT_ENV = "COMPAT_RESULTS_PATH"
DEFAULT_ARTIFACT_PATH = "compat-results.json"
@ -62,11 +63,10 @@ DEFAULT_RATE_LIMIT_SUMMARY_PATH = "compat-rate-limit-summary.json"
# the rate limiter is supposed to back off from. False positives on a
# genuinely slow upstream are tolerable here because the worst case is
# the binary search runs at a slightly lower rate than necessary.
_RATE_LIMIT_RE = re.compile(
r"(?:\b429\b|rate[\s_-]?limit|too\s+many\s+requests|throttl(?:ed|ing)|"
r"claude\s+CLI\s+timed\s+out)",
re.IGNORECASE,
)
#
# The pattern lives in `cli_driver` so the driver's retry-on-rate-limit
# logic and this summary classify failures identically.
_RATE_LIMIT_RE = RATE_LIMIT_SHAPED_RE
@dataclass

View file

@ -27,6 +27,15 @@ VERTEXAI_LOCATION=global
AZURE_FOUNDRY_API_KEY=
AZURE_FOUNDRY_API_BASE=
# Azure cell of the `passthrough` row. Foundry-mode Claude Code sends
# the model in the request body, so the proxy's /azure passthrough
# cannot resolve a router alias and falls back to these env vars.
# AZURE_API_BASE is the Foundry resource's Anthropic surface, i.e.
# https://<resource>.services.ai.azure.com/anthropic ; AZURE_API_KEY
# is the same key as AZURE_FOUNDRY_API_KEY.
AZURE_API_BASE=
AZURE_API_KEY=
# REQUIRED for publishing: PAT for the `agent-shin` user, used to push
# the daily compat-matrix branch to its fork (agent-shin/litellm-docs)
# and open the cross-repo PR against BerriAI/litellm-docs. Scopes:

View file

@ -91,6 +91,23 @@ features:
# Code releases. The HTTP probe hits the bug surface LiteLLM
# has actually shipped fixes for (2.1.117, 2.1.72, 2.1.70 per
# the Claude Code release notes).
- id: passthrough
name: Native API passthrough
# Drives the CLI in each cloud's native mode against LiteLLM's
# passthrough routes instead of the /v1/messages translation
# layer -- the "LLM gateway" setup from
# https://code.claude.com/docs/en/gateway. anthropic uses
# ANTHROPIC_BASE_URL={proxy}/anthropic; bedrock_invoke uses
# CLAUDE_CODE_USE_BEDROCK=1 against {proxy}/bedrock (InvokeModel
# wire, alias resolved from the URL by the router); vertex_ai
# uses CLAUDE_CODE_USE_VERTEX=1 against {proxy}/vertex_ai/v1
# (rawPredict wire, alias + project + location resolved from the
# deployment, which therefore needs `use_in_pass_through: true`);
# azure uses CLAUDE_CODE_USE_FOUNDRY=1 against {proxy}/azure and
# needs AZURE_API_BASE/AZURE_API_KEY on the proxy (see
# passthrough/test_azure.py and the cron env example).
# bedrock_converse is structurally not_applicable: Claude Code
# has no Converse-wire client.
- id: long_context_1m
name: Long context (1M)
# Sends a ~210k-token padded prompt with the

View file

@ -0,0 +1,44 @@
"""passthrough x Anthropic.
Drive the real `claude` CLI in its default first-party mode, but with
ANTHROPIC_BASE_URL aimed at the proxy's `/anthropic` passthrough route
instead of the `/v1/messages` translation endpoint. The proxy forwards
the request verbatim to api.anthropic.com, swapping the virtual-key
bearer for its own ANTHROPIC_API_KEY.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/passthrough/test_anthropic.py
^^^^^^^^^^^ ^^^^^^^^^
feature_id provider
Because nothing is translated, the model ids are the real Anthropic API
ids (which happen to equal the proxy aliases for this column). A red
cell here means the passthrough route broke forwarding itself -- auth
header swap, streaming SSE relay, or beta-header propagation -- since
no per-provider transformation is involved.
"""
from __future__ import annotations
from claude_code._passthrough import (
ANTHROPIC_PASSTHROUGH_BASE_PATH,
run_passthrough_cell,
)
ANTHROPIC_MODELS = [
"claude-haiku-4-5",
"claude-sonnet-4-6",
"claude-opus-4-7",
]
def test_passthrough_anthropic(compat_result):
"""Drive the `claude` CLI through `{proxy}/anthropic` and assert a reply."""
run_passthrough_cell(
compat_result=compat_result,
models=ANTHROPIC_MODELS,
prompt="Reply with the single word 'pong' and nothing else.",
passthrough_base_path=ANTHROPIC_PASSTHROUGH_BASE_PATH,
)

View file

@ -0,0 +1,59 @@
"""passthrough x Azure (Microsoft Foundry).
Drive the real `claude` CLI in foundry mode (CLAUDE_CODE_USE_FOUNDRY=1)
with ANTHROPIC_FOUNDRY_BASE_URL aimed at the proxy's `/azure`
passthrough route. The CLI POSTs `/v1/messages` with the model in the
JSON body -- unlike the bedrock/vertex modes there is no model segment
in the URL, so the proxy's router-alias resolution cannot engage and
the `/azure` route falls back to its env-configured target: the proxy
must set AZURE_API_BASE to the Foundry resource's Anthropic surface
(`https://<resource>.services.ai.azure.com/anthropic`) and
AZURE_API_KEY to the Foundry key (see
cron_vm/litellm-compat-matrix.env.example). The model ids are the
Foundry deployment names, which this matrix provisions to match the
Anthropic ids.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/passthrough/test_azure.py
^^^^^^^^^^^ ^^^^^
feature_id provider
Wiring verified live at authoring time: through `{proxy}/azure` the
Foundry Anthropic surface accepted the `api-key` / `Authorization:
Bearer` headers the fallback sends (a bogus key 401s, the real key
proceeds to deployment lookup), so a red cell here means missing
AZURE_API_BASE/AZURE_API_KEY on the proxy, missing Foundry deployments
for the three tiers, or a genuine forwarding gap -- not an auth-scheme
mismatch.
Known-red at authoring time against a healthy Foundry resource: the
`/azure` fallback assembles only its own auth headers and drops the
rest of the client's headers, including the `anthropic-version` header
the CLI sends, and Foundry's Anthropic surface rejects the request
with 400 "anthropic-version: header is required" (the same request
sent directly to Foundry with that header succeeds). This cell stays
red until that forwarding gap is fixed, which is precisely the class
of bug the row exists to surface.
"""
from __future__ import annotations
from claude_code._passthrough import foundry_extra_env, run_passthrough_cell
AZURE_MODELS = [
"claude-haiku-4-5",
"claude-sonnet-4-6",
"claude-opus-4-7",
]
def test_passthrough_azure(compat_result):
"""Drive the `claude` CLI through `{proxy}/azure` and assert a reply."""
run_passthrough_cell(
compat_result=compat_result,
models=AZURE_MODELS,
prompt="Reply with the single word 'pong' and nothing else.",
build_extra_env=foundry_extra_env,
)

View file

@ -0,0 +1,34 @@
"""passthrough x Bedrock (Converse).
Structurally not applicable. In bedrock mode the `claude` CLI speaks
only the InvokeModel wire (`/model/{id}/invoke-with-response-stream`);
it has no Converse-wire client, so there is no Claude Code traffic a
Converse passthrough could serve. LiteLLM's `/bedrock` route does
accept `/model/{id}/converse-stream`, but exercising it would test a
wire no Claude Code user can produce, which is out of scope for this
matrix.
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/passthrough/test_bedrock_converse.py
^^^^^^^^^^^ ^^^^^^^^^^^^^^^^
feature_id provider
"""
from __future__ import annotations
def test_passthrough_bedrock_converse(compat_result):
"""Report not_applicable: Claude Code has no Converse-wire mode."""
compat_result.set(
{
"status": "not_applicable",
"reason": (
"Claude Code's bedrock mode speaks only the InvokeModel wire "
"(/model/{id}/invoke-with-response-stream); it has no "
"Converse-wire client, so there is no Claude Code surface "
"for Converse passthrough."
),
}
)

View file

@ -0,0 +1,42 @@
"""passthrough x Bedrock (Invoke).
Drive the real `claude` CLI in bedrock mode (CLAUDE_CODE_USE_BEDROCK=1)
with ANTHROPIC_BEDROCK_BASE_URL aimed at the proxy's `/bedrock`
passthrough route. The CLI speaks the native InvokeModel wire --
`POST /model/{model}/invoke-with-response-stream` -- with the proxy
alias in the model segment; the proxy resolves the alias through its
router, rewrites the path to the deployment's upstream model id, and
SigV4-signs the forwarded request with its own AWS credentials
(CLAUDE_CODE_SKIP_BEDROCK_AUTH=1 keeps the CLI from signing).
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/passthrough/test_bedrock_invoke.py
^^^^^^^^^^^ ^^^^^^^^^^^^^^
feature_id provider
The CLI also fires a best-effort `GET /bedrock/inference-profiles`
listing at startup; its failure is non-fatal and does not gate this
cell.
"""
from __future__ import annotations
from claude_code._passthrough import bedrock_extra_env, run_passthrough_cell
BEDROCK_INVOKE_MODELS = [
"claude-haiku-4-5-bedrock-invoke",
"claude-sonnet-4-6-bedrock-invoke",
"claude-opus-4-7-bedrock-invoke",
]
def test_passthrough_bedrock_invoke(compat_result):
"""Drive the `claude` CLI through `{proxy}/bedrock` and assert a reply."""
run_passthrough_cell(
compat_result=compat_result,
models=BEDROCK_INVOKE_MODELS,
prompt="Reply with the single word 'pong' and nothing else.",
build_extra_env=bedrock_extra_env,
)

View file

@ -0,0 +1,45 @@
"""passthrough x Vertex AI.
Drive the real `claude` CLI in vertex mode (CLAUDE_CODE_USE_VERTEX=1)
with ANTHROPIC_VERTEX_BASE_URL aimed at the proxy's `/vertex_ai`
passthrough route. The CLI speaks the native rawPredict wire --
`POST .../projects/{p}/locations/{l}/publishers/anthropic/models/{model}:streamRawPredict`
-- with the proxy alias in the model segment; the proxy resolves the
alias through its router, replaces the placeholder project/location
path segments with the deployment's `vertex_project` /
`vertex_location`, and attaches its own Google credentials
(CLAUDE_CODE_SKIP_VERTEX_AUTH=1 keeps the CLI from minting a token).
The (feature, provider) for this cell is inferred from the file path by
`tests/e2e/claude_code/conftest.py`:
tests/e2e/claude_code/passthrough/test_vertex_ai.py
^^^^^^^^^^^ ^^^^^^^^^
feature_id provider
This cell requires the vertex deployments in the proxy config to carry
`use_in_pass_through: true` (see test_config.yaml) -- that is what
registers their credentials with the passthrough router. Without it
the proxy forwards the CLI's own headers (the virtual-key bearer) to
Google and every tier fails with a 401.
"""
from __future__ import annotations
from claude_code._passthrough import run_passthrough_cell, vertex_extra_env
VERTEX_MODELS = [
"claude-haiku-4-5-vertex",
"claude-sonnet-4-6-vertex",
"claude-opus-4-7-vertex",
]
def test_passthrough_vertex_ai(compat_result):
"""Drive the `claude` CLI through `{proxy}/vertex_ai` and assert a reply."""
run_passthrough_cell(
compat_result=compat_result,
models=VERTEX_MODELS,
prompt="Reply with the single word 'pong' and nothing else.",
build_extra_env=vertex_extra_env,
)

View file

@ -59,21 +59,31 @@ model_list:
aws_region_name: us-east-1
# ---- Vertex AI ----
# `use_in_pass_through: true` registers each deployment's
# project/location/credentials with the /vertex_ai passthrough
# router, which the `passthrough` row needs to resolve
# .../models/{alias}:streamRawPredict URLs. That registration only
# reads the canonical `vertex_project`/`vertex_location` param names
# (not the `vertex_ai_*` aliases); the chat translation path accepts
# both.
- model_name: claude-haiku-4-5-vertex
litellm_params:
model: vertex_ai/claude-haiku-4-5
vertex_ai_project: os.environ/VERTEXAI_PROJECT
vertex_ai_location: os.environ/VERTEXAI_LOCATION
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: os.environ/VERTEXAI_LOCATION
use_in_pass_through: true
- model_name: claude-sonnet-4-6-vertex
litellm_params:
model: vertex_ai/claude-sonnet-4-6
vertex_ai_project: os.environ/VERTEXAI_PROJECT
vertex_ai_location: os.environ/VERTEXAI_LOCATION
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: os.environ/VERTEXAI_LOCATION
use_in_pass_through: true
- model_name: claude-opus-4-7-vertex
litellm_params:
model: vertex_ai/claude-opus-4-7
vertex_ai_project: os.environ/VERTEXAI_PROJECT
vertex_ai_location: os.environ/VERTEXAI_LOCATION
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: os.environ/VERTEXAI_LOCATION
use_in_pass_through: true
# ---- Microsoft Foundry (Anthropic deployments on Azure) ----
- model_name: claude-haiku-4-5-azure

View file

@ -17,6 +17,7 @@
- {id: reliability.routing.cost_based.picks_lowest_cost, module: reliability, tier: P1, behavior: routing, variant: cost_based, assertions: [picks_lowest_cost], exercised_on: [chat_completions, messages], source: "router_strategy/lowest_cost.py", rationale: "Spend-aware routing"}
- {id: reliability.routing.usage_based.picks_under_tpm, module: reliability, tier: P0, behavior: routing, variant: usage_based, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "router_strategy/lowest_tpm_rpm_v2.py", rationale: "Routes to lowest-TPM deployment; prevents over-allocation"}
- {id: reliability.routing.least_busy.picks_lowest_traffic, module: reliability, tier: P1, behavior: routing, variant: least_busy, assertions: [picks_lowest_traffic], exercised_on: [chat_completions, messages], source: "router_strategy/least_busy.py", rationale: "Fewest in-flight requests"}
- {id: reliability.routing.complexity_llm_classifier.routes_by_llm_tier, module: reliability, tier: P1, behavior: routing, variant: complexity_llm_classifier, assertions: [routes_by_llm_tier], exercised_on: [chat_completions], source: "router_strategy/complexity_router/complexity_router.py", fail_before_fix: proven, rationale: "v2 auto-router LLM complexity classifier runs over the proxy and routes by semantic tier instead of silently crashing on absent litellm_metadata and falling back to heuristic scoring"}
- {id: reliability.cache.exact.returns_cached, module: reliability, tier: P1, behavior: cache, variant: exact, assertions: [returns_cached], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/caching.py", rationale: "Response cache returns cached on exact match"}
- {id: reliability.cache.prompt_caching_model_select.returns_cached, module: reliability, tier: P1, behavior: cache, variant: prompt_caching_model_select, assertions: [returns_cached], exercised_on: [chat_completions], source: "router_utils/prompt_caching_cache.py", rationale: "Selects model supporting prompt caching for cacheable prefix"}
- {id: reliability.circuit_breaker.redis.trips_then_recovers, module: reliability, tier: P0, behavior: circuit_breaker, variant: redis, assertions: [trips_then_recovers], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/redis_cache.py:99", rationale: "Redis breaker CLOSED->OPEN->HALF_OPEN; guards all cache/rate-limit ops"}

View file

@ -66,6 +66,23 @@ configs:
model: openai/text-embedding-3-small
api_key: os.environ/OPENAI_API_KEY
# v2 auto-router with the LLM complexity classifier. SIMPLE stays on the
# openai backend; every higher tier routes to the anthropic backend, so the
# served deployment (read back from the spend log's model) reveals whether
# the LLM classifier actually ran or silently fell back to heuristic scoring.
- model_name: complexity-smart-router
litellm_params:
model: auto_router/complexity_router
complexity_router_config:
classifier_type: llm
classifier_llm_config:
model: gpt-5.5
tiers:
SIMPLE: gpt-5.5
MEDIUM: claude-haiku-4-5
COMPLEX: claude-haiku-4-5
REASONING: claude-haiku-4-5
services:
litellm:
image: ghcr.io/berriai/litellm:main-latest

View file

@ -0,0 +1,20 @@
"""Client for the complexity auto-router e2e tests.
The suite drives the shared /chat/completions and spend-log reads on the Gateway,
so this client only carries the Gateway the shared lifecycle needs for cleanup.
"""
from __future__ import annotations
from dataclasses import dataclass
from e2e_gateway import Gateway, build_gateway
@dataclass(frozen=True, slots=True)
class ComplexityRouterClient:
gateway: Gateway
def build_client() -> ComplexityRouterClient:
return ComplexityRouterClient(gateway=build_gateway())

View file

@ -0,0 +1,15 @@
"""Router suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
live in the parent tests/e2e/conftest.py. ComplexityRouterClient holds the shared
Gateway, so the `resources` fixture cleans up keys this suite creates.
"""
import pytest
from complexity_router_client import ComplexityRouterClient, build_client
@pytest.fixture(scope="session")
def client() -> ComplexityRouterClient:
return build_client()

View file

@ -0,0 +1,62 @@
"""Live e2e: the v2 auto-router's LLM complexity classifier actually runs over the
proxy and drives routing, instead of silently crashing and falling back to the
local heuristic scorer.
The regression this guards (complexity_router.py `_classifier_call_metadata`
returning None when the request carries no `litellm_metadata`, which the classifier
sub-call then fed into a `.update`, raising `'NoneType' object has no attribute
'update'`) was invisible from the outside: the router caught the error and answered
from heuristic scoring, so every request still returned 200. The only tell is which
tier, and therefore which backend, served the request.
`complexity-smart-router` (see the inline config in docker-compose.yml) pins SIMPLE
to the openai backend and every higher tier to the anthropic backend. "Is P equal
to NP?" is lexically trivial, so the heuristic scorer lands it in SIMPLE (openai),
but any competent LLM classifier reads it as a hard reasoning question and lands it
above SIMPLE (anthropic). The served deployment is read back from the spend log's
`model`, so anthropic proves the classifier ran and openai proves it silently fell
back - the exact failure before the fix.
"""
import pytest
from complexity_router_client import ComplexityRouterClient
from e2e_http import unwrap
from models import ChatBody, ChatMessage
pytestmark = pytest.mark.e2e
ROUTER_MODEL = "complexity-smart-router"
# Lexically simple (heuristic -> SIMPLE) but a hard reasoning question (LLM -> above SIMPLE).
LEXICALLY_SIMPLE_HARD_PROMPT = "Is P equal to NP?"
# SIMPLE tier backend; served only when the classifier silently falls back to heuristic.
HEURISTIC_TIER_MODEL = "openai/gpt-5.5"
# MEDIUM/COMPLEX/REASONING tier backend; served only when the LLM classifier runs.
LLM_TIER_MODEL = "anthropic/claude-haiku-4-5"
class TestComplexityRouterLlmClassifier:
@pytest.mark.covers("reliability.routing.complexity_llm_classifier.routes_by_llm_tier")
def test_llm_classifier_runs_and_routes_by_semantic_tier(
self, client: ComplexityRouterClient, scoped_key: str
) -> None:
chat = unwrap(
client.gateway.chat(
scoped_key,
ChatBody(
model=ROUTER_MODEL,
messages=[ChatMessage(role="user", content=LEXICALLY_SIMPLE_HARD_PROMPT)],
max_tokens=16,
),
)
)
assert chat.choices, f"router returned no choices: {chat}"
rows = client.gateway.poll_logs_for_key(scoped_key, min_rows=1)
served = [row.model for row in rows]
assert served == [LLM_TIER_MODEL], (
f"expected the request to be served by {LLM_TIER_MODEL!r} (the higher-tier "
f"backend the LLM classifier picks for a hard prompt), but the spend log shows "
f"{served!r}. {HEURISTIC_TIER_MODEL!r} means the LLM classifier silently failed "
f"and the router fell back to heuristic scoring (SIMPLE) - the pre-fix regression"
)

View file

@ -28,8 +28,11 @@ from litellm.caching.caching import DualCache
@pytest.mark.asyncio
async def test_llm_guard_valid_response():
"""
Tests to see llm guard raises an error for a flagged response
A valid (is_valid=True) LLM Guard response must apply the returned
sanitized_prompt back onto the request data so the provider receives the
redacted content.
"""
litellm.llm_guard_mode = "all"
input_a_anonymizer_results = {
"sanitized_prompt": "hello world",
"is_valid": True,
@ -44,21 +47,65 @@ async def test_llm_guard_valid_response():
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key)
local_cache = DualCache()
try:
await llm_guard.async_moderation_hook(
data={
"messages": [
{
"role": "user",
"content": "hello world, my name is Jane Doe. My number is: 23r323r23r2wwkl",
}
]
},
user_api_key_dict=user_api_key_dict,
call_type="completion",
)
except Exception as e:
pytest.fail(f"An exception occurred - {str(e)}")
data = {
"messages": [
{
"role": "user",
"content": "hello world, my name is Jane Doe. My number is: 23r323r23r2wwkl",
}
]
}
result = await llm_guard.async_moderation_hook(
data=data,
user_api_key_dict=user_api_key_dict,
call_type="completion",
)
assert result is data
assert data["messages"][0]["content"] == "hello world"
@pytest.mark.asyncio
async def test_llm_guard_sanitizes_multimodal_and_input():
"""
Sanitization must reach text parts of multimodal message content and the
``input`` field (embeddings/moderation) while leaving non-text parts intact.
"""
litellm.llm_guard_mode = "all"
llm_guard = _ENTERPRISE_LLMGuard(
mock_testing=True,
mock_redacted_text={
"sanitized_prompt": "email: [REDACTED]",
"is_valid": True,
"scanners": {"Regex": 0.0},
},
)
user_api_key_dict = UserAPIKeyAuth(api_key=hash_token("sk-12345"))
image_part = {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}}
data = {
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": "email: person@example.com"},
image_part,
],
}
]
}
result = await llm_guard.async_moderation_hook(
data=data, user_api_key_dict=user_api_key_dict, call_type="completion"
)
assert result["messages"][0]["content"][0]["text"] == "email: [REDACTED]"
assert result["messages"][0]["content"][1] == image_part
input_data = {"input": ["email: person@example.com", "another prompt"]}
input_result = await llm_guard.async_moderation_hook(
data=input_data, user_api_key_dict=user_api_key_dict, call_type="embeddings"
)
assert input_result["input"] == ["email: [REDACTED]", "email: [REDACTED]"]
@pytest.mark.asyncio

View file

@ -556,3 +556,31 @@ async def test_embedding_cache_falls_back_to_token_counter_for_legacy_entries():
assert cache_hit
# token_counter over "hello world" yields a nonzero count — fallback path still runs
assert response.usage.prompt_tokens > 0
def test_request_kwargs_does_not_retain_logging_obj():
"""
The caching handler lives on logging_obj._llm_caching_handler, so keeping
litellm_logging_obj inside request_kwargs closes a reference cycle
(Logging -> LLMCachingHandler -> kwargs -> Logging). That cycle keeps the
full request payload alive until a generational GC pass instead of being
freed by refcount when the request finishes; under bursts of large-token
requests this presents as stepwise RSS growth that never returns to
baseline. Other kwargs (messages included) must be preserved.
"""
logging_obj = MagicMock()
kwargs = {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hello"}],
"litellm_logging_obj": logging_obj,
}
handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs=kwargs,
start_time=datetime.now(),
)
assert "litellm_logging_obj" not in handler.request_kwargs
assert handler.request_kwargs["messages"] == kwargs["messages"]
assert handler.request_kwargs["model"] == "gpt-4o"

View file

@ -32,6 +32,13 @@ def logging_obj():
)
def test_get_combined_callback_list_preserves_insertion_order(logging_obj):
assert logging_obj.get_combined_callback_list(
dynamic_success_callbacks=["prometheus", "langfuse", "datadog", "otel", "s3"],
global_callbacks=["langfuse", "gcs_bucket", "arize", "logfire"],
) == ["prometheus", "langfuse", "datadog", "otel", "s3", "gcs_bucket", "arize", "logfire"]
def test_get_masked_api_base(logging_obj):
api_base = "https://api.openai.com/v1"
masked_api_base = logging_obj._get_masked_api_base(api_base)
@ -3773,3 +3780,19 @@ def test_zero_token_video_usage_preserves_duration_seconds(logging_obj):
assert payload["metadata"]["usage_object"]["duration_seconds"] == 4.0
assert payload["total_tokens"] == 0
assert payload["completion_tokens"] == 0
def test_pre_call_does_not_pin_request_in_module_state(logging_obj):
"""
pre_call/post_call must not stash their locals (full messages, the Logging
object, complete_input_dict) into module-level state. That pinned the most
recent request's entire payload in memory for the life of the worker,
which with multi-hundred-KB requests is a permanent per-worker leak.
"""
litellm.error_logs.clear()
big_input = [{"role": "user", "content": "x" * 10_000}]
logging_obj.pre_call(input=big_input, api_key="sk-test")
logging_obj.post_call(original_response='{"ok": true}', input=big_input, api_key="sk-test")
assert litellm.error_logs == {}

View file

@ -320,6 +320,176 @@ class TestPerformRedaction:
assert choice.message.content == "redacted-by-litellm"
assert choice.message.reasoning_content == "redacted-by-litellm"
def test_redacts_tool_call_arguments_in_model_response_dict(self):
"""Assistant tool call arguments must not leak when redaction is on."""
result = {
"choices": [
{
"message": {
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"city": "sensitive-city"}',
},
}
],
"function_call": {
"name": "get_weather",
"arguments": '{"city": "sensitive-city"}',
},
}
}
]
}
redacted = perform_redaction({}, result)
message = redacted["choices"][0]["message"]
assert message["content"] == "redacted-by-litellm"
tool_call = message["tool_calls"][0]
assert tool_call["function"]["arguments"] == "redacted-by-litellm"
assert tool_call["function"]["name"] == "get_weather"
assert message["function_call"]["arguments"] == "redacted-by-litellm"
def test_redacts_tool_call_arguments_in_streaming_delta_dict(self):
result = {
"choices": [
{
"delta": {
"content": None,
"tool_calls": [
{
"index": 0,
"function": {
"name": "get_weather",
"arguments": '{"city": "sensitive-city"}',
},
}
],
}
}
]
}
redacted = perform_redaction({}, result)
delta = redacted["choices"][0]["delta"]
assert delta["tool_calls"][0]["function"]["arguments"] == "redacted-by-litellm"
def test_redacts_tool_call_arguments_on_model_response_object(self):
result = litellm.ModelResponse(
id="resp-1",
choices=[
litellm.Choices(
message=litellm.Message(
content=None,
role="assistant",
tool_calls=[
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"city": "sensitive-city"}',
},
}
],
)
)
],
model="gpt-4o",
)
redacted = perform_redaction({}, result)
tool_call = redacted.choices[0].message.tool_calls[0]
assert tool_call.function.arguments == "redacted-by-litellm"
assert tool_call.function.name == "get_weather"
assert result.choices[0].message.tool_calls[0].function.arguments == (
'{"city": "sensitive-city"}'
)
def test_redacts_tool_call_arguments_on_streaming_response_object(self):
"""Reproduces the Stream=True path where tool calls arrive as deltas."""
streaming_choice = litellm.utils.StreamingChoices(
delta=litellm.utils.Delta(
content=None,
role="assistant",
tool_calls=[
{
"index": 0,
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"city": "sensitive-city"}',
},
}
],
)
)
streaming_response = SimpleNamespace(choices=[streaming_choice])
details = {
"stream": True,
"complete_streaming_response": streaming_response,
}
perform_redaction(details, None)
tool_call = streaming_response.choices[0].delta.tool_calls[0]
assert tool_call.function.arguments == "redacted-by-litellm"
def test_redacts_tool_call_arguments_in_standard_logging_object(self):
details = {
"standard_logging_object": {
"response": {
"choices": [
{
"message": {
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"city": "sensitive-city"}',
},
}
],
}
}
]
}
}
}
perform_redaction(details, None)
message = details["standard_logging_object"]["response"]["choices"][0]["message"]
assert message["tool_calls"][0]["function"]["arguments"] == "redacted-by-litellm"
def test_redacts_responses_api_function_call_arguments_dict(self):
result = {
"output": [
{
"type": "function_call",
"name": "get_weather",
"arguments": '{"city": "sensitive-city"}',
"call_id": "call_1",
}
]
}
redacted = perform_redaction({}, result)
assert redacted["output"][0]["arguments"] == "redacted-by-litellm"
assert redacted["output"][0]["name"] == "get_weather"
def test_redacts_response_output_objects_with_top_level_text(self):
output_items = [
SimpleNamespace(text="top-level output"),

View file

@ -7,7 +7,7 @@ with guardrail transformations, specifically testing edge cases with empty choic
import os
import sys
from typing import Any, List, Literal, Optional
from typing import Any, Literal, Optional
from unittest.mock import MagicMock, patch
import pytest
@ -295,6 +295,81 @@ class TestAnthropicMessagesHandlerInputProcessing:
assert "input_schema" in tools[1]
class ToolAppendingGuardrail(CustomGuardrail):
"""Guardrail that appends a new OpenAI-format function tool, mimicking a
guardrail that injects a retrieval/recovery tool the model can later call."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
tools = list(inputs.get("tools") or [])
tools.append(
{
"type": "function",
"function": {
"name": "injected_tool",
"description": "injected by guardrail",
"parameters": {"type": "object", "properties": {}},
},
}
)
inputs["tools"] = tools
return inputs
class TestAnthropicMessagesHandlerToolInjection:
"""A tool a guardrail injects in OpenAI format must survive the write-back
to Anthropic format alongside the request's original tools."""
@pytest.mark.asyncio
async def test_injected_tool_survives_when_request_already_has_tools(self):
handler = AnthropicMessagesHandler()
guardrail = ToolAppendingGuardrail(guardrail_name="test")
data = {
"model": "claude-opus-4-6",
"messages": [{"role": "user", "content": "hi"}],
"tools": [
{
"name": "get_weather",
"description": "Get the weather at a specific location",
"input_schema": {
"type": "object",
"properties": {"location": {"type": "string"}},
},
}
],
}
result = await handler.process_input_messages(
data=data, guardrail_to_apply=guardrail, litellm_logging_obj=MagicMock()
)
names = [t.get("name") for t in result["tools"]]
assert "get_weather" in names
assert "injected_tool" in names
@pytest.mark.asyncio
async def test_injected_tool_survives_when_request_has_no_tools(self):
handler = AnthropicMessagesHandler()
guardrail = ToolAppendingGuardrail(guardrail_name="test")
data = {
"model": "claude-opus-4-6",
"messages": [{"role": "user", "content": "hi"}],
}
result = await handler.process_input_messages(
data=data, guardrail_to_apply=guardrail, litellm_logging_obj=MagicMock()
)
assert [t.get("name") for t in result["tools"]] == ["injected_tool"]
if __name__ == "__main__":
# Run the tests
pytest.main([__file__, "-v"])

View file

@ -474,7 +474,12 @@ class TestBedrockMantleResponsesRegistry:
model="xai.grok-4.3",
)
assert isinstance(cfg, BedrockMantleResponsesAPIConfig)
assert cfg.use_openai_path is False
# grok-4.3 is a third-party frontier model on Bedrock Mantle, served on the
# /openai/v1 base (like gpt-5.x / gemma-4), not the standard /v1 path used by
# open-weights models such as gpt-oss. The standard /v1 base returns
# "Berm is not enabled for this account", so the price-map entry carries
# use_openai_responses_path=true.
assert cfg.use_openai_path is True
def test_unmapped_frontier_model_falls_through_to_none(self, restore_model_cost):
# The gate is data-driven, not name-based: an unseen model not yet in the

View file

@ -1084,6 +1084,147 @@ def test_sync_delete_responses_sets_json_content_type():
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"litellm_params_kwargs, stream, global_timeout, expected",
[
({"timeout": 12.0}, False, None, 12.0),
({"request_timeout": 30.0}, False, None, 30.0),
({}, False, 1500.0, 1500.0),
({"timeout": 5.0, "stream_timeout": 50.0}, True, None, 50.0),
({"timeout": 5.0, "stream_timeout": 50.0}, False, None, 5.0),
({"timeout": 5.0, "request_timeout": 30.0}, False, None, 5.0),
({}, False, None, None),
({}, True, None, None),
],
)
def test_resolve_anthropic_messages_timeout(
monkeypatch, litellm_params_kwargs, stream, global_timeout, expected
):
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
if global_timeout is None:
monkeypatch.setattr(
"litellm.request_timeout",
float(DEFAULT_REQUEST_TIMEOUT_SECONDS),
raising=False,
)
monkeypatch.setattr(
"litellm.request_timeout_explicitly_set",
False,
raising=False,
)
else:
monkeypatch.setattr("litellm.request_timeout", global_timeout, raising=False)
monkeypatch.setattr(
"litellm.request_timeout_explicitly_set", True, raising=False
)
resolved = BaseLLMHTTPHandler._resolve_anthropic_messages_timeout(
litellm_params=GenericLiteLLMParams(**litellm_params_kwargs),
stream=stream,
custom_llm_provider="anthropic",
)
assert resolved == expected
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_forwards_request_timeout(monkeypatch):
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "request_timeout", float(DEFAULT_REQUEST_TIMEOUT_SECONDS))
monkeypatch.setattr(litellm, "request_timeout_explicitly_set", False)
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
)
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude", "messages": []}
)
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
mock_config.max_retry_on_anthropic_messages_http_error = 1
expected_response = {"id": "msg_1", "content": []}
mock_config.transform_anthropic_messages_response = Mock(return_value=expected_response)
ok_response = Mock()
ok_response.raise_for_status = Mock(return_value=None)
mock_client = AsyncMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=ok_response)
logging_obj = Mock()
logging_obj.model_call_details = {}
logging_obj.dynamic_success_callbacks = []
result = await handler.async_anthropic_messages_handler(
model="claude",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_provider_config=mock_config,
anthropic_messages_optional_request_params={},
custom_llm_provider="anthropic",
litellm_params=GenericLiteLLMParams(request_timeout=0.3),
logging_obj=logging_obj,
client=mock_client,
kwargs={},
)
assert result is expected_response
assert mock_client.post.await_args.kwargs["timeout"] == 0.3
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_forwards_stream_timeout(monkeypatch):
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "request_timeout", float(DEFAULT_REQUEST_TIMEOUT_SECONDS))
monkeypatch.setattr(litellm, "request_timeout_explicitly_set", False)
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
)
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude", "messages": []}
)
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
mock_config.max_retry_on_anthropic_messages_http_error = 1
mock_config.get_async_streaming_response_iterator = Mock(return_value=Mock())
ok_response = Mock()
ok_response.raise_for_status = Mock(return_value=None)
ok_response.headers = httpx.Headers({})
mock_client = AsyncMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=ok_response)
logging_obj = Mock()
logging_obj.model_call_details = {}
logging_obj.dynamic_success_callbacks = []
await handler.async_anthropic_messages_handler(
model="claude",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_provider_config=mock_config,
anthropic_messages_optional_request_params={},
custom_llm_provider="anthropic",
litellm_params=GenericLiteLLMParams(timeout=9.0, stream_timeout=0.7),
logging_obj=logging_obj,
client=mock_client,
stream=True,
kwargs={},
)
assert mock_client.post.await_args.kwargs["stream"] is True
assert mock_client.post.await_args.kwargs["timeout"] == 0.7
@pytest.mark.asyncio
async def test_anthropic_post_uses_prebuilt_body_without_redumping():
"""When the caller passes a pre-serialized (unsigned) body, attempt 0 must
@ -1894,7 +2035,9 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques
ok_response = httpx.Response(200, json={"id": "msg_1"}, request=httpx.Request("POST", request_url))
class FakeAsyncClient:
async def post(self, url, headers, data, stream=False, logging_obj=None):
async def post(
self, url, headers, data, stream=False, logging_obj=None, timeout=None
):
posts.append({"headers": dict(headers), "data": data})
return invalid_signature_response if len(posts) == 1 else ok_response

View file

@ -1168,3 +1168,69 @@ class TestGetStructuredMessages:
data = {"input": None}
result = handler.get_structured_messages(data)
assert result is None
class ToolAppendingGuardrail(CustomGuardrail):
"""Guardrail that appends a new function tool, mimicking a guardrail that
injects a retrieval/recovery tool the model can later call."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
tools = list(inputs.get("tools") or [])
tools.append(
{
"type": "function",
"function": {
"name": "injected_tool",
"description": "injected by guardrail",
"parameters": {"type": "object", "properties": {}},
},
}
)
inputs["tools"] = tools
return inputs
class TestOpenAIResponsesHandlerToolInjection:
"""A tool a guardrail injects must survive the write-back to Responses format."""
def test_merge_keeps_guardrail_appended_tool(self):
"""_merge_tools_after_guardrail must not drop the extra appended tool."""
handler = OpenAIResponsesHandler()
original = [{"type": "function", "name": "a"}]
remapped = [
{"type": "function", "name": "a"},
{"type": "function", "name": "b"},
]
merged = handler._merge_tools_after_guardrail(original, remapped)
assert [t["name"] for t in merged] == ["a", "b"]
@pytest.mark.asyncio
async def test_injected_tool_survives_when_request_already_has_tools(self):
"""Regression: the merge dropped the injected tool whenever the request
already carried tools, so the model never saw it."""
handler = OpenAIResponsesHandler()
guardrail = ToolAppendingGuardrail(guardrail_name="test")
data = {
"input": [{"role": "user", "content": "hi", "type": "message"}],
"tools": [
{
"type": "function",
"name": "get_weather",
"parameters": {"type": "object", "properties": {}},
}
],
"model": "gpt-4",
}
result = await handler.process_input_messages(data, guardrail)
names = [t.get("name") for t in result["tools"]]
assert "get_weather" in names
assert "injected_tool" in names

View file

@ -356,6 +356,86 @@ class TestMCPServerManager:
assert server.oauth2_flow == "authorization_code"
assert server.needs_user_oauth_token is True
@pytest.mark.asyncio
async def test_load_servers_from_config_rejects_uncorroborated_endpoints_but_keeps_resource_scopes(self):
"""A yaml server with a manual authorization_url has the same config-time mix-up exposure as a
DB row: a document advertising a different authorize endpoint has its token_url rejected. The
resource-driven scopes are kept, because scope selection is resource-driven (MCP Scope
Selection Strategy) and scope inflation is bounded by the authorization server at consent, not
by dropping scopes when an endpoint mismatches."""
manager = MCPServerManager()
metadata = MCPOAuthMetadata(
authorization_url="https://attacker.example.com/authorize",
token_url="https://attacker.example.com/token",
scopes=["read", "admin"],
)
config = self._oauth2_config(
oauth2_flow="authorization_code",
authorization_url="https://idp.example.com/authorize",
token_url=None,
)
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
await manager.load_servers_from_config(config)
server = next(iter(manager.config_mcp_servers.values()))
assert server.authorization_url == "https://idp.example.com/authorize"
assert server.token_url is None
assert server.scopes == ["read", "admin"]
@pytest.mark.asyncio
async def test_load_servers_from_config_fills_token_url_when_metadata_corroborates_manual_authorization_url(self):
"""Corroborated metadata keeps the self-heal on the config path: when the discovered document
advertises the same authorize endpoint the admin pinned, its token_url fills the blank field
and scopes come through resource-driven (the discovered document's resource-preferred scopes),
not the authorization server's own capability list."""
manager = MCPServerManager()
metadata = MCPOAuthMetadata(
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
scopes=["read", "admin"],
)
config = self._oauth2_config(
oauth2_flow="authorization_code",
authorization_url="https://idp.example.com/authorize/",
token_url=None,
)
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
await manager.load_servers_from_config(config)
server = next(iter(manager.config_mcp_servers.values()))
assert server.token_url == "https://idp.example.com/token"
assert server.scopes == ["read", "admin"]
@pytest.mark.asyncio
@pytest.mark.parametrize("blank_authorization_url", ["", " "])
async def test_load_servers_from_config_blank_authorization_url_is_not_a_pin(self, blank_authorization_url):
"""A blank authorization_url — empty or whitespace-only — is not a trust anchor, so discovery
backfills the whole set (authorize endpoint, token_url, and its resource-preferred scopes)
from the same chain, exactly as if the field had been omitted. The merge and the corroboration
gate must agree that blank means unpinned; a whitespace value that the merge kept for redirects
while the gate treated as unpinned would strand a broken half-discovered config."""
manager = MCPServerManager()
metadata = MCPOAuthMetadata(
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
scopes=["read"],
)
config = self._oauth2_config(
oauth2_flow="authorization_code",
authorization_url=blank_authorization_url,
token_url=None,
)
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
await manager.load_servers_from_config(config)
server = next(iter(manager.config_mcp_servers.values()))
assert server.authorization_url == "https://idp.example.com/authorize"
assert server.token_url == "https://idp.example.com/token"
assert server.scopes == ["read"]
@pytest.mark.asyncio
async def test_load_servers_from_config_non_oauth2_needs_no_flow(self):
manager = MCPServerManager()
@ -1026,7 +1106,6 @@ class TestMCPServerManager:
"""The gateway's relayed authorize flow (used by the browser-only Authorize) needs the
upstream's authorization_url on the registry entry, and these rows never persist one, so
the DB build must discover it the same way oauth2 rows do."""
from types import SimpleNamespace
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(
@ -1053,6 +1132,181 @@ class TestMCPServerManager:
assert built.authorization_url == "https://idp.example.com/authorize"
assert built.token_url == "https://idp.example.com/token"
@pytest.mark.asyncio
async def test_build_from_table_backfills_resource_driven_scopes_for_pinned_authorization_url(self):
"""When authorization_url is admin-pinned and corroborated, scopes backfill as the
resource-driven value (the WWW-Authenticate challenge scope, else the RFC 9728
protected-resource scopes_supported), per the MCP authorization spec Scope Selection Strategy.
The client does not restrict scopes to the authorization server's own scopes_supported; scope
minimization and inflation control are the authorization server's and user's job at consent
(RFC 6749 §3.3)."""
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(
server_id="manual-auth-url-1",
alias="manual_auth_url",
description="manual authorization_url, blank scopes",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://idp.example.com/authorize",
created_at=datetime.now(),
updated_at=datetime.now(),
)
metadata = MCPOAuthMetadata(
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
scopes=["read", "admin"],
)
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)) as mock_discovery:
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
mock_discovery.assert_awaited_once()
assert built.authorization_url == "https://idp.example.com/authorize"
assert built.token_url == "https://idp.example.com/token"
assert built.scopes == ["read", "admin"]
@pytest.mark.asyncio
async def test_build_from_table_whitespace_authorization_url_is_not_a_pin(self):
"""A whitespace-only authorization_url on the row must not be kept for redirects while the
gate treats it as unpinned. It is normalized to unpinned everywhere, so the built server
takes the discovered authorize endpoint, token_url, and scopes as one consistent group
rather than serving the whitespace value with half-discovered fields."""
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(
server_id="whitespace-auth-url",
alias="whitespace_auth_url",
description="whitespace authorization_url is not a pin",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url=" ",
created_at=datetime.now(),
updated_at=datetime.now(),
)
metadata = MCPOAuthMetadata(
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
scopes=["read"],
)
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
assert built.authorization_url == "https://idp.example.com/authorize"
assert built.token_url == "https://idp.example.com/token"
assert built.scopes == ["read"]
@pytest.mark.asyncio
async def test_build_from_table_fills_endpoints_when_metadata_corroborates_manual_authorization_url(self):
"""A discovered token_url is only trusted next to a manual authorization_url when the same
metadata document advertises that authorize endpoint, and the comparison must tolerate
formatting-only differences (host case, trailing slash, query params like ?prompt=consent)
so hand-copied URLs still self-heal."""
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(
server_id="manual-auth-url-2",
alias="manual_auth_url_match",
description="manual authorization_url matching discovery, blank token_url",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://IDP.example.com/authorize/?prompt=consent",
created_at=datetime.now(),
updated_at=datetime.now(),
)
metadata = MCPOAuthMetadata(
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
registration_url="https://idp.example.com/register",
scopes=["read"],
)
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
assert built.authorization_url == "https://IDP.example.com/authorize/?prompt=consent"
assert built.token_url == "https://idp.example.com/token"
assert built.registration_url == "https://idp.example.com/register"
assert built.scopes == ["read"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"advertised_authorization_url",
["https://attacker.example.com/authorize", None],
)
async def test_build_from_table_rejects_uncorroborated_endpoints_but_keeps_resource_scopes(
self, advertised_authorization_url
):
"""Resource-rooted discovery lets a compromised upstream advertise its own authorization
server. With a manual authorization_url pinned, a document that does not corroborate it has
its token_url and registration_url dropped: accepting them would send the code, client secret,
and PKCE verifier to the attacker (config-time RFC 9700 mix-up). The resource-driven scopes
are kept, because scope selection is resource-driven (MCP Scope Selection Strategy) and scope
inflation is bounded by the authorization server at consent (RFC 6749 §3.3), not by dropping
scopes on an endpoint mismatch. Both the in-memory merge and the persisted metadata drop only
the uncorroborated endpoints."""
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(
server_id="manual-auth-url-3",
alias="manual_auth_url_mismatch",
description="manual authorization_url, hostile discovery document",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://idp.example.com/authorize",
created_at=datetime.now(),
updated_at=datetime.now(),
)
metadata = MCPOAuthMetadata(
authorization_url=advertised_authorization_url,
token_url="https://attacker.example.com/token",
registration_url="https://attacker.example.com/register",
scopes=["read", "admin"],
)
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)),
patch.object(manager, "_persist_discovered_oauth_endpoints", new=AsyncMock()) as mock_persist,
):
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
assert built.authorization_url == "https://idp.example.com/authorize"
assert built.token_url is None
assert built.registration_url is None
assert built.scopes == ["read", "admin"]
persisted_metadata = mock_persist.await_args.kwargs["metadata"]
assert persisted_metadata.token_url is None
assert persisted_metadata.registration_url is None
assert persisted_metadata.scopes == ["read", "admin"]
@pytest.mark.asyncio
async def test_build_from_table_skips_discovery_when_all_upstream_oauth_fields_present(self):
"""A fully hand-configured server (authorization_url, token_url, and scopes all set) has
nothing left for discovery to fill, so the build must not fetch upstream metadata."""
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(
server_id="fully-manual-1",
alias="fully_manual",
description="all upstream oauth fields set by the admin",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://idp.example.com/manual-authorize",
token_url="https://idp.example.com/manual-token",
credentials={"scopes": ["calendar.read"]},
created_at=datetime.now(),
updated_at=datetime.now(),
)
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)) as mock_discovery:
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
mock_discovery.assert_not_awaited()
assert built.authorization_url == "https://idp.example.com/manual-authorize"
assert built.token_url == "https://idp.example.com/manual-token"
assert built.scopes == ["calendar.read"]
async def _capture_subject_token(self, call) -> Optional[str]:
"""Run a manager method (via ``call(manager)``) and return the subject_token it threaded
into ``_create_mcp_client``."""
@ -2127,6 +2381,46 @@ class TestMCPServerManager:
assert result.scopes == ["api://some-scope/.default"]
assert result.from_origin_fallback is False
@pytest.mark.asyncio
async def test_descovery_metadata_scopes_are_resource_driven(self):
"""The effective `scopes` are resource-driven: the RFC 9728 protected-resource advertisement
(or WWW-Authenticate challenge) overrides the authorization server's own scopes_supported. This
is the MCP Scope Selection Strategy: the client requests what the resource needs, not the AS's
full capability list."""
manager = MCPServerManager()
mock_response = MagicMock()
mock_response.raise_for_status = MagicMock()
mock_client = MagicMock()
mock_client.get = AsyncMock(return_value=mock_response)
authorization_server_metadata = MCPOAuthMetadata(
scopes=["as.read", "as.write"],
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
)
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=mock_client,
),
patch.object(
manager,
"_attempt_well_known_discovery",
AsyncMock(return_value=(["https://idp.example.com"], ["resource.only"])),
),
patch.object(
manager,
"_fetch_authorization_server_metadata",
AsyncMock(return_value=authorization_server_metadata),
),
):
result = await manager._descovery_metadata("https://up.example.com/mcp")
assert result is not None
assert result.scopes == ["resource.only"]
@pytest.mark.asyncio
async def test_fetch_single_authorization_server_metadata_supports_azure_issuer_path(
self,
@ -2252,6 +2546,10 @@ class TestMCPServerManager:
@pytest.mark.asyncio
async def test_load_servers_from_config_overrides_discovery_metadata(self):
"""Config values win per field. The discovered token_url/registration_url do NOT fill the
blanks here: the document advertises a different authorization_endpoint than the manually
configured one, so combining its endpoints with the pinned authorize URL would be the
config-time mix-up the discovery gate exists to prevent."""
manager = MCPServerManager()
discovered_metadata = MCPOAuthMetadata(
@ -2285,8 +2583,8 @@ class TestMCPServerManager:
server = next(iter(manager.config_mcp_servers.values()))
assert server.scopes == ["config"] # config overrides discovery
assert server.authorization_url == "https://config.example.com/auth"
assert server.token_url == "https://discovered.example.com/token"
assert server.registration_url == "https://discovered.example.com/register"
assert server.token_url is None
assert server.registration_url is None
@pytest.mark.asyncio
async def test_load_servers_from_config_filters_blank_scopes(self):
@ -5093,6 +5391,102 @@ class TestMCPServerTimestamps:
_carry_forward_resolved_oauth_endpoints(new_server=explicit, previous_server=previous)
assert explicit.authorization_url == "https://configured.example.com/auth"
def test_carry_forward_does_not_revive_token_url_across_authorization_url_change(self):
"""Carry-forward is a non-manual endpoint source, so it obeys the same trust rule as
discovery: a previous token_url/registration_url belongs to the previous authorization
server, so it must not be pinned to a NEW authorization_url the admin re-pointed to. Without
this, re-pointing authorize to server B while the same MCP url keeps serving A's token
endpoint recreates the RFC 9700 mix-up, durably, and the discovery gate alone cannot catch
it because the stale endpoint comes from the registry, not from discovery."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_carry_forward_resolved_oauth_endpoints,
)
previous = MCPServer(
server_id="s1",
name="s1",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://idp-a.example.com/authorize",
token_url="https://idp-a.example.com/token",
registration_url="https://idp-a.example.com/register",
)
repointed = MCPServer(
server_id="s1",
name="s1",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://idp-b.example.com/authorize",
)
_carry_forward_resolved_oauth_endpoints(new_server=repointed, previous_server=previous)
assert repointed.authorization_url == "https://idp-b.example.com/authorize"
assert repointed.token_url is None
assert repointed.registration_url is None
def test_carry_forward_restores_endpoints_when_authorization_url_unchanged(self):
"""The last-known-good path still works: a rebuild whose discovery blipped (no authorize
endpoint) adopts the previous authorize endpoint AND its token endpoint together as a
consistent group, and a rebuild that re-pins the same authorize endpoint (formatting aside)
keeps carrying the corroborated token endpoint."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_carry_forward_resolved_oauth_endpoints,
)
def previous() -> MCPServer:
return MCPServer(
server_id="s1",
name="s1",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
registration_url="https://idp.example.com/register",
)
blipped = MCPServer(
server_id="s1",
name="s1",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url=None,
)
_carry_forward_resolved_oauth_endpoints(new_server=blipped, previous_server=previous())
assert blipped.authorization_url == "https://idp.example.com/authorize"
assert blipped.token_url == "https://idp.example.com/token"
assert blipped.registration_url == "https://idp.example.com/register"
same_authorize = MCPServer(
server_id="s1",
name="s1",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://IDP.example.com:443/authorize/",
)
_carry_forward_resolved_oauth_endpoints(new_server=same_authorize, previous_server=previous())
assert same_authorize.token_url == "https://idp.example.com/token"
assert same_authorize.registration_url == "https://idp.example.com/register"
def test_normalized_authorize_endpoint_treats_default_port_and_slash_as_identity(self):
"""The corroboration check must not fail on formatting-only differences an IdP legitimately
emits: default port, trailing slash, host case, and query string are not identity, but a
non-default port is."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_normalized_authorize_endpoint,
)
canonical = _normalized_authorize_endpoint("https://idp.example.com/authorize")
assert _normalized_authorize_endpoint("https://idp.example.com:443/authorize") == canonical
assert _normalized_authorize_endpoint("https://IDP.example.com/authorize/") == canonical
assert _normalized_authorize_endpoint("https://idp.example.com/authorize?prompt=consent") == canonical
assert _normalized_authorize_endpoint("https://idp.example.com:8443/authorize") != canonical
def test_build_mcp_server_table_preserves_timestamps(self):
"""_build_mcp_server_table must use the MCPServer's stored timestamps, not datetime.now()."""
manager = MCPServerManager()

View file

@ -659,7 +659,16 @@ def test_get_model_from_request_ignores_session_model_on_non_realtime_routes():
def test_abbreviate_api_key():
assert abbreviate_api_key("sk-test-1234") == "sk-...1234"
assert abbreviate_api_key("sk-test-1234-abcdefgh") == "sk-...efgh"
assert abbreviate_api_key("sk-abcdefghijklm") == "sk-...jklm"
def test_abbreviate_api_key_short_key_is_fully_masked():
"""Regression test for LIT-4355: for keys shorter than the enforced minimum,
showing the last 4 characters can reveal the entire key (sk-1234 -> sk-...1234)."""
assert abbreviate_api_key("sk-1234") == "sk-..."
assert abbreviate_api_key("sk-test-1234") == "sk-..."
assert abbreviate_api_key("") == "sk-..."
def test_get_customer_user_header_returns_none_when_no_customer_role():

File diff suppressed because it is too large Load diff

View file

@ -1291,6 +1291,136 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker):
}
def _patch_apply_guardrail_env(mocker, guardrail_result):
mock_guardrail = mocker.Mock()
mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result)
mock_registry = mocker.Mock()
mock_registry.get_initialized_guardrail_callback.return_value = mock_guardrail
mocker.patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry
)
mock_logging_obj = mocker.Mock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_obj.model_call_details = {}
mock_processor = mocker.Mock()
mock_processor.common_processing_pre_call_logic = AsyncMock(
return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
)
mocker.patch(
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
return_value=mock_processor,
)
mock_proxy_logging = mocker.Mock()
mock_proxy_logging.post_call_success_hook = AsyncMock()
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
mocker.patch("litellm.proxy.proxy_server.general_settings", {})
mocker.patch("litellm.proxy.proxy_server.proxy_config", mocker.Mock())
mocker.patch("litellm.proxy.proxy_server.version", "test")
mocker.patch("litellm.litellm_core_utils.thread_pool_executor.executor")
return mock_guardrail
@pytest.mark.asyncio
async def test_apply_guardrail_forwards_metadata_to_guardrail(mocker):
"""Client-supplied metadata must reach apply_guardrail via request_data so
parameterized custom guardrails can read per-request configuration."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail",
text="What are tax loopholes?",
metadata={"forbidden_topics": ["tax"]},
)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["What are tax loopholes?"]},
request_data={"metadata": {"forbidden_topics": ["tax"]}},
input_type="request",
)
@pytest.mark.asyncio
async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker):
"""metadata and messages must coexist in request_data; the dict merge must
not clobber messages when both fields are sent."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
messages = [{"role": "user", "content": "What are tax loopholes?"}]
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail",
text="What are tax loopholes?",
messages=messages,
metadata={"forbidden_topics": ["tax"]},
)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["What are tax loopholes?"]},
request_data={
"messages": messages,
"metadata": {"forbidden_topics": ["tax"]},
},
input_type="request",
)
@pytest.mark.asyncio
async def test_apply_guardrail_omits_metadata_when_not_sent(mocker):
"""Without metadata, request_data stays empty (backward-compatible)."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello")
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["hello"]},
request_data={},
input_type="request",
)
@pytest.mark.asyncio
async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(mocker):
"""Explicitly-sent empty messages/metadata must be forwarded, not dropped;
only omitted fields stay out of request_data."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail",
text="hello",
messages=[],
metadata={},
)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["hello"]},
request_data={"messages": [], "metadata": {}},
input_type="request",
)
@pytest.mark.asyncio
async def test_get_guardrail_info_endpoint_config_guardrail(mocker):
"""

View file

@ -563,7 +563,7 @@ async def test_generate_key_debug_log_never_contains_raw_token(monkeypatch, capl
generate_key_fn,
)
raw_key = "sk-short-secret"
raw_key = "sk-short-secret-a1b2"
with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"):
await generate_key_fn(
data=GenerateKeyRequest(key=raw_key),
@ -1336,10 +1336,10 @@ async def test_get_new_token_with_valid_key(monkeypatch):
)
# Test with valid new_key
data = RegenerateKeyRequest(new_key="sk-test123456789")
data = RegenerateKeyRequest(new_key="sk-test1234567890abc")
result = await get_new_token(data)
assert result == "sk-test123456789"
assert result == "sk-test1234567890abc"
@pytest.mark.asyncio
@ -1370,6 +1370,110 @@ async def test_get_new_token_with_invalid_key(monkeypatch):
assert "New key must start with 'sk-'" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_get_new_token_rejects_short_new_key(monkeypatch):
"""Regression test for LIT-4355: a short custom key like sk-99 must be rejected,
otherwise the stored key_name (sk-...{last 4 chars}) reveals the entire key."""
from unittest.mock import AsyncMock
from fastapi import HTTPException
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
get_new_token,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached",
AsyncMock(return_value={}),
)
data = RegenerateKeyRequest(new_key="sk-99")
with pytest.raises(HTTPException) as exc_info:
await get_new_token(data)
assert exc_info.value.status_code == 400
assert "at least 16 characters" in str(exc_info.value.detail)
@pytest.mark.asyncio
@pytest.mark.parametrize("short_key", ["sk-1234", "sk-abcdefghijkl"])
async def test_generate_key_fn_rejects_short_custom_key(monkeypatch, short_key):
"""Regression test for LIT-4355: /key/generate must reject custom keys shorter
than the minimum length (including the 15-char boundary); sk-1234 used to be
accepted and fully exposed via key_name."""
mock_prisma_client = AsyncMock()
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_verificationtoken = MagicMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles, ProxyException
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_fn,
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached",
AsyncMock(return_value={}),
)
assert len(short_key) < 16
with pytest.raises(ProxyException) as exc_info:
await generate_key_fn(
data=GenerateKeyRequest(key=short_key),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234"
),
)
assert exc_info.value.code == "400"
assert "at least 16 characters" in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_generate_key_fn_accepts_custom_key_at_minimum_length(monkeypatch):
"""Custom keys at exactly the minimum length (16 chars) are still accepted."""
mock_prisma_client = AsyncMock()
mock_insert_data = AsyncMock(
return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None, object_permission=None)
)
mock_prisma_client.insert_data = mock_insert_data
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_verificationtoken = MagicMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0)
from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_fn,
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached",
AsyncMock(return_value={}),
)
custom_key = "sk-abcdefghijklm"
assert len(custom_key) == 16
response = await generate_key_fn(
data=GenerateKeyRequest(key=custom_key),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234"
),
)
assert response.key == custom_key
@pytest.mark.asyncio
async def test_check_custom_key_allowed_when_disabled(monkeypatch):
"""_check_custom_key_allowed raises 403 when disable_custom_api_keys is true."""

View file

@ -8,7 +8,7 @@ sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to
from unittest.mock import MagicMock
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_request
@pytest.mark.parametrize(
@ -42,6 +42,200 @@ async def test_route_request_dynamic_credentials(route_type):
getattr(llm_router, route_type).assert_called_once_with(**data)
@pytest.mark.asyncio
async def test_route_request_proxy_admin_can_call_all_team_scoped_deployments_without_team_id():
import litellm
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
router = litellm.Router(
model_list=[
{
"model_name": "internal-team-azure-east",
"litellm_params": {
"model": "azure/gpt-4o",
"api_key": "fake",
"api_base": "https://east.example.openai.azure.com",
"api_version": "2024-02-15-preview",
"mock_response": "east",
},
"model_info": {
"id": "team-azure-east",
"team_id": "team-a",
"team_public_model_name": "team-azure",
},
},
{
"model_name": "internal-team-azure-west",
"litellm_params": {
"model": "azure/gpt-4o",
"api_key": "fake",
"api_base": "https://west.example.openai.azure.com",
"api_version": "2024-02-15-preview",
"mock_response": "west",
},
"model_info": {
"id": "team-azure-west",
"team_id": "team-a",
"team_public_model_name": "team-azure",
},
},
]
)
admin_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
data = {
"model": "team-azure",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {"user_api_key_auth": admin_auth},
}
llm_call = await route_request(
data=data,
llm_router=router,
user_model=None,
route_type="acompletion",
user_api_key_dict=admin_auth,
)
response = await llm_call
deployments = await router.async_get_healthy_deployments(
model="team-azure",
request_kwargs=data,
)
assert response.choices[0].message.content in {"east", "west"}
assert {deployment["model_info"]["id"] for deployment in deployments} == {
"team-azure-east",
"team-azure-west",
}
non_admin_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER)
with pytest.raises(ProxyModelNotFoundError):
await route_request(
data={
**data,
"metadata": {"user_api_key_auth": non_admin_auth},
},
llm_router=router,
user_model=None,
route_type="acompletion",
user_api_key_dict=non_admin_auth,
)
from litellm.types.router import Deployment
router.add_deployment(
Deployment(
model_name="internal-team-only",
litellm_params={
"model": "azure/gpt-4o",
"api_key": "fake",
"api_base": "https://internal.example.openai.azure.com",
"api_version": "2024-02-15-preview",
},
model_info={
"id": "internal-team-only-id",
"team_id": "team-a",
},
)
)
internal_deployments = await router.async_get_healthy_deployments(
model="internal-team-only",
request_kwargs={
**data,
"model": "internal-team-only",
},
)
assert {deployment["model_info"]["id"] for deployment in internal_deployments} == {"internal-team-only-id"}
router.add_deployment(
Deployment(
model_name="internal-other-team-azure",
litellm_params={
"model": "azure/gpt-4o",
"api_key": "fake",
"api_base": "https://other.example.openai.azure.com",
"api_version": "2024-02-15-preview",
"mock_response": "other",
},
model_info={
"id": "other-team-azure",
"team_id": "team-b",
"team_public_model_name": "team-azure",
},
)
)
with pytest.raises(litellm.BadRequestError, match="multiple teams"):
ambiguous_call = await route_request(
data=data,
llm_router=router,
user_model=None,
route_type="acompletion",
user_api_key_dict=admin_auth,
)
await ambiguous_call
router.add_deployment(
Deployment(
model_name="team-azure",
litellm_params={
"model": "azure/gpt-4o",
"api_key": "fake",
"api_base": "https://legacy.example.openai.azure.com",
"api_version": "2024-02-15-preview",
},
model_info={
"id": "legacy-team-azure",
"team_id": "team-a",
"team_public_model_name": "team-azure",
},
)
)
router.add_deployment(
Deployment(
model_name="team-azure",
litellm_params={
"model": "azure/gpt-4o",
"api_key": "fake",
"api_base": "https://other-legacy.example.openai.azure.com",
"api_version": "2024-02-15-preview",
},
model_info={
"id": "other-legacy-team-azure",
"team_id": "team-b",
"team_public_model_name": "team-azure",
},
)
)
with pytest.raises(litellm.BadRequestError, match="multiple teams"):
await router.async_get_healthy_deployments(
model="team-azure",
request_kwargs=data,
)
router.add_deployment(
Deployment(
model_name="team-azure",
litellm_params={
"model": "azure/gpt-4o",
"api_key": "fake",
"api_base": "https://global.example.openai.azure.com",
"api_version": "2024-02-15-preview",
},
model_info={"id": "global-team-azure"},
)
)
collision_deployments = await router.async_get_healthy_deployments(
model="team-azure",
request_kwargs=data,
)
assert {deployment["model_info"]["id"] for deployment in collision_deployments} == {"global-team-azure"}
@pytest.mark.asyncio
async def test_route_request_no_model_required():
"""Test route types that don't require model parameter"""

View file

@ -330,6 +330,13 @@ def test_get_combined_callback_list_matrix(proxy_logging):
}
def test_get_combined_callback_list_preserves_insertion_order(proxy_logging):
assert proxy_logging.get_combined_callback_list(
dynamic_success_callbacks=["prometheus", "langfuse", "datadog", "otel", "s3"],
global_callbacks=["langfuse", "gcs_bucket", "arize", "logfire"],
) == ["prometheus", "langfuse", "datadog", "otel", "s3", "gcs_bucket", "arize", "logfire"]
def test_get_combined_callback_list_unhashable_dynamic_raises(proxy_logging):
with pytest.raises(TypeError):
proxy_logging.get_combined_callback_list(

View file

@ -2387,6 +2387,16 @@ class TestSubCallMetadataSanitization:
assert sanitized["user_api_key_auth"] is not None
assert _get_budget_reservation_from_metadata(sanitized) is None
def test_returns_empty_dict_for_missing_metadata(self):
from litellm.router_strategy.complexity_router.complexity_router import (
_classifier_call_metadata,
)
for absent in (None, {}):
result = _classifier_call_metadata(absent)
assert result == {}
assert isinstance(result, dict)
def test_sanitized_auth_keeps_access_group_fields_and_leaves_original_untouched(self):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.router_strategy.complexity_router.complexity_router import (

View file

@ -977,3 +977,65 @@ class TestRouterIOTokenIntegration:
assert info is not None
assert info.itpm == 100
assert info.otpm == 20
class TestContextSlotRetention:
def test_setter_stores_kwargs_only_for_io_limited_deployments(self):
"""
The context slot pins the entire request kwargs (messages included)
for the lifetime of the surrounding asyncio context, and pooled
resources created mid-request (e.g. redis connections) capture that
context, extending the pin far past the request. Only ITPM/OTPM
pre-call checks read the slot, so the setter must store None for
deployments without io token limits and still clear reservation
sentinels from kwargs either way.
"""
kwargs = {
"messages": [{"role": "user", "content": "x" * 1000}],
"metadata": {ITPM_RESERVED_KEY: 999, ITPM_CACHE_KEY: "forged"},
}
set_io_token_rate_limit_request_kwargs(kwargs, store_in_context=False)
assert get_io_token_rate_limit_request_kwargs() is None
assert ITPM_RESERVED_KEY not in kwargs["metadata"]
assert ITPM_CACHE_KEY not in kwargs["metadata"]
set_io_token_rate_limit_request_kwargs(kwargs, store_in_context=True)
assert get_io_token_rate_limit_request_kwargs() is kwargs
set_io_token_rate_limit_request_kwargs(kwargs, store_in_context=False)
assert get_io_token_rate_limit_request_kwargs() is None
@pytest.mark.asyncio
async def test_router_does_not_pin_kwargs_without_io_limits(self):
router = Router(
model_list=[
{
"model_name": "plain",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test"},
}
]
)
set_io_token_rate_limit_request_kwargs(None)
kwargs = {"messages": [{"role": "user", "content": "hello"}], "metadata": {}}
deployment = router.get_deployment_by_model_group_name("plain")
assert deployment is not None
router._update_kwargs_with_deployment(deployment=deployment.model_dump(), kwargs=kwargs)
assert get_io_token_rate_limit_request_kwargs() is None
@pytest.mark.asyncio
async def test_router_pins_kwargs_for_io_limited_deployment(self):
router = Router(
model_list=[
{
"model_name": "limited",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test", "itpm": 100},
}
],
optional_pre_call_checks=["enforce_model_rate_limits"],
)
set_io_token_rate_limit_request_kwargs(None)
kwargs = {"messages": [{"role": "user", "content": "hello"}], "metadata": {}}
deployment = router.get_deployment_by_model_group_name("limited")
assert deployment is not None
router._update_kwargs_with_deployment(deployment=deployment.model_dump(), kwargs=kwargs)
assert get_io_token_rate_limit_request_kwargs() is kwargs

View file

@ -65,6 +65,16 @@ def test_redact_string_catches_secret_patterns():
assert redact_string(normal) == normal
def test_redact_string_catches_minimum_length_virtual_key():
"""Regression test for LIT-4355: keys at the enforced 16-char minimum
(MINIMUM_CUSTOM_KEY_LENGTH) must be treated as key-shaped by the scrubber."""
minimum_length_key = "sk-abcdefghijklm"
assert len(minimum_length_key) == 16
result = redact_string("msg: " + minimum_length_key)
assert minimum_length_key not in result
assert "REDACTED" in result
def test_filter_redacts_secrets_in_logger_output():
def log_messages():
verbose_logger.debug("Key: " + SECRET)

View file

@ -1655,17 +1655,6 @@
"count": 1
}
},
"src/components/UsageIndicator.tsx": {
"no-nested-ternary": {
"count": 4
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/static-components": {
"count": 1
}
},
"src/components/activity_metrics.tsx": {
"no-nested-ternary": {
"count": 1

View file

@ -53,7 +53,72 @@ describe("GuardrailTestPanel", () => {
// Verify onSubmit was called with the correct text
await waitFor(() => {
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text");
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text", null);
});
});
it("should submit parsed metadata when a JSON object is provided", async () => {
/**
* Tests that a JSON object typed into the Metadata field is parsed and
* passed to onSubmit so it reaches the apply_guardrail request body.
*/
const user = userEvent.setup();
render(
<GuardrailTestPanel
guardrailNames={mockGuardrailNames}
onSubmit={mockOnSubmit}
isLoading={false}
results={null}
errors={null}
onClose={mockOnClose}
/>,
);
const textarea = screen.getByPlaceholderText("Enter text to test with guardrails...");
await user.type(textarea, "Test input text");
const metadataField = screen.getByPlaceholderText('{"forbidden_topics": ["tax", "finance"]}');
await user.click(metadataField);
await user.paste('{"forbidden_topics": ["tax"]}');
await user.click(screen.getByRole("button", { name: /Test 2 guardrails/ }));
await waitFor(() => {
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text", { forbidden_topics: ["tax"] });
});
});
it("should block submission and show an error for invalid metadata JSON", async () => {
/**
* Tests that invalid JSON in the Metadata field prevents submission
* instead of silently sending a request without metadata.
*/
const user = userEvent.setup();
render(
<GuardrailTestPanel
guardrailNames={mockGuardrailNames}
onSubmit={mockOnSubmit}
isLoading={false}
results={null}
errors={null}
onClose={mockOnClose}
/>,
);
const textarea = screen.getByPlaceholderText("Enter text to test with guardrails...");
await user.type(textarea, "Test input text");
const metadataField = screen.getByPlaceholderText('{"forbidden_topics": ["tax", "finance"]}');
await user.click(metadataField);
await user.paste("{not json");
await user.click(screen.getByRole("button", { name: /Test 2 guardrails/ }));
await waitFor(() => {
expect(screen.getByText("Invalid JSON")).toBeInTheDocument();
});
expect(mockOnSubmit).not.toHaveBeenCalled();
});
});

View file

@ -10,7 +10,7 @@ const { Text } = Typography;
interface GuardrailTestPanelProps {
guardrailNames: string[];
onSubmit: (text: string) => void;
onSubmit: (text: string, metadata?: Record<string, unknown> | null) => void;
isLoading: boolean;
results: Array<{ guardrailName: string; response_text: string; latency: number }> | null;
errors: Array<{ guardrailName: string; error: Error; latency: number }> | null;
@ -26,6 +26,23 @@ export function GuardrailTestPanel({
onClose,
}: GuardrailTestPanelProps) {
const [inputText, setInputText] = useState("");
const [metadataText, setMetadataText] = useState("");
const [metadataError, setMetadataError] = useState<string | null>(null);
const parseMetadata = (raw: string): { metadata: Record<string, unknown> | null; error: string | null } => {
if (!raw.trim()) {
return { metadata: null, error: null };
}
try {
const parsed = JSON.parse(raw);
if (parsed === null || typeof parsed !== "object" || Array.isArray(parsed)) {
return { metadata: null, error: "Metadata must be a JSON object" };
}
return { metadata: parsed, error: null };
} catch {
return { metadata: null, error: "Invalid JSON" };
}
};
const handleSubmit = () => {
if (!inputText.trim()) {
@ -33,7 +50,15 @@ export function GuardrailTestPanel({
return;
}
onSubmit(inputText);
const { metadata, error } = parseMetadata(metadataText);
if (error) {
setMetadataError(error);
NotificationsManager.fromBackend(`Metadata: ${error}`);
return;
}
setMetadataError(null);
onSubmit(inputText, metadata);
};
const handleKeyDown = (e: React.KeyboardEvent<HTMLTextAreaElement>) => {
@ -142,6 +167,33 @@ export function GuardrailTestPanel({
</div>
</div>
<div>
<div className="flex items-center gap-2 mb-2">
<label className="text-sm font-medium text-gray-700">Metadata (optional)</label>
<Tooltip title="JSON object forwarded to the guardrail as request_data['metadata']. Custom guardrails can read per-request configuration from it.">
<InfoCircleOutlined className="text-gray-400 cursor-help" />
</Tooltip>
</div>
<TextArea
value={metadataText}
onChange={(e) => {
setMetadataText(e.target.value);
if (metadataError) {
setMetadataError(parseMetadata(e.target.value).error);
}
}}
placeholder='{"forbidden_topics": ["tax", "finance"]}'
rows={3}
className="font-mono text-sm"
status={metadataError ? "error" : undefined}
/>
{metadataError && (
<Text type="danger" className="text-xs">
{metadataError}
</Text>
)}
</div>
<div className="pt-2">
<Button onClick={handleSubmit} loading={isLoading} disabled={!inputText.trim()} className="w-full">
{isLoading

View file

@ -63,7 +63,7 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
setSelectedGuardrails(newSelection);
};
const handleTestGuardrails = async (text: string) => {
const handleTestGuardrails = async (text: string, metadata?: Record<string, unknown> | null) => {
if (selectedGuardrails.size === 0 || !accessToken) {
return;
}
@ -79,7 +79,7 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
Array.from(selectedGuardrails).map(async (guardrailName) => {
const startTime = Date.now();
try {
const result = await applyGuardrail(accessToken, guardrailName, text, null, null);
const result = await applyGuardrail(accessToken, guardrailName, text, null, null, metadata);
const latency = Date.now() - startTime;
results.push({
guardrailName,

View file

@ -1,190 +0,0 @@
import { describe, it, expect, beforeEach, afterEach, vi } from "vitest";
import { act, renderHook, waitFor } from "@testing-library/react";
import { useDisableUsageIndicator } from "./useDisableUsageIndicator";
import { LOCAL_STORAGE_EVENT } from "@/utils/localStorageUtils";
describe("useDisableUsageIndicator", () => {
const STORAGE_KEY = "disableUsageIndicator";
beforeEach(() => {
localStorage.clear();
vi.clearAllMocks();
});
afterEach(() => {
localStorage.clear();
});
it("should return false when localStorage is empty", () => {
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(false);
});
it("should return false when localStorage value is not 'true'", () => {
localStorage.setItem(STORAGE_KEY, "false");
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(false);
});
it("should return true when localStorage value is 'true'", () => {
localStorage.setItem(STORAGE_KEY, "true");
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(true);
});
it("should return false when localStorage value is an empty string", () => {
localStorage.setItem(STORAGE_KEY, "");
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(false);
});
it("should update when storage event fires for the correct key", async () => {
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(false);
await act(async () => {
localStorage.setItem(STORAGE_KEY, "true");
const storageEvent = new StorageEvent("storage", {
key: STORAGE_KEY,
newValue: "true",
});
window.dispatchEvent(storageEvent);
});
await waitFor(() => {
expect(result.current).toBe(true);
});
});
it("should not update when storage event fires for a different key", () => {
localStorage.setItem(STORAGE_KEY, "false");
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(false);
const storageEvent = new StorageEvent("storage", {
key: "otherKey",
newValue: "true",
});
window.dispatchEvent(storageEvent);
expect(result.current).toBe(false);
});
it("should update when custom LOCAL_STORAGE_EVENT fires for the correct key", async () => {
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(false);
await act(async () => {
localStorage.setItem(STORAGE_KEY, "true");
const customEvent = new CustomEvent(LOCAL_STORAGE_EVENT, {
detail: { key: STORAGE_KEY },
});
window.dispatchEvent(customEvent);
});
await waitFor(() => {
expect(result.current).toBe(true);
});
});
it("should not update when custom LOCAL_STORAGE_EVENT fires for a different key", () => {
localStorage.setItem(STORAGE_KEY, "false");
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(false);
const customEvent = new CustomEvent(LOCAL_STORAGE_EVENT, {
detail: { key: "otherKey" },
});
window.dispatchEvent(customEvent);
expect(result.current).toBe(false);
});
it("should update when localStorage changes from false to true via custom event", async () => {
localStorage.setItem(STORAGE_KEY, "false");
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(false);
await act(async () => {
localStorage.setItem(STORAGE_KEY, "true");
const customEvent = new CustomEvent(LOCAL_STORAGE_EVENT, {
detail: { key: STORAGE_KEY },
});
window.dispatchEvent(customEvent);
});
await waitFor(() => {
expect(result.current).toBe(true);
});
});
it("should update when localStorage changes from true to false via storage event", async () => {
localStorage.setItem(STORAGE_KEY, "true");
const { result } = renderHook(() => useDisableUsageIndicator());
expect(result.current).toBe(true);
await act(async () => {
localStorage.setItem(STORAGE_KEY, "false");
const storageEvent = new StorageEvent("storage", {
key: STORAGE_KEY,
newValue: "false",
});
window.dispatchEvent(storageEvent);
});
await waitFor(() => {
expect(result.current).toBe(false);
});
});
it("should cleanup event listeners on unmount", () => {
const addEventListenerSpy = vi.spyOn(window, "addEventListener");
const removeEventListenerSpy = vi.spyOn(window, "removeEventListener");
const { unmount } = renderHook(() => useDisableUsageIndicator());
expect(addEventListenerSpy).toHaveBeenCalledTimes(2);
expect(addEventListenerSpy).toHaveBeenCalledWith("storage", expect.any(Function));
expect(addEventListenerSpy).toHaveBeenCalledWith(LOCAL_STORAGE_EVENT, expect.any(Function));
unmount();
expect(removeEventListenerSpy).toHaveBeenCalledTimes(2);
expect(removeEventListenerSpy).toHaveBeenCalledWith("storage", expect.any(Function));
expect(removeEventListenerSpy).toHaveBeenCalledWith(LOCAL_STORAGE_EVENT, expect.any(Function));
});
it("should handle multiple hooks independently", async () => {
const { result: result1 } = renderHook(() => useDisableUsageIndicator());
const { result: result2 } = renderHook(() => useDisableUsageIndicator());
expect(result1.current).toBe(false);
expect(result2.current).toBe(false);
await act(async () => {
localStorage.setItem(STORAGE_KEY, "true");
const customEvent = new CustomEvent(LOCAL_STORAGE_EVENT, {
detail: { key: STORAGE_KEY },
});
window.dispatchEvent(customEvent);
});
await waitFor(() => {
expect(result1.current).toBe(true);
expect(result2.current).toBe(true);
});
});
});

View file

@ -1,33 +0,0 @@
import { getLocalStorageItem, LOCAL_STORAGE_EVENT } from "@/utils/localStorageUtils";
import { useSyncExternalStore } from "react";
function subscribe(callback: () => void) {
const onStorage = (e: StorageEvent) => {
if (e.key === "disableUsageIndicator") {
callback();
}
};
const onCustom = (e: Event) => {
const { key } = (e as CustomEvent).detail;
if (key === "disableUsageIndicator") {
callback();
}
};
window.addEventListener("storage", onStorage);
window.addEventListener(LOCAL_STORAGE_EVENT, onCustom);
return () => {
window.removeEventListener("storage", onStorage);
window.removeEventListener(LOCAL_STORAGE_EVENT, onCustom);
};
}
function getSnapshot() {
return getLocalStorageItem("disableUsageIndicator") === "true";
}
export function useDisableUsageIndicator() {
return useSyncExternalStore(subscribe, getSnapshot);
}

View file

@ -10,7 +10,7 @@ import { Button } from "@/components/ui/button";
import MessageManager from "@/components/molecules/message_manager";
import { useRouter } from "next/navigation";
import { useChatShell } from "@/contexts/ChatShellContext";
import { CHAT_ROUTES } from "@/components/chat/ChatShell";
import { getChatRoutes } from "@/components/chat/ChatShell";
import ChatMessages from "@/components/chat/ChatMessages";
import MCPConnectPicker from "@/components/chat/MCPConnectPicker";
import { fetchAvailableModels } from "@/components/llm_calls/fetch_models";
@ -87,7 +87,7 @@ export default function ChatConversationPage() {
const streamScrollLock = useRef<number | null>(null);
useEffect(() => {
if (staleId) router.replace(CHAT_ROUTES.chats);
if (staleId) router.replace(getChatRoutes().chats);
}, [staleId, router]);
// Load models
@ -140,7 +140,7 @@ export default function ChatConversationPage() {
if (!convId) {
convId = createConversation(model);
setResponsesSessionId(null); // new conversation starts a fresh session
router.push(`${CHAT_ROUTES.chats}?id=${convId}`);
window.history.pushState(null, "", `${window.location.pathname}?id=${convId}`);
}
appendMessage(convId, { role: "user", content: trimmed });
@ -248,7 +248,6 @@ export default function ChatConversationPage() {
createConversation,
appendMessage,
updateLastAssistantMessage,
router,
isStreaming,
responsesSessionId,
],
@ -529,7 +528,7 @@ export default function ChatConversationPage() {
Chat with 100+ LLMs + MCP tools; authenticate once, use them here.{" "}
<Button
variant="link"
onClick={() => router.push(CHAT_ROUTES.integrations)}
onClick={() => router.push(getChatRoutes().integrations)}
className="h-auto p-0 text-sm font-medium"
>
Open Integrations -&gt;

View file

@ -0,0 +1,70 @@
import { render, screen } from "@testing-library/react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import BetaBadge from "./BetaBadge";
// Mock the hook directly
vi.mock("@/app/(dashboard)/hooks/useDisableShowNewBadge", () => ({
useDisableShowNewBadge: vi.fn(),
}));
import { useDisableShowNewBadge } from "@/app/(dashboard)/hooks/useDisableShowNewBadge";
const mockUseDisableShowNewBadge = vi.mocked(useDisableShowNewBadge);
describe("BetaBadge", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("should render the badge when disableShowNewBadge is false", () => {
mockUseDisableShowNewBadge.mockReturnValue(false);
render(<BetaBadge>Test Content</BetaBadge>);
expect(screen.getByText("Beta")).toBeInTheDocument();
expect(screen.getByText("Test Content")).toBeInTheDocument();
});
it("should render the badge when disableShowNewBadge is not set", () => {
mockUseDisableShowNewBadge.mockReturnValue(false);
render(<BetaBadge />);
expect(screen.getByText("Beta")).toBeInTheDocument();
});
it("should render only children when disableShowNewBadge is true", () => {
mockUseDisableShowNewBadge.mockReturnValue(true);
render(<BetaBadge>Test Content</BetaBadge>);
expect(screen.queryByText("Beta")).not.toBeInTheDocument();
expect(screen.getByText("Test Content")).toBeInTheDocument();
});
it("should render nothing when disableShowNewBadge is true and no children", () => {
mockUseDisableShowNewBadge.mockReturnValue(true);
const { container } = render(<BetaBadge />);
expect(container.firstChild).toBeNull();
});
it("should render badge with dot instead of text when dot prop is true", () => {
mockUseDisableShowNewBadge.mockReturnValue(false);
render(<BetaBadge dot={true}>Test Content</BetaBadge>);
expect(screen.queryByText("Beta")).not.toBeInTheDocument();
expect(screen.getByText("Test Content")).toBeInTheDocument();
});
it("should render badge with 'Beta' text when dot prop is not provided (defaults to false)", () => {
mockUseDisableShowNewBadge.mockReturnValue(false);
render(<BetaBadge>Test Content</BetaBadge>);
expect(screen.getByText("Beta")).toBeInTheDocument();
expect(screen.getByText("Test Content")).toBeInTheDocument();
});
});

View file

@ -0,0 +1,18 @@
import { Badge } from "antd";
import { useDisableShowNewBadge } from "@/app/(dashboard)/hooks/useDisableShowNewBadge";
export default function BetaBadge({ children, dot = false }: { children?: React.ReactNode; dot?: boolean }) {
const disableShowNewBadge = useDisableShowNewBadge();
if (disableShowNewBadge) {
return children ? <>{children}</> : null;
}
return children ? (
<Badge color="blue" count={dot ? undefined : "Beta"} dot={dot}>
{children}
</Badge>
) : (
<Badge color="blue" count={dot ? undefined : "Beta"} dot={dot} />
);
}

View file

@ -2,7 +2,6 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { useDisableBlogPosts } from "@/app/(dashboard)/hooks/useDisableBlogPosts";
import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBouncingIcon";
import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts";
import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator";
import {
emitLocalStorageChange,
getLocalStorageItem,
@ -72,7 +71,6 @@ interface UserDropdownProps {
const UserDropdown: React.FC<UserDropdownProps> = ({ onLogout, variant = "navbar", collapsed = false }) => {
const { userId, userEmail, userRole, premiumUser } = useAuthorized();
const disableShowPrompts = useDisableShowPrompts();
const disableUsageIndicator = useDisableUsageIndicator();
const disableBlogPosts = useDisableBlogPosts();
const disableBouncingIcon = useDisableBouncingIcon();
const [disableShowNewBadge, setDisableShowNewBadge] = useState(false);
@ -165,23 +163,6 @@ const UserDropdown: React.FC<UserDropdownProps> = ({ onLogout, variant = "navbar
aria-label="Toggle hide all prompts"
/>
</Space>
<Space style={{ width: "100%", justifyContent: "space-between" }}>
<Text type="secondary">Hide Usage Indicator</Text>
<Switch
size="small"
checked={disableUsageIndicator}
onChange={(checked) => {
if (checked) {
setLocalStorageItem("disableUsageIndicator", "true");
emitLocalStorageChange("disableUsageIndicator");
} else {
removeLocalStorageItem("disableUsageIndicator");
emitLocalStorageChange("disableUsageIndicator");
}
}}
aria-label="Toggle hide usage indicator"
/>
</Space>
<Space style={{ width: "100%", justifyContent: "space-between" }}>
<Text type="secondary">Hide Blog Posts</Text>
<Switch

View file

@ -37,10 +37,6 @@ vi.mock("@/app/(dashboard)/hooks/useDisableShowPrompts", () => ({
useDisableShowPrompts: () => mockUseDisableShowPromptsImpl(),
}));
vi.mock("@/app/(dashboard)/hooks/useDisableUsageIndicator", () => ({
useDisableUsageIndicator: () => false,
}));
vi.mock("@/app/(dashboard)/hooks/useDisableBlogPosts", () => ({
useDisableBlogPosts: () => false,
}));

View file

@ -4,7 +4,6 @@ import { useDisableBlogPosts } from "@/app/(dashboard)/hooks/useDisableBlogPosts
import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBouncingIcon";
import { useDisableShowNewBadge } from "@/app/(dashboard)/hooks/useDisableShowNewBadge";
import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts";
import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator";
import { emitLocalStorageChange, removeLocalStorageItem, setLocalStorageItem } from "@/utils/localStorageUtils";
import { navAccountDisplayName } from "@/components/Navbar/navDisplayName";
import CopyButton from "@/components/shared/CopyButton";
@ -86,7 +85,6 @@ const SidebarAccountMenu: React.FC<SidebarAccountMenuProps> = ({ onLogout, colla
const { data: healthData } = useHealthReadinessDetails(accessToken);
const version = healthData?.litellm_version;
const disableShowPrompts = useDisableShowPrompts();
const disableUsageIndicator = useDisableUsageIndicator();
const disableBlogPosts = useDisableBlogPosts();
const disableBouncingIcon = useDisableBouncingIcon();
const disableShowNewBadge = useDisableShowNewBadge();
@ -115,13 +113,6 @@ const SidebarAccountMenu: React.FC<SidebarAccountMenuProps> = ({ onLogout, colla
checked: disableShowPrompts,
onCheckedChange: (checked: boolean) => setFlag("disableShowPrompts", checked),
},
{
key: "disableUsageIndicator",
label: "Hide Usage Indicator",
ariaLabel: "Toggle hide usage indicator",
checked: disableUsageIndicator,
onCheckedChange: (checked: boolean) => setFlag("disableUsageIndicator", checked),
},
{
key: "disableBlogPosts",
label: "Hide Blog Posts",

View file

@ -8,10 +8,6 @@ import type { LicenseInfo } from "./networking";
vi.mock("./networking", () => ({ getRemainingUsers: vi.fn() }));
vi.mock("@/app/(dashboard)/hooks/useDisableUsageIndicator", () => ({
useDisableUsageIndicator: vi.fn(() => false),
}));
vi.mock("@/app/(dashboard)/hooks/license/useLicenseInfo", () => ({
useLicenseInfo: vi.fn(),
}));
@ -128,6 +124,29 @@ describe("SidebarUsageCard", () => {
expect(container.querySelector('[data-slot="meter"]')).toBeNull();
});
it("shows the exact license expiration date as the subtitle instead of time remaining", async () => {
mockUseLicenseInfo.mockReturnValue(licenseResult({ ...ACTIVE_LICENSE, expiration_date: "2099-12-31" }));
renderWithClient(<SidebarUsageCard accessToken="token" collapsed={false} onExpandRail={() => {}} />);
expect(await screen.findByText("Expires Dec 31, 2099")).toBeInTheDocument();
expect(screen.queryByText(/(day|days|month|months) remaining/)).not.toBeInTheDocument();
});
it("shows the exact date as the subtitle when the license is expired", async () => {
mockUseLicenseInfo.mockReturnValue(licenseResult({ ...ACTIVE_LICENSE, expiration_date: "2020-01-01" }));
renderWithClient(<SidebarUsageCard accessToken="token" collapsed={false} onExpandRail={() => {}} />);
expect(await screen.findByText("Expired Jan 1, 2020")).toBeInTheDocument();
});
it("falls back to Active plan when the license has no expiration date", async () => {
renderWithClient(<SidebarUsageCard accessToken="token" collapsed={false} onExpandRail={() => {}} />);
expect(await screen.findByText("Active plan")).toBeInTheDocument();
});
it("shows a collapsed rail button that expands the sidebar", async () => {
const onExpandRail = vi.fn();
const user = userEvent.setup();

View file

@ -1,6 +1,5 @@
import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator";
import { useLicenseInfo } from "@/app/(dashboard)/hooks/license/useLicenseInfo";
import { getDaysUntilExpiration } from "@/utils/licenseUtils";
import { formatExpirationStatus } from "@/utils/licenseUtils";
import { Button } from "@/components/ui/button";
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
import { Meter, MeterIndicator, MeterLabel, MeterTrack } from "@/components/ui/meter";
@ -20,16 +19,6 @@ interface MeterData {
total: number;
}
const formatExpiration = (daysRemaining: number | null): string => {
if (daysRemaining === null) return "No expiration";
if (daysRemaining < 0) return "Expired";
if (daysRemaining === 0) return "Expires today";
if (daysRemaining === 1) return "1 day remaining";
if (daysRemaining < 30) return `${daysRemaining} days remaining`;
if (daysRemaining < 60) return "1 month remaining";
return `${Math.floor(daysRemaining / 30)} months remaining`;
};
const meterTone = (pct: number): "default" | "warning" | "over" => {
if (pct > 100) return "over";
if (pct >= 80) return "warning";
@ -79,7 +68,6 @@ const buildMeters = (data: RemainingUsage | null): MeterData[] => {
* design's Spend / API-request meters are intentionally omitted.
*/
export default function SidebarUsageCard({ accessToken, collapsed, onExpandRail }: SidebarUsageCardProps) {
const disableUsageIndicator = useDisableUsageIndicator();
const licenseInfo = useLicenseInfo(accessToken).data ?? null;
const { data: usageData, isLoading } = useQuery(remainingUsersQuery(accessToken));
const data = usageData ?? null;
@ -87,7 +75,7 @@ export default function SidebarUsageCard({ accessToken, collapsed, onExpandRail
const hasData = data !== null && (data.total_users !== null || data.total_teams !== null);
const noUsableData = !isLoading && !hasData;
const noLicensedUsage = !licenseInfo?.has_license || noUsableData;
if (disableUsageIndicator || !accessToken || noLicensedUsage) {
if (!accessToken || noLicensedUsage) {
return null;
}
@ -104,8 +92,7 @@ export default function SidebarUsageCard({ accessToken, collapsed, onExpandRail
);
}
const daysUntilExpiration = licenseInfo?.expiration_date ? getDaysUntilExpiration(licenseInfo.expiration_date) : null;
const subtitle = licenseInfo?.expiration_date ? formatExpiration(daysUntilExpiration) : "Active plan";
const subtitle = licenseInfo?.expiration_date ? formatExpirationStatus(licenseInfo.expiration_date) : "Active plan";
const meters = buildMeters(data);
return (

View file

@ -1,193 +0,0 @@
import React from "react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import UsageIndicator from "./UsageIndicator";
vi.mock("./networking", () => ({
getRemainingUsers: vi.fn(),
getLicenseInfo: vi.fn().mockResolvedValue(null),
}));
vi.mock("@/app/(dashboard)/hooks/useDisableUsageIndicator", () => ({
useDisableUsageIndicator: vi.fn(() => false),
}));
import { getRemainingUsers } from "./networking";
const mockGetRemainingUsers = vi.mocked(getRemainingUsers);
const renderWithClient = (ui: React.ReactElement) => {
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
return render(<QueryClientProvider client={queryClient}>{ui}</QueryClientProvider>);
};
const DEFAULT_USAGE_DATA = {
total_users: 100,
total_users_used: 1,
total_users_remaining: 99,
total_teams: null,
total_teams_used: 0,
total_teams_remaining: null,
};
describe("UsageIndicator", () => {
beforeEach(() => {
vi.clearAllMocks();
mockGetRemainingUsers.mockResolvedValue(DEFAULT_USAGE_DATA);
});
it("should render when given access token and usage data loads", async () => {
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await screen.findByText("Usage");
expect(screen.getByText("Usage")).toBeInTheDocument();
});
it("should not show Near limit when users usage is below 80% (1/100 -> 1%)", async () => {
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await screen.findByText("Usage");
expect(screen.queryByText("Near limit")).not.toBeInTheDocument();
});
it("should render nothing when both total_users and total_teams are null", async () => {
mockGetRemainingUsers.mockResolvedValue({
total_users: null,
total_teams: null,
total_users_used: 520,
total_teams_used: 4,
total_teams_remaining: null,
total_users_remaining: null,
});
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await waitFor(() => {
expect(screen.queryByText("Usage")).not.toBeInTheDocument();
expect(screen.queryByText("Loading...")).not.toBeInTheDocument();
});
});
it("should show Near limit for Teams when at 80% usage (4/5)", async () => {
mockGetRemainingUsers.mockResolvedValue({
total_users: null,
total_users_used: 0,
total_users_remaining: null,
total_teams: 5,
total_teams_used: 4,
total_teams_remaining: 1,
});
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await screen.findByText("Usage");
expect(screen.getByText("Teams")).toBeInTheDocument();
expect(screen.getByText("Near limit")).toBeInTheDocument();
});
it("should show Over limit for Users when usage exceeds 100% (105/100)", async () => {
mockGetRemainingUsers.mockResolvedValue({
total_users: 100,
total_users_used: 105,
total_users_remaining: -5,
total_teams: null,
total_teams_used: 0,
total_teams_remaining: null,
});
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await screen.findByText("Usage");
expect(screen.getByText("Users")).toBeInTheDocument();
expect(screen.getByText("Over limit")).toBeInTheDocument();
});
it("should show Over limit for Teams when usage exceeds 100%", async () => {
mockGetRemainingUsers.mockResolvedValue({
total_users: null,
total_users_used: 0,
total_users_remaining: null,
total_teams: 10,
total_teams_used: 12,
total_teams_remaining: -2,
});
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await screen.findByText("Usage");
expect(screen.getByText("Teams")).toBeInTheDocument();
expect(screen.getByText("Over limit")).toBeInTheDocument();
});
it("should render nothing when accessToken is null", () => {
renderWithClient(<UsageIndicator accessToken={null} width={220} />);
expect(mockGetRemainingUsers).not.toHaveBeenCalled();
expect(screen.queryByText("Usage")).not.toBeInTheDocument();
});
it("should render nothing when disableUsageIndicator is true", async () => {
const { useDisableUsageIndicator } = await import("@/app/(dashboard)/hooks/useDisableUsageIndicator");
(useDisableUsageIndicator as ReturnType<typeof vi.fn>).mockReturnValue(true);
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await waitFor(() => {
expect(screen.queryByText("Usage")).not.toBeInTheDocument();
});
(useDisableUsageIndicator as ReturnType<typeof vi.fn>).mockReturnValue(false);
});
it("should show Loading while fetching", () => {
mockGetRemainingUsers.mockImplementation(() => new Promise(() => {}));
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
expect(screen.getByText("Loading...")).toBeInTheDocument();
});
it("should show error message when fetch fails", async () => {
const consoleSpy = vi.spyOn(console, "error").mockImplementation(() => {});
mockGetRemainingUsers.mockRejectedValue(new Error("Network error"));
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
expect(await screen.findByText("Failed to load usage data")).toBeInTheDocument();
consoleSpy.mockRestore();
});
it("should minimize when user clicks minimize button", async () => {
const user = userEvent.setup();
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await screen.findByText("Usage");
const minimizeButton = screen.getByTitle("Minimize");
await user.click(minimizeButton);
expect(screen.queryByText("Users")).not.toBeInTheDocument();
expect(screen.getByTitle("Show usage details")).toBeInTheDocument();
});
it("should restore from minimized when user clicks restore button", async () => {
const user = userEvent.setup();
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
await screen.findByText("Usage");
await user.click(screen.getByTitle("Minimize"));
await user.click(screen.getByTitle("Show usage details"));
expect(screen.getByText("Usage")).toBeInTheDocument();
expect(screen.getByText("Users")).toBeInTheDocument();
});
});

View file

@ -1,681 +0,0 @@
import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator";
import { Badge } from "@tremor/react";
import {
AlertTriangle,
Calendar,
ChevronDown,
ChevronUp,
Loader2,
Minus,
TrendingUp,
UserCheck,
Users,
} from "lucide-react";
import { useEffect, useState } from "react";
import { getRemainingUsers } from "./networking";
import { cn } from "@/lib/cva.config";
import { getDaysUntilExpiration } from "@/utils/licenseUtils";
import { useLicenseInfo } from "@/app/(dashboard)/hooks/license/useLicenseInfo";
interface UsageIndicatorProps {
accessToken: string | null;
width: number;
}
interface UsageData {
total_users: number | null;
total_users_used: number;
total_users_remaining: number | null;
total_teams: number | null;
total_teams_used: number;
total_teams_remaining: number | null;
}
// Format expiration for display
const formatExpirationDisplay = (daysRemaining: number | null): string => {
if (daysRemaining === null) return "No expiration";
if (daysRemaining < 0) return "Expired";
if (daysRemaining === 0) return "Expires today";
if (daysRemaining === 1) return "1 day remaining";
if (daysRemaining < 30) return `${daysRemaining} days remaining`;
if (daysRemaining < 60) return "1 month remaining";
const months = Math.floor(daysRemaining / 30);
return `${months} months remaining`;
};
export default function UsageIndicator({ accessToken, width = 220 }: UsageIndicatorProps) {
const disableUsageIndicator = useDisableUsageIndicator();
const [isExpanded, setIsExpanded] = useState(false);
const [isMinimized, setIsMinimized] = useState(false);
const [data, setData] = useState<UsageData | null>(null);
const [isLoading, setIsLoading] = useState(false);
const [error, setError] = useState<string | null>(null);
const licenseInfo = useLicenseInfo(accessToken).data ?? null;
useEffect(() => {
const fetchData = async () => {
if (!accessToken) return;
setIsLoading(true);
setError(null);
try {
const usageResult = await getRemainingUsers(accessToken);
setData(usageResult);
} catch (err) {
console.error("Failed to fetch usage data:", err);
setError("Failed to load usage data");
} finally {
setIsLoading(false);
}
};
fetchData();
}, [accessToken]);
// Calculate license expiration metrics
const daysUntilExpiration = licenseInfo?.expiration_date ? getDaysUntilExpiration(licenseInfo.expiration_date) : null;
const isLicenseExpired = daysUntilExpiration !== null && daysUntilExpiration < 0;
const isLicenseExpiringSoon = daysUntilExpiration !== null && daysUntilExpiration >= 0 && daysUntilExpiration < 30;
// Calculate derived values from data
const getUsageMetrics = (data: UsageData | null) => {
if (!data) {
return {
isOverLimit: false,
isNearLimit: false,
usagePercentage: 0,
userMetrics: {
isOverLimit: false,
isNearLimit: false,
usagePercentage: 0,
},
teamMetrics: {
isOverLimit: false,
isNearLimit: false,
usagePercentage: 0,
},
};
}
// User metrics
const userUsagePercentage = data.total_users ? (data.total_users_used / data.total_users) * 100 : 0;
const userIsOverLimit = userUsagePercentage > 100;
const userIsNearLimit = userUsagePercentage >= 80 && userUsagePercentage <= 100;
// Team metrics
const teamUsagePercentage = data.total_teams ? (data.total_teams_used / data.total_teams) * 100 : 0;
const teamIsOverLimit = teamUsagePercentage > 100;
const teamIsNearLimit = teamUsagePercentage >= 80 && teamUsagePercentage <= 100;
// Combined status (worst case scenario)
const isOverLimit = userIsOverLimit || teamIsOverLimit;
const isNearLimit = (userIsNearLimit || teamIsNearLimit) && !isOverLimit;
const usagePercentage = Math.max(userUsagePercentage, teamUsagePercentage);
return {
isOverLimit,
isNearLimit,
usagePercentage,
userMetrics: {
isOverLimit: userIsOverLimit,
isNearLimit: userIsNearLimit,
usagePercentage: userUsagePercentage,
},
teamMetrics: {
isOverLimit: teamIsOverLimit,
isNearLimit: teamIsNearLimit,
usagePercentage: teamUsagePercentage,
},
};
};
const { isOverLimit, isNearLimit, usagePercentage, userMetrics, teamMetrics } = getUsageMetrics(data);
// Include license status in overall status
const hasAnyIssue = isOverLimit || isNearLimit || isLicenseExpired || isLicenseExpiringSoon;
const hasError = isOverLimit || isLicenseExpired;
const hasWarning = (isNearLimit || isLicenseExpiringSoon) && !hasError;
const getStatusColor = () => {
if (hasError) return "red";
if (hasWarning) return "yellow";
return "green";
};
const getStatusIcon = () => {
if (hasError) return <AlertTriangle className="h-3 w-3" />;
if (hasWarning) return <TrendingUp className="h-3 w-3" />;
return null;
};
// Minimized view - just a small restore button
const MinimizedView = () => {
return (
<div className="px-3 py-1" style={{ maxWidth: `${width}px` }}>
<button
onClick={() => setIsMinimized(false)}
className={cn(
"flex items-center gap-2 text-xs text-gray-400 hover:text-gray-600 transition-colors p-1 rounded-sm w-full",
hasError && "text-red-400 hover:text-red-600",
hasWarning && "text-yellow-500 hover:text-yellow-700",
)}
title="Show usage details"
>
<Users className="h-3 w-3 shrink-0" />
{hasAnyIssue && <span className="shrink-0">{getStatusIcon()}</span>}
<div className="flex items-center gap-1 truncate">
{data && data.total_users !== null && (
<span className="shrink-0">
U:{data.total_users_used}/{data.total_users}
</span>
)}
{data && data.total_teams !== null && (
<span className="shrink-0">
T:{data.total_teams_used}/{data.total_teams}
</span>
)}
{licenseInfo?.expiration_date && daysUntilExpiration !== null && (
<span
className={cn(
"shrink-0",
isLicenseExpired && "text-red-500",
isLicenseExpiringSoon && "text-yellow-500",
)}
>
{daysUntilExpiration < 0 ? "Exp!" : `${daysUntilExpiration}d`}
</span>
)}
{!data ||
(data.total_users === null && data.total_teams === null && !licenseInfo && (
<span className="truncate">Usage</span>
))}
</div>
</button>
</div>
);
};
// Sidebar/nav style component
const NavStyleView = () => {
if (isMinimized) {
return <MinimizedView />;
}
if (isLoading) {
return (
<div className="flex items-center gap-3 px-3 py-2 text-gray-500" style={{ maxWidth: `${width}px` }}>
<Loader2 className="h-4 w-4 animate-spin shrink-0" />
<span className="text-sm truncate">Loading...</span>
</div>
);
}
if (error || !data) {
return (
<div
className="flex items-center justify-between gap-3 px-3 py-2 text-gray-400 group"
style={{ maxWidth: `${width}px` }}
>
<div className="flex items-center gap-3 min-w-0 flex-1">
<Users className="h-4 w-4 shrink-0" />
<span className="text-sm truncate">{error || "No data"}</span>
</div>
<button
onClick={() => setIsMinimized(true)}
className="opacity-0 group-hover:opacity-100 p-0.5 hover:bg-gray-100 rounded-sm transition-all shrink-0"
title="Minimize"
>
<Minus className="h-3 w-3" />
</button>
</div>
);
}
return (
<div className="px-3 py-2 group" style={{ maxWidth: `${width}px` }}>
{/* Main nav item style */}
<div className="flex items-center justify-between">
<button
onClick={() => setIsExpanded(!isExpanded)}
className={cn(
"flex items-center gap-3 text-left hover:bg-gray-50 rounded-md px-0 py-1 transition-colors flex-1 min-w-0",
hasError && "text-red-600",
hasWarning && "text-yellow-600",
)}
>
<Users className="h-4 w-4 shrink-0" />
<span className="text-sm font-medium truncate">Usage Status</span>
{hasAnyIssue && (
<Badge color={getStatusColor()} className="text-xs px-1.5 py-0.5 shrink-0">
{getStatusIcon()}
</Badge>
)}
{isExpanded ? (
<ChevronUp className="h-3 w-3 text-gray-400 ml-auto shrink-0" />
) : (
<ChevronDown className="h-3 w-3 text-gray-400 ml-auto shrink-0" />
)}
</button>
{/* Minimize button */}
<button
onClick={() => setIsMinimized(true)}
className="opacity-0 group-hover:opacity-100 p-0.5 hover:bg-gray-100 rounded-sm transition-all ml-1 shrink-0"
title="Minimize"
>
<Minus className="h-3 w-3 text-gray-400" />
</button>
</div>
{/* Expanded details - simple and compact */}
{isExpanded && (
<div className="mt-2 pl-7 text-xs text-gray-600 space-y-3">
{/* License expiration section */}
{licenseInfo?.has_license && licenseInfo.expiration_date && (
<div>
<div className="mb-1 flex items-center gap-1">
<Calendar className="h-3 w-3" />
<span className="font-medium">License</span>
</div>
<div
className={cn(
"flex items-center gap-1 text-xs",
isLicenseExpired && "text-red-600",
isLicenseExpiringSoon && "text-yellow-600",
)}
>
{isLicenseExpired ? (
<AlertTriangle className="h-3 w-3" />
) : isLicenseExpiringSoon ? (
<TrendingUp className="h-3 w-3" />
) : null}
<span className="truncate">{formatExpirationDisplay(daysUntilExpiration)}</span>
</div>
</div>
)}
{/* Users section */}
{data.total_users !== null && (
<div>
<div className="mb-1 flex items-center gap-1">
<Users className="h-3 w-3" />
<span className="font-medium">
{data.total_users_used}/{data.total_users}
</span>
<span className="text-gray-500">users</span>
</div>
{/* User progress bar */}
<div className="w-full bg-gray-200 rounded-full h-1 mb-1">
<div
className={cn(
"h-1 rounded-full transition-all duration-300",
userMetrics.isOverLimit && "bg-red-500",
userMetrics.isNearLimit && "bg-yellow-500",
!userMetrics.isOverLimit && !userMetrics.isNearLimit && "bg-green-500",
)}
style={{ width: `${Math.min(userMetrics.usagePercentage, 100)}%` }}
/>
</div>
{(userMetrics.isOverLimit || userMetrics.isNearLimit) && (
<div
className={cn(
"flex items-center gap-1 text-xs",
userMetrics.isOverLimit && "text-red-600",
userMetrics.isNearLimit && "text-yellow-600",
)}
>
{userMetrics.isOverLimit ? (
<AlertTriangle className="h-3 w-3" />
) : (
<TrendingUp className="h-3 w-3" />
)}
<span className="truncate">Users {userMetrics.isOverLimit ? "Over Limit" : "Near Limit"}</span>
</div>
)}
</div>
)}
{/* Teams section */}
{data.total_teams !== null && (
<div>
<div className="mb-1 flex items-center gap-1">
<UserCheck className="h-3 w-3" />
<span className="font-medium">
{data.total_teams_used}/{data.total_teams}
</span>
<span className="text-gray-500">teams</span>
</div>
{/* Team progress bar */}
<div className="w-full bg-gray-200 rounded-full h-1 mb-1">
<div
className={cn(
"h-1 rounded-full transition-all duration-300",
teamMetrics.isOverLimit && "bg-red-500",
teamMetrics.isNearLimit && "bg-yellow-500",
!teamMetrics.isOverLimit && !teamMetrics.isNearLimit && "bg-green-500",
)}
style={{ width: `${Math.min(teamMetrics.usagePercentage, 100)}%` }}
/>
</div>
{(teamMetrics.isOverLimit || teamMetrics.isNearLimit) && (
<div
className={cn(
"flex items-center gap-1 text-xs",
teamMetrics.isOverLimit && "text-red-600",
teamMetrics.isNearLimit && "text-yellow-600",
)}
>
{teamMetrics.isOverLimit ? (
<AlertTriangle className="h-3 w-3" />
) : (
<TrendingUp className="h-3 w-3" />
)}
<span className="truncate">Teams {teamMetrics.isOverLimit ? "Over Limit" : "Near Limit"}</span>
</div>
)}
</div>
)}
</div>
)}
</div>
);
};
// Optimized CardStyleView for 220px width
const CardStyleView = () => {
if (isMinimized) {
return (
<button
onClick={() => setIsMinimized(false)}
className={cn(
"bg-white border border-gray-200 rounded-lg shadow-xs p-3 hover:shadow-md transition-all w-full",
)}
title="Show usage details"
>
<div className="flex items-center gap-2">
<Users className="h-4 w-4 shrink-0" />
{hasAnyIssue && <span className="shrink-0">{getStatusIcon()}</span>}
<div className="flex items-center gap-2 text-sm font-medium truncate">
{data && data.total_users !== null && (
<span
className={cn(
"shrink-0 px-1.5 py-0.5 rounded-sm text-xs border",
userMetrics.isOverLimit && "bg-red-50 text-red-700 border-red-200",
userMetrics.isNearLimit && "bg-yellow-50 text-yellow-700 border-yellow-200",
!userMetrics.isOverLimit && !userMetrics.isNearLimit && "bg-gray-50 text-gray-700 border-gray-200",
)}
>
U: {data.total_users_used}/{data.total_users}
</span>
)}
{data && data.total_teams !== null && (
<span
className={cn(
"shrink-0 px-1.5 py-0.5 rounded-sm text-xs border",
teamMetrics.isOverLimit && "bg-red-50 text-red-700 border-red-200",
teamMetrics.isNearLimit && "bg-yellow-50 text-yellow-700 border-yellow-200",
!teamMetrics.isOverLimit && !teamMetrics.isNearLimit && "bg-gray-50 text-gray-700 border-gray-200",
)}
>
T: {data.total_teams_used}/{data.total_teams}
</span>
)}
{licenseInfo?.expiration_date && daysUntilExpiration !== null && (
<span
className={cn(
"shrink-0 px-1.5 py-0.5 rounded-sm text-xs border",
isLicenseExpired && "bg-red-50 text-red-700 border-red-200",
isLicenseExpiringSoon && "bg-yellow-50 text-yellow-700 border-yellow-200",
!isLicenseExpired && !isLicenseExpiringSoon && "bg-gray-50 text-gray-700 border-gray-200",
)}
>
{daysUntilExpiration < 0 ? "Exp!" : `${daysUntilExpiration}d`}
</span>
)}
{!data ||
(data.total_users === null && data.total_teams === null && !licenseInfo && (
<span className="truncate">Usage</span>
))}
</div>
</div>
</button>
);
}
if (isLoading) {
return (
<div className="bg-white border border-gray-200 rounded-lg shadow-xs p-4 w-full">
<div className="flex items-center justify-center gap-2 py-2">
<Loader2 className="h-4 w-4 animate-spin" />
<span className="text-sm text-gray-500 truncate">Loading...</span>
</div>
</div>
);
}
if (error || !data) {
return (
<div className="bg-white border border-gray-200 rounded-lg shadow-xs p-4 group w-full">
<div className="flex items-center justify-between gap-2">
<div className="flex-1 min-w-0">
<span className="text-sm text-gray-500 truncate block">{error || "No data"}</span>
</div>
<button
onClick={() => setIsMinimized(true)}
className="opacity-0 group-hover:opacity-100 p-1 hover:bg-gray-100 rounded-sm transition-all shrink-0"
title="Minimize"
>
<Minus className="h-3 w-3 text-gray-400" />
</button>
</div>
</div>
);
}
return (
<div className={cn("bg-white border rounded-lg shadow-xs p-3 transition-all duration-200 group w-full")}>
<div className="flex items-center justify-between gap-2 mb-3">
<div className="flex items-center gap-2 min-w-0 flex-1">
<Users className="h-4 w-4 shrink-0" />
<span className="font-medium text-sm truncate">Usage</span>
</div>
<button
onClick={() => setIsMinimized(true)}
className="opacity-0 group-hover:opacity-100 p-1 hover:bg-gray-100 rounded-sm transition-all shrink-0"
title="Minimize"
>
<Minus className="h-3 w-3 text-gray-400" />
</button>
</div>
{/* Compact stats optimized for 220px */}
<div className="space-y-3 text-sm">
{/* License expiration section */}
{licenseInfo?.has_license && licenseInfo.expiration_date && (
<div
className={cn(
"space-y-1 border rounded-md p-2",
isLicenseExpired && "border-red-200 bg-red-50",
isLicenseExpiringSoon && "border-yellow-200 bg-yellow-50",
)}
>
<div className="flex items-center gap-2 text-xs text-gray-600 mb-1">
<Calendar className="h-3 w-3" />
<span className="font-medium">License</span>
<span
className={cn(
"ml-1 px-1.5 py-0.5 rounded-sm border",
isLicenseExpired && "bg-red-50 text-red-700 border-red-200",
isLicenseExpiringSoon && "bg-yellow-50 text-yellow-700 border-yellow-200",
!isLicenseExpired && !isLicenseExpiringSoon && "bg-gray-50 text-gray-600 border-gray-200",
)}
>
{isLicenseExpired ? "Expired" : isLicenseExpiringSoon ? "Expiring soon" : "OK"}
</span>
</div>
<div className="flex justify-between items-center">
<span className="text-gray-600 text-xs">Status:</span>
<span
className={cn(
"font-medium text-right",
isLicenseExpired && "text-red-600",
isLicenseExpiringSoon && "text-yellow-600",
)}
>
{formatExpirationDisplay(daysUntilExpiration)}
</span>
</div>
{licenseInfo.license_type && (
<div className="flex justify-between items-center">
<span className="text-gray-600 text-xs">Type:</span>
<span className="font-medium text-right capitalize">{licenseInfo.license_type}</span>
</div>
)}
</div>
)}
{/* Users section */}
{data.total_users !== null && (
<div
className={cn(
"space-y-1 border rounded-md p-2",
userMetrics.isOverLimit && "border-red-200 bg-red-50",
userMetrics.isNearLimit && "border-yellow-200 bg-yellow-50",
)}
>
<div className="flex items-center gap-2 text-xs text-gray-600 mb-1">
<Users className="h-3 w-3" />
<span className="font-medium">Users</span>
<span
className={cn(
"ml-1 px-1.5 py-0.5 rounded-sm border",
userMetrics.isOverLimit && "bg-red-50 text-red-700 border-red-200",
userMetrics.isNearLimit && "bg-yellow-50 text-yellow-700 border-yellow-200",
!userMetrics.isOverLimit && !userMetrics.isNearLimit && "bg-gray-50 text-gray-600 border-gray-200",
)}
>
{userMetrics.isOverLimit ? "Over limit" : userMetrics.isNearLimit ? "Near limit" : "OK"}
</span>
</div>
<div className="flex justify-between items-center">
<span className="text-gray-600 text-xs">Used:</span>
<span className="font-medium text-right">
{data.total_users_used}/{data.total_users}
</span>
</div>
<div className="flex justify-between items-center">
<span className="text-gray-600 text-xs">Remaining:</span>
<span
className={cn(
"font-medium text-right",
userMetrics.isOverLimit && "text-red-600",
userMetrics.isNearLimit && "text-yellow-600",
)}
>
{data.total_users_remaining}
</span>
</div>
<div className="flex justify-between items-center">
<span className="text-gray-600 text-xs">Usage:</span>
<span className="font-medium text-right">{Math.round(userMetrics.usagePercentage)}%</span>
</div>
{/* User progress bar */}
<div className="w-full bg-gray-200 rounded-full h-2">
<div
className={cn(
"h-2 rounded-full transition-all duration-300",
userMetrics.isOverLimit && "bg-red-500",
userMetrics.isNearLimit && "bg-yellow-500",
!userMetrics.isOverLimit && !userMetrics.isNearLimit && "bg-green-500",
)}
style={{ width: `${Math.min(userMetrics.usagePercentage, 100)}%` }}
/>
</div>
</div>
)}
{/* Teams section */}
{data.total_teams !== null && (
<div
className={cn(
"space-y-1 border rounded-md p-2",
teamMetrics.isOverLimit && "border-red-200 bg-red-50",
teamMetrics.isNearLimit && "border-yellow-200 bg-yellow-50",
)}
>
<div className="flex items-center gap-2 text-xs text-gray-600 mb-1">
<UserCheck className="h-3 w-3" />
<span className="font-medium">Teams</span>
<span
className={cn(
"ml-1 px-1.5 py-0.5 rounded-sm border",
teamMetrics.isOverLimit && "bg-red-50 text-red-700 border-red-200",
teamMetrics.isNearLimit && "bg-yellow-50 text-yellow-700 border-yellow-200",
!teamMetrics.isOverLimit && !teamMetrics.isNearLimit && "bg-gray-50 text-gray-600 border-gray-200",
)}
>
{teamMetrics.isOverLimit ? "Over limit" : teamMetrics.isNearLimit ? "Near limit" : "OK"}
</span>
</div>
<div className="flex justify-between items-center">
<span className="text-gray-600 text-xs">Used:</span>
<span className="font-medium text-right">
{data.total_teams_used}/{data.total_teams}
</span>
</div>
<div className="flex justify-between items-center">
<span className="text-gray-600 text-xs">Remaining:</span>
<span
className={cn(
"font-medium text-right",
teamMetrics.isOverLimit && "text-red-600",
teamMetrics.isNearLimit && "text-yellow-600",
)}
>
{data.total_teams_remaining}
</span>
</div>
<div className="flex justify-between items-center">
<span className="text-gray-600 text-xs">Usage:</span>
<span className="font-medium text-right">{Math.round(teamMetrics.usagePercentage)}%</span>
</div>
{/* Team progress bar */}
<div className="w-full bg-gray-200 rounded-full h-2">
<div
className={cn(
"h-2 rounded-full transition-all duration-300",
teamMetrics.isOverLimit && "bg-red-500",
teamMetrics.isNearLimit && "bg-yellow-500",
!teamMetrics.isOverLimit && !teamMetrics.isNearLimit && "bg-green-500",
)}
style={{ width: `${Math.min(teamMetrics.usagePercentage, 100)}%` }}
/>
</div>
</div>
)}
</div>
</div>
);
};
// Don't render anything if disabled, no access token, or if both total_users and total_teams are null
if (disableUsageIndicator || !accessToken || (data?.total_users === null && data?.total_teams === null)) {
return null;
}
// Fixed positioning with proper spacing from edges
return (
<div className="fixed bottom-4 left-4 z-50" style={{ width: `${Math.min(width, 220)}px` }}>
<CardStyleView />
</div>
);
}

View file

@ -0,0 +1,43 @@
import { describe, it, expect, beforeEach, afterEach, vi } from "vitest";
// Regression for the chat sidebar / first-message navigation under SERVER_ROOT_PATH.
// getChatRoutes() must read the server root path at call time. The previous
// module-level `CHAT_ROUTES` captured it once at import, before the UI-config
// bootstrap runs setServerRootPath, so every chat route was permanently
// unprefixed and router.push() navigated to a 404 (which, mid-stream, also
// aborted the first message). These tests deliberately apply the root path
// AFTER importing the module so a frozen-at-import implementation fails.
describe("getChatRoutes under server_root_path", () => {
beforeEach(() => {
vi.resetModules();
vi.stubEnv("NODE_ENV", "test");
});
afterEach(() => {
vi.unstubAllEnvs();
});
it("reflects a server root path applied after the module is loaded", async () => {
const { getChatRoutes } = await import("./ChatShell");
const { setServerRootPath } = await import("@/lib/serverRootPath");
setServerRootPath("/gw");
const routes = getChatRoutes();
expect(routes.chats).toBe("/gw/ui/chat");
expect(routes.integrations).toBe("/gw/ui/chat/integrations");
expect(routes.credentials).toBe("/gw/ui/chat/credentials");
expect(routes.apiKeys).toBe("/gw/ui/chat/api-keys");
expect(routes.usage).toBe("/gw/ui/chat/usage");
});
it("builds /ui-rooted paths when no server root path is set", async () => {
const { getChatRoutes } = await import("./ChatShell");
const { setServerRootPath } = await import("@/lib/serverRootPath");
setServerRootPath("/");
expect(getChatRoutes().chats).toBe("/ui/chat");
expect(getChatRoutes().integrations).toBe("/ui/chat/integrations");
});
});

View file

@ -9,14 +9,16 @@ import { migratedHref } from "@/utils/migratedPages";
import { useChatShell } from "@/contexts/ChatShellContext";
import ConversationList from "./ConversationList";
const CHAT_BASE = migratedHref("chat");
export const CHAT_ROUTES = {
chats: CHAT_BASE,
integrations: `${CHAT_BASE}/integrations`,
credentials: `${CHAT_BASE}/credentials`,
apiKeys: `${CHAT_BASE}/api-keys`,
usage: `${CHAT_BASE}/usage`,
};
export function getChatRoutes() {
const base = migratedHref("chat");
return {
chats: base,
integrations: `${base}/integrations`,
credentials: `${base}/credentials`,
apiKeys: `${base}/api-keys`,
usage: `${base}/usage`,
};
}
function stripTrailingSlash(path: string): string {
return path.length > 1 ? path.replace(/\/+$/, "") : path;
@ -54,7 +56,8 @@ const ChatShell: React.FC<ChatShellProps> = ({ children }) => {
const pathname = stripTrailingSlash(usePathname() ?? "");
const { conversations, activeConversationId, deleteConversation, renameConversation } = useChatShell();
const isChatsRoute = pathname === CHAT_ROUTES.chats;
const routes = getChatRoutes();
const isChatsRoute = pathname === routes.chats;
return (
<div className="flex h-full w-full flex-col bg-background overflow-hidden">
@ -73,7 +76,7 @@ const ChatShell: React.FC<ChatShellProps> = ({ children }) => {
<div className="flex flex-1 min-h-0 overflow-hidden">
<div className="shrink-0 bg-sidebar border-sidebar-border border-r flex flex-col overflow-hidden w-[260px]">
<div className="px-2 pt-3 pb-1 shrink-0">
<Button onClick={() => router.push(CHAT_ROUTES.chats)} className="w-full justify-start gap-2.5">
<Button onClick={() => router.push(routes.chats)} className="w-full justify-start gap-2.5">
<Plus className="h-4 w-4" />
New Chat
</Button>
@ -85,32 +88,32 @@ const ChatShell: React.FC<ChatShellProps> = ({ children }) => {
<NavItem
icon={<MessageSquare className="h-4 w-4" />}
label="Chats"
onClick={() => router.push(CHAT_ROUTES.chats)}
onClick={() => router.push(routes.chats)}
active={isChatsRoute}
/>
<NavItem
icon={<LayoutGrid className="h-4 w-4" />}
label="Integrations"
onClick={() => router.push(CHAT_ROUTES.integrations)}
active={pathname === CHAT_ROUTES.integrations}
onClick={() => router.push(routes.integrations)}
active={pathname === routes.integrations}
/>
<NavItem
icon={<KeyRound className="h-4 w-4" />}
label="Credentials"
onClick={() => router.push(CHAT_ROUTES.credentials)}
active={pathname === CHAT_ROUTES.credentials}
onClick={() => router.push(routes.credentials)}
active={pathname === routes.credentials}
/>
<NavItem
icon={<Lock className="h-4 w-4" />}
label="API Keys"
onClick={() => router.push(CHAT_ROUTES.apiKeys)}
active={pathname === CHAT_ROUTES.apiKeys}
onClick={() => router.push(routes.apiKeys)}
active={pathname === routes.apiKeys}
/>
<NavItem
icon={<BarChart3 className="h-4 w-4" />}
label="Usage"
onClick={() => router.push(CHAT_ROUTES.usage)}
active={pathname === CHAT_ROUTES.usage}
onClick={() => router.push(routes.usage)}
active={pathname === routes.usage}
/>
</div>
@ -120,10 +123,10 @@ const ChatShell: React.FC<ChatShellProps> = ({ children }) => {
<ConversationList
conversations={conversations}
activeConversationId={activeConversationId}
onSelect={(id) => router.push(`${CHAT_ROUTES.chats}?id=${id}`)}
onSelect={(id) => router.push(`${routes.chats}?id=${id}`)}
onDelete={(id) => {
deleteConversation(id);
if (id === activeConversationId) router.push(CHAT_ROUTES.chats);
if (id === activeConversationId) router.push(routes.chats);
}}
onRename={renameConversation}
/>

View file

@ -52,6 +52,7 @@ export function useChatHistory(
): {
conversations: Conversation[];
activeConversation: Conversation | null;
currentActiveId: string | null;
storageUnavailable: boolean;
staleId: boolean;
createConversation: (model: string) => string;
@ -208,6 +209,7 @@ export function useChatHistory(
return {
conversations,
activeConversation,
currentActiveId,
storageUnavailable,
staleId,
createConversation,

View file

@ -72,6 +72,7 @@ import {
rolesAllowedToViewWriteScopedPages,
rolesWithWriteAccess,
} from "../utils/roles";
import BetaBadge from "./BetaBadge";
import NewBadge from "./common_components/NewBadge";
import type { Organization } from "./networking";
import SidebarAccountMenu from "./SidebarAccountMenu/SidebarAccountMenu";
@ -211,7 +212,7 @@ const menuGroups: MenuGroup[] = [
page: "projects",
label: (
<span className="flex items-center gap-2">
Projects <NewBadge />
Projects <BetaBadge />
</span>
),
icon: <Folder {...ICON} />,

View file

@ -252,7 +252,7 @@ describe("Navbar", () => {
expect(screen.queryByRole("button", { name: /^notifications$/i })).not.toBeInTheDocument();
});
it("should handle hide new features toggle", async () => {
it("should handle hide new feature indicators toggle", async () => {
const user = userEvent.setup();
// Initially disabled

View file

@ -179,7 +179,7 @@ export interface PromptSpec {
export interface PromptTemplateBase {
litellm_prompt_id: string;
content: string;
metadata?: Record<string, any> | null;
metadata?: Record<string, unknown> | null;
}
interface PromptInfoResponse {
@ -6227,6 +6227,7 @@ export const applyGuardrail = async (
text: string,
language?: string | null,
entities?: string[] | null,
metadata?: Record<string, unknown> | null,
) => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/guardrails/apply_guardrail` : `/guardrails/apply_guardrail`;
@ -6244,6 +6245,10 @@ export const applyGuardrail = async (
requestBody.entities = entities;
}
if (metadata != null) {
requestBody.metadata = metadata;
}
const response = await fetch(url, {
method: "POST",
headers: {

View file

@ -57,12 +57,13 @@ export function ChatShellProvider({
children,
}: ChatShellProviderProps) {
const searchParams = useSearchParams();
const activeConversationId = searchParams.get("id");
const urlConversationId = searchParams.get("id");
const [selectedMCPServers, setSelectedMCPServers] = useState<string[]>([]);
const {
conversations,
activeConversation,
currentActiveId,
storageUnavailable,
staleId,
createConversation,
@ -71,7 +72,7 @@ export function ChatShellProvider({
truncateFromMessage,
deleteConversation,
renameConversation,
} = useChatHistory(activeConversationId, userId);
} = useChatHistory(urlConversationId, userId);
return (
<ChatShellContext.Provider
@ -85,7 +86,7 @@ export function ChatShellProvider({
setSelectedMCPServers,
conversations,
activeConversation,
activeConversationId,
activeConversationId: currentActiveId,
storageUnavailable,
staleId,
createConversation,

View file

@ -6544,7 +6544,7 @@ export interface paths {
* Parameters:
* - duration: Optional[str] - Specify the length of time the token is valid for. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
* - key_alias: Optional[str] - User defined key alias
* - key: Optional[str] - User defined key value. If not set, a 16-digit unique sk-key is created for you.
* - key: Optional[str] - User defined key value. Must start with 'sk-' and be at least 16 characters long. If not set, a 16-digit unique sk-key is created for you.
* - team_id: Optional[str] - The team id of the key
* - user_id: Optional[str] - The user id of the key
* - agent_id: Optional[str] - The agent id associated with the key.
@ -6765,7 +6765,7 @@ export interface paths {
* - data: Optional[RegenerateKeyRequest] - Request body containing optional parameters to update
* - key: Optional[str] - The key to regenerate.
* - new_master_key: Optional[str] - The new master key to use, if key is the master key.
* - new_key: Optional[str] - The new key to use, if key is not the master key. If both set, new_master_key will be used.
* - new_key: Optional[str] - The new key to use, if key is not the master key. Must start with 'sk-' and be at least 16 characters long. If both set, new_master_key will be used.
* - key_alias: Optional[str] - User-friendly key alias
* - user_id: Optional[str] - User ID associated with key
* - team_id: Optional[str] - Team ID associated with key
@ -6834,7 +6834,7 @@ export interface paths {
* Parameters:
* - duration: Optional[str] - Specify the length of time the token is valid for. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
* - key_alias: Optional[str] - User defined key alias
* - key: Optional[str] - User defined key value. If not set, a 16-digit unique sk-key is created for you.
* - key: Optional[str] - User defined key value. Must start with 'sk-' and be at least 16 characters long. If not set, a 16-digit unique sk-key is created for you.
* - team_id: Optional[str] - The team id of the key
* - user_id: Optional[str] - [NON-FUNCTIONAL] THIS WILL BE IGNORED. The user id of the key
* - budget_id: Optional[str] - The budget id associated with the key. Created by calling `/budget/new`.
@ -7024,7 +7024,7 @@ export interface paths {
* - data: Optional[RegenerateKeyRequest] - Request body containing optional parameters to update
* - key: Optional[str] - The key to regenerate.
* - new_master_key: Optional[str] - The new master key to use, if key is the master key.
* - new_key: Optional[str] - The new key to use, if key is not the master key. If both set, new_master_key will be used.
* - new_key: Optional[str] - The new key to use, if key is not the master key. Must start with 'sk-' and be at least 16 characters long. If both set, new_master_key will be used.
* - key_alias: Optional[str] - User-friendly key alias
* - user_id: Optional[str] - User ID associated with key
* - team_id: Optional[str] - Team ID associated with key
@ -20659,6 +20659,10 @@ export interface components {
messages?: {
[key: string]: unknown;
}[] | null;
/** Metadata */
metadata?: {
[key: string]: unknown;
} | null;
/** Text */
text: string;
};

View file

@ -1,5 +1,11 @@
import { describe, it, expect } from "vitest";
import { type LicenseExpiryTier, formatExpiryDate, getDaysUntilExpiration, getLicenseExpiryTier } from "./licenseUtils";
import {
type LicenseExpiryTier,
formatExpirationStatus,
formatExpiryDate,
getDaysUntilExpiration,
getLicenseExpiryTier,
} from "./licenseUtils";
const NOW = new Date("2026-07-08T00:00:00Z");
@ -59,3 +65,25 @@ describe("formatExpiryDate", () => {
expect(formatExpiryDate("bogus")).toBe("bogus");
});
});
describe("formatExpirationStatus", () => {
it("shows the exact date for a future expiration", () => {
expect(formatExpirationStatus("2026-08-07", NOW)).toBe("Expires Aug 7, 2026");
});
it("still reads as upcoming on the expiration day itself", () => {
expect(formatExpirationStatus("2026-07-08", NOW)).toBe("Expires Jul 8, 2026");
});
it("shows the exact date for a past expiration", () => {
expect(formatExpirationStatus("2026-07-07", NOW)).toBe("Expired Jul 7, 2026");
});
it("returns No expiration for a null date", () => {
expect(formatExpirationStatus(null, NOW)).toBe("No expiration");
});
it("returns No expiration for an unparseable date", () => {
expect(formatExpirationStatus("not-a-date", NOW)).toBe("No expiration");
});
});

View file

@ -49,3 +49,11 @@ export const formatExpiryDate = (expirationDate: string): string => {
}
return date.toLocaleDateString("en-US", EXPIRY_DATE_FORMAT);
};
export const formatExpirationStatus = (expirationDate: string | null, now: Date = new Date()): string => {
const days = getDaysUntilExpiration(expirationDate, now);
if (expirationDate === null || days === null) {
return "No expiration";
}
return days < 0 ? `Expired ${formatExpiryDate(expirationDate)}` : `Expires ${formatExpiryDate(expirationDate)}`;
};

1340
uv.lock generated

File diff suppressed because it is too large Load diff