mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #33491 from BerriAI/litellm_internal_staging
Some checks are pending
CodeQL / Analyze (actions) (push) Waiting to run
CodeQL / Analyze (javascript-typescript) (push) Waiting to run
CodeQL / Analyze (python) (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Some checks are pending
CodeQL / Analyze (actions) (push) Waiting to run
CodeQL / Analyze (javascript-typescript) (push) Waiting to run
CodeQL / Analyze (python) (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
chore(ci): promote internal staging to main
This commit is contained in:
commit
229159c790
94 changed files with 8179 additions and 1526 deletions
|
|
@ -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" \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -186,6 +186,61 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
|
|||
)
|
||||
|
||||
|
||||
def _blank_to_none(value: str | None) -> str | None:
|
||||
"""Collapse an absent, empty, or whitespace-only string to ``None``.
|
||||
|
||||
OAuth endpoint fields are consumed by truthiness-based merges (``row or discovered``) and by the
|
||||
corroboration gate. A whitespace-only value is truthy to ``or`` but is not a usable endpoint, so
|
||||
without this the merge would keep the blank value for redirects while the gate treats it as
|
||||
unpinned and backfills the other fields, yielding a broken half-discovered config. Normalizing
|
||||
the pinned fields once, at each build entry point, gives every downstream consumer a single
|
||||
notion of "blank" so those code paths cannot disagree.
|
||||
"""
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
return value.strip() or None
|
||||
|
||||
|
||||
def _normalized_authorize_endpoint(url: str) -> str:
|
||||
"""Compare authorize endpoints on scheme, host, and path only. The default port is elided and
|
||||
the host is lowercased so ``https://IDP.example.com:443/authorize/`` and
|
||||
``https://idp.example.com/authorize`` are the same identity; query and trailing slash are not."""
|
||||
parsed = urlparse(url)
|
||||
scheme = parsed.scheme.lower()
|
||||
host = (parsed.hostname or "").lower()
|
||||
default_port = {"https": 443, "http": 80}.get(scheme)
|
||||
try:
|
||||
port = parsed.port
|
||||
except ValueError:
|
||||
port = None
|
||||
authority = host if port is None or port == default_port else f"{host}:{port}"
|
||||
return f"{scheme}://{authority}{parsed.path.rstrip('/')}"
|
||||
|
||||
|
||||
def _endpoints_corroborate_authorization_url(
|
||||
source_authorization_url: str | None,
|
||||
trusted_authorization_url: str | None,
|
||||
) -> bool:
|
||||
"""Whether a source's ``token_url``/``registration_url`` may be paired with a trusted authorize
|
||||
endpoint. This is the single trust rule for adopting OAuth endpoints from any non-manual source.
|
||||
|
||||
Discovery is rooted at the MCP resource (RFC 9728), so a compromised upstream can advertise an
|
||||
attacker-run authorization server. When ``authorization_url`` is admin-pinned, pairing it with a
|
||||
``token_url`` from a different source is the RFC 9700 authorization-server mix-up: the user signs
|
||||
in at the trusted authorize endpoint while the gateway redeems the code, with the stored client
|
||||
secret and PKCE verifier, at the attacker's token endpoint. Endpoints are trustworthy together
|
||||
only when they share an authorization server, so a source's endpoints are adopted only when the
|
||||
same source advertised an ``authorization_endpoint`` matching the pinned value. With no pinned
|
||||
value (``trusted_authorization_url is None``) there is nothing to protect: the authorize endpoint
|
||||
comes from the same source as the token endpoint, so they corroborate each other by construction.
|
||||
"""
|
||||
if not (trusted_authorization_url and trusted_authorization_url.strip()):
|
||||
return True
|
||||
return bool(source_authorization_url) and _normalized_authorize_endpoint(
|
||||
source_authorization_url
|
||||
) == _normalized_authorize_endpoint(trusted_authorization_url)
|
||||
|
||||
|
||||
def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_server: MCPServer | None) -> None:
|
||||
"""Keep the last known good OAuth endpoints when a rebuild's re-discovery comes back empty.
|
||||
|
||||
|
|
@ -193,26 +248,82 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv
|
|||
during re-discovery downgrades a working server (``authorization_url`` set) to a broken one
|
||||
(``None``, /authorize 400s) with no configuration change. Mirrors the ``short_prefix``
|
||||
carry-forward. Skipped when the server's ``url`` or ``auth_type`` changed, since the previous
|
||||
endpoints may then belong to a different upstream. ``registration_url`` IS carried here even
|
||||
though ``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only
|
||||
restores the same in-memory value the previous build already ran with, while persisting it
|
||||
would flip ``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for
|
||||
dcr_bridge servers that never had one configured.
|
||||
endpoints may then belong to a different upstream. ``registration_url`` IS carried even though
|
||||
``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only restores
|
||||
the same in-memory value the previous build already ran with, while persisting it would flip
|
||||
``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for dcr_bridge
|
||||
servers that never had one configured.
|
||||
|
||||
Carry-forward is a non-manual endpoint source, so the same trust rule as discovery applies: the
|
||||
previous ``token_url``/``registration_url``/``scopes`` are carried only when the previous
|
||||
``authorization_url`` corroborates the authorize endpoint this build will use, i.e. when the
|
||||
incoming build has no pinned authorize endpoint (``None`` -> we adopt the previous one too, a
|
||||
consistent group) or pins the same one. An admin re-pointing ``authorization_url`` to a different
|
||||
server must not keep serving the old server's token endpoint or granted scopes.
|
||||
"""
|
||||
if previous_server is None:
|
||||
return
|
||||
if previous_server.url != new_server.url or previous_server.auth_type != new_server.auth_type:
|
||||
return
|
||||
may_carry = _endpoints_corroborate_authorization_url(
|
||||
previous_server.authorization_url, new_server.authorization_url
|
||||
)
|
||||
if new_server.authorization_url is None and previous_server.authorization_url:
|
||||
new_server.authorization_url = previous_server.authorization_url
|
||||
if new_server.token_url is None and previous_server.token_url:
|
||||
if may_carry and new_server.token_url is None and previous_server.token_url:
|
||||
new_server.token_url = previous_server.token_url
|
||||
if new_server.registration_url is None and previous_server.registration_url:
|
||||
if may_carry and new_server.registration_url is None and previous_server.registration_url:
|
||||
new_server.registration_url = previous_server.registration_url
|
||||
if not new_server.scopes and previous_server.scopes:
|
||||
if may_carry and not new_server.scopes and previous_server.scopes:
|
||||
new_server.scopes = previous_server.scopes
|
||||
|
||||
|
||||
def _restrict_discovery_to_corroborated_authorization_server(
|
||||
metadata: MCPOAuthMetadata | None,
|
||||
manual_authorization_url: str | None,
|
||||
server_identifier: str,
|
||||
is_dcr_bridge: bool,
|
||||
) -> MCPOAuthMetadata | None:
|
||||
"""Reject discovered token/registration endpoints a manually pinned authorize endpoint cannot
|
||||
vouch for (the RFC 9700 authorization-server mix-up).
|
||||
|
||||
Discovery is rooted at the MCP resource, so a compromised upstream can advertise an attacker
|
||||
``token_endpoint``: with ``authorization_url`` admin-pinned but ``token_url`` blank, the merge
|
||||
would pair the trusted authorize endpoint with that attacker token endpoint, and the gateway would
|
||||
post the authorization code and client secret there. So the discovered ``token_url`` and
|
||||
``registration_url`` are kept only if the document corroborates the pin (its
|
||||
``authorization_endpoint`` matches). ``scopes`` are deliberately NOT gated here: per the MCP
|
||||
authorization spec Scope Selection Strategy and RFC 9700 §2.3, the scopes a client requests are
|
||||
resource-driven (the WWW-Authenticate challenge or the RFC 9728 protected-resource
|
||||
``scopes_supported``), and scope inflation by a compromised resource is bounded by the
|
||||
authorization server and user consent (RFC 6749 §3.3), not by the client second-guessing the
|
||||
request. With no pin there is no trust anchor to protect, so discovery is returned as-is.
|
||||
"""
|
||||
if metadata is None or not (manual_authorization_url and manual_authorization_url.strip()):
|
||||
return metadata
|
||||
if _endpoints_corroborate_authorization_url(metadata.authorization_url, manual_authorization_url):
|
||||
return metadata
|
||||
if not metadata.token_url and not metadata.registration_url:
|
||||
return metadata
|
||||
bridge_note = (
|
||||
" The discovered registration_url is rejected with it, so this dcr_bridge server stays on the"
|
||||
" short-circuit registration arm."
|
||||
if is_dcr_bridge and metadata.registration_url
|
||||
else ""
|
||||
)
|
||||
verbose_logger.warning(
|
||||
"MCP OAuth discovery for server %s advertised authorization_endpoint %s, which does not match the "
|
||||
"manually configured authorization_url %s; rejecting the discovered token_url/registration_url so "
|
||||
"authorization codes and client credentials only follow the configured authorization server. "
|
||||
"Configure Token URL manually if the mismatch is intentional.%s",
|
||||
server_identifier,
|
||||
_normalized_authorize_endpoint(metadata.authorization_url) if metadata.authorization_url else "<absent>",
|
||||
_normalized_authorize_endpoint(manual_authorization_url),
|
||||
bridge_note,
|
||||
)
|
||||
return metadata.model_copy(update={"token_url": None, "registration_url": None})
|
||||
|
||||
|
||||
def invalidate_user_env_vars_cache(user_id: str, server_id: str) -> None:
|
||||
"""Drop a cached entry after the user stores or clears their env var values
|
||||
so the next request reads the fresh value instead of a stale one."""
|
||||
|
|
@ -1026,12 +1137,15 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
auth_type = server_config.get("auth_type", None)
|
||||
manual_authorization_url = _blank_to_none(server_config.get("authorization_url"))
|
||||
manual_token_url = _blank_to_none(server_config.get("token_url"))
|
||||
manual_registration_url = _blank_to_none(server_config.get("registration_url"))
|
||||
if server_url and (
|
||||
auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
|
||||
or self._obo_needs_endpoint_discovery(
|
||||
auth_type,
|
||||
server_config.get("token_exchange_endpoint"),
|
||||
server_config.get("token_url"),
|
||||
manual_token_url,
|
||||
)
|
||||
):
|
||||
mcp_oauth_metadata = await self._descovery_metadata(
|
||||
|
|
@ -1041,20 +1155,29 @@ class MCPServerManager:
|
|||
else:
|
||||
mcp_oauth_metadata = None
|
||||
|
||||
gated_oauth_metadata = (
|
||||
_restrict_discovery_to_corroborated_authorization_server(
|
||||
mcp_oauth_metadata,
|
||||
manual_authorization_url,
|
||||
server_name or server_id,
|
||||
bool(server_config.get("dcr_bridge")),
|
||||
)
|
||||
if auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
|
||||
else mcp_oauth_metadata
|
||||
)
|
||||
|
||||
# Filter blank scopes (e.g. YAML ``scopes: [""]``) the same way the DB-build path does, so
|
||||
# an all-blank list normalizes to None rather than a ``("",)`` tuple that skips the
|
||||
# entra_obo fail-closed scope precondition and POSTs an empty scope to the IdP.
|
||||
resolved_scopes = self._extract_scopes(server_config.get("scopes")) or (
|
||||
mcp_oauth_metadata.scopes if mcp_oauth_metadata else None
|
||||
gated_oauth_metadata.scopes if gated_oauth_metadata else None
|
||||
)
|
||||
resolved_authorization_url = server_config.get("authorization_url") or (
|
||||
mcp_oauth_metadata.authorization_url if mcp_oauth_metadata else None
|
||||
resolved_authorization_url = manual_authorization_url or (
|
||||
gated_oauth_metadata.authorization_url if gated_oauth_metadata else None
|
||||
)
|
||||
resolved_token_url = server_config.get("token_url") or (
|
||||
mcp_oauth_metadata.token_url if mcp_oauth_metadata else None
|
||||
)
|
||||
resolved_registration_url = server_config.get("registration_url") or (
|
||||
mcp_oauth_metadata.registration_url if mcp_oauth_metadata else None
|
||||
resolved_token_url = manual_token_url or (gated_oauth_metadata.token_url if gated_oauth_metadata else None)
|
||||
resolved_registration_url = manual_registration_url or (
|
||||
gated_oauth_metadata.registration_url if gated_oauth_metadata else None
|
||||
)
|
||||
|
||||
config_oauth2_flow = server_config.get("oauth2_flow", None)
|
||||
|
|
@ -1447,13 +1570,17 @@ class MCPServerManager:
|
|||
|
||||
auth_type = cast(MCPAuthType, mcp_server.auth_type)
|
||||
server_url = mcp_server.url
|
||||
manual_authorization_url = _blank_to_none(mcp_server.authorization_url)
|
||||
manual_token_url = _blank_to_none(mcp_server.token_url)
|
||||
manual_registration_url = _blank_to_none(mcp_server.registration_url)
|
||||
has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes)
|
||||
needs_discovery = bool(server_url) and (
|
||||
(auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and not mcp_server.authorization_url)
|
||||
(auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and not has_all_upstream_oauth_fields)
|
||||
or self._obo_needs_endpoint_discovery(
|
||||
auth_type,
|
||||
mcp_server.token_exchange_endpoint
|
||||
or (credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
|
||||
mcp_server.token_url,
|
||||
manual_token_url,
|
||||
)
|
||||
)
|
||||
mcp_oauth_metadata = (
|
||||
|
|
@ -1467,12 +1594,22 @@ class MCPServerManager:
|
|||
if needs_discovery and mcp_oauth_metadata is None:
|
||||
verbose_logger.warning(
|
||||
"MCP OAuth discovery yielded no metadata for server %s (%s); "
|
||||
"OAuth endpoints stay unresolved until a rebuild succeeds",
|
||||
"OAuth endpoints/scopes stay unresolved until a rebuild succeeds",
|
||||
mcp_server.server_id,
|
||||
server_url,
|
||||
)
|
||||
gated_oauth_metadata = (
|
||||
_restrict_discovery_to_corroborated_authorization_server(
|
||||
mcp_oauth_metadata,
|
||||
manual_authorization_url,
|
||||
mcp_server.server_id,
|
||||
bool(getattr(mcp_server, "dcr_bridge", None)),
|
||||
)
|
||||
if auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
|
||||
else mcp_oauth_metadata
|
||||
)
|
||||
|
||||
resolved_scopes = scopes or (mcp_oauth_metadata.scopes if mcp_oauth_metadata else None)
|
||||
resolved_scopes = scopes or (gated_oauth_metadata.scopes if gated_oauth_metadata else None)
|
||||
|
||||
new_server = MCPServer(
|
||||
server_id=mcp_server.server_id,
|
||||
|
|
@ -1492,9 +1629,9 @@ class MCPServerManager:
|
|||
client_secret=client_secret_value or getattr(mcp_server, "client_secret", None),
|
||||
oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)),
|
||||
scopes=resolved_scopes,
|
||||
authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None),
|
||||
token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None),
|
||||
registration_url=mcp_server.registration_url or getattr(mcp_oauth_metadata, "registration_url", None),
|
||||
authorization_url=manual_authorization_url or getattr(gated_oauth_metadata, "authorization_url", None),
|
||||
token_url=manual_token_url or getattr(gated_oauth_metadata, "token_url", None),
|
||||
registration_url=manual_registration_url or getattr(gated_oauth_metadata, "registration_url", None),
|
||||
token_endpoint_auth_method=(
|
||||
credentials_dict.get("token_endpoint_auth_method") if credentials_dict else None
|
||||
),
|
||||
|
|
@ -1545,16 +1682,16 @@ class MCPServerManager:
|
|||
await self._persist_discovered_obo_token_url(
|
||||
server_id=mcp_server.server_id,
|
||||
auth_type=auth_type,
|
||||
existing_token_url=mcp_server.token_url,
|
||||
existing_token_url=manual_token_url,
|
||||
discovered_token_url=new_server.token_url,
|
||||
)
|
||||
await self._persist_discovered_oauth_endpoints(
|
||||
server_id=mcp_server.server_id,
|
||||
auth_type=auth_type,
|
||||
existing_authorization_url=mcp_server.authorization_url,
|
||||
existing_token_url=mcp_server.token_url,
|
||||
existing_authorization_url=manual_authorization_url,
|
||||
existing_token_url=manual_token_url,
|
||||
existing_scopes=scopes,
|
||||
metadata=mcp_oauth_metadata,
|
||||
metadata=gated_oauth_metadata,
|
||||
)
|
||||
return new_server
|
||||
|
||||
|
|
|
|||
|
|
@ -7613,6 +7613,18 @@
|
|||
],
|
||||
"title": "Messages"
|
||||
},
|
||||
"metadata": {
|
||||
"anyOf": [
|
||||
{
|
||||
"additionalProperties": true,
|
||||
"type": "object"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Metadata"
|
||||
},
|
||||
"text": {
|
||||
"title": "Text",
|
||||
"type": "string"
|
||||
|
|
|
|||
|
|
@ -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:]}"
|
||||
|
|
|
|||
|
|
@ -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]},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
1214
litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py
Normal file
1214
litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
135
litellm/types/proxy/guardrails/guardrail_hooks/compresr.py
Normal file
135
litellm/types/proxy/guardrails/guardrail_hooks/compresr.py
Normal file
|
|
@ -0,0 +1,135 @@
|
|||
from typing import Any, Dict, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
|
||||
class CompresrGuardrailOptionalParams(BaseModel):
|
||||
"""Optional tuning knobs for the Compresr guardrail."""
|
||||
|
||||
target_compression_ratio: float | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Compression strength. 0-1 is the fraction of tokens to remove "
|
||||
"(0.5 = remove ~50%, the default); a value >1 is an Nx reduction "
|
||||
"factor (e.g. 4 = ~4x smaller)."
|
||||
),
|
||||
)
|
||||
coarse: bool | None = Field(
|
||||
default=None,
|
||||
description=("Paragraph-level compression (default, faster) instead of token-level (finer-grained)."),
|
||||
)
|
||||
min_chars_to_compress: int | None = Field(
|
||||
default=None,
|
||||
description=("Skip messages whose text is shorter than this many characters. Defaults to 500."),
|
||||
)
|
||||
compress_tool_outputs: bool | None = Field(
|
||||
default=None,
|
||||
description=("Compress tool/function result messages (search hits, RAG chunks, API dumps). Defaults to True."),
|
||||
)
|
||||
compress_system: bool | None = Field(
|
||||
default=None,
|
||||
description="Also compress system messages. Defaults to False.",
|
||||
)
|
||||
compress_history: bool | None = Field(
|
||||
default=None,
|
||||
description="Also compress prior (non-last) user messages. Defaults to False.",
|
||||
)
|
||||
compress_last_user: bool | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Also compress the last user message. The query sent to Compresr "
|
||||
"is always the original verbatim text. Defaults to False."
|
||||
),
|
||||
)
|
||||
enable_retrieval: bool | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Make compression recoverable: inject a `compresr_retrieve` tool "
|
||||
"so the model can fetch the original content behind a compression "
|
||||
"marker via the agentic loop. Defaults to True. Set to False (or "
|
||||
"run the proxy with --workers 1) for multi-worker deployments: "
|
||||
"the recovery store is per-process, so pre-call and retrieval hooks "
|
||||
"on different workers cannot see each other's originals."
|
||||
),
|
||||
)
|
||||
max_bytes_per_call: int | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Cap on aggregate bytes of stored originals per litellm_call_id. "
|
||||
"When a call exceeds this, oldest entries are evicted so the "
|
||||
"in-process store cannot grow without bound. Defaults to 10 MiB."
|
||||
),
|
||||
)
|
||||
allow_bypass_header: bool | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Honor the `x-compresr-bypass: true` request header to skip "
|
||||
"compression for a single call. Off by default because the "
|
||||
"header is caller-settable; enable only on trusted deployments."
|
||||
),
|
||||
)
|
||||
dynamic: bool | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"latte_v2 only. Let the server choose the compression amount per input "
|
||||
"(Kneedle elbow) instead of using target_compression_ratio. Defaults to True."
|
||||
),
|
||||
)
|
||||
dynamic_min_ratio: float | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"latte_v2 only. Floor on the adaptive ratio when `dynamic` is on. "
|
||||
"Unset lets the server default apply (~1.5)."
|
||||
),
|
||||
)
|
||||
dynamic_max_ratio: float | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"latte_v2 only. Ceiling on the adaptive ratio when `dynamic` is on. "
|
||||
"Unset lets the server default apply (~10.0)."
|
||||
),
|
||||
)
|
||||
compression_params: Dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Passthrough of extra parameters forwarded verbatim in the Compresr "
|
||||
"compress payload (e.g. `heuristic_chunking`, or any newer knob), so "
|
||||
"a new Compresr feature works without a guardrail update. The named "
|
||||
"fields above take precedence on collision."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class CompresrGuardrailConfigModel(GuardrailConfigModel[CompresrGuardrailOptionalParams]):
|
||||
api_key: str | None = Field(
|
||||
default=None,
|
||||
description=("Compresr API key. Falls back to the COMPRESR_API_KEY env var."),
|
||||
)
|
||||
api_base: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Base URL of the Compresr API. Falls back to the COMPRESR_API_BASE "
|
||||
"env var, then https://api.compresr.ai. Point at your internal "
|
||||
"service URL for on-prem deployments."
|
||||
),
|
||||
)
|
||||
model: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Compresr compression model (not the LLM). Defaults to 'latte_v2', the query-aware compression model."
|
||||
),
|
||||
)
|
||||
unreachable_fallback: Literal["fail_closed", "fail_open"] = Field(
|
||||
default="fail_closed",
|
||||
description=(
|
||||
"Behavior when the Compresr compression service is unreachable or errors. "
|
||||
"'fail_closed' raises an error (default). 'fail_open' logs a critical error and "
|
||||
"forwards the request uncompressed instead of blocking it."
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "Compresr (context compression)"
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
195
tests/e2e/claude_code/_driver_unit_tests/test_passthrough.py
Normal file
195
tests/e2e/claude_code/_driver_unit_tests/test_passthrough.py
Normal file
|
|
@ -0,0 +1,195 @@
|
|||
"""Unit tests for the shared `run_passthrough_cell` helper.
|
||||
|
||||
These tests inject a fake `run_models` callable and an explicit `env`
|
||||
mapping (both are first-class parameters, no monkeypatching), so they
|
||||
exercise the helper's branching -- env-missing guard, base-URL
|
||||
assembly, extra-env forwarding, per-model pass/fail -- without
|
||||
spawning the real CLI.
|
||||
|
||||
The env-builder tests pin the provider-mode contract itself: the
|
||||
CLAUDE_CODE_USE_* / CLAUDE_CODE_SKIP_*_AUTH flags and the passthrough
|
||||
route each mode must target. Those values are the feature -- e.g.
|
||||
dropping the `/v1` from the vertex base URL produces a request Google
|
||||
404s on -- so a mutation to any of them must fail here before it burns
|
||||
a live matrix run.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Mapping, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
from claude_code._passthrough import (
|
||||
ANTHROPIC_PASSTHROUGH_BASE_PATH,
|
||||
CLIENT_SIDE_AWS_REGION,
|
||||
VERTEX_PLACEHOLDER_PROJECT,
|
||||
VERTEX_PLACEHOLDER_REGION,
|
||||
bedrock_extra_env,
|
||||
foundry_extra_env,
|
||||
run_passthrough_cell,
|
||||
vertex_extra_env,
|
||||
)
|
||||
from claude_code.cli_driver import ClaudeCLIError, DriverResult
|
||||
|
||||
PROXY_ENV = {
|
||||
"LITELLM_PROXY_BASE_URL": "http://localhost:4000",
|
||||
"LITELLM_PROXY_API_KEY": "sk-test",
|
||||
}
|
||||
|
||||
|
||||
class _FakeResult:
|
||||
def __init__(self) -> None:
|
||||
self.rows: List[Dict[str, Any]] = []
|
||||
self.single: Optional[Dict[str, Any]] = None
|
||||
|
||||
def set(self, payload: Mapping[str, Any]) -> None:
|
||||
self.single = dict(payload)
|
||||
|
||||
def add(self, payload: Mapping[str, Any]) -> None:
|
||||
self.rows.append(dict(payload))
|
||||
|
||||
|
||||
def _fake_run_models(outcomes_by_model, captured: Dict[str, Any]):
|
||||
def fake(*, models, prompt, base_url, api_key, extra_env=None, **_kwargs):
|
||||
captured["models"] = list(models)
|
||||
captured["prompt"] = prompt
|
||||
captured["base_url"] = base_url
|
||||
captured["api_key"] = api_key
|
||||
captured["extra_env"] = dict(extra_env) if extra_env is not None else None
|
||||
return {model: outcomes_by_model[model] for model in models}
|
||||
|
||||
return fake
|
||||
|
||||
|
||||
def test_env_missing_guard_reports_fail_and_aborts():
|
||||
fake_result = _FakeResult()
|
||||
with pytest.raises(pytest.fail.Exception):
|
||||
run_passthrough_cell(
|
||||
compat_result=fake_result,
|
||||
models=["claude-haiku-4-5"],
|
||||
prompt="ping",
|
||||
env={},
|
||||
)
|
||||
assert fake_result.single is not None
|
||||
assert fake_result.single["status"] == "fail"
|
||||
assert "LITELLM_PROXY_BASE_URL" in fake_result.single["error"]
|
||||
|
||||
|
||||
def test_anthropic_base_path_appended_to_normalized_proxy_url():
|
||||
fake_result = _FakeResult()
|
||||
captured: Dict[str, Any] = {}
|
||||
outcome = DriverResult(text="pong")
|
||||
|
||||
run_passthrough_cell(
|
||||
compat_result=fake_result,
|
||||
models=["claude-haiku-4-5"],
|
||||
prompt="ping",
|
||||
passthrough_base_path=ANTHROPIC_PASSTHROUGH_BASE_PATH,
|
||||
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
|
||||
env={**PROXY_ENV, "LITELLM_PROXY_BASE_URL": "http://localhost:4000/"},
|
||||
)
|
||||
|
||||
assert captured["base_url"] == "http://localhost:4000/anthropic"
|
||||
assert captured["extra_env"] is None
|
||||
assert fake_result.rows == [{"status": "pass"}]
|
||||
|
||||
|
||||
def test_extra_env_builder_receives_normalized_base_and_is_forwarded():
|
||||
fake_result = _FakeResult()
|
||||
captured: Dict[str, Any] = {}
|
||||
outcome = DriverResult(text="pong")
|
||||
seen_bases: List[str] = []
|
||||
|
||||
def build(proxy_base: str) -> Dict[str, str]:
|
||||
seen_bases.append(proxy_base)
|
||||
return {"SOME_FLAG": "1"}
|
||||
|
||||
run_passthrough_cell(
|
||||
compat_result=fake_result,
|
||||
models=["claude-haiku-4-5"],
|
||||
prompt="ping",
|
||||
build_extra_env=build,
|
||||
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
|
||||
env={**PROXY_ENV, "LITELLM_PROXY_BASE_URL": "http://localhost:4000/"},
|
||||
)
|
||||
|
||||
assert seen_bases == ["http://localhost:4000"]
|
||||
assert captured["extra_env"] == {"SOME_FLAG": "1"}
|
||||
assert captured["base_url"] == "http://localhost:4000"
|
||||
|
||||
|
||||
def test_per_model_failures_reported_individually():
|
||||
fake_result = _FakeResult()
|
||||
captured: Dict[str, Any] = {}
|
||||
outcomes = {
|
||||
"claude-haiku-4-5": DriverResult(text="pong"),
|
||||
"claude-sonnet-4-6": ClaudeCLIError("claude CLI timed out after 120s"),
|
||||
"claude-opus-4-7": DriverResult(text="", exit_code=1),
|
||||
}
|
||||
|
||||
with pytest.raises(pytest.fail.Exception):
|
||||
run_passthrough_cell(
|
||||
compat_result=fake_result,
|
||||
models=list(outcomes.keys()),
|
||||
prompt="ping",
|
||||
run_models=_fake_run_models(outcomes, captured),
|
||||
env=PROXY_ENV,
|
||||
)
|
||||
|
||||
statuses = [row["status"] for row in fake_result.rows]
|
||||
assert statuses == ["pass", "fail", "fail"]
|
||||
assert "timed out" in fake_result.rows[1]["error"]
|
||||
assert "claude CLI failed" in fake_result.rows[2]["error"]
|
||||
|
||||
|
||||
def test_empty_assistant_text_is_a_fail():
|
||||
fake_result = _FakeResult()
|
||||
captured: Dict[str, Any] = {}
|
||||
outcomes = {"claude-haiku-4-5": DriverResult(text=" ")}
|
||||
|
||||
with pytest.raises(pytest.fail.Exception):
|
||||
run_passthrough_cell(
|
||||
compat_result=fake_result,
|
||||
models=["claude-haiku-4-5"],
|
||||
prompt="ping",
|
||||
run_models=_fake_run_models(outcomes, captured),
|
||||
env=PROXY_ENV,
|
||||
)
|
||||
|
||||
assert fake_result.rows == [
|
||||
{
|
||||
"status": "fail",
|
||||
"error": "[claude-haiku-4-5] claude returned empty assistant text",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_bedrock_extra_env_targets_proxy_bedrock_route():
|
||||
env = bedrock_extra_env("http://localhost:4000")
|
||||
assert env == {
|
||||
"CLAUDE_CODE_USE_BEDROCK": "1",
|
||||
"CLAUDE_CODE_SKIP_BEDROCK_AUTH": "1",
|
||||
"ANTHROPIC_BEDROCK_BASE_URL": "http://localhost:4000/bedrock",
|
||||
"AWS_REGION": CLIENT_SIDE_AWS_REGION,
|
||||
}
|
||||
|
||||
|
||||
def test_vertex_extra_env_keeps_the_api_version_in_the_base_url():
|
||||
env = vertex_extra_env("http://localhost:4000")
|
||||
assert env == {
|
||||
"CLAUDE_CODE_USE_VERTEX": "1",
|
||||
"CLAUDE_CODE_SKIP_VERTEX_AUTH": "1",
|
||||
"ANTHROPIC_VERTEX_BASE_URL": "http://localhost:4000/vertex_ai/v1",
|
||||
"ANTHROPIC_VERTEX_PROJECT_ID": VERTEX_PLACEHOLDER_PROJECT,
|
||||
"CLOUD_ML_REGION": VERTEX_PLACEHOLDER_REGION,
|
||||
}
|
||||
|
||||
|
||||
def test_foundry_extra_env_targets_proxy_azure_route():
|
||||
env = foundry_extra_env("http://localhost:4000")
|
||||
assert env == {
|
||||
"CLAUDE_CODE_USE_FOUNDRY": "1",
|
||||
"CLAUDE_CODE_SKIP_FOUNDRY_AUTH": "1",
|
||||
"ANTHROPIC_FOUNDRY_BASE_URL": "http://localhost:4000/azure",
|
||||
}
|
||||
196
tests/e2e/claude_code/_passthrough.py
Normal file
196
tests/e2e/claude_code/_passthrough.py
Normal file
|
|
@ -0,0 +1,196 @@
|
|||
"""Shared body for the `passthrough` × <provider> compat cells.
|
||||
|
||||
Every other matrix row drives the proxy's `/v1/messages` translation
|
||||
layer: Claude Code speaks the first-party Anthropic wire and LiteLLM
|
||||
transforms the request per provider. This row instead exercises
|
||||
LiteLLM's *native passthrough* routes -- the "LLM gateway"
|
||||
configuration documented at https://code.claude.com/docs/en/gateway --
|
||||
where Claude Code speaks each cloud's own wire format and the proxy
|
||||
forwards it, attaching provider credentials on the way out:
|
||||
|
||||
anthropic ANTHROPIC_BASE_URL={proxy}/anthropic. The CLI's
|
||||
first-party wire, forwarded verbatim to
|
||||
api.anthropic.com, so the model ids are real
|
||||
Anthropic ids rather than proxy aliases.
|
||||
bedrock_invoke CLAUDE_CODE_USE_BEDROCK=1 +
|
||||
ANTHROPIC_BEDROCK_BASE_URL={proxy}/bedrock. The
|
||||
CLI POSTs /model/{model}/invoke-with-response-stream;
|
||||
the proxy recognizes a router alias in the model
|
||||
segment, rewrites it to the deployment's upstream
|
||||
model id, and SigV4-signs with its own AWS creds.
|
||||
vertex_ai CLAUDE_CODE_USE_VERTEX=1 +
|
||||
ANTHROPIC_VERTEX_BASE_URL={proxy}/vertex_ai/v1.
|
||||
The CLI POSTs
|
||||
.../models/{model}:streamRawPredict; the proxy
|
||||
resolves a router alias in the model segment and
|
||||
takes project, location, and credentials from the
|
||||
deployment (which is why the deployment must set
|
||||
`use_in_pass_through: true` -- see
|
||||
test_config.yaml).
|
||||
azure CLAUDE_CODE_USE_FOUNDRY=1 +
|
||||
ANTHROPIC_FOUNDRY_BASE_URL={proxy}/azure. Foundry
|
||||
mode sends the model in the JSON body, not the
|
||||
URL, so the proxy's /azure route cannot resolve a
|
||||
router alias and falls back to the env-configured
|
||||
AZURE_API_BASE / AZURE_API_KEY target.
|
||||
bedrock_converse not applicable -- Claude Code's bedrock mode is
|
||||
InvokeModel-only; no Converse-wire client exists.
|
||||
|
||||
Auth is the same in every mode: the CLI's provider-native signing is
|
||||
disabled via CLAUDE_CODE_SKIP_<PROVIDER>_AUTH, and the LiteLLM virtual
|
||||
key travels as `Authorization: Bearer` (ANTHROPIC_AUTH_TOKEN), exactly
|
||||
like the translation rows. The proxy holds the real provider
|
||||
credentials.
|
||||
|
||||
The per-mode env vars and URL shapes above were captured from a real
|
||||
`claude` CLI (2.1.210) run against a request-logging sink, not from
|
||||
docs; if a CLI release changes them, the cells fail with the CLI's own
|
||||
diagnostic rather than silently testing the wrong wire.
|
||||
|
||||
`run_models` and `env` are injection seams for
|
||||
`_driver_unit_tests/test_passthrough.py`; production callers leave
|
||||
them unset.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Any, Callable, Dict, Mapping, Optional, Sequence
|
||||
|
||||
import pytest
|
||||
|
||||
from claude_code.cli_driver import (
|
||||
ClaudeCLIError,
|
||||
failure_diagnostic,
|
||||
run_claude_models_parallel,
|
||||
)
|
||||
|
||||
PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL"
|
||||
PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY"
|
||||
|
||||
ANTHROPIC_PASSTHROUGH_BASE_PATH = "/anthropic"
|
||||
|
||||
CLIENT_SIDE_AWS_REGION = "us-east-1"
|
||||
"""Satisfies the CLI's embedded AWS SDK, which refuses to construct a
|
||||
client without a region. The value never influences routing: the proxy
|
||||
signs the upstream request with its own credentials and region."""
|
||||
|
||||
VERTEX_PLACEHOLDER_PROJECT = "proxy-resolved-project"
|
||||
VERTEX_PLACEHOLDER_REGION = "us-east5"
|
||||
"""The CLI refuses to build a Vertex URL without a project id and
|
||||
region, but the proxy replaces both path segments with the resolved
|
||||
deployment's `vertex_project` / `vertex_location` before forwarding,
|
||||
so deliberately-fake values prove the resolution actually happened."""
|
||||
|
||||
|
||||
def bedrock_extra_env(proxy_base_url: str) -> Dict[str, str]:
|
||||
return {
|
||||
"CLAUDE_CODE_USE_BEDROCK": "1",
|
||||
"CLAUDE_CODE_SKIP_BEDROCK_AUTH": "1",
|
||||
"ANTHROPIC_BEDROCK_BASE_URL": f"{proxy_base_url}/bedrock",
|
||||
"AWS_REGION": CLIENT_SIDE_AWS_REGION,
|
||||
}
|
||||
|
||||
|
||||
def vertex_extra_env(proxy_base_url: str) -> Dict[str, str]:
|
||||
"""Vertex-mode CLI env pointed at the proxy's /vertex_ai route.
|
||||
|
||||
The `/v1` suffix on ANTHROPIC_VERTEX_BASE_URL is load-bearing: the
|
||||
CLI's Vertex SDK ships its API version inside its *default* base
|
||||
URL (`https://{region}-aiplatform.googleapis.com/v1`), so
|
||||
overriding the base drops the version from the request path unless
|
||||
the override carries it. LiteLLM's /vertex_ai route reuses the
|
||||
incoming path verbatim when it contains `/projects/.../locations/...`,
|
||||
so a version-less path would reach Google as
|
||||
`aiplatform.googleapis.com/projects/...` and 404.
|
||||
"""
|
||||
return {
|
||||
"CLAUDE_CODE_USE_VERTEX": "1",
|
||||
"CLAUDE_CODE_SKIP_VERTEX_AUTH": "1",
|
||||
"ANTHROPIC_VERTEX_BASE_URL": f"{proxy_base_url}/vertex_ai/v1",
|
||||
"ANTHROPIC_VERTEX_PROJECT_ID": VERTEX_PLACEHOLDER_PROJECT,
|
||||
"CLOUD_ML_REGION": VERTEX_PLACEHOLDER_REGION,
|
||||
}
|
||||
|
||||
|
||||
def foundry_extra_env(proxy_base_url: str) -> Dict[str, str]:
|
||||
return {
|
||||
"CLAUDE_CODE_USE_FOUNDRY": "1",
|
||||
"CLAUDE_CODE_SKIP_FOUNDRY_AUTH": "1",
|
||||
"ANTHROPIC_FOUNDRY_BASE_URL": f"{proxy_base_url}/azure",
|
||||
}
|
||||
|
||||
|
||||
def run_passthrough_cell(
|
||||
*,
|
||||
compat_result,
|
||||
models: Sequence[str],
|
||||
prompt: str,
|
||||
passthrough_base_path: str = "",
|
||||
build_extra_env: Optional[Callable[[str], Mapping[str, str]]] = None,
|
||||
run_models: Callable[..., Mapping[str, Any]] = run_claude_models_parallel,
|
||||
env: Optional[Mapping[str, str]] = None,
|
||||
) -> None:
|
||||
"""Run the shared `passthrough` × <provider> cell body.
|
||||
|
||||
`passthrough_base_path` is appended to the proxy base URL and
|
||||
becomes the CLI's ANTHROPIC_BASE_URL (only the anthropic column
|
||||
uses it; the cloud columns ignore ANTHROPIC_BASE_URL entirely once
|
||||
their CLAUDE_CODE_USE_* flag is set). `build_extra_env` receives
|
||||
the trailing-slash-normalized proxy base URL and returns the
|
||||
provider-mode env for the CLI subprocess.
|
||||
"""
|
||||
environ = env if env is not None else os.environ
|
||||
base_url = environ.get(PROXY_BASE_URL_ENV)
|
||||
api_key = environ.get(PROXY_API_KEY_ENV)
|
||||
if not base_url or not api_key:
|
||||
compat_result.set(
|
||||
{
|
||||
"status": "fail",
|
||||
"error": (
|
||||
f"missing required env: set {PROXY_BASE_URL_ENV} and "
|
||||
f"{PROXY_API_KEY_ENV} to point at a running LiteLLM proxy"
|
||||
),
|
||||
}
|
||||
)
|
||||
pytest.fail(
|
||||
f"{PROXY_BASE_URL_ENV} / {PROXY_API_KEY_ENV} not configured",
|
||||
pytrace=False,
|
||||
)
|
||||
|
||||
proxy_base = base_url.rstrip("/")
|
||||
extra_env = dict(build_extra_env(proxy_base)) if build_extra_env else None
|
||||
|
||||
outcomes = run_models(
|
||||
models=models,
|
||||
prompt=prompt,
|
||||
base_url=proxy_base + passthrough_base_path,
|
||||
api_key=api_key,
|
||||
extra_env=extra_env,
|
||||
)
|
||||
|
||||
failures = []
|
||||
for model in models:
|
||||
outcome = outcomes[model]
|
||||
if isinstance(outcome, ClaudeCLIError):
|
||||
error = f"[{model}] {outcome}"
|
||||
compat_result.add({"status": "fail", "error": error})
|
||||
failures.append(error)
|
||||
continue
|
||||
|
||||
if outcome.exit_code != 0:
|
||||
error = f"[{model}] claude CLI failed: {failure_diagnostic(outcome)}"
|
||||
compat_result.add({"status": "fail", "error": error})
|
||||
failures.append(error)
|
||||
continue
|
||||
|
||||
if not outcome.text.strip():
|
||||
error = f"[{model}] claude returned empty assistant text"
|
||||
compat_result.add({"status": "fail", "error": error})
|
||||
failures.append(error)
|
||||
continue
|
||||
|
||||
compat_result.add({"status": "pass"})
|
||||
|
||||
if failures:
|
||||
pytest.fail("; ".join(failures), pytrace=False)
|
||||
|
|
@ -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] = {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -27,6 +27,15 @@ VERTEXAI_LOCATION=global
|
|||
AZURE_FOUNDRY_API_KEY=
|
||||
AZURE_FOUNDRY_API_BASE=
|
||||
|
||||
# Azure cell of the `passthrough` row. Foundry-mode Claude Code sends
|
||||
# the model in the request body, so the proxy's /azure passthrough
|
||||
# cannot resolve a router alias and falls back to these env vars.
|
||||
# AZURE_API_BASE is the Foundry resource's Anthropic surface, i.e.
|
||||
# https://<resource>.services.ai.azure.com/anthropic ; AZURE_API_KEY
|
||||
# is the same key as AZURE_FOUNDRY_API_KEY.
|
||||
AZURE_API_BASE=
|
||||
AZURE_API_KEY=
|
||||
|
||||
# REQUIRED for publishing: PAT for the `agent-shin` user, used to push
|
||||
# the daily compat-matrix branch to its fork (agent-shin/litellm-docs)
|
||||
# and open the cross-repo PR against BerriAI/litellm-docs. Scopes:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
tests/e2e/claude_code/passthrough/__init__.py
Normal file
0
tests/e2e/claude_code/passthrough/__init__.py
Normal file
44
tests/e2e/claude_code/passthrough/test_anthropic.py
Normal file
44
tests/e2e/claude_code/passthrough/test_anthropic.py
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
"""passthrough x Anthropic.
|
||||
|
||||
Drive the real `claude` CLI in its default first-party mode, but with
|
||||
ANTHROPIC_BASE_URL aimed at the proxy's `/anthropic` passthrough route
|
||||
instead of the `/v1/messages` translation endpoint. The proxy forwards
|
||||
the request verbatim to api.anthropic.com, swapping the virtual-key
|
||||
bearer for its own ANTHROPIC_API_KEY.
|
||||
|
||||
The (feature, provider) for this cell is inferred from the file path by
|
||||
`tests/e2e/claude_code/conftest.py`:
|
||||
|
||||
tests/e2e/claude_code/passthrough/test_anthropic.py
|
||||
^^^^^^^^^^^ ^^^^^^^^^
|
||||
feature_id provider
|
||||
|
||||
Because nothing is translated, the model ids are the real Anthropic API
|
||||
ids (which happen to equal the proxy aliases for this column). A red
|
||||
cell here means the passthrough route broke forwarding itself -- auth
|
||||
header swap, streaming SSE relay, or beta-header propagation -- since
|
||||
no per-provider transformation is involved.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from claude_code._passthrough import (
|
||||
ANTHROPIC_PASSTHROUGH_BASE_PATH,
|
||||
run_passthrough_cell,
|
||||
)
|
||||
|
||||
ANTHROPIC_MODELS = [
|
||||
"claude-haiku-4-5",
|
||||
"claude-sonnet-4-6",
|
||||
"claude-opus-4-7",
|
||||
]
|
||||
|
||||
|
||||
def test_passthrough_anthropic(compat_result):
|
||||
"""Drive the `claude` CLI through `{proxy}/anthropic` and assert a reply."""
|
||||
run_passthrough_cell(
|
||||
compat_result=compat_result,
|
||||
models=ANTHROPIC_MODELS,
|
||||
prompt="Reply with the single word 'pong' and nothing else.",
|
||||
passthrough_base_path=ANTHROPIC_PASSTHROUGH_BASE_PATH,
|
||||
)
|
||||
59
tests/e2e/claude_code/passthrough/test_azure.py
Normal file
59
tests/e2e/claude_code/passthrough/test_azure.py
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
"""passthrough x Azure (Microsoft Foundry).
|
||||
|
||||
Drive the real `claude` CLI in foundry mode (CLAUDE_CODE_USE_FOUNDRY=1)
|
||||
with ANTHROPIC_FOUNDRY_BASE_URL aimed at the proxy's `/azure`
|
||||
passthrough route. The CLI POSTs `/v1/messages` with the model in the
|
||||
JSON body -- unlike the bedrock/vertex modes there is no model segment
|
||||
in the URL, so the proxy's router-alias resolution cannot engage and
|
||||
the `/azure` route falls back to its env-configured target: the proxy
|
||||
must set AZURE_API_BASE to the Foundry resource's Anthropic surface
|
||||
(`https://<resource>.services.ai.azure.com/anthropic`) and
|
||||
AZURE_API_KEY to the Foundry key (see
|
||||
cron_vm/litellm-compat-matrix.env.example). The model ids are the
|
||||
Foundry deployment names, which this matrix provisions to match the
|
||||
Anthropic ids.
|
||||
|
||||
The (feature, provider) for this cell is inferred from the file path by
|
||||
`tests/e2e/claude_code/conftest.py`:
|
||||
|
||||
tests/e2e/claude_code/passthrough/test_azure.py
|
||||
^^^^^^^^^^^ ^^^^^
|
||||
feature_id provider
|
||||
|
||||
Wiring verified live at authoring time: through `{proxy}/azure` the
|
||||
Foundry Anthropic surface accepted the `api-key` / `Authorization:
|
||||
Bearer` headers the fallback sends (a bogus key 401s, the real key
|
||||
proceeds to deployment lookup), so a red cell here means missing
|
||||
AZURE_API_BASE/AZURE_API_KEY on the proxy, missing Foundry deployments
|
||||
for the three tiers, or a genuine forwarding gap -- not an auth-scheme
|
||||
mismatch.
|
||||
|
||||
Known-red at authoring time against a healthy Foundry resource: the
|
||||
`/azure` fallback assembles only its own auth headers and drops the
|
||||
rest of the client's headers, including the `anthropic-version` header
|
||||
the CLI sends, and Foundry's Anthropic surface rejects the request
|
||||
with 400 "anthropic-version: header is required" (the same request
|
||||
sent directly to Foundry with that header succeeds). This cell stays
|
||||
red until that forwarding gap is fixed, which is precisely the class
|
||||
of bug the row exists to surface.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from claude_code._passthrough import foundry_extra_env, run_passthrough_cell
|
||||
|
||||
AZURE_MODELS = [
|
||||
"claude-haiku-4-5",
|
||||
"claude-sonnet-4-6",
|
||||
"claude-opus-4-7",
|
||||
]
|
||||
|
||||
|
||||
def test_passthrough_azure(compat_result):
|
||||
"""Drive the `claude` CLI through `{proxy}/azure` and assert a reply."""
|
||||
run_passthrough_cell(
|
||||
compat_result=compat_result,
|
||||
models=AZURE_MODELS,
|
||||
prompt="Reply with the single word 'pong' and nothing else.",
|
||||
build_extra_env=foundry_extra_env,
|
||||
)
|
||||
34
tests/e2e/claude_code/passthrough/test_bedrock_converse.py
Normal file
34
tests/e2e/claude_code/passthrough/test_bedrock_converse.py
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
"""passthrough x Bedrock (Converse).
|
||||
|
||||
Structurally not applicable. In bedrock mode the `claude` CLI speaks
|
||||
only the InvokeModel wire (`/model/{id}/invoke-with-response-stream`);
|
||||
it has no Converse-wire client, so there is no Claude Code traffic a
|
||||
Converse passthrough could serve. LiteLLM's `/bedrock` route does
|
||||
accept `/model/{id}/converse-stream`, but exercising it would test a
|
||||
wire no Claude Code user can produce, which is out of scope for this
|
||||
matrix.
|
||||
|
||||
The (feature, provider) for this cell is inferred from the file path by
|
||||
`tests/e2e/claude_code/conftest.py`:
|
||||
|
||||
tests/e2e/claude_code/passthrough/test_bedrock_converse.py
|
||||
^^^^^^^^^^^ ^^^^^^^^^^^^^^^^
|
||||
feature_id provider
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def test_passthrough_bedrock_converse(compat_result):
|
||||
"""Report not_applicable: Claude Code has no Converse-wire mode."""
|
||||
compat_result.set(
|
||||
{
|
||||
"status": "not_applicable",
|
||||
"reason": (
|
||||
"Claude Code's bedrock mode speaks only the InvokeModel wire "
|
||||
"(/model/{id}/invoke-with-response-stream); it has no "
|
||||
"Converse-wire client, so there is no Claude Code surface "
|
||||
"for Converse passthrough."
|
||||
),
|
||||
}
|
||||
)
|
||||
42
tests/e2e/claude_code/passthrough/test_bedrock_invoke.py
Normal file
42
tests/e2e/claude_code/passthrough/test_bedrock_invoke.py
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
"""passthrough x Bedrock (Invoke).
|
||||
|
||||
Drive the real `claude` CLI in bedrock mode (CLAUDE_CODE_USE_BEDROCK=1)
|
||||
with ANTHROPIC_BEDROCK_BASE_URL aimed at the proxy's `/bedrock`
|
||||
passthrough route. The CLI speaks the native InvokeModel wire --
|
||||
`POST /model/{model}/invoke-with-response-stream` -- with the proxy
|
||||
alias in the model segment; the proxy resolves the alias through its
|
||||
router, rewrites the path to the deployment's upstream model id, and
|
||||
SigV4-signs the forwarded request with its own AWS credentials
|
||||
(CLAUDE_CODE_SKIP_BEDROCK_AUTH=1 keeps the CLI from signing).
|
||||
|
||||
The (feature, provider) for this cell is inferred from the file path by
|
||||
`tests/e2e/claude_code/conftest.py`:
|
||||
|
||||
tests/e2e/claude_code/passthrough/test_bedrock_invoke.py
|
||||
^^^^^^^^^^^ ^^^^^^^^^^^^^^
|
||||
feature_id provider
|
||||
|
||||
The CLI also fires a best-effort `GET /bedrock/inference-profiles`
|
||||
listing at startup; its failure is non-fatal and does not gate this
|
||||
cell.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from claude_code._passthrough import bedrock_extra_env, run_passthrough_cell
|
||||
|
||||
BEDROCK_INVOKE_MODELS = [
|
||||
"claude-haiku-4-5-bedrock-invoke",
|
||||
"claude-sonnet-4-6-bedrock-invoke",
|
||||
"claude-opus-4-7-bedrock-invoke",
|
||||
]
|
||||
|
||||
|
||||
def test_passthrough_bedrock_invoke(compat_result):
|
||||
"""Drive the `claude` CLI through `{proxy}/bedrock` and assert a reply."""
|
||||
run_passthrough_cell(
|
||||
compat_result=compat_result,
|
||||
models=BEDROCK_INVOKE_MODELS,
|
||||
prompt="Reply with the single word 'pong' and nothing else.",
|
||||
build_extra_env=bedrock_extra_env,
|
||||
)
|
||||
45
tests/e2e/claude_code/passthrough/test_vertex_ai.py
Normal file
45
tests/e2e/claude_code/passthrough/test_vertex_ai.py
Normal file
|
|
@ -0,0 +1,45 @@
|
|||
"""passthrough x Vertex AI.
|
||||
|
||||
Drive the real `claude` CLI in vertex mode (CLAUDE_CODE_USE_VERTEX=1)
|
||||
with ANTHROPIC_VERTEX_BASE_URL aimed at the proxy's `/vertex_ai`
|
||||
passthrough route. The CLI speaks the native rawPredict wire --
|
||||
`POST .../projects/{p}/locations/{l}/publishers/anthropic/models/{model}:streamRawPredict`
|
||||
-- with the proxy alias in the model segment; the proxy resolves the
|
||||
alias through its router, replaces the placeholder project/location
|
||||
path segments with the deployment's `vertex_project` /
|
||||
`vertex_location`, and attaches its own Google credentials
|
||||
(CLAUDE_CODE_SKIP_VERTEX_AUTH=1 keeps the CLI from minting a token).
|
||||
|
||||
The (feature, provider) for this cell is inferred from the file path by
|
||||
`tests/e2e/claude_code/conftest.py`:
|
||||
|
||||
tests/e2e/claude_code/passthrough/test_vertex_ai.py
|
||||
^^^^^^^^^^^ ^^^^^^^^^
|
||||
feature_id provider
|
||||
|
||||
This cell requires the vertex deployments in the proxy config to carry
|
||||
`use_in_pass_through: true` (see test_config.yaml) -- that is what
|
||||
registers their credentials with the passthrough router. Without it
|
||||
the proxy forwards the CLI's own headers (the virtual-key bearer) to
|
||||
Google and every tier fails with a 401.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from claude_code._passthrough import run_passthrough_cell, vertex_extra_env
|
||||
|
||||
VERTEX_MODELS = [
|
||||
"claude-haiku-4-5-vertex",
|
||||
"claude-sonnet-4-6-vertex",
|
||||
"claude-opus-4-7-vertex",
|
||||
]
|
||||
|
||||
|
||||
def test_passthrough_vertex_ai(compat_result):
|
||||
"""Drive the `claude` CLI through `{proxy}/vertex_ai` and assert a reply."""
|
||||
run_passthrough_cell(
|
||||
compat_result=compat_result,
|
||||
models=VERTEX_MODELS,
|
||||
prompt="Reply with the single word 'pong' and nothing else.",
|
||||
build_extra_env=vertex_extra_env,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
20
tests/e2e/router/complexity_router_client.py
Normal file
20
tests/e2e/router/complexity_router_client.py
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
"""Client for the complexity auto-router e2e tests.
|
||||
|
||||
The suite drives the shared /chat/completions and spend-log reads on the Gateway,
|
||||
so this client only carries the Gateway the shared lifecycle needs for cleanup.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from e2e_gateway import Gateway, build_gateway
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ComplexityRouterClient:
|
||||
gateway: Gateway
|
||||
|
||||
|
||||
def build_client() -> ComplexityRouterClient:
|
||||
return ComplexityRouterClient(gateway=build_gateway())
|
||||
15
tests/e2e/router/conftest.py
Normal file
15
tests/e2e/router/conftest.py
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
"""Router suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. ComplexityRouterClient holds the shared
|
||||
Gateway, so the `resources` fixture cleans up keys this suite creates.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from complexity_router_client import ComplexityRouterClient, build_client
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def client() -> ComplexityRouterClient:
|
||||
return build_client()
|
||||
62
tests/e2e/router/test_complexity_router_e2e.py
Normal file
62
tests/e2e/router/test_complexity_router_e2e.py
Normal file
|
|
@ -0,0 +1,62 @@
|
|||
"""Live e2e: the v2 auto-router's LLM complexity classifier actually runs over the
|
||||
proxy and drives routing, instead of silently crashing and falling back to the
|
||||
local heuristic scorer.
|
||||
|
||||
The regression this guards (complexity_router.py `_classifier_call_metadata`
|
||||
returning None when the request carries no `litellm_metadata`, which the classifier
|
||||
sub-call then fed into a `.update`, raising `'NoneType' object has no attribute
|
||||
'update'`) was invisible from the outside: the router caught the error and answered
|
||||
from heuristic scoring, so every request still returned 200. The only tell is which
|
||||
tier, and therefore which backend, served the request.
|
||||
|
||||
`complexity-smart-router` (see the inline config in docker-compose.yml) pins SIMPLE
|
||||
to the openai backend and every higher tier to the anthropic backend. "Is P equal
|
||||
to NP?" is lexically trivial, so the heuristic scorer lands it in SIMPLE (openai),
|
||||
but any competent LLM classifier reads it as a hard reasoning question and lands it
|
||||
above SIMPLE (anthropic). The served deployment is read back from the spend log's
|
||||
`model`, so anthropic proves the classifier ran and openai proves it silently fell
|
||||
back - the exact failure before the fix.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from complexity_router_client import ComplexityRouterClient
|
||||
from e2e_http import unwrap
|
||||
from models import ChatBody, ChatMessage
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
ROUTER_MODEL = "complexity-smart-router"
|
||||
# Lexically simple (heuristic -> SIMPLE) but a hard reasoning question (LLM -> above SIMPLE).
|
||||
LEXICALLY_SIMPLE_HARD_PROMPT = "Is P equal to NP?"
|
||||
# SIMPLE tier backend; served only when the classifier silently falls back to heuristic.
|
||||
HEURISTIC_TIER_MODEL = "openai/gpt-5.5"
|
||||
# MEDIUM/COMPLEX/REASONING tier backend; served only when the LLM classifier runs.
|
||||
LLM_TIER_MODEL = "anthropic/claude-haiku-4-5"
|
||||
|
||||
|
||||
class TestComplexityRouterLlmClassifier:
|
||||
@pytest.mark.covers("reliability.routing.complexity_llm_classifier.routes_by_llm_tier")
|
||||
def test_llm_classifier_runs_and_routes_by_semantic_tier(
|
||||
self, client: ComplexityRouterClient, scoped_key: str
|
||||
) -> None:
|
||||
chat = unwrap(
|
||||
client.gateway.chat(
|
||||
scoped_key,
|
||||
ChatBody(
|
||||
model=ROUTER_MODEL,
|
||||
messages=[ChatMessage(role="user", content=LEXICALLY_SIMPLE_HARD_PROMPT)],
|
||||
max_tokens=16,
|
||||
),
|
||||
)
|
||||
)
|
||||
assert chat.choices, f"router returned no choices: {chat}"
|
||||
|
||||
rows = client.gateway.poll_logs_for_key(scoped_key, min_rows=1)
|
||||
served = [row.model for row in rows]
|
||||
assert served == [LLM_TIER_MODEL], (
|
||||
f"expected the request to be served by {LLM_TIER_MODEL!r} (the higher-tier "
|
||||
f"backend the LLM classifier picks for a hard prompt), but the spend log shows "
|
||||
f"{served!r}. {HEURISTIC_TIER_MODEL!r} means the LLM classifier silently failed "
|
||||
f"and the router fell back to heuristic scoring (SIMPLE) - the pre-fix regression"
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 == {}
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
2091
tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py
Normal file
2091
tests/test_litellm/proxy/guardrails/guardrail_hooks/test_compresr.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -53,7 +53,72 @@ describe("GuardrailTestPanel", () => {
|
|||
|
||||
// Verify onSubmit was called with the correct text
|
||||
await waitFor(() => {
|
||||
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text");
|
||||
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text", null);
|
||||
});
|
||||
});
|
||||
|
||||
it("should submit parsed metadata when a JSON object is provided", async () => {
|
||||
/**
|
||||
* Tests that a JSON object typed into the Metadata field is parsed and
|
||||
* passed to onSubmit so it reaches the apply_guardrail request body.
|
||||
*/
|
||||
const user = userEvent.setup();
|
||||
|
||||
render(
|
||||
<GuardrailTestPanel
|
||||
guardrailNames={mockGuardrailNames}
|
||||
onSubmit={mockOnSubmit}
|
||||
isLoading={false}
|
||||
results={null}
|
||||
errors={null}
|
||||
onClose={mockOnClose}
|
||||
/>,
|
||||
);
|
||||
|
||||
const textarea = screen.getByPlaceholderText("Enter text to test with guardrails...");
|
||||
await user.type(textarea, "Test input text");
|
||||
|
||||
const metadataField = screen.getByPlaceholderText('{"forbidden_topics": ["tax", "finance"]}');
|
||||
await user.click(metadataField);
|
||||
await user.paste('{"forbidden_topics": ["tax"]}');
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Test 2 guardrails/ }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text", { forbidden_topics: ["tax"] });
|
||||
});
|
||||
});
|
||||
|
||||
it("should block submission and show an error for invalid metadata JSON", async () => {
|
||||
/**
|
||||
* Tests that invalid JSON in the Metadata field prevents submission
|
||||
* instead of silently sending a request without metadata.
|
||||
*/
|
||||
const user = userEvent.setup();
|
||||
|
||||
render(
|
||||
<GuardrailTestPanel
|
||||
guardrailNames={mockGuardrailNames}
|
||||
onSubmit={mockOnSubmit}
|
||||
isLoading={false}
|
||||
results={null}
|
||||
errors={null}
|
||||
onClose={mockOnClose}
|
||||
/>,
|
||||
);
|
||||
|
||||
const textarea = screen.getByPlaceholderText("Enter text to test with guardrails...");
|
||||
await user.type(textarea, "Test input text");
|
||||
|
||||
const metadataField = screen.getByPlaceholderText('{"forbidden_topics": ["tax", "finance"]}');
|
||||
await user.click(metadataField);
|
||||
await user.paste("{not json");
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Test 2 guardrails/ }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Invalid JSON")).toBeInTheDocument();
|
||||
});
|
||||
expect(mockOnSubmit).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ const { Text } = Typography;
|
|||
|
||||
interface GuardrailTestPanelProps {
|
||||
guardrailNames: string[];
|
||||
onSubmit: (text: string) => void;
|
||||
onSubmit: (text: string, metadata?: Record<string, unknown> | null) => void;
|
||||
isLoading: boolean;
|
||||
results: Array<{ guardrailName: string; response_text: string; latency: number }> | null;
|
||||
errors: Array<{ guardrailName: string; error: Error; latency: number }> | null;
|
||||
|
|
@ -26,6 +26,23 @@ export function GuardrailTestPanel({
|
|||
onClose,
|
||||
}: GuardrailTestPanelProps) {
|
||||
const [inputText, setInputText] = useState("");
|
||||
const [metadataText, setMetadataText] = useState("");
|
||||
const [metadataError, setMetadataError] = useState<string | null>(null);
|
||||
|
||||
const parseMetadata = (raw: string): { metadata: Record<string, unknown> | null; error: string | null } => {
|
||||
if (!raw.trim()) {
|
||||
return { metadata: null, error: null };
|
||||
}
|
||||
try {
|
||||
const parsed = JSON.parse(raw);
|
||||
if (parsed === null || typeof parsed !== "object" || Array.isArray(parsed)) {
|
||||
return { metadata: null, error: "Metadata must be a JSON object" };
|
||||
}
|
||||
return { metadata: parsed, error: null };
|
||||
} catch {
|
||||
return { metadata: null, error: "Invalid JSON" };
|
||||
}
|
||||
};
|
||||
|
||||
const handleSubmit = () => {
|
||||
if (!inputText.trim()) {
|
||||
|
|
@ -33,7 +50,15 @@ export function GuardrailTestPanel({
|
|||
return;
|
||||
}
|
||||
|
||||
onSubmit(inputText);
|
||||
const { metadata, error } = parseMetadata(metadataText);
|
||||
if (error) {
|
||||
setMetadataError(error);
|
||||
NotificationsManager.fromBackend(`Metadata: ${error}`);
|
||||
return;
|
||||
}
|
||||
setMetadataError(null);
|
||||
|
||||
onSubmit(inputText, metadata);
|
||||
};
|
||||
|
||||
const handleKeyDown = (e: React.KeyboardEvent<HTMLTextAreaElement>) => {
|
||||
|
|
@ -142,6 +167,33 @@ export function GuardrailTestPanel({
|
|||
</div>
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<div className="flex items-center gap-2 mb-2">
|
||||
<label className="text-sm font-medium text-gray-700">Metadata (optional)</label>
|
||||
<Tooltip title="JSON object forwarded to the guardrail as request_data['metadata']. Custom guardrails can read per-request configuration from it.">
|
||||
<InfoCircleOutlined className="text-gray-400 cursor-help" />
|
||||
</Tooltip>
|
||||
</div>
|
||||
<TextArea
|
||||
value={metadataText}
|
||||
onChange={(e) => {
|
||||
setMetadataText(e.target.value);
|
||||
if (metadataError) {
|
||||
setMetadataError(parseMetadata(e.target.value).error);
|
||||
}
|
||||
}}
|
||||
placeholder='{"forbidden_topics": ["tax", "finance"]}'
|
||||
rows={3}
|
||||
className="font-mono text-sm"
|
||||
status={metadataError ? "error" : undefined}
|
||||
/>
|
||||
{metadataError && (
|
||||
<Text type="danger" className="text-xs">
|
||||
{metadataError}
|
||||
</Text>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="pt-2">
|
||||
<Button onClick={handleSubmit} loading={isLoading} disabled={!inputText.trim()} className="w-full">
|
||||
{isLoading
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
|
|||
setSelectedGuardrails(newSelection);
|
||||
};
|
||||
|
||||
const handleTestGuardrails = async (text: string) => {
|
||||
const handleTestGuardrails = async (text: string, metadata?: Record<string, unknown> | null) => {
|
||||
if (selectedGuardrails.size === 0 || !accessToken) {
|
||||
return;
|
||||
}
|
||||
|
|
@ -79,7 +79,7 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
|
|||
Array.from(selectedGuardrails).map(async (guardrailName) => {
|
||||
const startTime = Date.now();
|
||||
try {
|
||||
const result = await applyGuardrail(accessToken, guardrailName, text, null, null);
|
||||
const result = await applyGuardrail(accessToken, guardrailName, text, null, null, metadata);
|
||||
const latency = Date.now() - startTime;
|
||||
results.push({
|
||||
guardrailName,
|
||||
|
|
|
|||
|
|
@ -1,190 +0,0 @@
|
|||
import { describe, it, expect, beforeEach, afterEach, vi } from "vitest";
|
||||
import { act, renderHook, waitFor } from "@testing-library/react";
|
||||
import { useDisableUsageIndicator } from "./useDisableUsageIndicator";
|
||||
import { LOCAL_STORAGE_EVENT } from "@/utils/localStorageUtils";
|
||||
|
||||
describe("useDisableUsageIndicator", () => {
|
||||
const STORAGE_KEY = "disableUsageIndicator";
|
||||
|
||||
beforeEach(() => {
|
||||
localStorage.clear();
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
localStorage.clear();
|
||||
});
|
||||
|
||||
it("should return false when localStorage is empty", () => {
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
});
|
||||
|
||||
it("should return false when localStorage value is not 'true'", () => {
|
||||
localStorage.setItem(STORAGE_KEY, "false");
|
||||
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
});
|
||||
|
||||
it("should return true when localStorage value is 'true'", () => {
|
||||
localStorage.setItem(STORAGE_KEY, "true");
|
||||
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(true);
|
||||
});
|
||||
|
||||
it("should return false when localStorage value is an empty string", () => {
|
||||
localStorage.setItem(STORAGE_KEY, "");
|
||||
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
});
|
||||
|
||||
it("should update when storage event fires for the correct key", async () => {
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
|
||||
await act(async () => {
|
||||
localStorage.setItem(STORAGE_KEY, "true");
|
||||
const storageEvent = new StorageEvent("storage", {
|
||||
key: STORAGE_KEY,
|
||||
newValue: "true",
|
||||
});
|
||||
window.dispatchEvent(storageEvent);
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current).toBe(true);
|
||||
});
|
||||
});
|
||||
|
||||
it("should not update when storage event fires for a different key", () => {
|
||||
localStorage.setItem(STORAGE_KEY, "false");
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
|
||||
const storageEvent = new StorageEvent("storage", {
|
||||
key: "otherKey",
|
||||
newValue: "true",
|
||||
});
|
||||
window.dispatchEvent(storageEvent);
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
});
|
||||
|
||||
it("should update when custom LOCAL_STORAGE_EVENT fires for the correct key", async () => {
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
|
||||
await act(async () => {
|
||||
localStorage.setItem(STORAGE_KEY, "true");
|
||||
const customEvent = new CustomEvent(LOCAL_STORAGE_EVENT, {
|
||||
detail: { key: STORAGE_KEY },
|
||||
});
|
||||
window.dispatchEvent(customEvent);
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current).toBe(true);
|
||||
});
|
||||
});
|
||||
|
||||
it("should not update when custom LOCAL_STORAGE_EVENT fires for a different key", () => {
|
||||
localStorage.setItem(STORAGE_KEY, "false");
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
|
||||
const customEvent = new CustomEvent(LOCAL_STORAGE_EVENT, {
|
||||
detail: { key: "otherKey" },
|
||||
});
|
||||
window.dispatchEvent(customEvent);
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
});
|
||||
|
||||
it("should update when localStorage changes from false to true via custom event", async () => {
|
||||
localStorage.setItem(STORAGE_KEY, "false");
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(false);
|
||||
|
||||
await act(async () => {
|
||||
localStorage.setItem(STORAGE_KEY, "true");
|
||||
const customEvent = new CustomEvent(LOCAL_STORAGE_EVENT, {
|
||||
detail: { key: STORAGE_KEY },
|
||||
});
|
||||
window.dispatchEvent(customEvent);
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current).toBe(true);
|
||||
});
|
||||
});
|
||||
|
||||
it("should update when localStorage changes from true to false via storage event", async () => {
|
||||
localStorage.setItem(STORAGE_KEY, "true");
|
||||
const { result } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result.current).toBe(true);
|
||||
|
||||
await act(async () => {
|
||||
localStorage.setItem(STORAGE_KEY, "false");
|
||||
const storageEvent = new StorageEvent("storage", {
|
||||
key: STORAGE_KEY,
|
||||
newValue: "false",
|
||||
});
|
||||
window.dispatchEvent(storageEvent);
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current).toBe(false);
|
||||
});
|
||||
});
|
||||
|
||||
it("should cleanup event listeners on unmount", () => {
|
||||
const addEventListenerSpy = vi.spyOn(window, "addEventListener");
|
||||
const removeEventListenerSpy = vi.spyOn(window, "removeEventListener");
|
||||
|
||||
const { unmount } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(addEventListenerSpy).toHaveBeenCalledTimes(2);
|
||||
expect(addEventListenerSpy).toHaveBeenCalledWith("storage", expect.any(Function));
|
||||
expect(addEventListenerSpy).toHaveBeenCalledWith(LOCAL_STORAGE_EVENT, expect.any(Function));
|
||||
|
||||
unmount();
|
||||
|
||||
expect(removeEventListenerSpy).toHaveBeenCalledTimes(2);
|
||||
expect(removeEventListenerSpy).toHaveBeenCalledWith("storage", expect.any(Function));
|
||||
expect(removeEventListenerSpy).toHaveBeenCalledWith(LOCAL_STORAGE_EVENT, expect.any(Function));
|
||||
});
|
||||
|
||||
it("should handle multiple hooks independently", async () => {
|
||||
const { result: result1 } = renderHook(() => useDisableUsageIndicator());
|
||||
const { result: result2 } = renderHook(() => useDisableUsageIndicator());
|
||||
|
||||
expect(result1.current).toBe(false);
|
||||
expect(result2.current).toBe(false);
|
||||
|
||||
await act(async () => {
|
||||
localStorage.setItem(STORAGE_KEY, "true");
|
||||
const customEvent = new CustomEvent(LOCAL_STORAGE_EVENT, {
|
||||
detail: { key: STORAGE_KEY },
|
||||
});
|
||||
window.dispatchEvent(customEvent);
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result1.current).toBe(true);
|
||||
expect(result2.current).toBe(true);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
import { getLocalStorageItem, LOCAL_STORAGE_EVENT } from "@/utils/localStorageUtils";
|
||||
import { useSyncExternalStore } from "react";
|
||||
|
||||
function subscribe(callback: () => void) {
|
||||
const onStorage = (e: StorageEvent) => {
|
||||
if (e.key === "disableUsageIndicator") {
|
||||
callback();
|
||||
}
|
||||
};
|
||||
|
||||
const onCustom = (e: Event) => {
|
||||
const { key } = (e as CustomEvent).detail;
|
||||
if (key === "disableUsageIndicator") {
|
||||
callback();
|
||||
}
|
||||
};
|
||||
|
||||
window.addEventListener("storage", onStorage);
|
||||
window.addEventListener(LOCAL_STORAGE_EVENT, onCustom);
|
||||
|
||||
return () => {
|
||||
window.removeEventListener("storage", onStorage);
|
||||
window.removeEventListener(LOCAL_STORAGE_EVENT, onCustom);
|
||||
};
|
||||
}
|
||||
|
||||
function getSnapshot() {
|
||||
return getLocalStorageItem("disableUsageIndicator") === "true";
|
||||
}
|
||||
|
||||
export function useDisableUsageIndicator() {
|
||||
return useSyncExternalStore(subscribe, getSnapshot);
|
||||
}
|
||||
|
|
@ -10,7 +10,7 @@ import { Button } from "@/components/ui/button";
|
|||
import MessageManager from "@/components/molecules/message_manager";
|
||||
import { useRouter } from "next/navigation";
|
||||
import { useChatShell } from "@/contexts/ChatShellContext";
|
||||
import { CHAT_ROUTES } from "@/components/chat/ChatShell";
|
||||
import { getChatRoutes } from "@/components/chat/ChatShell";
|
||||
import ChatMessages from "@/components/chat/ChatMessages";
|
||||
import MCPConnectPicker from "@/components/chat/MCPConnectPicker";
|
||||
import { fetchAvailableModels } from "@/components/llm_calls/fetch_models";
|
||||
|
|
@ -87,7 +87,7 @@ export default function ChatConversationPage() {
|
|||
const streamScrollLock = useRef<number | null>(null);
|
||||
|
||||
useEffect(() => {
|
||||
if (staleId) router.replace(CHAT_ROUTES.chats);
|
||||
if (staleId) router.replace(getChatRoutes().chats);
|
||||
}, [staleId, router]);
|
||||
|
||||
// Load models
|
||||
|
|
@ -140,7 +140,7 @@ export default function ChatConversationPage() {
|
|||
if (!convId) {
|
||||
convId = createConversation(model);
|
||||
setResponsesSessionId(null); // new conversation starts a fresh session
|
||||
router.push(`${CHAT_ROUTES.chats}?id=${convId}`);
|
||||
window.history.pushState(null, "", `${window.location.pathname}?id=${convId}`);
|
||||
}
|
||||
|
||||
appendMessage(convId, { role: "user", content: trimmed });
|
||||
|
|
@ -248,7 +248,6 @@ export default function ChatConversationPage() {
|
|||
createConversation,
|
||||
appendMessage,
|
||||
updateLastAssistantMessage,
|
||||
router,
|
||||
isStreaming,
|
||||
responsesSessionId,
|
||||
],
|
||||
|
|
@ -529,7 +528,7 @@ export default function ChatConversationPage() {
|
|||
Chat with 100+ LLMs + MCP tools; authenticate once, use them here.{" "}
|
||||
<Button
|
||||
variant="link"
|
||||
onClick={() => router.push(CHAT_ROUTES.integrations)}
|
||||
onClick={() => router.push(getChatRoutes().integrations)}
|
||||
className="h-auto p-0 text-sm font-medium"
|
||||
>
|
||||
Open Integrations ->
|
||||
|
|
|
|||
70
ui/litellm-dashboard/src/components/BetaBadge.test.tsx
Normal file
70
ui/litellm-dashboard/src/components/BetaBadge.test.tsx
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import BetaBadge from "./BetaBadge";
|
||||
|
||||
// Mock the hook directly
|
||||
vi.mock("@/app/(dashboard)/hooks/useDisableShowNewBadge", () => ({
|
||||
useDisableShowNewBadge: vi.fn(),
|
||||
}));
|
||||
|
||||
import { useDisableShowNewBadge } from "@/app/(dashboard)/hooks/useDisableShowNewBadge";
|
||||
|
||||
const mockUseDisableShowNewBadge = vi.mocked(useDisableShowNewBadge);
|
||||
|
||||
describe("BetaBadge", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should render the badge when disableShowNewBadge is false", () => {
|
||||
mockUseDisableShowNewBadge.mockReturnValue(false);
|
||||
|
||||
render(<BetaBadge>Test Content</BetaBadge>);
|
||||
|
||||
expect(screen.getByText("Beta")).toBeInTheDocument();
|
||||
expect(screen.getByText("Test Content")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render the badge when disableShowNewBadge is not set", () => {
|
||||
mockUseDisableShowNewBadge.mockReturnValue(false);
|
||||
|
||||
render(<BetaBadge />);
|
||||
|
||||
expect(screen.getByText("Beta")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render only children when disableShowNewBadge is true", () => {
|
||||
mockUseDisableShowNewBadge.mockReturnValue(true);
|
||||
|
||||
render(<BetaBadge>Test Content</BetaBadge>);
|
||||
|
||||
expect(screen.queryByText("Beta")).not.toBeInTheDocument();
|
||||
expect(screen.getByText("Test Content")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render nothing when disableShowNewBadge is true and no children", () => {
|
||||
mockUseDisableShowNewBadge.mockReturnValue(true);
|
||||
|
||||
const { container } = render(<BetaBadge />);
|
||||
|
||||
expect(container.firstChild).toBeNull();
|
||||
});
|
||||
|
||||
it("should render badge with dot instead of text when dot prop is true", () => {
|
||||
mockUseDisableShowNewBadge.mockReturnValue(false);
|
||||
|
||||
render(<BetaBadge dot={true}>Test Content</BetaBadge>);
|
||||
|
||||
expect(screen.queryByText("Beta")).not.toBeInTheDocument();
|
||||
expect(screen.getByText("Test Content")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render badge with 'Beta' text when dot prop is not provided (defaults to false)", () => {
|
||||
mockUseDisableShowNewBadge.mockReturnValue(false);
|
||||
|
||||
render(<BetaBadge>Test Content</BetaBadge>);
|
||||
|
||||
expect(screen.getByText("Beta")).toBeInTheDocument();
|
||||
expect(screen.getByText("Test Content")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
18
ui/litellm-dashboard/src/components/BetaBadge.tsx
Normal file
18
ui/litellm-dashboard/src/components/BetaBadge.tsx
Normal file
|
|
@ -0,0 +1,18 @@
|
|||
import { Badge } from "antd";
|
||||
import { useDisableShowNewBadge } from "@/app/(dashboard)/hooks/useDisableShowNewBadge";
|
||||
|
||||
export default function BetaBadge({ children, dot = false }: { children?: React.ReactNode; dot?: boolean }) {
|
||||
const disableShowNewBadge = useDisableShowNewBadge();
|
||||
|
||||
if (disableShowNewBadge) {
|
||||
return children ? <>{children}</> : null;
|
||||
}
|
||||
|
||||
return children ? (
|
||||
<Badge color="blue" count={dot ? undefined : "Beta"} dot={dot}>
|
||||
{children}
|
||||
</Badge>
|
||||
) : (
|
||||
<Badge color="blue" count={dot ? undefined : "Beta"} dot={dot} />
|
||||
);
|
||||
}
|
||||
|
|
@ -2,7 +2,6 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
|||
import { useDisableBlogPosts } from "@/app/(dashboard)/hooks/useDisableBlogPosts";
|
||||
import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBouncingIcon";
|
||||
import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts";
|
||||
import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator";
|
||||
import {
|
||||
emitLocalStorageChange,
|
||||
getLocalStorageItem,
|
||||
|
|
@ -72,7 +71,6 @@ interface UserDropdownProps {
|
|||
const UserDropdown: React.FC<UserDropdownProps> = ({ onLogout, variant = "navbar", collapsed = false }) => {
|
||||
const { userId, userEmail, userRole, premiumUser } = useAuthorized();
|
||||
const disableShowPrompts = useDisableShowPrompts();
|
||||
const disableUsageIndicator = useDisableUsageIndicator();
|
||||
const disableBlogPosts = useDisableBlogPosts();
|
||||
const disableBouncingIcon = useDisableBouncingIcon();
|
||||
const [disableShowNewBadge, setDisableShowNewBadge] = useState(false);
|
||||
|
|
@ -165,23 +163,6 @@ const UserDropdown: React.FC<UserDropdownProps> = ({ onLogout, variant = "navbar
|
|||
aria-label="Toggle hide all prompts"
|
||||
/>
|
||||
</Space>
|
||||
<Space style={{ width: "100%", justifyContent: "space-between" }}>
|
||||
<Text type="secondary">Hide Usage Indicator</Text>
|
||||
<Switch
|
||||
size="small"
|
||||
checked={disableUsageIndicator}
|
||||
onChange={(checked) => {
|
||||
if (checked) {
|
||||
setLocalStorageItem("disableUsageIndicator", "true");
|
||||
emitLocalStorageChange("disableUsageIndicator");
|
||||
} else {
|
||||
removeLocalStorageItem("disableUsageIndicator");
|
||||
emitLocalStorageChange("disableUsageIndicator");
|
||||
}
|
||||
}}
|
||||
aria-label="Toggle hide usage indicator"
|
||||
/>
|
||||
</Space>
|
||||
<Space style={{ width: "100%", justifyContent: "space-between" }}>
|
||||
<Text type="secondary">Hide Blog Posts</Text>
|
||||
<Switch
|
||||
|
|
|
|||
|
|
@ -37,10 +37,6 @@ vi.mock("@/app/(dashboard)/hooks/useDisableShowPrompts", () => ({
|
|||
useDisableShowPrompts: () => mockUseDisableShowPromptsImpl(),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/useDisableUsageIndicator", () => ({
|
||||
useDisableUsageIndicator: () => false,
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/useDisableBlogPosts", () => ({
|
||||
useDisableBlogPosts: () => false,
|
||||
}));
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ import { useDisableBlogPosts } from "@/app/(dashboard)/hooks/useDisableBlogPosts
|
|||
import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBouncingIcon";
|
||||
import { useDisableShowNewBadge } from "@/app/(dashboard)/hooks/useDisableShowNewBadge";
|
||||
import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts";
|
||||
import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator";
|
||||
import { emitLocalStorageChange, removeLocalStorageItem, setLocalStorageItem } from "@/utils/localStorageUtils";
|
||||
import { navAccountDisplayName } from "@/components/Navbar/navDisplayName";
|
||||
import CopyButton from "@/components/shared/CopyButton";
|
||||
|
|
@ -86,7 +85,6 @@ const SidebarAccountMenu: React.FC<SidebarAccountMenuProps> = ({ onLogout, colla
|
|||
const { data: healthData } = useHealthReadinessDetails(accessToken);
|
||||
const version = healthData?.litellm_version;
|
||||
const disableShowPrompts = useDisableShowPrompts();
|
||||
const disableUsageIndicator = useDisableUsageIndicator();
|
||||
const disableBlogPosts = useDisableBlogPosts();
|
||||
const disableBouncingIcon = useDisableBouncingIcon();
|
||||
const disableShowNewBadge = useDisableShowNewBadge();
|
||||
|
|
@ -115,13 +113,6 @@ const SidebarAccountMenu: React.FC<SidebarAccountMenuProps> = ({ onLogout, colla
|
|||
checked: disableShowPrompts,
|
||||
onCheckedChange: (checked: boolean) => setFlag("disableShowPrompts", checked),
|
||||
},
|
||||
{
|
||||
key: "disableUsageIndicator",
|
||||
label: "Hide Usage Indicator",
|
||||
ariaLabel: "Toggle hide usage indicator",
|
||||
checked: disableUsageIndicator,
|
||||
onCheckedChange: (checked: boolean) => setFlag("disableUsageIndicator", checked),
|
||||
},
|
||||
{
|
||||
key: "disableBlogPosts",
|
||||
label: "Hide Blog Posts",
|
||||
|
|
|
|||
|
|
@ -8,10 +8,6 @@ import type { LicenseInfo } from "./networking";
|
|||
|
||||
vi.mock("./networking", () => ({ getRemainingUsers: vi.fn() }));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/useDisableUsageIndicator", () => ({
|
||||
useDisableUsageIndicator: vi.fn(() => false),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/license/useLicenseInfo", () => ({
|
||||
useLicenseInfo: vi.fn(),
|
||||
}));
|
||||
|
|
@ -128,6 +124,29 @@ describe("SidebarUsageCard", () => {
|
|||
expect(container.querySelector('[data-slot="meter"]')).toBeNull();
|
||||
});
|
||||
|
||||
it("shows the exact license expiration date as the subtitle instead of time remaining", async () => {
|
||||
mockUseLicenseInfo.mockReturnValue(licenseResult({ ...ACTIVE_LICENSE, expiration_date: "2099-12-31" }));
|
||||
|
||||
renderWithClient(<SidebarUsageCard accessToken="token" collapsed={false} onExpandRail={() => {}} />);
|
||||
|
||||
expect(await screen.findByText("Expires Dec 31, 2099")).toBeInTheDocument();
|
||||
expect(screen.queryByText(/(day|days|month|months) remaining/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("shows the exact date as the subtitle when the license is expired", async () => {
|
||||
mockUseLicenseInfo.mockReturnValue(licenseResult({ ...ACTIVE_LICENSE, expiration_date: "2020-01-01" }));
|
||||
|
||||
renderWithClient(<SidebarUsageCard accessToken="token" collapsed={false} onExpandRail={() => {}} />);
|
||||
|
||||
expect(await screen.findByText("Expired Jan 1, 2020")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("falls back to Active plan when the license has no expiration date", async () => {
|
||||
renderWithClient(<SidebarUsageCard accessToken="token" collapsed={false} onExpandRail={() => {}} />);
|
||||
|
||||
expect(await screen.findByText("Active plan")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("shows a collapsed rail button that expands the sidebar", async () => {
|
||||
const onExpandRail = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator";
|
||||
import { useLicenseInfo } from "@/app/(dashboard)/hooks/license/useLicenseInfo";
|
||||
import { getDaysUntilExpiration } from "@/utils/licenseUtils";
|
||||
import { formatExpirationStatus } from "@/utils/licenseUtils";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
|
||||
import { Meter, MeterIndicator, MeterLabel, MeterTrack } from "@/components/ui/meter";
|
||||
|
|
@ -20,16 +19,6 @@ interface MeterData {
|
|||
total: number;
|
||||
}
|
||||
|
||||
const formatExpiration = (daysRemaining: number | null): string => {
|
||||
if (daysRemaining === null) return "No expiration";
|
||||
if (daysRemaining < 0) return "Expired";
|
||||
if (daysRemaining === 0) return "Expires today";
|
||||
if (daysRemaining === 1) return "1 day remaining";
|
||||
if (daysRemaining < 30) return `${daysRemaining} days remaining`;
|
||||
if (daysRemaining < 60) return "1 month remaining";
|
||||
return `${Math.floor(daysRemaining / 30)} months remaining`;
|
||||
};
|
||||
|
||||
const meterTone = (pct: number): "default" | "warning" | "over" => {
|
||||
if (pct > 100) return "over";
|
||||
if (pct >= 80) return "warning";
|
||||
|
|
@ -79,7 +68,6 @@ const buildMeters = (data: RemainingUsage | null): MeterData[] => {
|
|||
* design's Spend / API-request meters are intentionally omitted.
|
||||
*/
|
||||
export default function SidebarUsageCard({ accessToken, collapsed, onExpandRail }: SidebarUsageCardProps) {
|
||||
const disableUsageIndicator = useDisableUsageIndicator();
|
||||
const licenseInfo = useLicenseInfo(accessToken).data ?? null;
|
||||
const { data: usageData, isLoading } = useQuery(remainingUsersQuery(accessToken));
|
||||
const data = usageData ?? null;
|
||||
|
|
@ -87,7 +75,7 @@ export default function SidebarUsageCard({ accessToken, collapsed, onExpandRail
|
|||
const hasData = data !== null && (data.total_users !== null || data.total_teams !== null);
|
||||
const noUsableData = !isLoading && !hasData;
|
||||
const noLicensedUsage = !licenseInfo?.has_license || noUsableData;
|
||||
if (disableUsageIndicator || !accessToken || noLicensedUsage) {
|
||||
if (!accessToken || noLicensedUsage) {
|
||||
return null;
|
||||
}
|
||||
|
||||
|
|
@ -104,8 +92,7 @@ export default function SidebarUsageCard({ accessToken, collapsed, onExpandRail
|
|||
);
|
||||
}
|
||||
|
||||
const daysUntilExpiration = licenseInfo?.expiration_date ? getDaysUntilExpiration(licenseInfo.expiration_date) : null;
|
||||
const subtitle = licenseInfo?.expiration_date ? formatExpiration(daysUntilExpiration) : "Active plan";
|
||||
const subtitle = licenseInfo?.expiration_date ? formatExpirationStatus(licenseInfo.expiration_date) : "Active plan";
|
||||
const meters = buildMeters(data);
|
||||
|
||||
return (
|
||||
|
|
|
|||
|
|
@ -1,193 +0,0 @@
|
|||
import React from "react";
|
||||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import { render, screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import UsageIndicator from "./UsageIndicator";
|
||||
|
||||
vi.mock("./networking", () => ({
|
||||
getRemainingUsers: vi.fn(),
|
||||
getLicenseInfo: vi.fn().mockResolvedValue(null),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/useDisableUsageIndicator", () => ({
|
||||
useDisableUsageIndicator: vi.fn(() => false),
|
||||
}));
|
||||
|
||||
import { getRemainingUsers } from "./networking";
|
||||
|
||||
const mockGetRemainingUsers = vi.mocked(getRemainingUsers);
|
||||
|
||||
const renderWithClient = (ui: React.ReactElement) => {
|
||||
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
|
||||
return render(<QueryClientProvider client={queryClient}>{ui}</QueryClientProvider>);
|
||||
};
|
||||
|
||||
const DEFAULT_USAGE_DATA = {
|
||||
total_users: 100,
|
||||
total_users_used: 1,
|
||||
total_users_remaining: 99,
|
||||
total_teams: null,
|
||||
total_teams_used: 0,
|
||||
total_teams_remaining: null,
|
||||
};
|
||||
|
||||
describe("UsageIndicator", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
mockGetRemainingUsers.mockResolvedValue(DEFAULT_USAGE_DATA);
|
||||
});
|
||||
|
||||
it("should render when given access token and usage data loads", async () => {
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await screen.findByText("Usage");
|
||||
|
||||
expect(screen.getByText("Usage")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not show Near limit when users usage is below 80% (1/100 -> 1%)", async () => {
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await screen.findByText("Usage");
|
||||
|
||||
expect(screen.queryByText("Near limit")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render nothing when both total_users and total_teams are null", async () => {
|
||||
mockGetRemainingUsers.mockResolvedValue({
|
||||
total_users: null,
|
||||
total_teams: null,
|
||||
total_users_used: 520,
|
||||
total_teams_used: 4,
|
||||
total_teams_remaining: null,
|
||||
total_users_remaining: null,
|
||||
});
|
||||
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.queryByText("Usage")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("Loading...")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should show Near limit for Teams when at 80% usage (4/5)", async () => {
|
||||
mockGetRemainingUsers.mockResolvedValue({
|
||||
total_users: null,
|
||||
total_users_used: 0,
|
||||
total_users_remaining: null,
|
||||
total_teams: 5,
|
||||
total_teams_used: 4,
|
||||
total_teams_remaining: 1,
|
||||
});
|
||||
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await screen.findByText("Usage");
|
||||
|
||||
expect(screen.getByText("Teams")).toBeInTheDocument();
|
||||
expect(screen.getByText("Near limit")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show Over limit for Users when usage exceeds 100% (105/100)", async () => {
|
||||
mockGetRemainingUsers.mockResolvedValue({
|
||||
total_users: 100,
|
||||
total_users_used: 105,
|
||||
total_users_remaining: -5,
|
||||
total_teams: null,
|
||||
total_teams_used: 0,
|
||||
total_teams_remaining: null,
|
||||
});
|
||||
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await screen.findByText("Usage");
|
||||
|
||||
expect(screen.getByText("Users")).toBeInTheDocument();
|
||||
expect(screen.getByText("Over limit")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show Over limit for Teams when usage exceeds 100%", async () => {
|
||||
mockGetRemainingUsers.mockResolvedValue({
|
||||
total_users: null,
|
||||
total_users_used: 0,
|
||||
total_users_remaining: null,
|
||||
total_teams: 10,
|
||||
total_teams_used: 12,
|
||||
total_teams_remaining: -2,
|
||||
});
|
||||
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await screen.findByText("Usage");
|
||||
|
||||
expect(screen.getByText("Teams")).toBeInTheDocument();
|
||||
expect(screen.getByText("Over limit")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render nothing when accessToken is null", () => {
|
||||
renderWithClient(<UsageIndicator accessToken={null} width={220} />);
|
||||
|
||||
expect(mockGetRemainingUsers).not.toHaveBeenCalled();
|
||||
expect(screen.queryByText("Usage")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render nothing when disableUsageIndicator is true", async () => {
|
||||
const { useDisableUsageIndicator } = await import("@/app/(dashboard)/hooks/useDisableUsageIndicator");
|
||||
(useDisableUsageIndicator as ReturnType<typeof vi.fn>).mockReturnValue(true);
|
||||
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.queryByText("Usage")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
(useDisableUsageIndicator as ReturnType<typeof vi.fn>).mockReturnValue(false);
|
||||
});
|
||||
|
||||
it("should show Loading while fetching", () => {
|
||||
mockGetRemainingUsers.mockImplementation(() => new Promise(() => {}));
|
||||
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
expect(screen.getByText("Loading...")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show error message when fetch fails", async () => {
|
||||
const consoleSpy = vi.spyOn(console, "error").mockImplementation(() => {});
|
||||
mockGetRemainingUsers.mockRejectedValue(new Error("Network error"));
|
||||
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
expect(await screen.findByText("Failed to load usage data")).toBeInTheDocument();
|
||||
|
||||
consoleSpy.mockRestore();
|
||||
});
|
||||
|
||||
it("should minimize when user clicks minimize button", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await screen.findByText("Usage");
|
||||
|
||||
const minimizeButton = screen.getByTitle("Minimize");
|
||||
await user.click(minimizeButton);
|
||||
|
||||
expect(screen.queryByText("Users")).not.toBeInTheDocument();
|
||||
expect(screen.getByTitle("Show usage details")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should restore from minimized when user clicks restore button", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithClient(<UsageIndicator accessToken="token" width={220} />);
|
||||
|
||||
await screen.findByText("Usage");
|
||||
|
||||
await user.click(screen.getByTitle("Minimize"));
|
||||
await user.click(screen.getByTitle("Show usage details"));
|
||||
|
||||
expect(screen.getByText("Usage")).toBeInTheDocument();
|
||||
expect(screen.getByText("Users")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -1,681 +0,0 @@
|
|||
import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator";
|
||||
import { Badge } from "@tremor/react";
|
||||
import {
|
||||
AlertTriangle,
|
||||
Calendar,
|
||||
ChevronDown,
|
||||
ChevronUp,
|
||||
Loader2,
|
||||
Minus,
|
||||
TrendingUp,
|
||||
UserCheck,
|
||||
Users,
|
||||
} from "lucide-react";
|
||||
import { useEffect, useState } from "react";
|
||||
import { getRemainingUsers } from "./networking";
|
||||
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { getDaysUntilExpiration } from "@/utils/licenseUtils";
|
||||
import { useLicenseInfo } from "@/app/(dashboard)/hooks/license/useLicenseInfo";
|
||||
|
||||
interface UsageIndicatorProps {
|
||||
accessToken: string | null;
|
||||
width: number;
|
||||
}
|
||||
|
||||
interface UsageData {
|
||||
total_users: number | null;
|
||||
total_users_used: number;
|
||||
total_users_remaining: number | null;
|
||||
total_teams: number | null;
|
||||
total_teams_used: number;
|
||||
total_teams_remaining: number | null;
|
||||
}
|
||||
|
||||
// Format expiration for display
|
||||
const formatExpirationDisplay = (daysRemaining: number | null): string => {
|
||||
if (daysRemaining === null) return "No expiration";
|
||||
if (daysRemaining < 0) return "Expired";
|
||||
if (daysRemaining === 0) return "Expires today";
|
||||
if (daysRemaining === 1) return "1 day remaining";
|
||||
if (daysRemaining < 30) return `${daysRemaining} days remaining`;
|
||||
if (daysRemaining < 60) return "1 month remaining";
|
||||
const months = Math.floor(daysRemaining / 30);
|
||||
return `${months} months remaining`;
|
||||
};
|
||||
|
||||
export default function UsageIndicator({ accessToken, width = 220 }: UsageIndicatorProps) {
|
||||
const disableUsageIndicator = useDisableUsageIndicator();
|
||||
const [isExpanded, setIsExpanded] = useState(false);
|
||||
const [isMinimized, setIsMinimized] = useState(false);
|
||||
const [data, setData] = useState<UsageData | null>(null);
|
||||
const [isLoading, setIsLoading] = useState(false);
|
||||
const [error, setError] = useState<string | null>(null);
|
||||
|
||||
const licenseInfo = useLicenseInfo(accessToken).data ?? null;
|
||||
|
||||
useEffect(() => {
|
||||
const fetchData = async () => {
|
||||
if (!accessToken) return;
|
||||
|
||||
setIsLoading(true);
|
||||
setError(null);
|
||||
|
||||
try {
|
||||
const usageResult = await getRemainingUsers(accessToken);
|
||||
setData(usageResult);
|
||||
} catch (err) {
|
||||
console.error("Failed to fetch usage data:", err);
|
||||
setError("Failed to load usage data");
|
||||
} finally {
|
||||
setIsLoading(false);
|
||||
}
|
||||
};
|
||||
|
||||
fetchData();
|
||||
}, [accessToken]);
|
||||
|
||||
// Calculate license expiration metrics
|
||||
const daysUntilExpiration = licenseInfo?.expiration_date ? getDaysUntilExpiration(licenseInfo.expiration_date) : null;
|
||||
const isLicenseExpired = daysUntilExpiration !== null && daysUntilExpiration < 0;
|
||||
const isLicenseExpiringSoon = daysUntilExpiration !== null && daysUntilExpiration >= 0 && daysUntilExpiration < 30;
|
||||
|
||||
// Calculate derived values from data
|
||||
const getUsageMetrics = (data: UsageData | null) => {
|
||||
if (!data) {
|
||||
return {
|
||||
isOverLimit: false,
|
||||
isNearLimit: false,
|
||||
usagePercentage: 0,
|
||||
userMetrics: {
|
||||
isOverLimit: false,
|
||||
isNearLimit: false,
|
||||
usagePercentage: 0,
|
||||
},
|
||||
teamMetrics: {
|
||||
isOverLimit: false,
|
||||
isNearLimit: false,
|
||||
usagePercentage: 0,
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
// User metrics
|
||||
const userUsagePercentage = data.total_users ? (data.total_users_used / data.total_users) * 100 : 0;
|
||||
const userIsOverLimit = userUsagePercentage > 100;
|
||||
const userIsNearLimit = userUsagePercentage >= 80 && userUsagePercentage <= 100;
|
||||
|
||||
// Team metrics
|
||||
const teamUsagePercentage = data.total_teams ? (data.total_teams_used / data.total_teams) * 100 : 0;
|
||||
const teamIsOverLimit = teamUsagePercentage > 100;
|
||||
const teamIsNearLimit = teamUsagePercentage >= 80 && teamUsagePercentage <= 100;
|
||||
|
||||
// Combined status (worst case scenario)
|
||||
const isOverLimit = userIsOverLimit || teamIsOverLimit;
|
||||
const isNearLimit = (userIsNearLimit || teamIsNearLimit) && !isOverLimit;
|
||||
const usagePercentage = Math.max(userUsagePercentage, teamUsagePercentage);
|
||||
|
||||
return {
|
||||
isOverLimit,
|
||||
isNearLimit,
|
||||
usagePercentage,
|
||||
userMetrics: {
|
||||
isOverLimit: userIsOverLimit,
|
||||
isNearLimit: userIsNearLimit,
|
||||
usagePercentage: userUsagePercentage,
|
||||
},
|
||||
teamMetrics: {
|
||||
isOverLimit: teamIsOverLimit,
|
||||
isNearLimit: teamIsNearLimit,
|
||||
usagePercentage: teamUsagePercentage,
|
||||
},
|
||||
};
|
||||
};
|
||||
|
||||
const { isOverLimit, isNearLimit, usagePercentage, userMetrics, teamMetrics } = getUsageMetrics(data);
|
||||
|
||||
// Include license status in overall status
|
||||
const hasAnyIssue = isOverLimit || isNearLimit || isLicenseExpired || isLicenseExpiringSoon;
|
||||
const hasError = isOverLimit || isLicenseExpired;
|
||||
const hasWarning = (isNearLimit || isLicenseExpiringSoon) && !hasError;
|
||||
|
||||
const getStatusColor = () => {
|
||||
if (hasError) return "red";
|
||||
if (hasWarning) return "yellow";
|
||||
return "green";
|
||||
};
|
||||
|
||||
const getStatusIcon = () => {
|
||||
if (hasError) return <AlertTriangle className="h-3 w-3" />;
|
||||
if (hasWarning) return <TrendingUp className="h-3 w-3" />;
|
||||
return null;
|
||||
};
|
||||
|
||||
// Minimized view - just a small restore button
|
||||
const MinimizedView = () => {
|
||||
return (
|
||||
<div className="px-3 py-1" style={{ maxWidth: `${width}px` }}>
|
||||
<button
|
||||
onClick={() => setIsMinimized(false)}
|
||||
className={cn(
|
||||
"flex items-center gap-2 text-xs text-gray-400 hover:text-gray-600 transition-colors p-1 rounded-sm w-full",
|
||||
hasError && "text-red-400 hover:text-red-600",
|
||||
hasWarning && "text-yellow-500 hover:text-yellow-700",
|
||||
)}
|
||||
title="Show usage details"
|
||||
>
|
||||
<Users className="h-3 w-3 shrink-0" />
|
||||
{hasAnyIssue && <span className="shrink-0">{getStatusIcon()}</span>}
|
||||
<div className="flex items-center gap-1 truncate">
|
||||
{data && data.total_users !== null && (
|
||||
<span className="shrink-0">
|
||||
U:{data.total_users_used}/{data.total_users}
|
||||
</span>
|
||||
)}
|
||||
{data && data.total_teams !== null && (
|
||||
<span className="shrink-0">
|
||||
T:{data.total_teams_used}/{data.total_teams}
|
||||
</span>
|
||||
)}
|
||||
{licenseInfo?.expiration_date && daysUntilExpiration !== null && (
|
||||
<span
|
||||
className={cn(
|
||||
"shrink-0",
|
||||
isLicenseExpired && "text-red-500",
|
||||
isLicenseExpiringSoon && "text-yellow-500",
|
||||
)}
|
||||
>
|
||||
{daysUntilExpiration < 0 ? "Exp!" : `${daysUntilExpiration}d`}
|
||||
</span>
|
||||
)}
|
||||
{!data ||
|
||||
(data.total_users === null && data.total_teams === null && !licenseInfo && (
|
||||
<span className="truncate">Usage</span>
|
||||
))}
|
||||
</div>
|
||||
</button>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
// Sidebar/nav style component
|
||||
const NavStyleView = () => {
|
||||
if (isMinimized) {
|
||||
return <MinimizedView />;
|
||||
}
|
||||
|
||||
if (isLoading) {
|
||||
return (
|
||||
<div className="flex items-center gap-3 px-3 py-2 text-gray-500" style={{ maxWidth: `${width}px` }}>
|
||||
<Loader2 className="h-4 w-4 animate-spin shrink-0" />
|
||||
<span className="text-sm truncate">Loading...</span>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (error || !data) {
|
||||
return (
|
||||
<div
|
||||
className="flex items-center justify-between gap-3 px-3 py-2 text-gray-400 group"
|
||||
style={{ maxWidth: `${width}px` }}
|
||||
>
|
||||
<div className="flex items-center gap-3 min-w-0 flex-1">
|
||||
<Users className="h-4 w-4 shrink-0" />
|
||||
<span className="text-sm truncate">{error || "No data"}</span>
|
||||
</div>
|
||||
<button
|
||||
onClick={() => setIsMinimized(true)}
|
||||
className="opacity-0 group-hover:opacity-100 p-0.5 hover:bg-gray-100 rounded-sm transition-all shrink-0"
|
||||
title="Minimize"
|
||||
>
|
||||
<Minus className="h-3 w-3" />
|
||||
</button>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="px-3 py-2 group" style={{ maxWidth: `${width}px` }}>
|
||||
{/* Main nav item style */}
|
||||
<div className="flex items-center justify-between">
|
||||
<button
|
||||
onClick={() => setIsExpanded(!isExpanded)}
|
||||
className={cn(
|
||||
"flex items-center gap-3 text-left hover:bg-gray-50 rounded-md px-0 py-1 transition-colors flex-1 min-w-0",
|
||||
hasError && "text-red-600",
|
||||
hasWarning && "text-yellow-600",
|
||||
)}
|
||||
>
|
||||
<Users className="h-4 w-4 shrink-0" />
|
||||
<span className="text-sm font-medium truncate">Usage Status</span>
|
||||
{hasAnyIssue && (
|
||||
<Badge color={getStatusColor()} className="text-xs px-1.5 py-0.5 shrink-0">
|
||||
{getStatusIcon()}
|
||||
</Badge>
|
||||
)}
|
||||
{isExpanded ? (
|
||||
<ChevronUp className="h-3 w-3 text-gray-400 ml-auto shrink-0" />
|
||||
) : (
|
||||
<ChevronDown className="h-3 w-3 text-gray-400 ml-auto shrink-0" />
|
||||
)}
|
||||
</button>
|
||||
|
||||
{/* Minimize button */}
|
||||
<button
|
||||
onClick={() => setIsMinimized(true)}
|
||||
className="opacity-0 group-hover:opacity-100 p-0.5 hover:bg-gray-100 rounded-sm transition-all ml-1 shrink-0"
|
||||
title="Minimize"
|
||||
>
|
||||
<Minus className="h-3 w-3 text-gray-400" />
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{/* Expanded details - simple and compact */}
|
||||
{isExpanded && (
|
||||
<div className="mt-2 pl-7 text-xs text-gray-600 space-y-3">
|
||||
{/* License expiration section */}
|
||||
{licenseInfo?.has_license && licenseInfo.expiration_date && (
|
||||
<div>
|
||||
<div className="mb-1 flex items-center gap-1">
|
||||
<Calendar className="h-3 w-3" />
|
||||
<span className="font-medium">License</span>
|
||||
</div>
|
||||
<div
|
||||
className={cn(
|
||||
"flex items-center gap-1 text-xs",
|
||||
isLicenseExpired && "text-red-600",
|
||||
isLicenseExpiringSoon && "text-yellow-600",
|
||||
)}
|
||||
>
|
||||
{isLicenseExpired ? (
|
||||
<AlertTriangle className="h-3 w-3" />
|
||||
) : isLicenseExpiringSoon ? (
|
||||
<TrendingUp className="h-3 w-3" />
|
||||
) : null}
|
||||
<span className="truncate">{formatExpirationDisplay(daysUntilExpiration)}</span>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Users section */}
|
||||
{data.total_users !== null && (
|
||||
<div>
|
||||
<div className="mb-1 flex items-center gap-1">
|
||||
<Users className="h-3 w-3" />
|
||||
<span className="font-medium">
|
||||
{data.total_users_used}/{data.total_users}
|
||||
</span>
|
||||
<span className="text-gray-500">users</span>
|
||||
</div>
|
||||
|
||||
{/* User progress bar */}
|
||||
<div className="w-full bg-gray-200 rounded-full h-1 mb-1">
|
||||
<div
|
||||
className={cn(
|
||||
"h-1 rounded-full transition-all duration-300",
|
||||
userMetrics.isOverLimit && "bg-red-500",
|
||||
userMetrics.isNearLimit && "bg-yellow-500",
|
||||
!userMetrics.isOverLimit && !userMetrics.isNearLimit && "bg-green-500",
|
||||
)}
|
||||
style={{ width: `${Math.min(userMetrics.usagePercentage, 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{(userMetrics.isOverLimit || userMetrics.isNearLimit) && (
|
||||
<div
|
||||
className={cn(
|
||||
"flex items-center gap-1 text-xs",
|
||||
userMetrics.isOverLimit && "text-red-600",
|
||||
userMetrics.isNearLimit && "text-yellow-600",
|
||||
)}
|
||||
>
|
||||
{userMetrics.isOverLimit ? (
|
||||
<AlertTriangle className="h-3 w-3" />
|
||||
) : (
|
||||
<TrendingUp className="h-3 w-3" />
|
||||
)}
|
||||
<span className="truncate">Users {userMetrics.isOverLimit ? "Over Limit" : "Near Limit"}</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Teams section */}
|
||||
{data.total_teams !== null && (
|
||||
<div>
|
||||
<div className="mb-1 flex items-center gap-1">
|
||||
<UserCheck className="h-3 w-3" />
|
||||
<span className="font-medium">
|
||||
{data.total_teams_used}/{data.total_teams}
|
||||
</span>
|
||||
<span className="text-gray-500">teams</span>
|
||||
</div>
|
||||
|
||||
{/* Team progress bar */}
|
||||
<div className="w-full bg-gray-200 rounded-full h-1 mb-1">
|
||||
<div
|
||||
className={cn(
|
||||
"h-1 rounded-full transition-all duration-300",
|
||||
teamMetrics.isOverLimit && "bg-red-500",
|
||||
teamMetrics.isNearLimit && "bg-yellow-500",
|
||||
!teamMetrics.isOverLimit && !teamMetrics.isNearLimit && "bg-green-500",
|
||||
)}
|
||||
style={{ width: `${Math.min(teamMetrics.usagePercentage, 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{(teamMetrics.isOverLimit || teamMetrics.isNearLimit) && (
|
||||
<div
|
||||
className={cn(
|
||||
"flex items-center gap-1 text-xs",
|
||||
teamMetrics.isOverLimit && "text-red-600",
|
||||
teamMetrics.isNearLimit && "text-yellow-600",
|
||||
)}
|
||||
>
|
||||
{teamMetrics.isOverLimit ? (
|
||||
<AlertTriangle className="h-3 w-3" />
|
||||
) : (
|
||||
<TrendingUp className="h-3 w-3" />
|
||||
)}
|
||||
<span className="truncate">Teams {teamMetrics.isOverLimit ? "Over Limit" : "Near Limit"}</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
// Optimized CardStyleView for 220px width
|
||||
const CardStyleView = () => {
|
||||
if (isMinimized) {
|
||||
return (
|
||||
<button
|
||||
onClick={() => setIsMinimized(false)}
|
||||
className={cn(
|
||||
"bg-white border border-gray-200 rounded-lg shadow-xs p-3 hover:shadow-md transition-all w-full",
|
||||
)}
|
||||
title="Show usage details"
|
||||
>
|
||||
<div className="flex items-center gap-2">
|
||||
<Users className="h-4 w-4 shrink-0" />
|
||||
{hasAnyIssue && <span className="shrink-0">{getStatusIcon()}</span>}
|
||||
<div className="flex items-center gap-2 text-sm font-medium truncate">
|
||||
{data && data.total_users !== null && (
|
||||
<span
|
||||
className={cn(
|
||||
"shrink-0 px-1.5 py-0.5 rounded-sm text-xs border",
|
||||
userMetrics.isOverLimit && "bg-red-50 text-red-700 border-red-200",
|
||||
userMetrics.isNearLimit && "bg-yellow-50 text-yellow-700 border-yellow-200",
|
||||
!userMetrics.isOverLimit && !userMetrics.isNearLimit && "bg-gray-50 text-gray-700 border-gray-200",
|
||||
)}
|
||||
>
|
||||
U: {data.total_users_used}/{data.total_users}
|
||||
</span>
|
||||
)}
|
||||
{data && data.total_teams !== null && (
|
||||
<span
|
||||
className={cn(
|
||||
"shrink-0 px-1.5 py-0.5 rounded-sm text-xs border",
|
||||
teamMetrics.isOverLimit && "bg-red-50 text-red-700 border-red-200",
|
||||
teamMetrics.isNearLimit && "bg-yellow-50 text-yellow-700 border-yellow-200",
|
||||
!teamMetrics.isOverLimit && !teamMetrics.isNearLimit && "bg-gray-50 text-gray-700 border-gray-200",
|
||||
)}
|
||||
>
|
||||
T: {data.total_teams_used}/{data.total_teams}
|
||||
</span>
|
||||
)}
|
||||
{licenseInfo?.expiration_date && daysUntilExpiration !== null && (
|
||||
<span
|
||||
className={cn(
|
||||
"shrink-0 px-1.5 py-0.5 rounded-sm text-xs border",
|
||||
isLicenseExpired && "bg-red-50 text-red-700 border-red-200",
|
||||
isLicenseExpiringSoon && "bg-yellow-50 text-yellow-700 border-yellow-200",
|
||||
!isLicenseExpired && !isLicenseExpiringSoon && "bg-gray-50 text-gray-700 border-gray-200",
|
||||
)}
|
||||
>
|
||||
{daysUntilExpiration < 0 ? "Exp!" : `${daysUntilExpiration}d`}
|
||||
</span>
|
||||
)}
|
||||
{!data ||
|
||||
(data.total_users === null && data.total_teams === null && !licenseInfo && (
|
||||
<span className="truncate">Usage</span>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
</button>
|
||||
);
|
||||
}
|
||||
|
||||
if (isLoading) {
|
||||
return (
|
||||
<div className="bg-white border border-gray-200 rounded-lg shadow-xs p-4 w-full">
|
||||
<div className="flex items-center justify-center gap-2 py-2">
|
||||
<Loader2 className="h-4 w-4 animate-spin" />
|
||||
<span className="text-sm text-gray-500 truncate">Loading...</span>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (error || !data) {
|
||||
return (
|
||||
<div className="bg-white border border-gray-200 rounded-lg shadow-xs p-4 group w-full">
|
||||
<div className="flex items-center justify-between gap-2">
|
||||
<div className="flex-1 min-w-0">
|
||||
<span className="text-sm text-gray-500 truncate block">{error || "No data"}</span>
|
||||
</div>
|
||||
<button
|
||||
onClick={() => setIsMinimized(true)}
|
||||
className="opacity-0 group-hover:opacity-100 p-1 hover:bg-gray-100 rounded-sm transition-all shrink-0"
|
||||
title="Minimize"
|
||||
>
|
||||
<Minus className="h-3 w-3 text-gray-400" />
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<div className={cn("bg-white border rounded-lg shadow-xs p-3 transition-all duration-200 group w-full")}>
|
||||
<div className="flex items-center justify-between gap-2 mb-3">
|
||||
<div className="flex items-center gap-2 min-w-0 flex-1">
|
||||
<Users className="h-4 w-4 shrink-0" />
|
||||
<span className="font-medium text-sm truncate">Usage</span>
|
||||
</div>
|
||||
<button
|
||||
onClick={() => setIsMinimized(true)}
|
||||
className="opacity-0 group-hover:opacity-100 p-1 hover:bg-gray-100 rounded-sm transition-all shrink-0"
|
||||
title="Minimize"
|
||||
>
|
||||
<Minus className="h-3 w-3 text-gray-400" />
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{/* Compact stats optimized for 220px */}
|
||||
<div className="space-y-3 text-sm">
|
||||
{/* License expiration section */}
|
||||
{licenseInfo?.has_license && licenseInfo.expiration_date && (
|
||||
<div
|
||||
className={cn(
|
||||
"space-y-1 border rounded-md p-2",
|
||||
isLicenseExpired && "border-red-200 bg-red-50",
|
||||
isLicenseExpiringSoon && "border-yellow-200 bg-yellow-50",
|
||||
)}
|
||||
>
|
||||
<div className="flex items-center gap-2 text-xs text-gray-600 mb-1">
|
||||
<Calendar className="h-3 w-3" />
|
||||
<span className="font-medium">License</span>
|
||||
<span
|
||||
className={cn(
|
||||
"ml-1 px-1.5 py-0.5 rounded-sm border",
|
||||
isLicenseExpired && "bg-red-50 text-red-700 border-red-200",
|
||||
isLicenseExpiringSoon && "bg-yellow-50 text-yellow-700 border-yellow-200",
|
||||
!isLicenseExpired && !isLicenseExpiringSoon && "bg-gray-50 text-gray-600 border-gray-200",
|
||||
)}
|
||||
>
|
||||
{isLicenseExpired ? "Expired" : isLicenseExpiringSoon ? "Expiring soon" : "OK"}
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex justify-between items-center">
|
||||
<span className="text-gray-600 text-xs">Status:</span>
|
||||
<span
|
||||
className={cn(
|
||||
"font-medium text-right",
|
||||
isLicenseExpired && "text-red-600",
|
||||
isLicenseExpiringSoon && "text-yellow-600",
|
||||
)}
|
||||
>
|
||||
{formatExpirationDisplay(daysUntilExpiration)}
|
||||
</span>
|
||||
</div>
|
||||
{licenseInfo.license_type && (
|
||||
<div className="flex justify-between items-center">
|
||||
<span className="text-gray-600 text-xs">Type:</span>
|
||||
<span className="font-medium text-right capitalize">{licenseInfo.license_type}</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Users section */}
|
||||
{data.total_users !== null && (
|
||||
<div
|
||||
className={cn(
|
||||
"space-y-1 border rounded-md p-2",
|
||||
userMetrics.isOverLimit && "border-red-200 bg-red-50",
|
||||
userMetrics.isNearLimit && "border-yellow-200 bg-yellow-50",
|
||||
)}
|
||||
>
|
||||
<div className="flex items-center gap-2 text-xs text-gray-600 mb-1">
|
||||
<Users className="h-3 w-3" />
|
||||
<span className="font-medium">Users</span>
|
||||
<span
|
||||
className={cn(
|
||||
"ml-1 px-1.5 py-0.5 rounded-sm border",
|
||||
userMetrics.isOverLimit && "bg-red-50 text-red-700 border-red-200",
|
||||
userMetrics.isNearLimit && "bg-yellow-50 text-yellow-700 border-yellow-200",
|
||||
!userMetrics.isOverLimit && !userMetrics.isNearLimit && "bg-gray-50 text-gray-600 border-gray-200",
|
||||
)}
|
||||
>
|
||||
{userMetrics.isOverLimit ? "Over limit" : userMetrics.isNearLimit ? "Near limit" : "OK"}
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex justify-between items-center">
|
||||
<span className="text-gray-600 text-xs">Used:</span>
|
||||
<span className="font-medium text-right">
|
||||
{data.total_users_used}/{data.total_users}
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex justify-between items-center">
|
||||
<span className="text-gray-600 text-xs">Remaining:</span>
|
||||
<span
|
||||
className={cn(
|
||||
"font-medium text-right",
|
||||
userMetrics.isOverLimit && "text-red-600",
|
||||
userMetrics.isNearLimit && "text-yellow-600",
|
||||
)}
|
||||
>
|
||||
{data.total_users_remaining}
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex justify-between items-center">
|
||||
<span className="text-gray-600 text-xs">Usage:</span>
|
||||
<span className="font-medium text-right">{Math.round(userMetrics.usagePercentage)}%</span>
|
||||
</div>
|
||||
|
||||
{/* User progress bar */}
|
||||
<div className="w-full bg-gray-200 rounded-full h-2">
|
||||
<div
|
||||
className={cn(
|
||||
"h-2 rounded-full transition-all duration-300",
|
||||
userMetrics.isOverLimit && "bg-red-500",
|
||||
userMetrics.isNearLimit && "bg-yellow-500",
|
||||
!userMetrics.isOverLimit && !userMetrics.isNearLimit && "bg-green-500",
|
||||
)}
|
||||
style={{ width: `${Math.min(userMetrics.usagePercentage, 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Teams section */}
|
||||
{data.total_teams !== null && (
|
||||
<div
|
||||
className={cn(
|
||||
"space-y-1 border rounded-md p-2",
|
||||
teamMetrics.isOverLimit && "border-red-200 bg-red-50",
|
||||
teamMetrics.isNearLimit && "border-yellow-200 bg-yellow-50",
|
||||
)}
|
||||
>
|
||||
<div className="flex items-center gap-2 text-xs text-gray-600 mb-1">
|
||||
<UserCheck className="h-3 w-3" />
|
||||
<span className="font-medium">Teams</span>
|
||||
<span
|
||||
className={cn(
|
||||
"ml-1 px-1.5 py-0.5 rounded-sm border",
|
||||
teamMetrics.isOverLimit && "bg-red-50 text-red-700 border-red-200",
|
||||
teamMetrics.isNearLimit && "bg-yellow-50 text-yellow-700 border-yellow-200",
|
||||
!teamMetrics.isOverLimit && !teamMetrics.isNearLimit && "bg-gray-50 text-gray-600 border-gray-200",
|
||||
)}
|
||||
>
|
||||
{teamMetrics.isOverLimit ? "Over limit" : teamMetrics.isNearLimit ? "Near limit" : "OK"}
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex justify-between items-center">
|
||||
<span className="text-gray-600 text-xs">Used:</span>
|
||||
<span className="font-medium text-right">
|
||||
{data.total_teams_used}/{data.total_teams}
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex justify-between items-center">
|
||||
<span className="text-gray-600 text-xs">Remaining:</span>
|
||||
<span
|
||||
className={cn(
|
||||
"font-medium text-right",
|
||||
teamMetrics.isOverLimit && "text-red-600",
|
||||
teamMetrics.isNearLimit && "text-yellow-600",
|
||||
)}
|
||||
>
|
||||
{data.total_teams_remaining}
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex justify-between items-center">
|
||||
<span className="text-gray-600 text-xs">Usage:</span>
|
||||
<span className="font-medium text-right">{Math.round(teamMetrics.usagePercentage)}%</span>
|
||||
</div>
|
||||
|
||||
{/* Team progress bar */}
|
||||
<div className="w-full bg-gray-200 rounded-full h-2">
|
||||
<div
|
||||
className={cn(
|
||||
"h-2 rounded-full transition-all duration-300",
|
||||
teamMetrics.isOverLimit && "bg-red-500",
|
||||
teamMetrics.isNearLimit && "bg-yellow-500",
|
||||
!teamMetrics.isOverLimit && !teamMetrics.isNearLimit && "bg-green-500",
|
||||
)}
|
||||
style={{ width: `${Math.min(teamMetrics.usagePercentage, 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
// Don't render anything if disabled, no access token, or if both total_users and total_teams are null
|
||||
if (disableUsageIndicator || !accessToken || (data?.total_users === null && data?.total_teams === null)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
// Fixed positioning with proper spacing from edges
|
||||
return (
|
||||
<div className="fixed bottom-4 left-4 z-50" style={{ width: `${Math.min(width, 220)}px` }}>
|
||||
<CardStyleView />
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
@ -0,0 +1,43 @@
|
|||
import { describe, it, expect, beforeEach, afterEach, vi } from "vitest";
|
||||
|
||||
// Regression for the chat sidebar / first-message navigation under SERVER_ROOT_PATH.
|
||||
// getChatRoutes() must read the server root path at call time. The previous
|
||||
// module-level `CHAT_ROUTES` captured it once at import, before the UI-config
|
||||
// bootstrap runs setServerRootPath, so every chat route was permanently
|
||||
// unprefixed and router.push() navigated to a 404 (which, mid-stream, also
|
||||
// aborted the first message). These tests deliberately apply the root path
|
||||
// AFTER importing the module so a frozen-at-import implementation fails.
|
||||
describe("getChatRoutes under server_root_path", () => {
|
||||
beforeEach(() => {
|
||||
vi.resetModules();
|
||||
vi.stubEnv("NODE_ENV", "test");
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
vi.unstubAllEnvs();
|
||||
});
|
||||
|
||||
it("reflects a server root path applied after the module is loaded", async () => {
|
||||
const { getChatRoutes } = await import("./ChatShell");
|
||||
const { setServerRootPath } = await import("@/lib/serverRootPath");
|
||||
|
||||
setServerRootPath("/gw");
|
||||
|
||||
const routes = getChatRoutes();
|
||||
expect(routes.chats).toBe("/gw/ui/chat");
|
||||
expect(routes.integrations).toBe("/gw/ui/chat/integrations");
|
||||
expect(routes.credentials).toBe("/gw/ui/chat/credentials");
|
||||
expect(routes.apiKeys).toBe("/gw/ui/chat/api-keys");
|
||||
expect(routes.usage).toBe("/gw/ui/chat/usage");
|
||||
});
|
||||
|
||||
it("builds /ui-rooted paths when no server root path is set", async () => {
|
||||
const { getChatRoutes } = await import("./ChatShell");
|
||||
const { setServerRootPath } = await import("@/lib/serverRootPath");
|
||||
|
||||
setServerRootPath("/");
|
||||
|
||||
expect(getChatRoutes().chats).toBe("/ui/chat");
|
||||
expect(getChatRoutes().integrations).toBe("/ui/chat/integrations");
|
||||
});
|
||||
});
|
||||
|
|
@ -9,14 +9,16 @@ import { migratedHref } from "@/utils/migratedPages";
|
|||
import { useChatShell } from "@/contexts/ChatShellContext";
|
||||
import ConversationList from "./ConversationList";
|
||||
|
||||
const CHAT_BASE = migratedHref("chat");
|
||||
export const CHAT_ROUTES = {
|
||||
chats: CHAT_BASE,
|
||||
integrations: `${CHAT_BASE}/integrations`,
|
||||
credentials: `${CHAT_BASE}/credentials`,
|
||||
apiKeys: `${CHAT_BASE}/api-keys`,
|
||||
usage: `${CHAT_BASE}/usage`,
|
||||
};
|
||||
export function getChatRoutes() {
|
||||
const base = migratedHref("chat");
|
||||
return {
|
||||
chats: base,
|
||||
integrations: `${base}/integrations`,
|
||||
credentials: `${base}/credentials`,
|
||||
apiKeys: `${base}/api-keys`,
|
||||
usage: `${base}/usage`,
|
||||
};
|
||||
}
|
||||
|
||||
function stripTrailingSlash(path: string): string {
|
||||
return path.length > 1 ? path.replace(/\/+$/, "") : path;
|
||||
|
|
@ -54,7 +56,8 @@ const ChatShell: React.FC<ChatShellProps> = ({ children }) => {
|
|||
const pathname = stripTrailingSlash(usePathname() ?? "");
|
||||
const { conversations, activeConversationId, deleteConversation, renameConversation } = useChatShell();
|
||||
|
||||
const isChatsRoute = pathname === CHAT_ROUTES.chats;
|
||||
const routes = getChatRoutes();
|
||||
const isChatsRoute = pathname === routes.chats;
|
||||
|
||||
return (
|
||||
<div className="flex h-full w-full flex-col bg-background overflow-hidden">
|
||||
|
|
@ -73,7 +76,7 @@ const ChatShell: React.FC<ChatShellProps> = ({ children }) => {
|
|||
<div className="flex flex-1 min-h-0 overflow-hidden">
|
||||
<div className="shrink-0 bg-sidebar border-sidebar-border border-r flex flex-col overflow-hidden w-[260px]">
|
||||
<div className="px-2 pt-3 pb-1 shrink-0">
|
||||
<Button onClick={() => router.push(CHAT_ROUTES.chats)} className="w-full justify-start gap-2.5">
|
||||
<Button onClick={() => router.push(routes.chats)} className="w-full justify-start gap-2.5">
|
||||
<Plus className="h-4 w-4" />
|
||||
New Chat
|
||||
</Button>
|
||||
|
|
@ -85,32 +88,32 @@ const ChatShell: React.FC<ChatShellProps> = ({ children }) => {
|
|||
<NavItem
|
||||
icon={<MessageSquare className="h-4 w-4" />}
|
||||
label="Chats"
|
||||
onClick={() => router.push(CHAT_ROUTES.chats)}
|
||||
onClick={() => router.push(routes.chats)}
|
||||
active={isChatsRoute}
|
||||
/>
|
||||
<NavItem
|
||||
icon={<LayoutGrid className="h-4 w-4" />}
|
||||
label="Integrations"
|
||||
onClick={() => router.push(CHAT_ROUTES.integrations)}
|
||||
active={pathname === CHAT_ROUTES.integrations}
|
||||
onClick={() => router.push(routes.integrations)}
|
||||
active={pathname === routes.integrations}
|
||||
/>
|
||||
<NavItem
|
||||
icon={<KeyRound className="h-4 w-4" />}
|
||||
label="Credentials"
|
||||
onClick={() => router.push(CHAT_ROUTES.credentials)}
|
||||
active={pathname === CHAT_ROUTES.credentials}
|
||||
onClick={() => router.push(routes.credentials)}
|
||||
active={pathname === routes.credentials}
|
||||
/>
|
||||
<NavItem
|
||||
icon={<Lock className="h-4 w-4" />}
|
||||
label="API Keys"
|
||||
onClick={() => router.push(CHAT_ROUTES.apiKeys)}
|
||||
active={pathname === CHAT_ROUTES.apiKeys}
|
||||
onClick={() => router.push(routes.apiKeys)}
|
||||
active={pathname === routes.apiKeys}
|
||||
/>
|
||||
<NavItem
|
||||
icon={<BarChart3 className="h-4 w-4" />}
|
||||
label="Usage"
|
||||
onClick={() => router.push(CHAT_ROUTES.usage)}
|
||||
active={pathname === CHAT_ROUTES.usage}
|
||||
onClick={() => router.push(routes.usage)}
|
||||
active={pathname === routes.usage}
|
||||
/>
|
||||
</div>
|
||||
|
||||
|
|
@ -120,10 +123,10 @@ const ChatShell: React.FC<ChatShellProps> = ({ children }) => {
|
|||
<ConversationList
|
||||
conversations={conversations}
|
||||
activeConversationId={activeConversationId}
|
||||
onSelect={(id) => router.push(`${CHAT_ROUTES.chats}?id=${id}`)}
|
||||
onSelect={(id) => router.push(`${routes.chats}?id=${id}`)}
|
||||
onDelete={(id) => {
|
||||
deleteConversation(id);
|
||||
if (id === activeConversationId) router.push(CHAT_ROUTES.chats);
|
||||
if (id === activeConversationId) router.push(routes.chats);
|
||||
}}
|
||||
onRename={renameConversation}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -52,6 +52,7 @@ export function useChatHistory(
|
|||
): {
|
||||
conversations: Conversation[];
|
||||
activeConversation: Conversation | null;
|
||||
currentActiveId: string | null;
|
||||
storageUnavailable: boolean;
|
||||
staleId: boolean;
|
||||
createConversation: (model: string) => string;
|
||||
|
|
@ -208,6 +209,7 @@ export function useChatHistory(
|
|||
return {
|
||||
conversations,
|
||||
activeConversation,
|
||||
currentActiveId,
|
||||
storageUnavailable,
|
||||
staleId,
|
||||
createConversation,
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ import {
|
|||
rolesAllowedToViewWriteScopedPages,
|
||||
rolesWithWriteAccess,
|
||||
} from "../utils/roles";
|
||||
import BetaBadge from "./BetaBadge";
|
||||
import NewBadge from "./common_components/NewBadge";
|
||||
import type { Organization } from "./networking";
|
||||
import SidebarAccountMenu from "./SidebarAccountMenu/SidebarAccountMenu";
|
||||
|
|
@ -211,7 +212,7 @@ const menuGroups: MenuGroup[] = [
|
|||
page: "projects",
|
||||
label: (
|
||||
<span className="flex items-center gap-2">
|
||||
Projects <NewBadge />
|
||||
Projects <BetaBadge />
|
||||
</span>
|
||||
),
|
||||
icon: <Folder {...ICON} />,
|
||||
|
|
|
|||
|
|
@ -252,7 +252,7 @@ describe("Navbar", () => {
|
|||
expect(screen.queryByRole("button", { name: /^notifications$/i })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should handle hide new features toggle", async () => {
|
||||
it("should handle hide new feature indicators toggle", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
||||
// Initially disabled
|
||||
|
|
|
|||
|
|
@ -179,7 +179,7 @@ export interface PromptSpec {
|
|||
export interface PromptTemplateBase {
|
||||
litellm_prompt_id: string;
|
||||
content: string;
|
||||
metadata?: Record<string, any> | null;
|
||||
metadata?: Record<string, unknown> | null;
|
||||
}
|
||||
|
||||
interface PromptInfoResponse {
|
||||
|
|
@ -6227,6 +6227,7 @@ export const applyGuardrail = async (
|
|||
text: string,
|
||||
language?: string | null,
|
||||
entities?: string[] | null,
|
||||
metadata?: Record<string, unknown> | null,
|
||||
) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/guardrails/apply_guardrail` : `/guardrails/apply_guardrail`;
|
||||
|
|
@ -6244,6 +6245,10 @@ export const applyGuardrail = async (
|
|||
requestBody.entities = entities;
|
||||
}
|
||||
|
||||
if (metadata != null) {
|
||||
requestBody.metadata = metadata;
|
||||
}
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
|
|
|
|||
|
|
@ -57,12 +57,13 @@ export function ChatShellProvider({
|
|||
children,
|
||||
}: ChatShellProviderProps) {
|
||||
const searchParams = useSearchParams();
|
||||
const activeConversationId = searchParams.get("id");
|
||||
const urlConversationId = searchParams.get("id");
|
||||
const [selectedMCPServers, setSelectedMCPServers] = useState<string[]>([]);
|
||||
|
||||
const {
|
||||
conversations,
|
||||
activeConversation,
|
||||
currentActiveId,
|
||||
storageUnavailable,
|
||||
staleId,
|
||||
createConversation,
|
||||
|
|
@ -71,7 +72,7 @@ export function ChatShellProvider({
|
|||
truncateFromMessage,
|
||||
deleteConversation,
|
||||
renameConversation,
|
||||
} = useChatHistory(activeConversationId, userId);
|
||||
} = useChatHistory(urlConversationId, userId);
|
||||
|
||||
return (
|
||||
<ChatShellContext.Provider
|
||||
|
|
@ -85,7 +86,7 @@ export function ChatShellProvider({
|
|||
setSelectedMCPServers,
|
||||
conversations,
|
||||
activeConversation,
|
||||
activeConversationId,
|
||||
activeConversationId: currentActiveId,
|
||||
storageUnavailable,
|
||||
staleId,
|
||||
createConversation,
|
||||
|
|
|
|||
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -6544,7 +6544,7 @@ export interface paths {
|
|||
* Parameters:
|
||||
* - duration: Optional[str] - Specify the length of time the token is valid for. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
|
||||
* - key_alias: Optional[str] - User defined key alias
|
||||
* - key: Optional[str] - User defined key value. If not set, a 16-digit unique sk-key is created for you.
|
||||
* - key: Optional[str] - User defined key value. Must start with 'sk-' and be at least 16 characters long. If not set, a 16-digit unique sk-key is created for you.
|
||||
* - team_id: Optional[str] - The team id of the key
|
||||
* - user_id: Optional[str] - The user id of the key
|
||||
* - agent_id: Optional[str] - The agent id associated with the key.
|
||||
|
|
@ -6765,7 +6765,7 @@ export interface paths {
|
|||
* - data: Optional[RegenerateKeyRequest] - Request body containing optional parameters to update
|
||||
* - key: Optional[str] - The key to regenerate.
|
||||
* - new_master_key: Optional[str] - The new master key to use, if key is the master key.
|
||||
* - new_key: Optional[str] - The new key to use, if key is not the master key. If both set, new_master_key will be used.
|
||||
* - new_key: Optional[str] - The new key to use, if key is not the master key. Must start with 'sk-' and be at least 16 characters long. If both set, new_master_key will be used.
|
||||
* - key_alias: Optional[str] - User-friendly key alias
|
||||
* - user_id: Optional[str] - User ID associated with key
|
||||
* - team_id: Optional[str] - Team ID associated with key
|
||||
|
|
@ -6834,7 +6834,7 @@ export interface paths {
|
|||
* Parameters:
|
||||
* - duration: Optional[str] - Specify the length of time the token is valid for. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
|
||||
* - key_alias: Optional[str] - User defined key alias
|
||||
* - key: Optional[str] - User defined key value. If not set, a 16-digit unique sk-key is created for you.
|
||||
* - key: Optional[str] - User defined key value. Must start with 'sk-' and be at least 16 characters long. If not set, a 16-digit unique sk-key is created for you.
|
||||
* - team_id: Optional[str] - The team id of the key
|
||||
* - user_id: Optional[str] - [NON-FUNCTIONAL] THIS WILL BE IGNORED. The user id of the key
|
||||
* - budget_id: Optional[str] - The budget id associated with the key. Created by calling `/budget/new`.
|
||||
|
|
@ -7024,7 +7024,7 @@ export interface paths {
|
|||
* - data: Optional[RegenerateKeyRequest] - Request body containing optional parameters to update
|
||||
* - key: Optional[str] - The key to regenerate.
|
||||
* - new_master_key: Optional[str] - The new master key to use, if key is the master key.
|
||||
* - new_key: Optional[str] - The new key to use, if key is not the master key. If both set, new_master_key will be used.
|
||||
* - new_key: Optional[str] - The new key to use, if key is not the master key. Must start with 'sk-' and be at least 16 characters long. If both set, new_master_key will be used.
|
||||
* - key_alias: Optional[str] - User-friendly key alias
|
||||
* - user_id: Optional[str] - User ID associated with key
|
||||
* - team_id: Optional[str] - Team ID associated with key
|
||||
|
|
@ -20659,6 +20659,10 @@ export interface components {
|
|||
messages?: {
|
||||
[key: string]: unknown;
|
||||
}[] | null;
|
||||
/** Metadata */
|
||||
metadata?: {
|
||||
[key: string]: unknown;
|
||||
} | null;
|
||||
/** Text */
|
||||
text: string;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,5 +1,11 @@
|
|||
import { describe, it, expect } from "vitest";
|
||||
import { type LicenseExpiryTier, formatExpiryDate, getDaysUntilExpiration, getLicenseExpiryTier } from "./licenseUtils";
|
||||
import {
|
||||
type LicenseExpiryTier,
|
||||
formatExpirationStatus,
|
||||
formatExpiryDate,
|
||||
getDaysUntilExpiration,
|
||||
getLicenseExpiryTier,
|
||||
} from "./licenseUtils";
|
||||
|
||||
const NOW = new Date("2026-07-08T00:00:00Z");
|
||||
|
||||
|
|
@ -59,3 +65,25 @@ describe("formatExpiryDate", () => {
|
|||
expect(formatExpiryDate("bogus")).toBe("bogus");
|
||||
});
|
||||
});
|
||||
|
||||
describe("formatExpirationStatus", () => {
|
||||
it("shows the exact date for a future expiration", () => {
|
||||
expect(formatExpirationStatus("2026-08-07", NOW)).toBe("Expires Aug 7, 2026");
|
||||
});
|
||||
|
||||
it("still reads as upcoming on the expiration day itself", () => {
|
||||
expect(formatExpirationStatus("2026-07-08", NOW)).toBe("Expires Jul 8, 2026");
|
||||
});
|
||||
|
||||
it("shows the exact date for a past expiration", () => {
|
||||
expect(formatExpirationStatus("2026-07-07", NOW)).toBe("Expired Jul 7, 2026");
|
||||
});
|
||||
|
||||
it("returns No expiration for a null date", () => {
|
||||
expect(formatExpirationStatus(null, NOW)).toBe("No expiration");
|
||||
});
|
||||
|
||||
it("returns No expiration for an unparseable date", () => {
|
||||
expect(formatExpirationStatus("not-a-date", NOW)).toBe("No expiration");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -49,3 +49,11 @@ export const formatExpiryDate = (expirationDate: string): string => {
|
|||
}
|
||||
return date.toLocaleDateString("en-US", EXPIRY_DATE_FORMAT);
|
||||
};
|
||||
|
||||
export const formatExpirationStatus = (expirationDate: string | null, now: Date = new Date()): string => {
|
||||
const days = getDaysUntilExpiration(expirationDate, now);
|
||||
if (expirationDate === null || days === null) {
|
||||
return "No expiration";
|
||||
}
|
||||
return days < 0 ? `Expired ${formatExpiryDate(expirationDate)}` : `Expires ${formatExpiryDate(expirationDate)}`;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue