diff --git a/docker/build_from_pip/Dockerfile.build_from_pip b/docker/build_from_pip/Dockerfile.build_from_pip index bda742c71a9..372606a5b0f 100644 --- a/docker/build_from_pip/Dockerfile.build_from_pip +++ b/docker/build_from_pip/Dockerfile.build_from_pip @@ -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" \ diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py b/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py index ad8aabf77b6..d10b5a2ab09 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py @@ -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 diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index c860f8e540d..b17e055c7ea 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -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): diff --git a/litellm/constants.py b/litellm/constants.py index 715d57e594d..8e0a5cfe50f 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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", diff --git a/litellm/litellm_core_utils/dd_tracing.py b/litellm/litellm_core_utils/dd_tracing.py index ae4f46c38bd..3a1bd72e1a5 100644 --- a/litellm/litellm_core_utils/dd_tracing.py +++ b/litellm/litellm_core_utils/dd_tracing.py @@ -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 diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 461ab62b815..268dcc3df78 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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: """ diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index 6e8429839ad..43181e7f5ff 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -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) diff --git a/litellm/litellm_core_utils/secret_redaction.py b/litellm/litellm_core_utils/secret_redaction.py index b526068589d..455d0f00c35 100644 --- a/litellm/litellm_core_utils/secret_redaction.py +++ b/litellm/litellm_core_utils/secret_redaction.py @@ -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) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 96d0ad48b79..b47fc50e196 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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 diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 093dffccac0..d90703d1544 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -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( diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4988160f19b..dedb9bbf40a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index e6e265abb61..8b1d00c2855 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 "", + _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 diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 27cdc483d4a..67d935e34e3 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -7613,6 +7613,18 @@ ], "title": "Messages" }, + "metadata": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Metadata" + }, "text": { "title": "Text", "type": "string" diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 893e09ece6e..a610e44e69c 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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:]}" diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index b6a2d8d9069..1ed67a93d94 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -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]}, diff --git a/litellm/proxy/guardrails/guardrail_hooks/compresr/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/compresr/__init__.py new file mode 100644 index 00000000000..73d31f7aec0 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/compresr/__init__.py @@ -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, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py b/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py new file mode 100644 index 00000000000..a95bdb670c3 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py @@ -0,0 +1,1214 @@ +"""Compresr guardrail — query-aware, recoverable context compression. + +Compresses bulky message content (tool outputs by default) through the +Compresr API before the request reaches the LLM. Each compressed message +carries a hash marker; a ``compresr_retrieve`` tool is injected so the model +can fetch the original content back through the agentic loop when the +compressed version is not enough — making compression recoverable instead +of lossy. + +Unlike gateway-side compressors that operate on whole message lists, each +target is compressed *query-aware*: the query sent to Compresr is the intent +of the tool call that produced the message (``name + arguments``, resolved +via ``tool_call_id``), falling back to the last user message. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import ipaddress +import json +import time +from collections import Counter, OrderedDict +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Literal +from urllib.parse import urlparse + +import httpx +from fastapi import HTTPException +from httpx import Response as HttpxResponse +from typing_extensions import TypeGuard + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.litellm_core_utils.prompt_templates.factory import ( + get_attribute_or_key, + get_tool_calls_from_response, + has_tool_with_name, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # helper is untyped in http_handler + httpxSpecialProvider, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.secret_managers.main import get_secret_str +from litellm.types.guardrails import GuardrailEventHooks, Mode +from litellm.types.integrations.custom_logger import ( + AgenticLoopPlan, + AgenticLoopRequestPatch, +) +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import ( + Logging as LiteLLMLoggingObj, + ) + from litellm.types.proxy.guardrails.guardrail_hooks.base import ( + GuardrailConfigModel, + ) + +BYPASS_HEADER = "x-compresr-bypass" +COMPRESR_RETRIEVE_TOOL_NAME = "compresr_retrieve" +DEFAULT_API_BASE = "https://api.compresr.ai" +DEFAULT_COMPRESSION_MODEL = "latte_v2" +DEFAULT_TARGET_COMPRESSION_RATIO = 0.5 +DEFAULT_MIN_CHARS_TO_COMPRESS = 500 +_ORIGINALS_TTL_SECONDS = 15 * 60 +_NO_SCOPE_WARNING_INTERVAL_SECONDS = 15 * 60 +_MAX_TRACKED_CALLS = 256 +_DEFAULT_MAX_BYTES_PER_CALL = 10 * 1024 * 1024 +# Aggregate ceiling across all recovery-store entries. max_bytes_per_call only +# bounds a single call; this caps the whole store so many calls cannot exhaust it. +_MAX_TOTAL_STORE_BYTES = 256 * 1024 * 1024 +# Max compresr_retrieve calls expanded into a single follow-up (repeats deduped). +_MAX_RETRIEVALS_PER_LOOP = 8 +# The shared client's 600s read timeout is far too long for an on-request +# guardrail; bound the compress call so a stall hits the fail policy quickly. +_COMPRESS_TIMEOUT_SECONDS = 60.0 +_SOURCE_TAG = "integration:litellm" +# Request-content fields the compression_params passthrough must never +# override — they carry the actual message content/queries being compressed. +_RESERVED_COMPRESSION_PARAM_KEYS = frozenset({"context", "query", "inputs"}) +_BLOCKED_METADATA_HOSTS = frozenset( + { + "metadata.google.internal", + "metadata.goog", + "metadata.azure.com", + "metadata.azure.internal", + } +) +_BLOCKED_METADATA_IPS = frozenset( + ipaddress.ip_address(ip) for ip in ("169.254.169.254", "fd00:ec2::254", "100.100.100.200", "168.63.129.16") +) + + +def _parse_ip_literal(host: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address | None: + """Parse ``host`` as an IP literal, covering the alternate spellings the + socket layer accepts (decimal/hex single-integer IPv4, IPv4-mapped IPv6) + so a blocked address cannot be smuggled past a string comparison.""" + try: + addr = ipaddress.ip_address(host) + except ValueError: + try: + addr = ipaddress.ip_address(int(host, 0)) + except (TypeError, ValueError): + return None + if isinstance(addr, ipaddress.IPv6Address) and addr.ipv4_mapped is not None: + return addr.ipv4_mapped + return addr + + +def _validate_api_base(url: str) -> str: + """Return ``url`` if it passes basic outbound-target checks, else raise. + + Best-effort defense in depth for a mis/maliciously-configured ``api_base``: + rejects non-http(s) schemes and cloud-metadata IPs/hosts (incl. alternate IP + encodings); private ranges are allowed for on-prem deployments. NOT a complete + SSRF control — no DNS resolution, and the shared client follows redirects and + re-resolves DNS (TOCTOU / rebinding); ``api_base`` is trusted operator config, + so this is an accepted limitation. + """ + parsed = urlparse(url) + if parsed.scheme not in ("http", "https"): + raise ValueError(f"Compresr guardrail api_base must be http or https, got scheme={parsed.scheme!r}") + host = (parsed.hostname or "").lower() + if not host: + raise ValueError("Compresr guardrail api_base has no host") + ip_literal = _parse_ip_literal(host) + if host in _BLOCKED_METADATA_HOSTS or (ip_literal is not None and ip_literal in _BLOCKED_METADATA_IPS): + raise ValueError(f"Compresr guardrail api_base {host!r} is a blocked cloud-metadata host") + return url + + +def _is_str_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip + return isinstance(value, dict) + + +def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip + return isinstance(value, list) + + +def _content_to_text(content: object) -> str: + """Collapse a message ``content`` (str or list-of-parts) to plain text. + + For the multimodal list shape, joins ``{type: "text", text: ...}`` parts + with blank-line separators; non-text parts are ignored. + """ + if isinstance(content, str): + return content + if isinstance(content, list): + parts: list[str] = [] + for part in content: + if isinstance(part, dict) and part.get("type") == "text": + text = part.get("text") + if isinstance(text, str): + parts.append(text) + return "\n\n".join(parts) + return "" + + +def _replace_text_in_content(content: object, new_text: str) -> object: + """Write ``new_text`` back into a ``content`` value, preserving shape. + + ``str`` content is replaced directly. For list-of-parts content the first + text part carries ``new_text``, later text parts are dropped, and + non-text parts (images, audio, files) pass through untouched. + """ + if isinstance(content, str): + return new_text + if isinstance(content, list): + out: list[object] = [] + replaced = False + for part in content: + if isinstance(part, dict) and part.get("type") == "text": + if not replaced: + out.append({**part, "text": new_text}) + replaced = True + continue + out.append(part) + if not replaced: + out.insert(0, {"type": "text", "text": new_text}) + return out + return new_text + + +def _render_tool_intent(fn: dict[str, object]) -> str: + name = str(fn.get("name") or "").strip() + args = fn.get("arguments") + if isinstance(args, dict): + try: + args_str = json.dumps(args, separators=(",", ":")) + except (TypeError, ValueError): + args_str = str(args) + else: + args_str = str(args).strip() if args is not None else "" + if name and args_str: + return f"{name}: {args_str}" + return name or args_str + + +def _query_for_target(messages: list[dict[str, object]], target_idx: int, fallback: str) -> str: + """Query used to compress ``messages[target_idx]``. + + Tool/function outputs are compressed against the intent of the tool call + that produced them (found via ``tool_call_id`` on a prior assistant + message); everything else uses the last user message. + """ + msg = messages[target_idx] + if msg.get("role") not in ("tool", "function"): + return fallback + + tool_call_id = msg.get("tool_call_id") + fn_name = msg.get("name") + for j in range(target_idx - 1, -1, -1): + prev = messages[j] + if prev.get("role") != "assistant": + continue + tool_calls = prev.get("tool_calls") + if isinstance(tool_calls, list): + for tc in tool_calls: + if not isinstance(tc, dict): + continue + if tool_call_id and tc.get("id") == tool_call_id: + fn = tc.get("function") + intent = _render_tool_intent(fn if isinstance(fn, dict) else {}) + if intent: + return intent + # Legacy function_call fallback: require a name match, else an earlier + # function_call turn would attribute the wrong intent. + fc = prev.get("function_call") + if isinstance(fc, dict) and fn_name and fc.get("name") == fn_name: + intent = _render_tool_intent(fc) + if intent: + return intent + return fallback + + +def _safe_int(value: object) -> int: + """Parse a token-stat field defensively. + + A malformed-but-200 response must not raise here: ``_call_compress`` has + already returned successfully, so the fail_open/fail_closed decision is + behind us. A bare ``int()`` on a non-numeric field would surface as an + unhandled 500 even when ``fail_open`` is configured. + """ + try: + return int(value) if value is not None else 0 + except (TypeError, ValueError): + return 0 + + +def _safe_response_text(response: object, limit: int = 500) -> str: + """Read a response body for error logging without letting the read itself + raise. A corrupt ``Content-Encoding`` makes ``httpx``'s ``.text`` raise a + ``DecodingError``; if that happened while building a failure detail it would + turn an already-handled error into an unhandled 500.""" + try: + text = getattr(response, "text", "") + except httpx.DecodingError: + return "" + return (text or "")[:limit] + + +def _content_hash(text: str) -> str: + # surrogatepass so a lone surrogate in untrusted content (valid via a JSON + # \uXXXX escape) hashes instead of raising past the fail policy. + return hashlib.sha256(text.encode("utf-8", "surrogatepass")).hexdigest()[:24] + + +def _entry_bytes(originals: dict[str, str]) -> int: + """UTF-8 byte size of one recovery-store entry (surrogatepass, like _content_hash).""" + return sum(len(value.encode("utf-8", "surrogatepass")) for value in originals.values()) + + +def _display_hash(hash_value: str) -> str: + """Bound a model-supplied hash for logs/fallback text. A real marker hash is + 24 hex chars; a prompt-injected ``compresr_retrieve`` call could pass a huge + or control-character-laden string, so strip non-printables (no forged log + lines / ANSI escapes) and cap length before echoing into logs and the + conversation.""" + printable = "".join(ch for ch in hash_value if ch.isprintable()) + return printable if len(printable) <= 32 else f"{printable[:32]}…" + + +def _recovery_marker(hash_value: str) -> str: + return ( + f"\n\n[compresr hash={hash_value}: parts of this content were compressed " + f"away. If you need the full original, call the " + f"{COMPRESR_RETRIEVE_TOOL_NAME} tool with this hash.]" + ) + + +def _build_compresr_retrieve_tool() -> dict[str, object]: + return { + "type": "function", + "function": { + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "description": ( + "Retrieve the original, uncompressed content behind a Compresr " + "compression marker. Call this when a compression marker's hash " + "points at content you need in full." + ), + "parameters": { + "type": "object", + "properties": { + "hash": { + "type": "string", + "description": "The 24-character hex hash from the compression marker.", + }, + }, + "required": ["hash"], + }, + }, + } + + +def has_compresr_retrieve_tool(tools: object) -> bool: + return has_tool_with_name(tools, COMPRESR_RETRIEVE_TOOL_NAME) + + +def _merge_retrieve_tool(existing_tools: object) -> list[object] | None: + """The request's tools plus the retrieve tool, or None when the incoming + shape is not a list (leave the caller's tools untouched; markers stay + inert text).""" + if existing_tools is not None and not isinstance(existing_tools, list): + return None + retrieve_tool = _build_compresr_retrieve_tool() + if existing_tools is None: + return [retrieve_tool] + if has_compresr_retrieve_tool(existing_tools): + return list(existing_tools) + return list(existing_tools) + [retrieve_tool] + + +def _extract_compresr_tool_calls(response: object) -> list[dict[str, object]]: + return [ + {"id": tc.get("id"), "type": "function", "name": tc.get("name"), "arguments": tc.get("arguments", {})} + for tc in get_tool_calls_from_response(response) + if tc.get("name") == COMPRESR_RETRIEVE_TOOL_NAME + ] + + +def _resolve_call_id(logging_obj: object) -> str | None: + """The call id from the framework logging object. + + This value ultimately derives from the client-settable ``x-litellm-call-id`` + header and is echoed back in responses, so it is NOT a trust boundary on its + own — ``_scoped_store_key`` prefixes it with the caller's virtual-key hash to + partition the recovery store per tenant. Request-body/kwargs call ids are + deliberately not consulted here. + """ + logging_call_id = getattr(logging_obj, "litellm_call_id", None) + if isinstance(logging_call_id, str) and logging_call_id: + return logging_call_id + return None + + +def _caller_scope(logging_obj: object) -> str: + """The caller's virtual-key hash, used to partition the recovery store. + + Trust is anchored on the ``UserAPIKeyAuth`` object the proxy sets + server-side (``metadata.user_api_key_auth``, litellm_pre_call_utils). Its + ``api_key`` is the hash of the authenticated key. Both metadata spellings + are scanned (``/v1/messages`` and ``/v1/responses`` carry it under + ``litellm_metadata``), but the bare ``user_api_key`` *string* is never + trusted on its own: a JSON request body can place one in the client-supplied + ``metadata`` field, which is only sanitized on the route's canonical + container. Returns "" when the proxy runs without per-key auth, in which case + all traffic is a single trust domain and the call id alone suffices. + """ + details = getattr(logging_obj, "model_call_details", None) + if not _is_str_object_dict(details): + return "" + litellm_params = details.get("litellm_params") + for container in (litellm_params, details): + if not _is_str_object_dict(container): + continue + for meta_key in ("metadata", "litellm_metadata"): + metadata = container.get(meta_key) + if not _is_str_object_dict(metadata): + continue + auth = metadata.get("user_api_key_auth") + if isinstance(auth, UserAPIKeyAuth) and isinstance(auth.api_key, str) and auth.api_key: + return auth.api_key + return "" + + +def _scoped_store_key(logging_obj: object) -> str | None: + """Key for the recovery store: caller identity plus framework call id. + + Keying on the call id alone is unsafe: it comes from the client-settable + ``x-litellm-call-id`` header and is echoed back in responses, so one caller + could read or evict another's originals by reusing the id. Prefixing the + unforgeable virtual-key hash binds each entry to the tenant that created it. + Returns None when there is no call id, which disables recovery for the call. + """ + call_id = _resolve_call_id(logging_obj) + if call_id is None: + return None + scope = _caller_scope(logging_obj) + return f"{scope}\x00{call_id}" if scope else call_id + + +def _is_responses_api_response(response: object) -> bool: + return isinstance(get_attribute_or_key(response, "output", None), list) + + +def _is_anthropic_messages_response(response: object) -> bool: + return isinstance(get_attribute_or_key(response, "content", None), list) + + +def _assistant_text_from_response(response: object) -> str | None: + """The assistant's natural-language text from a model response, across chat, + Anthropic, and Responses shapes. Preserved when the turn is rebuilt for the + retrieval follow-up so the model's reasoning is not lost.""" + choices = get_attribute_or_key(response, "choices", None) + if isinstance(choices, list) and choices: + message = get_attribute_or_key(choices[0], "message", None) + if message is not None: + text = _content_to_text(get_attribute_or_key(message, "content", None)) + if text: + return text + content = get_attribute_or_key(response, "content", None) + if isinstance(content, list): + parts = [ + text + for block in content + if get_attribute_or_key(block, "type", None) == "text" + for text in (get_attribute_or_key(block, "text", None),) + if isinstance(text, str) and text + ] + if parts: + return "".join(parts) + output = get_attribute_or_key(response, "output", None) + if isinstance(output, list): + parts = [] + for item in output: + if get_attribute_or_key(item, "type", None) != "message": + continue + item_content = get_attribute_or_key(item, "content", None) + if not isinstance(item_content, list): + continue + for chunk in item_content: + if get_attribute_or_key(chunk, "type", None) == "output_text": + text = get_attribute_or_key(chunk, "text", None) + if isinstance(text, str) and text: + parts.append(text) + if parts: + return "".join(parts) + return None + + +def _build_assistant_message_from_response( + response: object, + retrieved: list[tuple[dict[str, object], str]], +) -> dict[str, object]: + """Rebuild the chat-completions assistant turn for the retrieval follow-up. + + Only the ``compresr_retrieve`` calls are echoed, each answered by a tool + result below. Other tool calls made in the same turn are omitted on purpose: + the follow-up re-runs the model with the recovered content so it re-plans + them. Echoing them would leave tool_calls with no matching tool result and + the provider would reject the request. + """ + return { + "role": "assistant", + "content": _assistant_text_from_response(response), + "tool_calls": [ + { + "id": tool_call.get("id"), + "type": "function", + "function": { + "name": tool_call.get("name"), + "arguments": json.dumps(tool_call.get("arguments", {})), + }, + } + for tool_call, _ in retrieved + ], + } + + +def _build_anthropic_followup_messages( + response: object, + retrieved: list[tuple[dict[str, object], str]], +) -> list[dict[str, object]]: + """Anthropic requires the tool_use block echoed back in an assistant + message paired with a tool_result block keyed by the same tool_use_id. The + assistant text is preserved; non-retrieve tool calls are re-planned by the + follow-up (see _build_assistant_message_from_response).""" + assistant_content: list[dict[str, object]] = [] + text = _assistant_text_from_response(response) + if text: + assistant_content.append({"type": "text", "text": text}) + assistant_content.extend( + { + "type": "tool_use", + "id": tool_call.get("id"), + "name": tool_call.get("name"), + "input": tool_call.get("arguments", {}), + } + for tool_call, _ in retrieved + ) + assistant_message: dict[str, object] = {"role": "assistant", "content": assistant_content} + user_message: dict[str, object] = { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": tool_call.get("id"), "content": content} + for tool_call, content in retrieved + ], + } + return [assistant_message, user_message] + + +def _build_responses_followup_items( + response: object, + retrieved: list[tuple[dict[str, object], str]], +) -> list[dict[str, object]]: + """The Responses API requires the model's function_call echoed back paired + with a function_call_output keyed by the same call_id. The assistant text is + preserved; non-retrieve tool calls are re-planned by the follow-up.""" + items: list[dict[str, object]] = [] + text = _assistant_text_from_response(response) + if text: + items.append({"role": "assistant", "content": text}) + for tool_call, content in retrieved: + call_id = tool_call.get("id") + items.append( + { + "type": "function_call", + "call_id": call_id, + "name": tool_call.get("name"), + "arguments": json.dumps(tool_call.get("arguments", {})), + } + ) + items.append({"type": "function_call_output", "call_id": call_id, "output": content}) + return items + + +@dataclass +class _CompressionResult: + """Outcome of applying compression results to a message list.""" + + compressed_messages: list[dict[str, object]] + originals: dict[str, str] = field(default_factory=dict) + # original text -> compressed text, plus the machinery the Responses `texts` + # mirror needs to replace only where it is unambiguous. + text_replacements: dict[str, str] = field(default_factory=dict) + replaced_text_counts: dict[str, int] = field(default_factory=dict) + ambiguous_texts: set[str] = field(default_factory=set) + messages_compressed: int = 0 + tokens_before: int = 0 + tokens_after: int = 0 + + +class CompresrGuardrail(CustomGuardrail): + def __init__( + self, + api_base: str | None = None, + api_key: str | None = None, + model: str | None = None, + target_compression_ratio: float | None = None, + coarse: bool | None = None, + min_chars_to_compress: int | None = None, + compress_tool_outputs: bool | None = None, + compress_system: bool | None = None, + compress_history: bool | None = None, + compress_last_user: bool | None = None, + enable_retrieval: bool | None = None, + guardrail_name: str | None = None, + event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, + default_on: bool = False, + unreachable_fallback: str | None = None, + max_bytes_per_call: int | None = None, + allow_bypass_header: bool | None = None, + dynamic: bool | None = None, + dynamic_min_ratio: float | None = None, + dynamic_max_ratio: float | None = None, + compression_params: dict[str, object] | None = None, + ): + raw_api_base = (api_base or get_secret_str("COMPRESR_API_BASE") or DEFAULT_API_BASE).rstrip("/") + self.compresr_api_base = _validate_api_base(raw_api_base) + self.compresr_api_key = api_key or get_secret_str("COMPRESR_API_KEY") + if not self.compresr_api_key: + raise ValueError( + "Compresr guardrail requires an API key. Set `api_key` in the " + "guardrail config or the COMPRESR_API_KEY env var." + ) + self.compression_model = model or DEFAULT_COMPRESSION_MODEL + self.target_compression_ratio = ( + DEFAULT_TARGET_COMPRESSION_RATIO if target_compression_ratio is None else target_compression_ratio + ) + self.coarse = True if coarse is None else coarse + self.min_chars_to_compress = ( + DEFAULT_MIN_CHARS_TO_COMPRESS if min_chars_to_compress is None else min_chars_to_compress + ) + self.compress_tool_outputs = True if compress_tool_outputs is None else compress_tool_outputs + self.compress_system = False if compress_system is None else compress_system + self.compress_history = False if compress_history is None else compress_history + self.compress_last_user = False if compress_last_user is None else compress_last_user + self.enable_retrieval = True if enable_retrieval is None else enable_retrieval + self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( + "fail_open" if unreachable_fallback == "fail_open" else "fail_closed" + ) + self.max_bytes_per_call = _DEFAULT_MAX_BYTES_PER_CALL if max_bytes_per_call is None else max_bytes_per_call + if self.max_bytes_per_call < 0: + raise ValueError("max_bytes_per_call must be >= 0 (0 disables the cap; positive values enforce it)") + self.allow_bypass_header = False if allow_bypass_header is None else allow_bypass_header + # Dynamic (adaptive) compression — latte_v2 only, on by default: the server + # picks the ratio per input instead of honoring target_compression_ratio. + self.dynamic = True if dynamic is None else dynamic + self.dynamic_min_ratio = dynamic_min_ratio + self.dynamic_max_ratio = dynamic_max_ratio + # Passthrough of extra compression params forwarded verbatim, so a new + # Compresr feature works without changing this guardrail. Named fields win; + # request-content fields are stripped. + reserved_keys = _RESERVED_COMPRESSION_PARAM_KEYS.intersection(compression_params or {}) + if reserved_keys: + verbose_proxy_logger.warning( + "Compresr: ignoring reserved compression_params keys %s", sorted(reserved_keys) + ) + self.compression_params: dict[str, object] = { + k: v for k, v in (compression_params or {}).items() if k not in _RESERVED_COMPRESSION_PARAM_KEYS + } + self.async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback, + ) + self._originals_by_call_id: OrderedDict[str, tuple[dict[str, str], float]] = OrderedDict() + # Running byte size of the store, kept in sync to enforce the global cap cheaply. + self._store_total_bytes = 0 + # Rate-limits the "recovery skipped, no auth scope" warning so an ongoing + # misconfiguration stays visible without flooding hot-path logs. + self._no_scope_warning_expiry = 0.0 + if self.enable_retrieval: + verbose_proxy_logger.warning( + "Compresr: enable_retrieval is on; the recovery store is per-process. " + "For multi-worker deployments, set enable_retrieval=false or run with --workers 1." + ) + super().__init__( # pyright: ignore[reportUnknownMemberType] # CustomGuardrail.__init__ is untyped + guardrail_name=guardrail_name, + event_hook=event_hook, + default_on=default_on, + ) + + def _should_bypass(self, request_data: dict) -> bool: + if not self.allow_bypass_header: + return False + psr = request_data.get("proxy_server_request") + if not _is_str_object_dict(psr): + return False + headers = psr.get("headers") + if not _is_str_object_dict(headers): + return False + return str(headers.get(BYPASS_HEADER)).lower() == "true" + + def _request_headers(self) -> dict[str, str]: + return { + "Content-Type": "application/json", + "X-API-Key": self.compresr_api_key or "", + } + + def _handle_compress_failure(self, error: str, log_detail: dict[str, object]) -> None: + """fail_open logs and returns (caller forwards uncompressed); + fail_closed raises. ``log_detail`` may include upstream response bodies + and is written only to server logs; the raised ``HTTPException`` carries + a generic message so a malicious ``api_base`` cannot exfiltrate response + bytes through the client-visible error.""" + if self.unreachable_fallback == "fail_open": + verbose_proxy_logger.warning( + "Compresr: %s; fail_open configured, forwarding request uncompressed. detail=%s", + error, + log_detail, + ) + return + verbose_proxy_logger.error("Compresr: %s. detail=%s", error, log_detail) + raise HTTPException(status_code=502, detail={"error": error}) + + def _evict_oldest(self) -> None: + """Drop the front (oldest) entry and decrement the running byte total.""" + _key, (evicted, _expiry) = self._originals_by_call_id.popitem(last=False) + self._store_total_bytes -= _entry_bytes(evicted) + + def _prune_originals(self) -> None: + # Insertion order == expiry order (shared TTL); prune from the front. + now = time.monotonic() + store = self._originals_by_call_id + while store and store[next(iter(store))][1] <= now: + self._evict_oldest() + while len(store) > _MAX_TRACKED_CALLS: + self._evict_oldest() + # Global byte budget; keep the most-recent entry so the current call's + # originals survive (a single call is already bounded by max_bytes_per_call). + while len(store) > 1 and self._store_total_bytes > _MAX_TOTAL_STORE_BYTES: + self._evict_oldest() + + def _existing_originals(self, store_key: str | None) -> dict[str, str]: + """Originals already stored under this key, so the per-call byte budget + can account for an earlier turn that reused the store key.""" + if store_key is None: + return {} + return self._originals_by_call_id.get(store_key, ({}, 0.0))[0] + + def _store_originals(self, store_key: str, originals: dict[str, str]) -> None: + existing, _ = self._originals_by_call_id.get(store_key, ({}, 0.0)) + merged = self._bound_call_bytes({**existing, **originals}) + # Keep the running total in sync: drop the overwritten entry, add the new one. + self._store_total_bytes += _entry_bytes(merged) - _entry_bytes(existing) + self._originals_by_call_id[store_key] = ( + merged, + time.monotonic() + _ORIGINALS_TTL_SECONDS, + ) + self._originals_by_call_id.move_to_end(store_key) + self._prune_originals() + + def _bound_call_bytes(self, merged: dict[str, str]) -> dict[str, str]: + """Drop oldest entries (dict insertion order) until the aggregate byte + size fits ``self.max_bytes_per_call``. Prevents one call with many + large tool outputs from growing proxy memory without bound.""" + if self.max_bytes_per_call <= 0: + return merged + total = _entry_bytes(merged) + if total <= self.max_bytes_per_call: + return merged + bounded = dict(merged) + for key in list(bounded.keys()): + if total <= self.max_bytes_per_call: + break + total -= len(bounded[key].encode("utf-8", "surrogatepass")) + del bounded[key] + verbose_proxy_logger.warning("Compresr: originals-store byte cap hit, evicted hash=%s", key) + return bounded + + def _retrieve_original(self, store_key: str | None, hash_value: str) -> str | None: + """Stored original for a marker hash, or None if not issued for this + request (unknown, expired, or from another caller's scope).""" + if store_key: + originals, expiry = self._originals_by_call_id.get(store_key, ({}, 0.0)) + if expiry > time.monotonic() and hash_value in originals: + return originals[hash_value] + verbose_proxy_logger.warning( + "Compresr retrieve: rejecting hash=%s (not issued for this request, or expired)", + _display_hash(hash_value), + ) + return None + + def _resolve_retrievals( + self, store_key: str | None, tool_calls: list[dict[str, object]] + ) -> tuple[list[tuple[dict[str, object], str]], bool]: + """Resolve compresr_retrieve calls to (call, result_text) pairs, deduping + repeated hashes and capping the count so the follow-up cannot be amplified. + The bool is True iff at least one call resolved to real stored content.""" + retrieved: list[tuple[dict[str, object], str]] = [] + seen: set[str] = set() + resolved_any = False + for idx, tc in enumerate(tool_calls): + arguments = tc.get("arguments", {}) + hash_value = str(arguments.get("hash", "")) if isinstance(arguments, dict) else "" + if idx >= _MAX_RETRIEVALS_PER_LOOP: + result = "[compresr: retrieval limit reached for this turn]" + elif hash_value in seen: + result = "[compresr: already retrieved above for this hash]" + else: + content = self._retrieve_original(store_key, hash_value) + if content is None: + result = f"[compresr: hash={_display_hash(hash_value)} not found, expired, or not issued for this request]" + else: + seen.add(hash_value) + resolved_any = True + result = content + verbose_proxy_logger.debug("Compresr retrieve: hash=%s -> %d chars", _display_hash(hash_value), len(result)) + retrieved.append((tc, result)) + return retrieved, resolved_any + + async def _call_compress( + self, + contexts: list[str], + queries: list[str], + ) -> list[dict[str, object]] | None: + """Compress ``contexts`` (query-aware). Returns one result dict per + context, or None when the service failed and fail_open applies.""" + common: dict[str, object] = { + # Passthrough first so the named fields below always win on collision. + **self.compression_params, + "compression_model_name": self.compression_model, + "target_compression_ratio": self.target_compression_ratio, + "coarse": self.coarse, + "dynamic": self.dynamic, + "source": _SOURCE_TAG, + } + # Only send the bounds the operator actually set; otherwise let the + # server apply its own floor/ceiling. + if self.dynamic_min_ratio is not None: + common["dynamic_min_ratio"] = self.dynamic_min_ratio + if self.dynamic_max_ratio is not None: + common["dynamic_max_ratio"] = self.dynamic_max_ratio + if len(contexts) == 1: + url = f"{self.compresr_api_base}/api/compress/question-specific/" + payload: dict[str, object] = { + "context": contexts[0], + "query": queries[0], + **common, + } + else: + url = f"{self.compresr_api_base}/api/compress/question-specific/batch" + payload = { + "inputs": [{"context": ctx, "query": q} for ctx, q in zip(contexts, queries)], + **common, + } + + try: + raw_response: HttpxResponse | None = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped + url=url, + json=payload, + headers=self._request_headers(), + timeout=_COMPRESS_TIMEOUT_SECONDS, + ) + except asyncio.CancelledError: + raise + except httpx.HTTPStatusError as e: + # The shared handler calls raise_for_status(), so a non-2xx reply arrives + # here as an error carrying the upstream body + our API key header; route + # it through the fail policy so none of that reaches the client. + resp = getattr(e, "response", None) + self._handle_compress_failure( + "Compresr compression service returned an error", + { + "status_code": getattr(resp, "status_code", None), + "body": _safe_response_text(resp), + }, + ) + return None + except (httpx.RequestError, litellm.Timeout) as e: + # Every request-side httpx failure is a RequestError; route the whole + # class through the fail policy so none escapes as a 500 under fail_open. + # (HTTPStatusError is handled above and is not a RequestError.) + self._handle_compress_failure( + "Compresr compression service request failed", + {"detail": str(e)}, + ) + return None + if raw_response is None or not 200 <= raw_response.status_code < 300: + self._handle_compress_failure( + "Compresr compression service returned an error", + { + "status_code": getattr(raw_response, "status_code", None), + "body": _safe_response_text(raw_response), + }, + ) + return None + + try: + body: object = raw_response.json() + except (ValueError, httpx.DecodingError, RecursionError): + # RecursionError: a deeply nested JSON body overflows the parser; + # route it through the fail policy rather than let it escape as a 500. + self._handle_compress_failure( + "Compresr compression service returned an unreadable response", + {"body": _safe_response_text(raw_response)}, + ) + return None + if not _is_str_object_dict(body) or not _is_str_object_dict(body.get("data")): + self._handle_compress_failure( + "Compresr compression service returned unexpected response shape", + {"body": _safe_response_text(raw_response)}, + ) + return None + data: dict[str, object] = body["data"] # pyright: ignore[reportAssignmentType] # dict-guarded above; subscript does not narrow + + if len(contexts) == 1: + return [data] + results = data.get("results") + if ( + not _is_object_list(results) + or len(results) != len(contexts) + or not all(_is_str_object_dict(r) for r in results) + ): + # Anything but a 1:1 dict-per-context mapping would misalign + # results with their target messages. + self._handle_compress_failure( + "Compresr batch response missing or mismatched 'results'", + {"expected": len(contexts), "got": len(results) if _is_object_list(results) else None}, + ) + return None + return results # pyright: ignore[reportReturnType] # every element dict-checked above; list[object] does not narrow + + def _select_targets(self, messages: list[dict[str, object]], query_idx: int | None) -> list[int]: + """Indices of messages whose text content should be compressed.""" + targets: list[int] = [] + for idx, msg in enumerate(messages): + if idx == query_idx and not self.compress_last_user: + continue + role = msg.get("role") + if role in ("tool", "function"): + if not self.compress_tool_outputs: + continue + elif role == "system": + if not self.compress_system: + continue + elif role == "user": + if idx != query_idx and not self.compress_history: + continue + else: + continue + if len(_content_to_text(msg.get("content"))) < self.min_chars_to_compress: + continue + targets.append(idx) + return targets + + @staticmethod + def _extract_fallback_query( + messages: list[dict[str, object]], + ) -> tuple[str, int | None]: + for idx in range(len(messages) - 1, -1, -1): + if messages[idx].get("role") == "user": + return _content_to_text(messages[idx].get("content")), idx + return "", None + + def _apply_compression_results( + self, + messages: list[dict[str, object]], + targets: list[int], + contexts: list[str], + results: list[dict[str, object]], + recovery_enabled: bool, + existing_originals: dict[str, str] | None = None, + ) -> _CompressionResult: + """Write each compression result into a copy of ``messages``. + + A result is a real compression only when it is a non-empty string that + differs from the original; identical text is treated as a no-op so an + untouched request is not needlessly rewritten downstream. + """ + out = _CompressionResult(compressed_messages=list(messages)) + existing = existing_originals or {} + cap = self.max_bytes_per_call + # Seed with what is already stored under this store key: markers are + # attached only while the store (existing + this call's originals) stays + # within the cap, so _store_originals never has to evict a hash this call + # just shipped a marker for -- including on a later turn that reuses the + # store key. A hash already stored (or repeated here) costs no new bytes. + recovery_bytes = _entry_bytes(existing) + for target_idx, original_text, result in zip(targets, contexts, results): + compressed_text = result.get("compressed_context") + if not isinstance(compressed_text, str) or not compressed_text or compressed_text == original_text: + continue + out.messages_compressed += 1 + if recovery_enabled: + hash_value = _content_hash(original_text) + already_stored = hash_value in existing or hash_value in out.originals + new_bytes = 0 if already_stored else len(original_text.encode("utf-8", "surrogatepass")) + if cap <= 0 or recovery_bytes + new_bytes <= cap: + recovery_bytes += new_bytes + out.originals[hash_value] = original_text + compressed_text += _recovery_marker(hash_value) + previous = out.text_replacements.get(original_text) + if previous is not None and previous != compressed_text: + # Two targets with identical text but different query-specific + # compressions; a value-keyed replacement cannot tell them apart. + out.ambiguous_texts.add(original_text) + else: + out.text_replacements[original_text] = compressed_text + out.replaced_text_counts[original_text] = out.replaced_text_counts.get(original_text, 0) + 1 + original_msg = out.compressed_messages[target_idx] + out.compressed_messages[target_idx] = { + **original_msg, + "content": _replace_text_in_content(original_msg.get("content"), compressed_text), + } + out.tokens_before += _safe_int(result.get("original_tokens")) + out.tokens_after += _safe_int(result.get("compressed_tokens")) + return out + + @staticmethod + def _mirror_texts_channel(input_texts: object, applied: _CompressionResult) -> list[object] | None: + """Compressed content mirrored into the Responses `texts` channel. + + The chat/Anthropic handlers round-trip ``structured_messages``; the + Responses translation cannot rebuild its input from chat messages and + instead writes back through ``texts``. This matches by value, so a + replacement is applied only when it is unambiguous: one compression per + text, and every occurrence in ``texts`` accounted for by a compressed + target. Anything else is left uncompressed rather than risk a wrong or + out-of-policy replacement. Returns None when nothing safe applies. + """ + if not applied.text_replacements or not isinstance(input_texts, list): + return None + counts = Counter(text for text in input_texts if isinstance(text, str)) + safe = { + text: replacement + for text, replacement in applied.text_replacements.items() + if text not in applied.ambiguous_texts and counts.get(text) == applied.replaced_text_counts.get(text) + } + if not safe: + return None + return [safe.get(text, text) if isinstance(text, str) else text for text in input_texts] + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: LiteLLMLoggingObj | None = None, + ) -> GenericGuardrailAPIInputs: + if input_type != "request": + return inputs + + if self._should_bypass(request_data): + verbose_proxy_logger.debug("Compresr: %s header set; skipping compression", BYPASS_HEADER) + return inputs + + structured_messages = inputs.get("structured_messages") + if not _is_object_list(structured_messages) or not structured_messages: + return inputs + messages = [m for m in structured_messages if _is_str_object_dict(m)] + if len(messages) != len(structured_messages): + return inputs + + fallback_query, query_idx = self._extract_fallback_query(messages) + targets: list[int] = [] + queries: list[str] = [] + for idx in self._select_targets(messages, query_idx): + query = _query_for_target(messages, idx, fallback_query) + # latte models require a non-empty query; leave targets we cannot + # derive one for uncompressed rather than erroring. + if not query.strip(): + continue + targets.append(idx) + queries.append(query) + if not targets: + verbose_proxy_logger.debug("Compresr: no messages eligible for compression") + return inputs + + contexts = [_content_to_text(messages[idx].get("content")) for idx in targets] + + start_time = time.monotonic() + results = await self._call_compress(contexts=contexts, queries=queries) + end_time = time.monotonic() + if results is None: # service failed, fail_open configured + return inputs + + # Recovery needs a per-tenant scope; without per-key auth the key would fall + # back to the client-settable call id (cross-tenant reads), so skip it. + store_key = _scoped_store_key(logging_obj) + scope = _caller_scope(logging_obj) + recovery_enabled = self.enable_retrieval and store_key is not None and bool(scope) + if self.enable_retrieval and not scope and time.monotonic() >= self._no_scope_warning_expiry: + # Surface the silent no-recovery case (compressed, but no auth scope + # to inject the retrieve tool), re-warning once per interval. + self._no_scope_warning_expiry = time.monotonic() + _NO_SCOPE_WARNING_INTERVAL_SECONDS + verbose_proxy_logger.warning( + "Compresr: enable_retrieval is on but this request has no per-key auth scope; " + "compressing without recovery (compresr_retrieve tool not injected). " + "Configure virtual-key auth to enable recovery." + ) + + existing_originals = self._existing_originals(store_key) + applied = self._apply_compression_results( + messages, targets, contexts, results, recovery_enabled, existing_originals + ) + if applied.messages_compressed == 0: + # Nothing replaced: return the original inputs object (handlers detect + # edits by identity; a fresh list forces write-back that strips Anthropic + # cache_control from thinking blocks). + verbose_proxy_logger.debug("Compresr: service returned no compressed content; request unchanged") + return inputs + + stats: dict[str, object] = { + "messages_compressed": applied.messages_compressed, + "tokens_before": applied.tokens_before, + "tokens_after": applied.tokens_after, + "tokens_saved": applied.tokens_before - applied.tokens_after, + "compression_model": self.compression_model, + } + verbose_proxy_logger.debug( + "Compresr: compressed %s message(s), %s -> %s tokens", + applied.messages_compressed, + applied.tokens_before, + applied.tokens_after, + ) + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=stats, + request_data=request_data, + guardrail_status="success", + guardrail_provider="compresr", + start_time=start_time, + end_time=end_time, + duration=end_time - start_time, + ) + + compressed_inputs: dict[str, object] = {**inputs, "structured_messages": applied.compressed_messages} + mirrored_texts = self._mirror_texts_channel(inputs.get("texts"), applied) + if mirrored_texts is not None: + compressed_inputs["texts"] = mirrored_texts + + originals = applied.originals + if not recovery_enabled or not originals or store_key is None: + return compressed_inputs # pyright: ignore[reportReturnType] # plain dicts satisfy AllMessageValues at runtime + + self._store_originals(store_key, originals) + + merged_tools = _merge_retrieve_tool(inputs.get("tools")) + if merged_tools is not None: + compressed_inputs["tools"] = merged_tools + return compressed_inputs # pyright: ignore[reportReturnType] # plain dicts satisfy AllMessageValues at runtime + + async def async_should_run_agentic_loop( + self, + response: Any, + model: str, + messages: list[dict], + tools: list[dict] | None, + stream: bool, + custom_llm_provider: str, + kwargs: dict, + ) -> tuple[bool, dict]: + if not has_compresr_retrieve_tool(tools): + return False, {} + tool_calls = _extract_compresr_tool_calls(response) + if not tool_calls: + return False, {} + return True, {"tool_calls": tool_calls} + + async def async_build_agentic_loop_plan( + self, + tools: dict, + model: str, + messages: list[dict], + response: Any, + anthropic_messages_provider_config: Any, + anthropic_messages_optional_request_params: dict, + logging_obj: Any, + stream: bool, + kwargs: dict, + ) -> AgenticLoopPlan: + tool_calls: list[dict[str, object]] = tools.get("tool_calls", []) # pyright: ignore[reportAssignmentType] # gate hook builds this dict with list values only + + self._prune_originals() + store_key = _scoped_store_key(logging_obj) + retrieved, resolved_any = self._resolve_retrievals(store_key, tool_calls) + if not resolved_any: + # Nothing this guardrail stored resolved; skip the extra provider round-trip. + return AgenticLoopPlan(run_agentic_loop=False) + + if _is_responses_api_response(response): + follow_up_messages = list(messages) + _build_responses_followup_items(response, retrieved) + elif _is_anthropic_messages_response(response): + follow_up_messages = list(messages) + _build_anthropic_followup_messages(response, retrieved) + else: + assistant_message = _build_assistant_message_from_response(response, retrieved) + tool_results = [ + {"role": "tool", "tool_call_id": tc.get("id"), "content": content} for tc, content in retrieved + ] + follow_up_messages = list(messages) + [assistant_message] + tool_results + + anthropic_max = anthropic_messages_optional_request_params.get("max_tokens") + max_tokens: int | None = anthropic_max if anthropic_max is not None else kwargs.get("max_tokens") + optional_params_without_max_tokens = { + k: v for k, v in anthropic_messages_optional_request_params.items() if k != "max_tokens" + } + + full_model_name = model + if logging_obj is not None: + agentic_params = getattr(logging_obj, "model_call_details", {}).get("agentic_loop_params", {}) + candidate = agentic_params.get("model", model) + if isinstance(candidate, str) and candidate: + full_model_name = candidate + + return AgenticLoopPlan( + run_agentic_loop=True, + request_patch=AgenticLoopRequestPatch( + model=full_model_name, + messages=follow_up_messages, + max_tokens=max_tokens, + optional_params=optional_params_without_max_tokens, + kwargs=self._sanitized_follow_up_kwargs(kwargs), + ), + metadata={"tool_type": "compresr_retrieve"}, + ) + + def _sanitized_follow_up_kwargs(self, kwargs: dict) -> dict[str, object]: + """Copy of the request kwargs for the retrieval follow-up with other + guardrails' pre-call-executed markers stripped, so input guardrails + re-inspect the restored originals; only this guardrail's own marker is + kept, to avoid recompressing what it just retrieved.""" + out: dict[str, object] = { + k: v for k, v in kwargs.items() if not k.startswith("_compresr") and k != "litellm_logging_obj" + } + own_marker = self._pre_call_marker() + for meta_key in ("metadata", "litellm_metadata"): + meta = out.get(meta_key) + if not isinstance(meta, dict): + continue + executed = meta.get(PRE_CALL_EXECUTED_GUARDRAILS_KEY) + if not isinstance(executed, list): + continue + kept = [m for m in executed if own_marker is not None and m == own_marker] + out[meta_key] = ( + {**meta, PRE_CALL_EXECUTED_GUARDRAILS_KEY: kept} + if kept + else {k: v for k, v in meta.items() if k != PRE_CALL_EXECUTED_GUARDRAILS_KEY} + ) + return out + + @staticmethod + def get_config_model() -> type[GuardrailConfigModel[object]] | None: + from litellm.types.proxy.guardrails.guardrail_hooks.compresr import ( + CompresrGuardrailConfigModel, + ) + + return CompresrGuardrailConfigModel diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index df311bed7b2..01f4e040e58 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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 diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 8c980f33b01..94871ff072d 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 62fc28256cd..9f36e729330 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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, diff --git a/litellm/router.py b/litellm/router.py index 3d49d7048e4..f408d030b8b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index eb0f74a58e7..e85987870e1 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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] diff --git a/litellm/router_utils/common_utils.py b/litellm/router_utils/common_utils.py index 18bce2348f6..5cfea5e3bf2 100644 --- a/litellm/router_utils/common_utils.py +++ b/litellm/router_utils/common_utils.py @@ -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 diff --git a/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py b/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py index 62c99a1c9f6..803fdc4b353 100644 --- a/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py +++ b/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py @@ -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]]: diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 03b6b36ec2e..f3410935ec7 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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): diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index ced6dfe5e2e..801436c774a 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/compresr.py b/litellm/types/proxy/guardrails/guardrail_hooks/compresr.py new file mode 100644 index 00000000000..dad61f83b7d --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/compresr.py @@ -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)" diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 8b273174d99..e10dde793d1 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/pyproject.toml b/pyproject.toml index f890cc976f3..d70e99c5775 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/tests/e2e/claude_code/_builder_unit_tests/test_v0_layout.py b/tests/e2e/claude_code/_builder_unit_tests/test_v0_layout.py index b1745008fac..a3569ebdb49 100644 --- a/tests/e2e/claude_code/_builder_unit_tests/test_v0_layout.py +++ b/tests/e2e/claude_code/_builder_unit_tests/test_v0_layout.py @@ -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 " diff --git a/tests/e2e/claude_code/_driver_unit_tests/test_cli_driver.py b/tests/e2e/claude_code/_driver_unit_tests/test_cli_driver.py index 786f6029993..f1a0534906e 100644 --- a/tests/e2e/claude_code/_driver_unit_tests/test_cli_driver.py +++ b/tests/e2e/claude_code/_driver_unit_tests/test_cli_driver.py @@ -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 diff --git a/tests/e2e/claude_code/_driver_unit_tests/test_passthrough.py b/tests/e2e/claude_code/_driver_unit_tests/test_passthrough.py new file mode 100644 index 00000000000..2d17a84d418 --- /dev/null +++ b/tests/e2e/claude_code/_driver_unit_tests/test_passthrough.py @@ -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", + } diff --git a/tests/e2e/claude_code/_passthrough.py b/tests/e2e/claude_code/_passthrough.py new file mode 100644 index 00000000000..24b6d694d3f --- /dev/null +++ b/tests/e2e/claude_code/_passthrough.py @@ -0,0 +1,196 @@ +"""Shared body for the `passthrough` × 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__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` × 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) diff --git a/tests/e2e/claude_code/cli_driver.py b/tests/e2e/claude_code/cli_driver.py index 97eaa0e6847..5b18c1c291a 100644 --- a/tests/e2e/claude_code/cli_driver.py +++ b/tests/e2e/claude_code/cli_driver.py @@ -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] = {} diff --git a/tests/e2e/claude_code/conftest.py b/tests/e2e/claude_code/conftest.py index d2bfa1a54bf..6ee8b940648 100644 --- a/tests/e2e/claude_code/conftest.py +++ b/tests/e2e/claude_code/conftest.py @@ -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 diff --git a/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.env.example b/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.env.example index 11633810533..5ca7937a426 100644 --- a/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.env.example +++ b/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.env.example @@ -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://.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: diff --git a/tests/e2e/claude_code/manifest.yaml b/tests/e2e/claude_code/manifest.yaml index f7cccf0cef2..e5a956991cb 100644 --- a/tests/e2e/claude_code/manifest.yaml +++ b/tests/e2e/claude_code/manifest.yaml @@ -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 diff --git a/tests/e2e/claude_code/passthrough/__init__.py b/tests/e2e/claude_code/passthrough/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/e2e/claude_code/passthrough/test_anthropic.py b/tests/e2e/claude_code/passthrough/test_anthropic.py new file mode 100644 index 00000000000..aa0443e0625 --- /dev/null +++ b/tests/e2e/claude_code/passthrough/test_anthropic.py @@ -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, + ) diff --git a/tests/e2e/claude_code/passthrough/test_azure.py b/tests/e2e/claude_code/passthrough/test_azure.py new file mode 100644 index 00000000000..09b0824047a --- /dev/null +++ b/tests/e2e/claude_code/passthrough/test_azure.py @@ -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://.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, + ) diff --git a/tests/e2e/claude_code/passthrough/test_bedrock_converse.py b/tests/e2e/claude_code/passthrough/test_bedrock_converse.py new file mode 100644 index 00000000000..d1093a7a958 --- /dev/null +++ b/tests/e2e/claude_code/passthrough/test_bedrock_converse.py @@ -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." + ), + } + ) diff --git a/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py b/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py new file mode 100644 index 00000000000..6e84dea6779 --- /dev/null +++ b/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py @@ -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, + ) diff --git a/tests/e2e/claude_code/passthrough/test_vertex_ai.py b/tests/e2e/claude_code/passthrough/test_vertex_ai.py new file mode 100644 index 00000000000..5e3c6bce419 --- /dev/null +++ b/tests/e2e/claude_code/passthrough/test_vertex_ai.py @@ -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, + ) diff --git a/tests/e2e/claude_code/test_config.yaml b/tests/e2e/claude_code/test_config.yaml index eec68d11dcf..e9253da2b3c 100644 --- a/tests/e2e/claude_code/test_config.yaml +++ b/tests/e2e/claude_code/test_config.yaml @@ -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 diff --git a/tests/e2e/coverage_registry/reliability.yaml b/tests/e2e/coverage_registry/reliability.yaml index 3746b029331..ba192c2912e 100644 --- a/tests/e2e/coverage_registry/reliability.yaml +++ b/tests/e2e/coverage_registry/reliability.yaml @@ -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"} diff --git a/tests/e2e/docker-compose.yml b/tests/e2e/docker-compose.yml index e2bb6ca8933..b64f3d8dbfd 100644 --- a/tests/e2e/docker-compose.yml +++ b/tests/e2e/docker-compose.yml @@ -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 diff --git a/tests/e2e/router/complexity_router_client.py b/tests/e2e/router/complexity_router_client.py new file mode 100644 index 00000000000..929acbb3461 --- /dev/null +++ b/tests/e2e/router/complexity_router_client.py @@ -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()) diff --git a/tests/e2e/router/conftest.py b/tests/e2e/router/conftest.py new file mode 100644 index 00000000000..e8c05520b10 --- /dev/null +++ b/tests/e2e/router/conftest.py @@ -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() diff --git a/tests/e2e/router/test_complexity_router_e2e.py b/tests/e2e/router/test_complexity_router_e2e.py new file mode 100644 index 00000000000..88d79a9cac0 --- /dev/null +++ b/tests/e2e/router/test_complexity_router_e2e.py @@ -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" + ) diff --git a/tests/local_testing/test_llm_guard.py b/tests/local_testing/test_llm_guard.py index da881e19d3c..78bbd1c0af8 100644 --- a/tests/local_testing/test_llm_guard.py +++ b/tests/local_testing/test_llm_guard.py @@ -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 diff --git a/tests/test_litellm/caching/test_caching_handler.py b/tests/test_litellm/caching/test_caching_handler.py index 01327529410..1136a0b7e7b 100644 --- a/tests/test_litellm/caching/test_caching_handler.py +++ b/tests/test_litellm/caching/test_caching_handler.py @@ -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" diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index ade2c677745..f99cfb953ba 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -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 == {} diff --git a/tests/test_litellm/litellm_core_utils/test_redact_messages.py b/tests/test_litellm/litellm_core_utils/test_redact_messages.py index 0f7f492ddb6..e1ffabb3515 100644 --- a/tests/test_litellm/litellm_core_utils/test_redact_messages.py +++ b/tests/test_litellm/litellm_core_utils/test_redact_messages.py @@ -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"), diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 807f1fe95f5..c5422e0d70f 100644 --- a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -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"]) diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index 19cb8e4c07c..816a025e11a 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 926f40a6c67..fddd8d09dfc 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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 diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py index 49cd1b71ef2..4c45eaac7b9 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index adcfff6fe9d..bba00ed1819 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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() diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 042fc107f40..17ff700791f 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -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(): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py new file mode 100644 index 00000000000..f6f29eee5bc --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py @@ -0,0 +1,2091 @@ +""" +Unit tests for the Compresr guardrail. + +Tests cover: +- apply_guardrail compresses eligible messages query-aware (tool-call intent + resolved via tool_call_id, falling back to the last user message) +- target selection: tool outputs by default, system/history opt-in, min-chars + threshold, targets without a derivable query are left uncompressed +- multimodal content: text parts replaced, non-text parts preserved +- recovery: hash marker appended, compresr_retrieve tool injected, originals + stored per litellm_call_id, agentic loop returns the original content and + rejects hashes not issued for the current request +- x-compresr-bypass header, response-type passthrough +- fail_closed raises HTTPException; fail_open forwards uncompressed +""" + +import hashlib +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, create_autospec, patch + +import httpx +import pytest +from fastapi import HTTPException + +from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY +from litellm.proxy.guardrails.guardrail_hooks.compresr.compresr import ( + COMPRESR_RETRIEVE_TOOL_NAME, + CompresrGuardrail, + _content_hash, + _extract_compresr_tool_calls, + _scoped_store_key, + has_compresr_retrieve_tool, +) +from litellm.types.utils import GenericGuardrailAPIInputs + +FAKE_API_BASE = "https://compresr.example.com" +FAKE_API_KEY = "cmp_test-key" + +TOOL_OUTPUT = "Result 1: EV range comparison. " * 40 # > 500 chars +USER_QUESTION = "Which 2026 EV has the longest range?" + +AGENT_MESSAGES = [ + {"role": "system", "content": "You are a research assistant."}, + {"role": "user", "content": USER_QUESTION}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "web_search", "arguments": '{"query": "2026 EV range"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": TOOL_OUTPUT}, +] + + +def _make_guardrail(**kwargs) -> CompresrGuardrail: + defaults = dict( + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + guardrail_name="compresr", + default_on=True, + ) + defaults.update(kwargs) + return CompresrGuardrail(**defaults) + + +def _make_single_compress_response( + compressed_context: str = "compressed summary", + original_tokens: int = 1000, + compressed_tokens: int = 400, + status: int = 200, +) -> MagicMock: + mock = MagicMock() + mock.status_code = status + mock.json.return_value = { + "success": True, + "data": { + "compressed_context": compressed_context, + "original_tokens": original_tokens, + "compressed_tokens": compressed_tokens, + "actual_compression_ratio": 0.6, + "tokens_saved": original_tokens - compressed_tokens, + "duration_ms": 42, + }, + } + mock.text = "" + return mock + + +def _make_batch_compress_response(compressed_contexts: list, status: int = 200) -> MagicMock: + mock = MagicMock() + mock.status_code = status + mock.json.return_value = { + "success": True, + "data": { + "results": [ + { + "compressed_context": ctx, + "original_tokens": 1000, + "compressed_tokens": 400, + "actual_compression_ratio": 0.6, + "tokens_saved": 600, + "duration_ms": 42, + } + for ctx in compressed_contexts + ], + "count": len(compressed_contexts), + }, + } + mock.text = "" + return mock + + +def _make_openai_response_with_tool_call(tool_name: str, arguments: dict, tool_id: str = "call_abc123") -> MagicMock: + fn = MagicMock() + fn.name = tool_name + fn.arguments = json.dumps(arguments) + + tc = MagicMock() + tc.id = tool_id + tc.type = "function" + tc.function = fn + + message = MagicMock() + message.content = None + message.tool_calls = [tc] + + choice = MagicMock() + choice.message = message + + response = MagicMock() + response.choices = [choice] + # Plain chat-completion shape: no responses-API `output` list, no + # anthropic `content` list. + response.output = None + response.content = None + return response + + +def _make_openai_response_with_tool_calls(tool_calls: list, content: object = None) -> MagicMock: + """Chat-completion response carrying several tool calls in one turn + (parallel tool calling). ``tool_calls`` items are (name, arguments, id).""" + tcs = [] + for name, arguments, tool_id in tool_calls: + fn = MagicMock() + fn.name = name + fn.arguments = json.dumps(arguments) + tc = MagicMock() + tc.id = tool_id + tc.type = "function" + tc.function = fn + tcs.append(tc) + + message = MagicMock() + message.content = content + message.tool_calls = tcs + + choice = MagicMock() + choice.message = message + + response = MagicMock() + response.choices = [choice] + response.output = None + response.content = None + return response + + +def _apply_inputs(messages: list) -> GenericGuardrailAPIInputs: + return GenericGuardrailAPIInputs(structured_messages=[dict(m) for m in messages]) + + +def _logging_obj(call_id: str) -> SimpleNamespace: + # Default fixture models a proxy with per-key auth enabled (the production + # shape). Recovery requires a caller scope; tests that need the no-auth + # path should build the object explicitly. + from litellm.proxy._types import UserAPIKeyAuth + + return SimpleNamespace( + litellm_call_id=call_id, + model_call_details={ + "litellm_params": {"metadata": {"user_api_key_auth": UserAPIKeyAuth(api_key="hash-default")}} + }, + ) + + +def _logging_obj_with_key(call_id: str, user_api_key: str, meta_key: str = "metadata") -> SimpleNamespace: + """Logging object carrying the server-set UserAPIKeyAuth object, the way the + proxy populates it for an authenticated request (the bare user_api_key + string alone is never trusted — a client could forge that).""" + from litellm.proxy._types import UserAPIKeyAuth + + return SimpleNamespace( + litellm_call_id=call_id, + model_call_details={"litellm_params": {meta_key: {"user_api_key_auth": UserAPIKeyAuth(api_key=user_api_key)}}}, + ) + + +def _retrieve_tool_call(hash_value: str, tool_id: str) -> dict: + return { + "id": tool_id, + "type": "function", + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "arguments": {"hash": hash_value}, + } + + +def _retrieve_tool_stub() -> dict: + return { + "type": "function", + "function": {"name": COMPRESR_RETRIEVE_TOOL_NAME, "parameters": {}}, + } + + +@pytest.fixture +def guardrail() -> CompresrGuardrail: + return _make_guardrail() + + +# ── init ────────────────────────────────────────────────────────────── + + +def test_init_raises_without_api_key(monkeypatch): + monkeypatch.delenv("COMPRESR_API_KEY", raising=False) + with pytest.raises(ValueError, match="API key"): + CompresrGuardrail(guardrail_name="compresr") + + +def test_init_defaults(): + g = _make_guardrail() + assert g.compresr_api_base == FAKE_API_BASE + assert g.compression_model == "latte_v2" + assert g.target_compression_ratio == 0.5 + assert g.coarse is True + assert g.min_chars_to_compress == 500 + assert g.compress_tool_outputs is True + assert g.compress_system is False + assert g.compress_history is False + assert g.compress_last_user is False + assert g.enable_retrieval is True + assert g.unreachable_fallback == "fail_closed" + + +def test_init_coerces_unknown_unreachable_fallback_to_fail_closed(): + g = _make_guardrail(unreachable_fallback="banana") + assert g.unreachable_fallback == "fail_closed" + + +# ── compression core ───────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_apply_guardrail_compresses_tool_output_with_intent_query( + guardrail: CompresrGuardrail, +): + inputs = _apply_inputs(AGENT_MESSAGES) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + _, call_kwargs = mock_post.call_args + assert call_kwargs["url"] == f"{FAKE_API_BASE}/api/compress/question-specific/" + assert call_kwargs["headers"]["X-API-Key"] == FAKE_API_KEY + payload = call_kwargs["json"] + assert payload["context"] == TOOL_OUTPUT + # Query is the tool call's intent, not the user question. + assert payload["query"] == 'web_search: {"query": "2026 EV range"}' + assert payload["compression_model_name"] == "latte_v2" + assert payload["target_compression_ratio"] == 0.5 + + out = result["structured_messages"] + assert out[3]["content"].startswith("compressed summary") + # Untouched messages pass through byte-identical. + assert out[0] == AGENT_MESSAGES[0] + assert out[1] == AGENT_MESSAGES[1] + assert out[2] == AGENT_MESSAGES[2] + + +@pytest.mark.asyncio +async def test_apply_guardrail_mirrors_compression_into_texts_channel( + guardrail: CompresrGuardrail, +): + """The /v1/responses translation writes compressed output back through the + `texts` channel, not structured_messages. Compression must be mirrored there + or that surface silently forwards the original content uncompressed.""" + messages = [ + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "call_unknown", "content": TOOL_OUTPUT}, + ] + inputs = GenericGuardrailAPIInputs( + texts=[USER_QUESTION, TOOL_OUTPUT], + structured_messages=[dict(m) for m in messages], + ) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + texts = result["texts"] + # The compressed tool output replaces the original in the texts channel... + assert texts[1].startswith("compressed summary") + assert texts[1] != TOOL_OUTPUT + # ...while untouched text passes through byte-identical. + assert texts[0] == USER_QUESTION + + +@pytest.mark.asyncio +async def test_apply_guardrail_returns_inputs_unchanged_when_nothing_compressed( + guardrail: CompresrGuardrail, +): + """A 200 response whose compressed_context is empty is a functional no-op. + The exact inputs object must come back: handlers detect guardrail edits by + identity, and a fresh structured_messages list would force a full write-back + of an untouched request (on Anthropic, reconversion strips cache_control + from thinking blocks).""" + inputs = _apply_inputs(AGENT_MESSAGES) + mock_post = AsyncMock(return_value=_make_single_compress_response(compressed_context="")) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + assert result is inputs + + +@pytest.mark.asyncio +async def test_texts_mirror_skips_duplicate_content_with_diverging_compressions( + guardrail: CompresrGuardrail, +): + """Two targets with identical text but different query-specific compressions: + the value-keyed texts mirror cannot tell the occurrences apart, so it must + leave them uncompressed rather than apply an arbitrary variant to both.""" + messages = [ + {"role": "user", "content": USER_QUESTION}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "search_docs", "arguments": '{"q": "a"}'}}, + {"id": "call_2", "type": "function", "function": {"name": "search_web", "arguments": '{"q": "b"}'}}, + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": TOOL_OUTPUT}, + {"role": "tool", "tool_call_id": "call_2", "content": TOOL_OUTPUT}, + ] + inputs = GenericGuardrailAPIInputs( + texts=[USER_QUESTION, TOOL_OUTPUT, TOOL_OUTPUT], + structured_messages=[dict(m) for m in messages], + ) + mock_post = AsyncMock(return_value=_make_batch_compress_response(["compressed for docs", "compressed for web"])) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + # Each message position still gets its own query-specific compression... + out = result["structured_messages"] + assert out[2]["content"].startswith("compressed for docs") + assert out[3]["content"].startswith("compressed for web") + # ...but the texts mirror leaves the ambiguous occurrences untouched. + assert result["texts"] == [USER_QUESTION, TOOL_OUTPUT, TOOL_OUTPUT] + + +@pytest.mark.asyncio +async def test_texts_mirror_skips_text_that_also_appears_outside_targets( + guardrail: CompresrGuardrail, +): + """compress_system is off, so a system message whose text happens to equal + a compressed tool output must not be rewritten through the texts mirror.""" + messages = [ + {"role": "system", "content": TOOL_OUTPUT}, + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "call_x", "content": TOOL_OUTPUT}, + ] + inputs = GenericGuardrailAPIInputs( + texts=[TOOL_OUTPUT, USER_QUESTION, TOOL_OUTPUT], + structured_messages=[dict(m) for m in messages], + ) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + out = result["structured_messages"] + assert out[0]["content"] == TOOL_OUTPUT # system message untouched + assert out[2]["content"].startswith("compressed summary") + # One compressed target cannot account for two occurrences in texts. + assert result["texts"] == [TOOL_OUTPUT, USER_QUESTION, TOOL_OUTPUT] + + +@pytest.mark.asyncio +async def test_tool_output_without_matching_call_uses_user_question( + guardrail: CompresrGuardrail, +): + messages = [ + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "call_unknown", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + payload = mock_post.call_args.kwargs["json"] + assert payload["query"] == USER_QUESTION + + +@pytest.mark.asyncio +async def test_function_result_without_name_does_not_bind_unrelated_call( + guardrail: CompresrGuardrail, +): + """A legacy function-role result missing its name must not adopt the intent + of an arbitrary earlier assistant function_call; it falls back to the last + user message.""" + messages = [ + {"role": "user", "content": USER_QUESTION}, + { + "role": "assistant", + "content": None, + "function_call": {"name": "get_weather", "arguments": '{"city": "Paris"}'}, + }, + {"role": "function", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + payload = mock_post.call_args.kwargs["json"] + assert payload["query"] == USER_QUESTION + + +@pytest.mark.asyncio +async def test_target_without_derivable_query_left_uncompressed( + guardrail: CompresrGuardrail, +): + # No user message and no tool-call intent anywhere -> nothing to compress. + messages = [{"role": "tool", "tool_call_id": "call_x", "content": TOOL_OUTPUT}] + mock_post = AsyncMock() + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + mock_post.assert_not_called() + assert result["structured_messages"][0]["content"] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_system_and_history_not_compressed_by_default( + guardrail: CompresrGuardrail, +): + long_system = "Rules. " * 200 + messages = [ + {"role": "system", "content": long_system}, + {"role": "user", "content": "Old question? " * 100}, + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "c1", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + # Only one (single, non-batch) call: the tool output. + assert mock_post.call_count == 1 + assert mock_post.call_args.kwargs["json"]["context"] == TOOL_OUTPUT + out = result["structured_messages"] + assert out[0]["content"] == long_system + assert out[2]["content"] == USER_QUESTION + + +@pytest.mark.asyncio +async def test_opt_in_system_uses_batch_endpoint(): + guardrail = _make_guardrail(compress_system=True) + long_system = "Rules. " * 200 + messages = [ + {"role": "system", "content": long_system}, + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "c1", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_batch_compress_response(["short system", "short tool"])) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + call_kwargs = mock_post.call_args.kwargs + assert call_kwargs["url"].endswith("/api/compress/question-specific/batch") + batch_inputs = call_kwargs["json"]["inputs"] + assert [i["context"] for i in batch_inputs] == [long_system, TOOL_OUTPUT] + out = result["structured_messages"] + assert out[0]["content"].startswith("short system") + assert out[2]["content"].startswith("short tool") + + +@pytest.mark.asyncio +async def test_short_messages_skipped(guardrail: CompresrGuardrail): + messages = [ + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "c1", "content": "tiny result"}, + ] + mock_post = AsyncMock() + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + mock_post.assert_not_called() + assert result["structured_messages"][1]["content"] == "tiny result" + + +@pytest.mark.asyncio +async def test_multimodal_text_replaced_non_text_preserved( + guardrail: CompresrGuardrail, +): + image_part = {"type": "image_url", "image_url": {"url": "https://example.com/x.png"}} + messages = [ + {"role": "user", "content": USER_QUESTION}, + { + "role": "tool", + "tool_call_id": "c1", + "content": [{"type": "text", "text": TOOL_OUTPUT}, image_part], + }, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + content = result["structured_messages"][1]["content"] + assert isinstance(content, list) + assert content[0]["type"] == "text" + assert content[0]["text"].startswith("compressed summary") + assert content[1] == image_part + + +# ── passthrough / bypass ───────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_bypass_header_skips_compression_when_allowed(): + guardrail = _make_guardrail(allow_bypass_header=True) + inputs = _apply_inputs(AGENT_MESSAGES) + mock_post = AsyncMock() + request_data = { + "model": "gpt-4o", + "proxy_server_request": {"headers": {"x-compresr-bypass": "true"}}, + } + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request") + + mock_post.assert_not_called() + assert result is inputs + + +@pytest.mark.asyncio +async def test_bypass_header_ignored_by_default(guardrail: CompresrGuardrail): + inputs = _apply_inputs(AGENT_MESSAGES) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + request_data = { + "model": "gpt-4o", + "proxy_server_request": {"headers": {"x-compresr-bypass": "true"}}, + } + + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request") + + mock_post.assert_called_once() + + +@pytest.mark.asyncio +async def test_response_input_type_passthrough(guardrail: CompresrGuardrail): + inputs = _apply_inputs(AGENT_MESSAGES) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="response") + assert result is inputs + + +@pytest.mark.asyncio +async def test_missing_structured_messages_passthrough(guardrail: CompresrGuardrail): + inputs = GenericGuardrailAPIInputs(texts=["hello"]) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request") + assert result is inputs + + +# ── failure policy ──────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_transport_error_raises_when_fail_closed(guardrail: CompresrGuardrail): + with patch.object( + guardrail.async_handler, + "post", + AsyncMock(side_effect=httpx.ConnectError("boom")), + ): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert exc_info.value.status_code == 502 + + +@pytest.mark.asyncio +async def test_transport_error_fail_open_forwards_uncompressed(): + guardrail = _make_guardrail(unreachable_fallback="fail_open") + inputs = _apply_inputs(AGENT_MESSAGES) + + with patch.object( + guardrail.async_handler, + "post", + AsyncMock(side_effect=httpx.ConnectError("boom")), + ): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + assert result is inputs + assert result["structured_messages"][3]["content"] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_non_json_response_raises_when_fail_closed(guardrail: CompresrGuardrail): + mock = MagicMock() + mock.status_code = 200 + mock.json.side_effect = ValueError("not json") + mock.text = "gateway error" + + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock)): + with pytest.raises(HTTPException): + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + +@pytest.mark.asyncio +async def test_http_exception_does_not_reflect_upstream_body(guardrail: CompresrGuardrail): + mock = MagicMock() + mock.status_code = 500 + mock.json.side_effect = ValueError("not json") + mock.text = "SECRET_INSTANCE_METADATA_TOKEN=aws-imds-response" + + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=mock)): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert "SECRET_INSTANCE_METADATA_TOKEN" not in json.dumps(exc_info.value.detail) + + +def test_init_rejects_non_http_api_base(): + with pytest.raises(ValueError, match="scheme"): + CompresrGuardrail( + api_base="file:///etc/passwd", + api_key=FAKE_API_KEY, + guardrail_name="compresr", + ) + + +def test_init_rejects_cloud_metadata_api_base(): + with pytest.raises(ValueError, match="metadata"): + CompresrGuardrail( + api_base="http://169.254.169.254", + api_key=FAKE_API_KEY, + guardrail_name="compresr", + ) + + +@pytest.mark.parametrize( + "api_base", + [ + "http://2852039166", # decimal encoding of 169.254.169.254 + "http://0xa9fea9fe", # hex encoding + "http://[::ffff:169.254.169.254]", # IPv4-mapped IPv6 + "http://metadata.azure.com", + "http://metadata.azure.internal", + "http://168.63.129.16", # Azure WireServer + ], +) +def test_init_rejects_encoded_cloud_metadata_api_base(api_base): + with pytest.raises(ValueError, match="metadata"): + CompresrGuardrail( + api_base=api_base, + api_key=FAKE_API_KEY, + guardrail_name="compresr", + ) + + +@pytest.mark.asyncio +async def test_apply_guardrail_ignores_user_supplied_call_id(guardrail: CompresrGuardrail): + mock_post = AsyncMock(return_value=_make_single_compress_response()) + attacker_call_id = "victim-tenant-call-id" + + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o", "litellm_call_id": attacker_call_id}, + input_type="request", + logging_obj=_logging_obj("real-framework-call-id"), + ) + + assert not any(attacker_call_id in k for k in guardrail._originals_by_call_id) + assert any(k.endswith("real-framework-call-id") for k in guardrail._originals_by_call_id) + + +@pytest.mark.asyncio +async def test_agentic_plan_ignores_user_supplied_call_id(guardrail: CompresrGuardrail): + hash_value = "d" * 24 + guardrail._store_originals("victim-tenant-call-id", {hash_value: "victim-original"}) + + plan = await guardrail.async_build_agentic_loop_plan( + tools={ + "tool_calls": [ + { + "id": "call_abc", + "type": "function", + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "arguments": {"hash": hash_value}, + } + ] + }, + model="gpt-4o", + messages=[], + response=_make_openai_response_with_tool_call( + COMPRESR_RETRIEVE_TOOL_NAME, {"hash": hash_value}, tool_id="call_abc" + ), + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("attacker-call-id"), + stream=False, + kwargs={"litellm_call_id": "victim-tenant-call-id"}, + ) + + # Attacker's scope resolves nothing, so the loop is vetoed and the victim + # original never surfaces. + assert plan.run_agentic_loop is False + assert plan.request_patch is None + + +@pytest.mark.asyncio +async def test_recovery_store_partitioned_by_caller_identity(guardrail: CompresrGuardrail): + """Two tenants that set the SAME client-forgeable x-litellm-call-id must not + read each other's stored originals, and each still reads its own.""" + shared_call_id = "shared-call-id" + expected_hash = hashlib.sha256(TOOL_OUTPUT.encode()).hexdigest()[:24] + + async def _plan_for(user_api_key: str, tool_id: str): + return await guardrail.async_build_agentic_loop_plan( + tools={"tool_calls": [_retrieve_tool_call(expected_hash, tool_id)]}, + model="gpt-4o", + messages=[], + response=_make_openai_response_with_tool_call( + COMPRESR_RETRIEVE_TOOL_NAME, {"hash": expected_hash}, tool_id=tool_id + ), + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj_with_key(shared_call_id, user_api_key), + stream=False, + kwargs={}, + ) + + # Tenant A compresses and stores its original. + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=_make_single_compress_response())): + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj_with_key(shared_call_id, "hash-tenant-A"), + ) + + # Tenant B, same call id, different virtual-key hash → different bucket, so + # nothing resolves and the loop is vetoed (Tenant A's original never leaks). + plan_b = await _plan_for("hash-tenant-B", "call_b") + assert plan_b.run_agentic_loop is False + assert plan_b.request_patch is None + + # Tenant A retrieves its own content successfully. + plan_a = await _plan_for("hash-tenant-A", "call_a") + assert TOOL_OUTPUT in plan_a.request_patch.messages[-1]["content"] + + +@pytest.mark.asyncio +async def test_caller_scope_read_from_litellm_metadata(guardrail: CompresrGuardrail): + """/v1/messages and /v1/responses carry the auth object under + litellm_metadata rather than metadata; the store key must be scoped by it + there too, without relying on upstream's metadata backfill.""" + logging_obj = _logging_obj_with_key("call-lm", "hash-tenant-lm", meta_key="litellm_metadata") + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=logging_obj, + ) + + assert "hash-tenant-lm\x00call-lm" in guardrail._originals_by_call_id + assert "call-lm" not in guardrail._originals_by_call_id + + +@pytest.mark.asyncio +async def test_caller_scope_rejects_forged_user_api_key_string(guardrail: CompresrGuardrail): + """A client-supplied metadata.user_api_key STRING (no server-set + UserAPIKeyAuth object) must not be trusted as a tenant scope — otherwise a + caller could forge another tenant's recovery bucket on /v1/messages.""" + logging_obj = SimpleNamespace( + litellm_call_id="call-forge", + model_call_details={"litellm_params": {"metadata": {"user_api_key": "victim-tenant-hash"}}}, + ) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=logging_obj, + ) + + # Forged string is ignored: scope resolves to empty, so recovery is + # disabled entirely (no bucket keyed on victim-tenant-hash, no unscoped + # bucket that another caller could reuse). + assert not guardrail._originals_by_call_id + + +@pytest.mark.asyncio +async def test_compress_post_called_with_real_handler_signature(): + """AsyncHTTPHandler.post has a fixed signature; an autospec mock enforces it + (unlike AsyncMock(spec=...), which silently accepts any kwarg) so a kwarg the + real handler rejects — which would raise TypeError past the fail policy — + fails the test instead.""" + guardrail = _make_guardrail() + autospec_post = create_autospec(guardrail.async_handler.post, return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", autospec_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + assert result["structured_messages"][3]["content"].startswith("compressed summary") + + +@pytest.mark.asyncio +async def test_batch_result_count_mismatch_raises_when_fail_closed(): + guardrail = _make_guardrail(compress_system=True) + messages = [ + {"role": "system", "content": "Rules. " * 200}, + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "c1", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_batch_compress_response(["only one"])) + + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException): + await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + +# ── recovery (compresr_retrieve) ────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_recovery_marker_tool_injection_and_original_stored( + guardrail: CompresrGuardrail, +): + inputs = _apply_inputs(AGENT_MESSAGES) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + compressed_content = result["structured_messages"][3]["content"] + expected_hash = hashlib.sha256(TOOL_OUTPUT.encode()).hexdigest()[:24] + assert f"compresr hash={expected_hash}" in compressed_content + + tools = result.get("tools") + assert tools is not None and has_compresr_retrieve_tool(tools) + + scoped_key = next(k for k in guardrail._originals_by_call_id if k.endswith("call-id-1")) + originals, _expiry = guardrail._originals_by_call_id[scoped_key] + assert originals[expected_hash] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_enable_retrieval_false_no_marker_no_tool(): + guardrail = _make_guardrail(enable_retrieval=False) + inputs = _apply_inputs(AGENT_MESSAGES) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + assert result["structured_messages"][3]["content"] == "compressed summary" + assert not has_compresr_retrieve_tool(result.get("tools") or []) + assert guardrail._originals_by_call_id == {} + + +@pytest.mark.asyncio +async def test_existing_tools_preserved_when_injecting(guardrail: CompresrGuardrail): + existing_tool = {"type": "function", "function": {"name": "my_tool", "parameters": {}}} + inputs = GenericGuardrailAPIInputs( + structured_messages=[dict(m) for m in AGENT_MESSAGES], + tools=[existing_tool], + ) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + tools = result["tools"] + assert existing_tool in tools + assert has_compresr_retrieve_tool(tools) + assert len(tools) == 2 + + +@pytest.mark.asyncio +async def test_non_list_tools_left_untouched_when_injecting(guardrail: CompresrGuardrail): + # An unexpected non-list tools value must survive unchanged rather than be + # clobbered by the injected retrieve tool. + odd_tools = {"type": "function", "function": {"name": "my_tool"}} + inputs = GenericGuardrailAPIInputs( + structured_messages=[dict(m) for m in AGENT_MESSAGES], + tools=odd_tools, # type: ignore[typeddict-item] + ) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + + assert result["tools"] is odd_tools + + +def test_extract_compresr_tool_calls_tolerates_missing_keys(): + # A retrieve call missing id/arguments must not KeyError in the post-call + # hook; it extracts with safe defaults and resolves to a rejection later. + with patch( + "litellm.proxy.guardrails.guardrail_hooks.compresr.compresr.get_tool_calls_from_response", + return_value=[{"name": COMPRESR_RETRIEVE_TOOL_NAME}, {"id": "x"}], + ): + extracted = _extract_compresr_tool_calls(object()) + + assert extracted == [{"id": None, "type": "function", "name": COMPRESR_RETRIEVE_TOOL_NAME, "arguments": {}}] + + +# ── agentic loop ────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_async_should_run_agentic_loop_true_for_retrieve_call( + guardrail: CompresrGuardrail, +): + response = _make_openai_response_with_tool_call(COMPRESR_RETRIEVE_TOOL_NAME, {"hash": "a" * 24}) + tools = [dict(t) for t in [_retrieve_tool_stub()]] + + should_run, gate_tools = await guardrail.async_should_run_agentic_loop( + response=response, + model="gpt-4o", + messages=[], + tools=tools, + stream=False, + custom_llm_provider="openai", + kwargs={}, + ) + assert should_run is True + assert gate_tools["tool_calls"][0]["name"] == COMPRESR_RETRIEVE_TOOL_NAME + + +@pytest.mark.asyncio +async def test_async_should_run_agentic_loop_false_without_retrieve_tool( + guardrail: CompresrGuardrail, +): + response = _make_openai_response_with_tool_call("other_tool", {"x": 1}) + should_run, _ = await guardrail.async_should_run_agentic_loop( + response=response, + model="gpt-4o", + messages=[], + tools=[], + stream=False, + custom_llm_provider="openai", + kwargs={}, + ) + assert should_run is False + + +@pytest.mark.asyncio +async def test_agentic_plan_returns_stored_original(guardrail: CompresrGuardrail): + hash_value = hashlib.sha256(TOOL_OUTPUT.encode()).hexdigest()[:24] + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {hash_value: TOOL_OUTPUT}) + + plan = await guardrail.async_build_agentic_loop_plan( + tools={ + "tool_calls": [ + { + "id": "call_abc", + "type": "function", + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "arguments": {"hash": hash_value}, + } + ] + }, + model="gpt-4o", + messages=[{"role": "user", "content": "q"}], + response=_make_openai_response_with_tool_call( + COMPRESR_RETRIEVE_TOOL_NAME, {"hash": hash_value}, tool_id="call_abc" + ), + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + + assert plan.run_agentic_loop is True + follow_up = plan.request_patch.messages + tool_result = follow_up[-1] + assert tool_result["role"] == "tool" + assert tool_result["tool_call_id"] == "call_abc" + assert tool_result["content"] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_agentic_plan_preserves_list_shaped_assistant_text(guardrail: CompresrGuardrail): + """Some providers return chat assistant content as list-of-parts; the + retrieval follow-up must keep that text, not drop it to None.""" + hash_value = "a" * 24 + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {hash_value: "original"}) + response = _make_openai_response_with_tool_calls( + [(COMPRESR_RETRIEVE_TOOL_NAME, {"hash": hash_value}, "call_r")], + content=[{"type": "text", "text": "Let me fetch the original."}], + ) + + plan = await guardrail.async_build_agentic_loop_plan( + tools={"tool_calls": [_retrieve_tool_call(hash_value, "call_r")]}, + model="gpt-4o", + messages=[{"role": "user", "content": "q"}], + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + + assistant_message = plan.request_patch.messages[-2] + assert assistant_message["role"] == "assistant" + assert assistant_message["content"] == "Let me fetch the original." + + +@pytest.mark.asyncio +async def test_agentic_plan_strips_other_guardrails_executed_markers(guardrail: CompresrGuardrail): + # The retrieval follow-up restores content other pre-call guardrails may + # never have inspected, so their executed markers must not be replayed; + # only this guardrail's own marker survives (no recompression loop). + hash_value = "a" * 24 + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {hash_value: TOOL_OUTPUT}) + response = _make_openai_response_with_tool_calls([(COMPRESR_RETRIEVE_TOOL_NAME, {"hash": hash_value}, "call_r")]) + own_marker = guardrail._pre_call_marker() + assert own_marker is not None + kwargs = { + "metadata": { + "user_api_key": "key-hash", + PRE_CALL_EXECUTED_GUARDRAILS_KEY: [own_marker, "token:pii_guardrail"], + }, + "litellm_metadata": {PRE_CALL_EXECUTED_GUARDRAILS_KEY: ["token:other_guardrail"]}, + } + + plan = await guardrail.async_build_agentic_loop_plan( + tools={"tool_calls": [_retrieve_tool_call(hash_value, "call_r")]}, + model="gpt-4o", + messages=[{"role": "user", "content": "q"}], + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs=kwargs, + ) + + out = plan.request_patch.kwargs + assert out["metadata"][PRE_CALL_EXECUTED_GUARDRAILS_KEY] == [own_marker] + assert out["metadata"]["user_api_key"] == "key-hash" + assert PRE_CALL_EXECUTED_GUARDRAILS_KEY not in out["litellm_metadata"] + assert kwargs["metadata"][PRE_CALL_EXECUTED_GUARDRAILS_KEY] == [own_marker, "token:pii_guardrail"] + assert kwargs["litellm_metadata"][PRE_CALL_EXECUTED_GUARDRAILS_KEY] == ["token:other_guardrail"] + + +@pytest.mark.asyncio +async def test_agentic_plan_rejects_hash_from_other_request( + guardrail: CompresrGuardrail, +): + hash_value = "b" * 24 + guardrail._store_originals("someone-elses-call", {hash_value: "secret"}) + + plan = await guardrail.async_build_agentic_loop_plan( + tools={ + "tool_calls": [ + { + "id": "call_abc", + "type": "function", + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "arguments": {"hash": hash_value}, + } + ] + }, + model="gpt-4o", + messages=[], + response=_make_openai_response_with_tool_call( + COMPRESR_RETRIEVE_TOOL_NAME, {"hash": hash_value}, tool_id="call_abc" + ), + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("my-call"), + stream=False, + kwargs={}, + ) + + # Hash belongs to another caller's scope; the loop is vetoed and the secret + # never surfaces. + assert plan.run_agentic_loop is False + assert plan.request_patch is None + + +@pytest.mark.asyncio +async def test_agentic_loop_vetoed_when_no_recovery_state(guardrail: CompresrGuardrail): + # A caller-defined compresr_retrieve tool with no stored original must not + # trigger an extra provider round-trip. + hash_value = "f" * 24 + response = _make_openai_response_with_tool_call(COMPRESR_RETRIEVE_TOOL_NAME, {"hash": hash_value}, tool_id="call_x") + plan = await guardrail.async_build_agentic_loop_plan( + tools={"tool_calls": [_retrieve_tool_call(hash_value, "call_x")]}, + model="gpt-4o", + messages=[{"role": "user", "content": "q"}], + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + assert plan.run_agentic_loop is False + assert plan.request_patch is None + + +@pytest.mark.asyncio +async def test_agentic_loop_dedupes_repeated_retrievals(guardrail: CompresrGuardrail): + # Retrieving the same marker many times expands the original once; repeats + # get a short marker (no follow-up amplification). + hash_value = "a" * 24 + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {hash_value: TOOL_OUTPUT}) + calls = [(COMPRESR_RETRIEVE_TOOL_NAME, {"hash": hash_value}, f"call_{i}") for i in range(5)] + plan = await guardrail.async_build_agentic_loop_plan( + tools={"tool_calls": [_retrieve_tool_call(hash_value, f"call_{i}") for i in range(5)]}, + model="gpt-4o", + messages=[{"role": "user", "content": "q"}], + response=_make_openai_response_with_tool_calls(calls), + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + tool_results = [m for m in plan.request_patch.messages if m.get("role") == "tool"] + assert len(tool_results) == 5 + assert sum(1 for m in tool_results if m["content"] == TOOL_OUTPUT) == 1 + assert all("already retrieved" in m["content"] for m in tool_results if m["content"] != TOOL_OUTPUT) + + +@pytest.mark.asyncio +async def test_agentic_loop_caps_retrieval_count(guardrail: CompresrGuardrail): + # Beyond _MAX_RETRIEVALS_PER_LOOP retrievals, extra calls get a bounded marker. + n = 10 + hashes = [f"{i:024x}" for i in range(n)] + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {h: f"original-{h}" for h in hashes}) + calls = [(COMPRESR_RETRIEVE_TOOL_NAME, {"hash": h}, f"call_{i}") for i, h in enumerate(hashes)] + plan = await guardrail.async_build_agentic_loop_plan( + tools={"tool_calls": [_retrieve_tool_call(h, f"call_{i}") for i, h in enumerate(hashes)]}, + model="gpt-4o", + messages=[{"role": "user", "content": "q"}], + response=_make_openai_response_with_tool_calls(calls), + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + tool_results = [m for m in plan.request_patch.messages if m.get("role") == "tool"] + assert len(tool_results) == n + over_limit = [m for m in tool_results if "retrieval limit reached" in m["content"]] + assert len(over_limit) == n - 8 # only the first 8 expand + + +def test_display_hash_strips_control_characters(): + """The compresr_retrieve `hash` argument is model/tool-output-influenced, so + control characters (newlines, ANSI escapes) must be stripped — not just + length-capped — before it is echoed into logs or the fallback message.""" + from litellm.proxy.guardrails.guardrail_hooks.compresr.compresr import _display_hash + + assert _display_hash("a" * 24) == "a" * 24 # a real marker hash passes through + assert "\n" not in _display_hash("abc\ndef\rFORGED LOG LINE") + assert "\x1b" not in _display_hash("hash\x1b[31mred") + capped = _display_hash("z" * 100) + assert capped.endswith("…") and len(capped) <= 33 + + +@pytest.mark.asyncio +async def test_agentic_plan_builds_anthropic_followup_shape( + guardrail: CompresrGuardrail, +): + hash_value = "c" * 24 + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {hash_value: TOOL_OUTPUT}) + + response = MagicMock() + response.output = None + response.content = [{"type": "tool_use", "id": "toolu_1"}] + + plan = await guardrail.async_build_agentic_loop_plan( + tools={ + "tool_calls": [ + { + "id": "toolu_1", + "type": "function", + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "arguments": {"hash": hash_value}, + } + ] + }, + model="claude-sonnet-5", + messages=[{"role": "user", "content": "q"}], + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={"max_tokens": 1024}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + + follow_up = plan.request_patch.messages + assistant_msg, user_msg = follow_up[-2], follow_up[-1] + assert assistant_msg["role"] == "assistant" + assert assistant_msg["content"][0]["type"] == "tool_use" + assert user_msg["content"][0]["type"] == "tool_result" + assert user_msg["content"][0]["tool_use_id"] == "toolu_1" + assert user_msg["content"][0]["content"] == TOOL_OUTPUT + assert plan.request_patch.max_tokens == 1024 + + +@pytest.mark.asyncio +async def test_agentic_plan_builds_responses_followup_shape( + guardrail: CompresrGuardrail, +): + """The /v1/responses path echoes the function_call and pairs it with a + function_call_output keyed by the same call_id.""" + hash_value = "e" * 24 + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {hash_value: TOOL_OUTPUT}) + + response = MagicMock() + response.output = [{"type": "function_call", "call_id": "fc_1"}] # responses-API shape + response.content = None + + plan = await guardrail.async_build_agentic_loop_plan( + tools={ + "tool_calls": [ + { + "id": "fc_1", + "type": "function", + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "arguments": {"hash": hash_value}, + } + ] + }, + model="gpt-4o", + messages=[{"role": "user", "content": "q"}], + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + + call_item, output_item = plan.request_patch.messages[-2], plan.request_patch.messages[-1] + assert call_item["type"] == "function_call" + assert call_item["call_id"] == "fc_1" + assert output_item["type"] == "function_call_output" + assert output_item["call_id"] == "fc_1" + assert output_item["output"] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_agentic_plan_chat_parallel_tool_calls_echoes_only_retrieve( + guardrail: CompresrGuardrail, +): + """When the model calls a real tool alongside compresr_retrieve in one turn, + only the retrieve call may be echoed in the reconstructed assistant message: + every echoed tool_call must have a matching tool result or the provider 400s. + The real call is re-planned by the follow-up; the assistant text is kept.""" + hash_value = "f" * 24 + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {hash_value: TOOL_OUTPUT}) + + response = _make_openai_response_with_tool_calls( + [ + ("get_weather", {"city": "Paris"}, "call_weather"), + (COMPRESR_RETRIEVE_TOOL_NAME, {"hash": hash_value}, "call_retrieve"), + ], + content="Let me expand that note and check the weather.", + ) + + plan = await guardrail.async_build_agentic_loop_plan( + tools={"tool_calls": [_retrieve_tool_call(hash_value, "call_retrieve")]}, + model="gpt-4o", + messages=[{"role": "user", "content": "q"}], + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + + follow_up = plan.request_patch.messages + assistant_msg = follow_up[-2] + echoed_ids = {tc["id"] for tc in assistant_msg["tool_calls"]} + result_ids = {m["tool_call_id"] for m in follow_up if m.get("role") == "tool"} + # get_weather is not echoed; every echoed tool_call is answered. + assert echoed_ids == {"call_retrieve"} + assert echoed_ids == result_ids + assert assistant_msg["content"] == "Let me expand that note and check the weather." + assert follow_up[-1]["content"] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_agentic_plan_anthropic_parallel_preserves_text_and_balances( + guardrail: CompresrGuardrail, +): + """Anthropic parallel-tool-call turn: the assistant text is preserved, the + real tool_use is dropped (re-planned), and the reconstructed turn stays + balanced — one tool_result per echoed tool_use.""" + hash_value = "a" * 23 + "9" + guardrail._store_originals(_scoped_store_key(_logging_obj("call-id-1")), {hash_value: TOOL_OUTPUT}) + + response = MagicMock() + response.output = None + response.content = [ + {"type": "text", "text": "Checking the weather and expanding the note."}, + {"type": "tool_use", "id": "toolu_weather", "name": "get_weather", "input": {"city": "Paris"}}, + { + "type": "tool_use", + "id": "toolu_retrieve", + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "input": {"hash": hash_value}, + }, + ] + + plan = await guardrail.async_build_agentic_loop_plan( + tools={"tool_calls": [_retrieve_tool_call(hash_value, "toolu_retrieve")]}, + model="claude-sonnet-5", + messages=[{"role": "user", "content": "q"}], + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={"max_tokens": 1024}, + logging_obj=_logging_obj("call-id-1"), + stream=False, + kwargs={}, + ) + + assistant_msg, user_msg = plan.request_patch.messages[-2], plan.request_patch.messages[-1] + assert assistant_msg["content"][0] == { + "type": "text", + "text": "Checking the weather and expanding the note.", + } + echoed_ids = [b["id"] for b in assistant_msg["content"] if b["type"] == "tool_use"] + answered_ids = [b["tool_use_id"] for b in user_msg["content"]] + # get_weather dropped; balanced tool_use/tool_result pairing. + assert echoed_ids == ["toolu_retrieve"] + assert answered_ids == echoed_ids + + +# ── store hygiene ───────────────────────────────────────────────────── + + +def test_originals_store_prunes_expired(guardrail: CompresrGuardrail): + guardrail._originals_by_call_id["old"] = ({"a" * 24: "x"}, 0.0) # already expired + guardrail._store_originals("new", {"b" * 24: "y"}) + assert "old" not in guardrail._originals_by_call_id + assert "new" in guardrail._originals_by_call_id + + +def test_originals_store_caps_tracked_calls(guardrail: CompresrGuardrail): + for i in range(300): + guardrail._store_originals(f"call-{i}", {("%024x" % i): "x"}) + assert len(guardrail._originals_by_call_id) <= 256 + # Most recent entries survive. + assert "call-299" in guardrail._originals_by_call_id + + +def test_originals_store_caps_bytes_per_call(): + guardrail = _make_guardrail(max_bytes_per_call=1000) + hashes = tuple(f"{i:024x}" for i in range(5)) + values = tuple("x" * 400 for _ in range(5)) + guardrail._store_originals("c", dict(zip(hashes, values))) + + stored, _expiry = guardrail._originals_by_call_id["c"] + assert sum(len(v.encode("utf-8")) for v in stored.values()) <= 1000 + # Oldest entries are evicted first; newest survives. + assert hashes[-1] in stored + assert hashes[0] not in stored + + +def test_originals_store_byte_cap_survives_lone_surrogates(): + # Regression: eviction path must use surrogatepass to match the hash + # function; a bare encode("utf-8") crashed on lone surrogates. + guardrail = _make_guardrail(max_bytes_per_call=500) + surrogate_value = "\ud800" * 60 + hashes = tuple(f"{i:024x}" for i in range(3)) + guardrail._store_originals("c", dict(zip(hashes, (surrogate_value, surrogate_value, surrogate_value)))) + + stored, _expiry = guardrail._originals_by_call_id["c"] + assert hashes[-1] in stored + assert hashes[0] not in stored + + +def test_originals_store_caps_total_bytes_across_calls(monkeypatch: pytest.MonkeyPatch): + # Global byte budget: many distinct call ids must not retain unbounded memory. + monkeypatch.setattr( + "litellm.proxy.guardrails.guardrail_hooks.compresr.compresr._MAX_TOTAL_STORE_BYTES", + 10_000, + ) + guardrail = _make_guardrail(max_bytes_per_call=4_000) + for i in range(20): + guardrail._store_originals(f"call-{i}", {f"{i:024x}": "x" * 3_000}) + + total = sum( + len(v.encode("utf-8")) + for originals, _expiry in guardrail._originals_by_call_id.values() + for v in originals.values() + ) + assert total <= 10_000 + assert guardrail._store_total_bytes == total # running counter stays exact + # Oldest calls evicted; the most-recent call's originals survive. + assert "call-0" not in guardrail._originals_by_call_id + assert "call-19" in guardrail._originals_by_call_id + + +def test_originals_store_global_cap_keeps_current_when_single_call_is_large( + monkeypatch: pytest.MonkeyPatch, +): + # One call over the global cap is still kept (only max_bytes_per_call trims it); + # global eviction never empties the store. + monkeypatch.setattr( + "litellm.proxy.guardrails.guardrail_hooks.compresr.compresr._MAX_TOTAL_STORE_BYTES", + 1_000, + ) + guardrail = _make_guardrail(max_bytes_per_call=5_000) + guardrail._store_originals("solo", {f"{0:024x}": "x" * 4_000}) + assert "solo" in guardrail._originals_by_call_id + + +def test_recovery_markers_respect_per_call_byte_cap(): + # Regression: markers were built from every original before _store_originals + # applied the byte cap, so an evicted original left a dangling marker the + # model could never retrieve. Recovery must be attached only for originals + # that fit the cap, so every shipped marker stays retrievable. + guardrail = _make_guardrail(max_bytes_per_call=1000) + contexts = ["a" * 400, "b" * 400, "c" * 400] + messages = [{"role": "tool", "content": text} for text in contexts] + results = [{"compressed_context": f"small-{i}"} for i in range(3)] + + applied = guardrail._apply_compression_results( + messages, [0, 1, 2], contexts, results, recovery_enabled=True + ) + + # 400 + 400 fit under 1000; the third (which would reach 1200) is skipped. + assert applied.messages_compressed == 3 + assert len(applied.originals) == 2 + third_hash = _content_hash("c" * 400) + assert third_hash not in applied.originals + assert f"compresr hash={third_hash}" not in applied.compressed_messages[2]["content"] + + # Every marker still shipped must resolve to a stored original. + guardrail._store_originals("c", applied.originals) + for hash_value in applied.originals: + assert guardrail._retrieve_original("c", hash_value) is not None + assert f"compresr hash={hash_value}" in "".join( + str(m["content"]) for m in applied.compressed_messages + ) + + +def test_recovery_markers_respect_byte_cap_across_reused_store_key(): + # Regression: the per-call budget must also count bytes already stored under + # the same store key (a later turn reusing the call id). Otherwise merging + # this call's originals with the existing entry overflows the cap and + # _store_originals evicts an original this call just shipped a marker for. + guardrail = _make_guardrail(max_bytes_per_call=100) + old_hash = _content_hash("A" * 50) + guardrail._store_originals("k", {old_hash: "A" * 50}) + guardrail._store_originals("k", {_content_hash("C" * 40): "C" * 40}) + existing = guardrail._originals_by_call_id["k"][0] + + # This turn recompresses the same "A" (already stored) plus a new "D". + contexts = ["D" * 40, "A" * 50] + messages = [{"role": "tool", "content": text} for text in contexts] + results = [{"compressed_context": "dd"}, {"compressed_context": "aa"}] + applied = guardrail._apply_compression_results( + messages, [0, 1], contexts, results, recovery_enabled=True, existing_originals=existing + ) + + guardrail._store_originals("k", applied.originals) + # No marker shipped this turn may dangle after the store enforces the cap. + for hash_value in applied.originals: + assert guardrail._retrieve_original("k", hash_value) is not None + assert f"compresr hash={hash_value}" in "".join( + str(m["content"]) for m in applied.compressed_messages + ) + # The zero-cost repeat of an already-stored original stays retrievable. + assert old_hash in applied.originals + assert guardrail._retrieve_original("k", old_hash) is not None + + +# ── dynamic (adaptive) compression — latte_v2 Kneedle ───────────────── + + +@pytest.mark.asyncio +async def test_dynamic_flag_in_payload(): + """dynamic=True must appear in the compress payload; unset bounds omitted.""" + guardrail = _make_guardrail(dynamic=True) + messages = [ + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "call_x", "name": "search", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + payload = mock_post.call_args.kwargs["json"] + assert payload["dynamic"] is True + assert "dynamic_min_ratio" not in payload + assert "dynamic_max_ratio" not in payload + + +@pytest.mark.asyncio +async def test_dynamic_bounds_in_payload_when_set(): + guardrail = _make_guardrail(dynamic=True, dynamic_min_ratio=2.0, dynamic_max_ratio=8.0) + messages = [ + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "call_x", "name": "search", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + payload = mock_post.call_args.kwargs["json"] + assert payload["dynamic"] is True + assert payload["dynamic_min_ratio"] == 2.0 + assert payload["dynamic_max_ratio"] == 8.0 + + +@pytest.mark.asyncio +async def test_dynamic_on_by_default(): + guardrail = _make_guardrail() # dynamic defaults on (latte_v2) + messages = [ + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "call_x", "name": "search", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert mock_post.call_args.kwargs["json"]["dynamic"] is True + + +# ── generic passthrough compression params ──────────────────────────── + + +@pytest.mark.asyncio +async def test_compression_params_passthrough_in_payload(): + """Extra params in compression_params are forwarded verbatim; named fields + still win on collision.""" + guardrail = _make_guardrail(compression_params={"heuristic_chunking": True, "coarse": False}) + messages = [ + {"role": "user", "content": USER_QUESTION}, + {"role": "tool", "tool_call_id": "call_x", "name": "search", "content": TOOL_OUTPUT}, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + payload = mock_post.call_args.kwargs["json"] + assert payload["heuristic_chunking"] is True + # named `coarse` (default True) wins over the passthrough's coarse=False + assert payload["coarse"] is True + + +@pytest.mark.asyncio +async def test_compression_params_cannot_override_request_content_fields(): + """context/query/inputs carry the actual content being compressed; a + passthrough collision on them must be dropped, not silently win.""" + guardrail = _make_guardrail( + compression_params={ + "context": "injected", + "query": "injected", + "inputs": [], + "heuristic_chunking": True, + } + ) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + payload = mock_post.call_args.kwargs["json"] + assert payload["context"] == TOOL_OUTPUT + assert payload["query"] == 'web_search: {"query": "2026 EV range"}' + assert "inputs" not in payload + assert payload["heuristic_chunking"] is True + + +# ── compress_last_user ──────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_compress_last_user_compresses_with_verbatim_query(): + """compress_last_user=True compresses the last user message, but the query + sent to Compresr is still the original verbatim user text.""" + guardrail = _make_guardrail(compress_last_user=True) + long_question = "Which 2026 EV has the longest range? " * 20 # > 500 chars + messages = [{"role": "user", "content": long_question}] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + payload = mock_post.call_args.kwargs["json"] + assert payload["context"] == long_question + assert payload["query"] == long_question # verbatim, not the compressed text + assert result["structured_messages"][0]["content"] == "compressed summary" + + +# ── malformed-but-200 token stats (must not defeat fail policy) ─────── + + +@pytest.mark.asyncio +async def test_non_numeric_token_stats_do_not_raise(guardrail: CompresrGuardrail): + """A 200 response with non-numeric token counts must not raise: _call_compress + already succeeded, so a bare int() here would 500 even under fail policy.""" + resp = _make_single_compress_response() + resp.json.return_value["data"]["original_tokens"] = "not-a-number" + resp.json.return_value["data"]["compressed_tokens"] = None + + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=resp)): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + assert result["structured_messages"][3]["content"].startswith("compressed summary") + + +# ── HTTP status errors (non-2xx from the shared handler) ────────────── + + +def _http_status_error(status: int = 500, text: str = "upstream error body") -> httpx.HTTPStatusError: + # The shared AsyncHTTPHandler.post() raises HTTPStatusError on any non-2xx, + # carrying the upstream body and request headers; this simulates that. + request = httpx.Request("POST", f"{FAKE_API_BASE}/api/compress/question-specific/") + response = httpx.Response(status, text=text, request=request) + return httpx.HTTPStatusError(str(status), request=request, response=response) + + +@pytest.mark.asyncio +async def test_http_status_error_raises_when_fail_closed(guardrail: CompresrGuardrail): + with patch.object(guardrail.async_handler, "post", AsyncMock(side_effect=_http_status_error(500))): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert exc_info.value.status_code == 502 + + +@pytest.mark.asyncio +async def test_http_status_error_fail_open_forwards_uncompressed(): + guardrail = _make_guardrail(unreachable_fallback="fail_open") + inputs = _apply_inputs(AGENT_MESSAGES) + with patch.object(guardrail.async_handler, "post", AsyncMock(side_effect=_http_status_error(429))): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert result is inputs + assert result["structured_messages"][3]["content"] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_http_status_error_does_not_leak_upstream_body(guardrail: CompresrGuardrail): + secret = "SECRET_INSTANCE_METADATA_TOKEN=aws-imds-response" + with patch.object(guardrail.async_handler, "post", AsyncMock(side_effect=_http_status_error(500, text=secret))): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert secret not in json.dumps(exc_info.value.detail) + + +# ── non-transport httpx errors must still honor the fail policy ──────── +# TooManyRedirects and DecodingError are httpx.RequestError but NOT +# httpx.TransportError, so a narrow except would let them escape as a 500 +# even under fail_open. These lock in that they are routed through the policy. + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "error", + [ + httpx.TooManyRedirects("redirect loop"), + httpx.DecodingError("bad content-encoding"), + ], +) +async def test_request_errors_raise_when_fail_closed(guardrail: CompresrGuardrail, error): + with patch.object(guardrail.async_handler, "post", AsyncMock(side_effect=error)): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert exc_info.value.status_code == 502 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "error", + [ + httpx.TooManyRedirects("redirect loop"), + httpx.DecodingError("bad content-encoding"), + ], +) +async def test_request_errors_fail_open_forwards_uncompressed(error): + guardrail = _make_guardrail(unreachable_fallback="fail_open") + inputs = _apply_inputs(AGENT_MESSAGES) + with patch.object(guardrail.async_handler, "post", AsyncMock(side_effect=error)): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert result is inputs + assert result["structured_messages"][3]["content"] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_undecodable_body_on_200_forwards_uncompressed_when_fail_open(): + """A 200 whose body raises DecodingError on .json()/.text must not 500.""" + guardrail = _make_guardrail(unreachable_fallback="fail_open") + resp = MagicMock() + resp.status_code = 200 + resp.json.side_effect = httpx.DecodingError("bad content-encoding") + type(resp).text = property(lambda self: (_ for _ in ()).throw(httpx.DecodingError("bad content-encoding"))) + inputs = _apply_inputs(AGENT_MESSAGES) + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=resp)): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert result is inputs + assert result["structured_messages"][3]["content"] == TOOL_OUTPUT + + +@pytest.mark.asyncio +async def test_recursion_error_on_json_forwards_when_fail_open(): + """A deeply nested JSON body can raise RecursionError while parsing; it must + route through the fail policy, not escape as a 500.""" + guardrail = _make_guardrail(unreachable_fallback="fail_open") + resp = MagicMock() + resp.status_code = 200 + resp.json.side_effect = RecursionError("maximum recursion depth exceeded") + resp.text = "" + inputs = _apply_inputs(AGENT_MESSAGES) + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=resp)): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert result is inputs + + +@pytest.mark.asyncio +async def test_lone_surrogate_in_content_does_not_crash(guardrail: CompresrGuardrail): + """A lone Unicode surrogate (reachable via a JSON \\uXXXX escape) in content + must not crash hashing/byte-accounting after the fail-policy decision.""" + surrogate_output = ("x" * 600) + "\ud800" + messages = [ + {"role": "user", "content": USER_QUESTION}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "web_search", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "call_1", "content": surrogate_output}, + ] + mock_post = AsyncMock(return_value=_make_single_compress_response()) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(messages), + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + assert result["structured_messages"][2]["content"].startswith("compressed summary") + # The original (surrogate included) is recoverable by its hash. + stored = next(iter(guardrail._originals_by_call_id.values()))[0] + assert surrogate_output in stored.values() + + +@pytest.mark.asyncio +async def test_identical_compressed_text_treated_as_noop(guardrail: CompresrGuardrail): + """If the service returns text byte-identical to the original, nothing + changed: the exact inputs object is returned so no write-back is forced.""" + inputs = _apply_inputs(AGENT_MESSAGES) + mock_post = AsyncMock(return_value=_make_single_compress_response(compressed_context=TOOL_OUTPUT)) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=_logging_obj("call-id-1"), + ) + assert result is inputs + + +# ── recovery requires a framework-issued call id ────────────────────── + + +@pytest.mark.asyncio +async def test_recovery_disabled_without_call_id(): + # enable_retrieval defaults True, but with no framework litellm_call_id we + # cannot scope stored originals to the request, so compression proceeds + # without markers, the retrieve tool, or any stored originals. + guardrail = _make_guardrail() + inputs = _apply_inputs(AGENT_MESSAGES) + mock_post = AsyncMock(return_value=_make_single_compress_response()) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"model": "gpt-4o"}, + input_type="request", + ) + assert result["structured_messages"][3]["content"] == "compressed summary" + assert not has_compresr_retrieve_tool(result.get("tools") or []) + assert guardrail._originals_by_call_id == {} + + +def test_config_model_exposes_unreachable_fallback(): + from litellm.types.proxy.guardrails.guardrail_hooks.compresr import ( + CompresrGuardrailConfigModel, + ) + + field = CompresrGuardrailConfigModel.model_fields.get("unreachable_fallback") + assert field is not None + assert field.default == "fail_closed" + + +# ── audit fixes ─────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_cancelled_error_propagates_not_swallowed(): + # Regression: CancelledError is a BaseException, not caught by + # (RequestError, Timeout). It must re-raise so cooperative cancellation + # (asyncio.wait_for, client disconnect) still fires. + import asyncio as _asyncio + + guardrail = _make_guardrail(unreachable_fallback="fail_open") + with patch.object(guardrail.async_handler, "post", AsyncMock(side_effect=_asyncio.CancelledError())): + with pytest.raises(_asyncio.CancelledError): + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + ) + + +def test_max_bytes_per_call_negative_rejected(): + # Regression: a negative value silently disabled the byte cap (< 0 behaves + # like 0 in _bound_call_bytes). Validate at construction so the footgun + # surfaces as a ValueError at startup, not silent unbounded storage. + with pytest.raises(ValueError, match="max_bytes_per_call"): + _make_guardrail(max_bytes_per_call=-1) + + +@pytest.mark.asyncio +async def test_max_tokens_zero_from_optional_params_wins_over_kwargs(): + # Regression: `or` short-circuits on falsy values, so an explicit + # max_tokens=0 from optional_params fell through to kwargs["max_tokens"]. + # Must use `is not None`. + guardrail = _make_guardrail() + hash_value = "deadbeef" + guardrail._store_originals(_scoped_store_key(_logging_obj("call-1")), {hash_value: TOOL_OUTPUT}) + response = MagicMock() + response.content = [{"type": "tool_use", "id": "toolu_1"}] + plan = await guardrail.async_build_agentic_loop_plan( + tools={ + "tool_calls": [ + { + "id": "toolu_1", + "type": "function", + "name": COMPRESR_RETRIEVE_TOOL_NAME, + "arguments": {"hash": hash_value}, + } + ] + }, + model="claude-sonnet-5", + messages=[{"role": "user", "content": "q"}], + response=response, + anthropic_messages_provider_config=None, + anthropic_messages_optional_request_params={"max_tokens": 0}, + logging_obj=_logging_obj("call-1"), + stream=False, + kwargs={"max_tokens": 999}, + ) + assert plan.request_patch.max_tokens == 0 + + +@pytest.mark.asyncio +async def test_recovery_disabled_when_no_caller_scope(): + # Regression: on a no-auth deployment (no UserAPIKeyAuth in metadata) the + # store key would fall back to the client-settable call id alone, letting + # any caller retrieve any other caller's originals. Recovery must be off. + guardrail = _make_guardrail() + logging_obj = MagicMock() + logging_obj.litellm_call_id = "call-abc" + logging_obj.model_call_details = {"litellm_params": {"metadata": {}}} + mock_post = AsyncMock(return_value=_make_single_compress_response()) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=logging_obj, + ) + assert not has_compresr_retrieve_tool(result.get("tools") or []) + assert guardrail._originals_by_call_id == {} + + +@pytest.mark.asyncio +async def test_warns_once_per_interval_when_recovery_skipped_without_scope(): + # enable_retrieval is on but the request has no per-key auth scope: recovery + # is silently skipped, so a call-time warning must surface it, rate-limited + # within the interval but re-arming after it so an ongoing misconfiguration + # stays visible. + guardrail = _make_guardrail() + logging_obj = MagicMock() + logging_obj.litellm_call_id = "call-abc" + logging_obj.model_call_details = {"litellm_params": {"metadata": {}}} + mock_post = AsyncMock(return_value=_make_single_compress_response()) + + def _no_scope_warnings(mock_log): + return [c for c in mock_log.warning.call_args_list if "no per-key auth scope" in str(c)] + + with patch.object(guardrail.async_handler, "post", mock_post): + with patch("litellm.proxy.guardrails.guardrail_hooks.compresr.compresr.verbose_proxy_logger") as mock_log: + for _ in range(3): + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=logging_obj, + ) + assert len(_no_scope_warnings(mock_log)) == 1 + + guardrail._no_scope_warning_expiry = 0.0 + await guardrail.apply_guardrail( + inputs=_apply_inputs(AGENT_MESSAGES), + request_data={"model": "gpt-4o"}, + input_type="request", + logging_obj=logging_obj, + ) + assert len(_no_scope_warnings(mock_log)) == 2 diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 21e7186fca3..359b1807344 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -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): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index aa24b0199ab..dffca3093fa 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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.""" diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index 74a0efba43d..303871e3981 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -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""" diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py index 45b81acbce1..9452e8042bd 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py @@ -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( diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 41dc7269372..3404b55f0db 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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 ( diff --git a/tests/test_litellm/test_router/test_io_token_rate_limits.py b/tests/test_litellm/test_router/test_io_token_rate_limits.py index 939e3189596..a5a68271111 100644 --- a/tests/test_litellm/test_router/test_io_token_rate_limits.py +++ b/tests/test_litellm/test_router/test_io_token_rate_limits.py @@ -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 diff --git a/tests/test_litellm/test_secret_redaction.py b/tests/test_litellm/test_secret_redaction.py index 85430ba752b..7152eea7c9f 100644 --- a/tests/test_litellm/test_secret_redaction.py +++ b/tests/test_litellm/test_secret_redaction.py @@ -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) diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 05ad536b260..9b55bf6ca0d 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -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 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.test.tsx index 94ed707f31d..d6f5f22616b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.test.tsx @@ -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( + , + ); + + 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( + , + ); + + 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(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.tsx index a65b980d8b8..e310da55063 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.tsx @@ -10,7 +10,7 @@ const { Text } = Typography; interface GuardrailTestPanelProps { guardrailNames: string[]; - onSubmit: (text: string) => void; + onSubmit: (text: string, metadata?: Record | 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(null); + + const parseMetadata = (raw: string): { metadata: Record | 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) => { @@ -142,6 +167,33 @@ export function GuardrailTestPanel({ +
+
+ + + + +
+