diff --git a/litellm/integrations/datadog/datadog.py b/litellm/integrations/datadog/datadog.py index bd62d1e303a..20239d831cc 100644 --- a/litellm/integrations/datadog/datadog.py +++ b/litellm/integrations/datadog/datadog.py @@ -19,7 +19,7 @@ import os import time import traceback from datetime import datetime as datetimeObj -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Sequence, Union import httpx from httpx import Response @@ -50,6 +50,7 @@ from litellm.types.integrations.base_health_check import IntegrationHealthCheckS from litellm.types.integrations.datadog import ( DD_ERRORS, DD_MAX_BATCH_SIZE, + DD_MAX_PAYLOAD_SIZE_BYTES, DataDogStatus, DatadogInitParams, DatadogPayload, @@ -384,8 +385,10 @@ class DataDogLogger( async def _send_with_413_split(self, batch: List) -> List: """ - Send a batch, halving any sub-batch that 413s (payload too large) and retrying the - halves, since Datadog enforces a 5MB uncompressed limit per request. + Send a batch, halving any sub-batch that exceeds Datadog's intake limits before + sending, and halving again on a 413 (payload too large) response, since Datadog + enforces a 5MB uncompressed limit per request. The proactive split avoids paying + a serialize + gzip + round trip for a payload the intake is guaranteed to reject. A 413 surfaces as a raised MaskedHTTPStatusError (httpx raise_for_status), not a returned response, so both paths are handled. A lone event that still 413s is @@ -398,6 +401,11 @@ class DataDogLogger( chunk = pending.pop() if not chunk: continue + if len(chunk) > 1 and self._exceeds_intake_limits(chunk): + mid = len(chunk) // 2 + pending.append(chunk[mid:]) + pending.append(chunk[:mid]) + continue try: response = await self.async_send_compressed_data(chunk) except Exception as e: @@ -436,6 +444,21 @@ class DataDogLogger( def _undelivered(chunk: List, pending: List[List]) -> List: return chunk + [event for remaining in reversed(pending) for event in remaining] + @staticmethod + def _exceeds_intake_limits(chunk: Sequence[DatadogPayload]) -> bool: + """ + True when a chunk would breach Datadog's log intake limits: more than + DD_MAX_BATCH_SIZE events per payload, or a serialized size above + DD_MAX_PAYLOAD_SIZE_BYTES (held under Datadog's 5MB uncompressed cap so + the batch is split before the intake rejects it with a 413). + """ + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + + if len(chunk) > DD_MAX_BATCH_SIZE: + return True + payload_size_bytes = len(safe_dumps(chunk).encode("utf-8")) + return payload_size_bytes > DD_MAX_PAYLOAD_SIZE_BYTES + async def flush_queue(self): if self.flush_lock is None: return diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 60fc021a6a8..f575372fc3d 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -1618,6 +1618,14 @@ class PrometheusLogger(CustomLogger): user_id: Optional[str] = None, user_api_key_org_id: Optional[str] = None, ): + if ( + isinstance(self.litellm_remaining_team_budget_metric, NoOpMetric) + and isinstance(self.litellm_remaining_api_key_budget_metric, NoOpMetric) + and isinstance(self.litellm_remaining_user_budget_metric, NoOpMetric) + and isinstance(self.litellm_remaining_org_budget_metric, NoOpMetric) + ): + return + _metadata = litellm_params.get("metadata") or {} _team_spend = _metadata.get("user_api_key_team_spend", None) _team_max_budget = _metadata.get("user_api_key_team_max_budget", None) @@ -3332,6 +3340,9 @@ class PrometheusLogger(CustomLogger): - looks up team info from db if not available in metadata - Set team budget metrics """ + if isinstance(self.litellm_remaining_team_budget_metric, NoOpMetric): + return + if user_api_team: team_object = await self._assemble_team_object( team_id=user_api_team, @@ -3453,6 +3464,9 @@ class PrometheusLogger(CustomLogger): - Fetches org info via cache (get_org_object) - Sets org budget metrics """ + if isinstance(self.litellm_remaining_org_budget_metric, NoOpMetric): + return + if not org_id: return @@ -3582,6 +3596,9 @@ class PrometheusLogger(CustomLogger): key_max_budget: Optional[float], key_spend: Optional[float], ): + if isinstance(self.litellm_remaining_api_key_budget_metric, NoOpMetric): + return + if user_api_key: user_api_key_dict = await self._assemble_key_object( user_api_key=user_api_key, @@ -3642,6 +3659,9 @@ class PrometheusLogger(CustomLogger): - looks up user info from db if not available in metadata - Set user budget metrics """ + if isinstance(self.litellm_remaining_user_budget_metric, NoOpMetric): + return + if user_id: user_object = await self._assemble_user_object( user_id=user_id, diff --git a/litellm/litellm_core_utils/fallback_generalizations.py b/litellm/litellm_core_utils/fallback_generalizations.py index abc171f900a..8b5a797cbd7 100644 --- a/litellm/litellm_core_utils/fallback_generalizations.py +++ b/litellm/litellm_core_utils/fallback_generalizations.py @@ -8,9 +8,12 @@ The metadata is a partial cost-map entry: ``litellm_provider`` drives provider routing, and the remaining fields (``mode``, ``supports_*``, context window, pricing, ...) drive ``get_model_info`` / ``supports_*``. -Precedence: rules are evaluated in file order and the first match wins. They are -consulted only after exact and case-insensitive lookups miss, so an exact entry -always takes precedence over a rule. +Precedence: rules are evaluated in file order and the first match wins. Callers +with extra constraints (model-info resolution checks the provider) use +``match_all_fallback_generalizations`` to skip inapplicable earlier rules instead +of discarding the model name. Rules are consulted only after exact and +case-insensitive lookups miss, so an exact entry always takes precedence over a +rule. Patterns are matched case-insensitively with ``re.search`` and are not implicitly anchored: a rule must include ``^`` and ``$`` (as the shipped rules do) to bind to @@ -105,15 +108,15 @@ class _FallbackGeneralizations: ) return compiled - def match(self, model: str) -> Optional[dict]: + def matches(self, model: str) -> list[dict]: if not model: - return None + return [] if self._compiled is None: self._compiled = self._compile() - for pattern, model_info in self._compiled: - if pattern.search(model) is not None: - return dict(model_info) - return None + return [dict(model_info) for pattern, model_info in self._compiled if pattern.search(model) is not None] + + def match(self, model: str) -> Optional[dict]: + return next(iter(self.matches(model)), None) _registry = _FallbackGeneralizations() @@ -139,3 +142,12 @@ def match_fallback_generalization(model: str) -> Optional[dict]: O(number of rules). Only call this once exact lookups have missed. """ return _registry.match(model) + + +def match_all_fallback_generalizations(model: str) -> list[dict]: + """Return the ``model_info`` of every rule whose regex matches ``model``, in rule order. + + Lets a caller with extra constraints (e.g. a provider match) skip an + inapplicable earlier rule instead of discarding the whole candidate. + """ + return _registry.matches(model) diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index db540e5441d..964186fd76a 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -289,6 +289,13 @@ class AnthropicModelInfo(BaseLLMModelInfo): status_code=400, ) + @staticmethod + def _strip_version_suffix(model: str) -> str: + at = model.rfind("@") + if at > 0: + return model[:at] + return model + @staticmethod def _model_map_lookup_candidates(model: str) -> List[str]: """Model-map keys to try for ``model``: the id itself, the same id with a @@ -324,6 +331,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): _DATED_RELEASE_SUFFIX_RE.sub("", cand), _DOTTED_VERSION_RE.sub(r"\1-\2", cand), _strip_bedrock_id_suffixes(cand), + AnthropicModelInfo._strip_version_suffix(cand), ) ) return list(dict.fromkeys((*primary, *normalized))) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index e78802a1587..b74bbda0b66 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -3,6 +3,7 @@ from typing import Any, AsyncIterator, Dict, List, Optional, Tuple import httpx from litellm.constants import ( + ANTHROPIC_MIN_THINKING_BUDGET_TOKENS, DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET, @@ -32,6 +33,12 @@ from ...common_utils import ( DEFAULT_ANTHROPIC_API_VERSION = "2023-06-01" +DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING = ( + "Dropping adaptive `thinking`/`output_config.effort` for model=%s: the model " + "does not support extended thinking, or max_tokens is too small to fit the " + "minimum thinking budget." +) + class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): def get_supported_anthropic_messages_params(self, model: str) -> list: @@ -253,6 +260,111 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): existing_output_config.setdefault("effort", effort) optional_params["output_config"] = existing_output_config + @staticmethod + def _translate_adaptive_effort_for_non_adaptive_model( + model: str, optional_params: Dict, max_tokens: Optional[int] + ) -> None: + """Translate the 4.6+ adaptive-thinking interface (``thinking.type=adaptive`` + and/or ``output_config.effort``) down to what an older Anthropic model + supports. Clients like Claude Code send this interface unconditionally, so + without translation it reaches a pre-4.6 model and Anthropic rejects it with + "This model does not support the effort parameter". + + The reshape is silent, matching how the messages path already strips + unsupported ``output_config`` for older models (bedrock invoke, issue + #22797): the goal is to keep the request working, not to fail it. + + ``thinking.type=adaptive`` and ``output_config.effort`` are independent + capabilities. Adaptive thinking needs ``supports_adaptive_thinking`` (4.6+); + ``output_config.effort`` needs ``supports_output_config``, which some + non-adaptive models (e.g. Claude Opus 4.5) advertise on its own. So the two + are handled separately: + + - Adaptive-thinking models (4.6+): both are native, left untouched. + - ``supports_output_config`` but non-adaptive (Opus 4.5): keep + ``output_config.effort`` (native), only drop the unsupported adaptive + ``thinking`` block. When adaptive thinking is being dropped and the + effort level itself isn't supported by the model (e.g. ``xhigh``/``max`` + on Opus 4.5, which only accepts low/medium/high, while ``xhigh`` is + Claude Code's default), fall through to the legacy translation below + instead of forwarding a level Anthropic would reject. Effort-only + requests are always left untouched: provider subclasses own their level + normalization (bedrock clamps ``xhigh`` to the model's ceiling after + this base transform runs). + - Thinking-capable but neither (``supports_reasoning``, e.g. Haiku/Sonnet + 4.5): map effort to legacy ``thinking={type: enabled, budget_tokens}`` via + ``AnthropicConfig._map_reasoning_effort``, capped below ``max_tokens`` + (Anthropic requires ``max_tokens > budget_tokens``) and dropped when + ``max_tokens`` can't fit even the minimum budget. + - No reasoning support: ``thinking`` is dropped. + + For the last two, only the consumed ``effort`` key is removed from + ``output_config``; any residual (e.g. ``format``) is left for provider + subclasses to handle. + """ + from litellm.exceptions import BadRequestError as _BadRequestError + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + if AnthropicConfig._is_adaptive_thinking_model(model): + return + + output_config = optional_params.get("output_config") + thinking = optional_params.get("thinking") + effort = output_config.get("effort") if isinstance(output_config, dict) else None + adaptive_thinking = isinstance(thinking, dict) and thinking.get("type") == "adaptive" + if effort is None and not adaptive_thinking: + return + + if AnthropicConfig._model_supports_effort_param(model) and ( + not adaptive_thinking or AnthropicConfig._validate_effort_for_model(model, effort) is None + ): + if adaptive_thinking: + optional_params.pop("thinking", None) + return + + supports_thinking = AnthropicModelInfo._supports_model_capability(model, "supports_reasoning") + try: + legacy_thinking = ( + AnthropicConfig._map_reasoning_effort(reasoning_effort=effort or "medium", model=model) + if supports_thinking + else None + ) + except _BadRequestError as e: + raise AnthropicError(message=str(e.message), status_code=400) + capped_thinking = ( + AnthropicMessagesConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens) + if legacy_thinking is not None + else None + ) + + if capped_thinking is not None: + optional_params["thinking"] = capped_thinking + else: + verbose_logger.warning(DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING, model) + optional_params.pop("thinking", None) + + if isinstance(output_config, dict) and "effort" in output_config: + residual = {k: v for k, v in output_config.items() if k != "effort"} + if residual: + optional_params["output_config"] = residual + else: + optional_params.pop("output_config", None) + + @staticmethod + def _cap_thinking_budget_to_max_tokens(thinking: Dict, max_tokens: Optional[int]) -> Optional[Dict]: + """Cap a legacy ``thinking.budget_tokens`` below ``max_tokens`` (Anthropic + requires ``max_tokens > budget_tokens``). Returns the (possibly capped) + thinking dict, or ``None`` when ``max_tokens`` is too small to fit even the + minimum thinking budget and thinking should be dropped.""" + budget = thinking.get("budget_tokens") + if max_tokens is None or not isinstance(budget, int): + return thinking + if max_tokens <= ANTHROPIC_MIN_THINKING_BUDGET_TOKENS: + return None + if budget < max_tokens: + return thinking + return {**thinking, "budget_tokens": max_tokens - 1} + def transform_anthropic_messages_request( self, model: str, @@ -284,6 +396,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): optional_params=anthropic_messages_optional_request_params, ) + self._translate_adaptive_effort_for_non_adaptive_model( + model=model, + optional_params=anthropic_messages_optional_request_params, + max_tokens=max_tokens, + ) + system_param = anthropic_messages_optional_request_params.get("system") if self.should_strip_billing_metadata() and system_param is not None: filtered_system = self._filter_billing_headers_from_system(system_param) diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 155fba008d4..56a47432997 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -93,31 +93,48 @@ class AmazonAnthropicClaudeMessagesConfig( return [{"type": "text", "text": value}] return [value] - def _normalize_system_role_messages_for_bedrock(self, anthropic_messages_request: dict) -> None: - """Bedrock Invoke rejects a conversation that opens with ``role: "system"`` - entries inside ``messages`` ("messages.0: use the top-level 'system' - parameter for the initial system prompt"); Anthropic Messages carries that - content in the top-level ``system`` field, so hoist the leading run of - system entries there. Mid-conversation system entries (e.g. Claude Code's - ``mid-conversation-system-2026-04-07`` reminders) are accepted by Invoke in - place and MUST stay in place: hoisting one mutates the ``system`` prefix - and invalidates the prompt cache for the entire message history. + @staticmethod + def _is_system_role_message(message: Any) -> bool: + return isinstance(message, dict) and message.get("role") == "system" + + def _normalize_system_role_messages_for_bedrock(self, anthropic_messages_request: dict, model: str) -> None: + """Bedrock Invoke validates ``role: "system"`` entries inside ``messages`` + per model. Models carrying ``supports_mid_conversation_system`` in the + cost map (the Opus 4.8 family) only reject a leading run ("messages.0: + use the top-level 'system' parameter for the initial system prompt") and + accept mid-conversation entries (e.g. Claude Code's + ``mid-conversation-system-2026-04-07`` reminders) in place, where they + MUST stay: hoisting one mutates the ``system`` prefix and invalidates the + prompt cache for the entire message history. Older Claude models (Opus + 4.7, Sonnet 4.6, Haiku 4.5, ...) reject the role in every position + ("role 'system' is not supported on this model"), so without the flag + every system entry is hoisted into the top-level ``system`` field. Billing-header system blocks are stripped from the top-level ``system`` field regardless of whether anything was hoisted.""" messages = anthropic_messages_request.get("messages") if not isinstance(messages, list): return - leading_count = next( - (i for i, m in enumerate(messages) if not (isinstance(m, dict) and m.get("role") == "system")), - len(messages), - ) - if leading_count: - anthropic_messages_request["messages"] = messages[leading_count:] + if _supports_factory( + model=model, + custom_llm_provider="bedrock", + key="supports_mid_conversation_system", + ): + leading_count = next( + (i for i, m in enumerate(messages) if not self._is_system_role_message(m)), + len(messages), + ) + hoisted = messages[:leading_count] + remaining = messages[leading_count:] + else: + hoisted = [m for m in messages if self._is_system_role_message(m)] + remaining = [m for m in messages if not self._is_system_role_message(m)] + if hoisted: + anthropic_messages_request["messages"] = remaining system_content = [ block for source in ( anthropic_messages_request.get("system"), - *(m.get("content") for m in messages[:leading_count]), + *(m.get("content") for m in hoisted), ) for block in self._as_system_content_blocks(source) ] @@ -674,7 +691,7 @@ class AmazonAnthropicClaudeMessagesConfig( litellm_params=litellm_params, headers=headers, ) - self._normalize_system_role_messages_for_bedrock(anthropic_messages_request) + self._normalize_system_role_messages_for_bedrock(anthropic_messages_request, model=model) ######################################################### ############## BEDROCK Invoke SPECIFIC TRANSFORMATION ### ######################################################### diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2a269d693ec..6ecadaf1fe6 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -1359,6 +1359,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1393,6 +1394,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1427,6 +1429,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1461,6 +1464,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1481,6 +1485,7 @@ "anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, @@ -1516,6 +1521,7 @@ "global.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, @@ -1551,6 +1557,7 @@ "us.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.875e-06, "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, @@ -1586,6 +1593,7 @@ "eu.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.875e-06, "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, @@ -1621,6 +1629,43 @@ "au.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, + "cache_read_input_token_cost": 5.5e-07, + "input_cost_per_token": 5.5e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.75e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true + }, + "jp.anthropic.claude-opus-4-8": { + "bedrock_converse_supports_strict_tools": false, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.875e-06, "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, @@ -1703,6 +1748,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1737,6 +1783,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1771,6 +1818,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1805,6 +1853,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1839,6 +1888,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1873,6 +1923,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -44998,6 +45049,17 @@ }, "fallback_generalizations": { "rules": [ + { + "name": "bedrock-anthropic-claude-mid-conversation-system", + "pattern": "anthropic\\.claude-[a-z]+-(?:4[-._](?:[89]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d)(?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)", + "description": "Bedrock Invoke ids for Claude 4.8 or higher: anthropic.claude- with minor 4.8 through 4.99, any 5.x or later major-minor, or a bare 5+ major, which also admits new families such as fable. These models accept mid-conversation role system messages in place (verified live on Opus 4.8, Sonnet 5 and Fable 5), so unmapped future Bedrock Claudes keep the cache-preserving in-place handling instead of the hoist-all default. Listed first so bare-id provider inference, which takes the first pattern hit, resolves these Bedrock ids to bedrock; model-info resolution skips provider-mismatched rules either way.", + "extends": "anthropic-claude", + "model_info": { + "litellm_provider": "bedrock", + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true + } + }, { "name": "anthropic-claude-adaptive-thinking", "pattern": "(?:opus|sonnet|haiku)[-._](?:4[-._](?:[6-9]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d{1,})[-._]\\d{1,2}(?!\\d))", diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 1aaa38cea14..6b1ca3564e3 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -2020,8 +2020,6 @@ class LiteLLMCompletionResponsesConfig: output_details_dict: dict[str, int] = {} if hasattr(completion_details, "reasoning_tokens") and completion_details.reasoning_tokens is not None: output_details_dict["reasoning_tokens"] = completion_details.reasoning_tokens - else: - output_details_dict["reasoning_tokens"] = 0 if hasattr(completion_details, "text_tokens") and completion_details.text_tokens is not None: output_details_dict["text_tokens"] = completion_details.text_tokens diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 8e3be2bc12d..92dd5a513b9 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -127,7 +127,7 @@ def mock_responses_api_response( "input_tokens": 36, "input_tokens_details": {"cached_tokens": 0}, "output_tokens": 87, - "output_tokens_details": {"reasoning_tokens": 0}, + "output_tokens_details": {}, "total_tokens": 123, }, "user": None, diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 3618331f0f5..eb78e6f9c8d 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -17,6 +17,7 @@ from litellm.constants import ( LITELLM_MAX_STREAMING_DURATION_SECONDS, STREAM_SSE_DONE_STRING, ) +from litellm.exceptions import MidStreamFallbackError from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -26,7 +27,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( ) from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig -from litellm.responses.utils import ResponsesAPIRequestUtils +from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import CallTypes from litellm.utils import async_post_call_success_deployment_hook @@ -47,6 +48,44 @@ def _log_background_task_failure(task: "asyncio.Task[Any]", *, task_name: str) - verbose_logger.error("%s failed: %s", task_name, exception) +_CLIENT_ERROR_CODES: frozenset[str] = frozenset( + ( + "invalid_request_error", + "context_length_exceeded", + "content_policy_violation", + "model_not_found", + ) +) + + +def _error_event_fields(error_obj: object) -> tuple[str, Optional[str], Optional[str]]: + if isinstance(error_obj, dict): + raw_message = error_obj.get("message") + raw_type = error_obj.get("type") + raw_code = error_obj.get("code") + elif error_obj is not None: + raw_message = getattr(error_obj, "message", None) + raw_type = getattr(error_obj, "type", None) + raw_code = getattr(error_obj, "code", None) + else: + raw_message = None + raw_type = None + raw_code = None + message = str(raw_message) if raw_message is not None else "Response API in-stream error" + error_type = raw_type if isinstance(raw_type, str) else None + code = raw_code if isinstance(raw_code, str) else None + return message, error_type, code + + +def _status_code_for_error_fields(error_type: Optional[str], error_code: Optional[str]) -> int: + fields = tuple(field for field in (error_type, error_code) if field is not None) + if any(field.startswith("rate_limit") or field == "insufficient_quota" for field in fields): + return 429 + if any(field in _CLIENT_ERROR_CODES for field in fields): + return 400 + return 500 + + class BaseResponsesAPIStreamingIterator: """ Base class for streaming iterators that process responses from the Responses API. @@ -73,6 +112,8 @@ class BaseResponsesAPIStreamingIterator: self.completed_response: Optional[Any] = None self.start_time = getattr(logging_obj, "start_time", datetime.now()) self._failure_handled = False # Track if failure handler has been called + self._yielded_first_chunk = False + self._generated_content = "" self._completed_response_cached = False self._completed_response_logged = False self._completed_response_cache_hit: Optional[bool] = None @@ -160,6 +201,10 @@ class BaseResponsesAPIStreamingIterator: # Encode container_id on streaming events so proxy/UI follow-ups route correctly _event_type = getattr(openai_responses_api_chunk, "type", None) + if _event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA: + _delta = getattr(openai_responses_api_chunk, "delta", None) + if isinstance(_delta, str): + self._generated_content += _delta _stream_model_id = ( self.litellm_metadata.get("model_info", {}).get("id") if self.litellm_metadata else None ) @@ -327,17 +372,66 @@ class BaseResponsesAPIStreamingIterator: """ response_obj = getattr(self.completed_response, "response", None) if self.completed_response else None error_info = getattr(response_obj, "error", None) if response_obj else None - error_message = "Response failed" - if isinstance(error_info, dict): - error_message = error_info.get("message", str(error_info)) + error_message, error_type, error_code = _error_event_fields(error_info) + self._record_failed_response_usage(response_obj) exception = litellm.APIError( - status_code=500, + status_code=_status_code_for_error_fields(error_type, error_code), message=error_message, llm_provider=self.custom_llm_provider or "", model=self.model or "", ) self._handle_failure(exception) + def _record_failed_response_usage(self, response_obj: Optional[Any]) -> None: + if response_obj is None or self.logging_obj is None: + return + usage_obj = getattr(response_obj, "usage", None) + if usage_obj is None: + return + try: + self.logging_obj.model_call_details["combined_usage_object"] = ( + ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage_obj) + ) + except (TypeError, ValueError) as usage_error: + verbose_logger.debug( + "could not record usage for failed responses stream: %s", + usage_error, + ) + return + self.logging_obj.model_call_details["response_cost"] = ( + self.logging_obj._response_cost_calculator(result=response_obj) or 0.0 + ) + + def _maybe_raise_for_error_event(self, result: object) -> None: + chunk_type = getattr(result, "type", None) + if chunk_type not in ("error", "response.failed"): + return + + error_obj: object = ( + getattr(getattr(result, "response", None), "error", None) + if chunk_type == "response.failed" + else getattr(result, "error", None) + ) + + error_message, error_type, error_code = _error_event_fields(error_obj) + status_code = _status_code_for_error_fields(error_type, error_code) + mapped_exception = litellm.APIError( + status_code=status_code, + message=error_message, + llm_provider=self.custom_llm_provider or "", + model=self.model or "", + ) + if 400 <= status_code < 500 and status_code != 429: + raise mapped_exception + raise MidStreamFallbackError( + message=str(mapped_exception), + model=self.model or "", + llm_provider=self.custom_llm_provider or "", + original_exception=mapped_exception, + generated_content=self._generated_content, + is_pre_first_chunk=not self._yielded_first_chunk, + ) + def _get_completed_response_object(self) -> Optional[Any]: openai_types = _get_openai_response_types() completed_response = self.completed_response @@ -611,11 +705,13 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): if self.finished: raise StopAsyncIteration elif result is not None: + self._maybe_raise_for_error_event(result) # Await hook directly instead of run_async_function # (which spawns a thread + event loop per call) result = await self._call_post_streaming_deployment_hook( chunk=result, ) + self._yielded_first_chunk = True return result # If result is None, continue the loop to get the next chunk @@ -685,11 +781,13 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): if self.finished: raise StopIteration elif result is not None: + self._maybe_raise_for_error_event(result) # Sync path: use run_async_function for the hook result = run_async_function( async_function=self._call_post_streaming_deployment_hook, chunk=result, ) + self._yielded_first_chunk = True return result # If result is None, continue the loop to get the next chunk diff --git a/litellm/router.py b/litellm/router.py index 5ffe60c2da0..6e773a06c7f 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -2344,6 +2344,8 @@ class Router: self.completed_response = None self.start_time = getattr(source_iterator, "start_time", datetime.now()) self._failure_handled = False + self._yielded_first_chunk = False + self._generated_content = "" self._completed_response_cached = False self._completed_response_logged = False self._completed_response_cache_hit = None diff --git a/litellm/types/integrations/datadog.py b/litellm/types/integrations/datadog.py index 89faac27830..e0f43519b3d 100644 --- a/litellm/types/integrations/datadog.py +++ b/litellm/types/integrations/datadog.py @@ -6,6 +6,7 @@ from typing_extensions import NotRequired, TypedDict from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams DD_MAX_BATCH_SIZE = 1000 +DD_MAX_PAYLOAD_SIZE_BYTES = 4_000_000 class DataDogStatus(str, Enum): diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 3ab5a7b736e..daac1e4506f 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1185,7 +1185,7 @@ class ResponsesAPIRequestParams(ResponsesAPIOptionalRequestParams, total=False): class OutputTokensDetails(BaseLiteLLMOpenAIResponseObject): - reasoning_tokens: int = 0 + reasoning_tokens: Optional[int] = None text_tokens: Optional[int] = None @@ -1720,7 +1720,7 @@ class ErrorEventError(BaseLiteLLMOpenAIResponseObject): type: str # e.g., 'invalid_request_error' code: str # e.g., 'context_length_exceeded' message: str - param: Optional[str] = None + param: Optional[Union[str, Dict[str, Any]]] = None class ErrorEvent(BaseLiteLLMOpenAIResponseObject): diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 8a5acc24d19..e33e2335525 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -143,6 +143,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False): supports_web_search: Optional[bool] supports_reasoning: Optional[bool] supports_adaptive_thinking: Optional[bool] + supports_mid_conversation_system: Optional[bool] supports_url_context: Optional[bool] supports_none_reasoning_effort: Optional[bool] supports_minimal_reasoning_effort: Optional[bool] diff --git a/litellm/utils.py b/litellm/utils.py index 5af7b62b332..731be992af0 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -61,7 +61,7 @@ from litellm._lazy_imports import ( ) from litellm._uuid import uuid from litellm.litellm_core_utils.fallback_generalizations import ( - match_fallback_generalization, + match_all_fallback_generalizations, ) from litellm.constants import ( DEFAULT_CHAT_COMPLETION_PARAM_VALUES, @@ -5046,9 +5046,10 @@ def _get_model_info_from_generalization( """Resolve an unmapped model via a declarative fallback-generalization rule. Tries the same name candidates as the exact lookups, in the same order, and - returns ``(matched_name, model_info)`` for the first candidate whose rule also - satisfies the provider constraint. O(number of rules); only call after the - exact lookups have missed. + returns ``(matched_name, model_info)`` for the first matching rule that also + satisfies the provider constraint; a rule scoped to another provider is + skipped in favor of later rules rather than discarding the candidate. + O(number of rules); only call after the exact lookups have missed. """ candidates = [ potential_model_names["combined_model_name"], @@ -5058,11 +5059,9 @@ def _get_model_info_from_generalization( potential_model_names["stripped_model_name"], ] for candidate in candidates: - generalized_info = match_fallback_generalization(candidate) - if generalized_info is not None and _check_provider_match( - model_info=generalized_info, custom_llm_provider=custom_llm_provider - ): - return candidate, generalized_info + for generalized_info in match_all_fallback_generalizations(candidate): + if _check_provider_match(model_info=generalized_info, custom_llm_provider=custom_llm_provider): + return candidate, generalized_info return None @@ -5472,6 +5471,7 @@ def _get_model_info_helper( supports_url_context=_model_info.get("supports_url_context", None), supports_reasoning=_model_info.get("supports_reasoning", None), supports_adaptive_thinking=_model_info.get("supports_adaptive_thinking", None), + supports_mid_conversation_system=_model_info.get("supports_mid_conversation_system", None), supports_none_reasoning_effort=_model_info.get("supports_none_reasoning_effort", None), supports_minimal_reasoning_effort=_model_info.get("supports_minimal_reasoning_effort", None), supports_low_reasoning_effort=_model_info.get("supports_low_reasoning_effort", None), diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index e79ddbe35d2..31a6146cbd6 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -1359,6 +1359,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1393,6 +1394,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1427,6 +1429,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1461,6 +1464,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1481,6 +1485,7 @@ "anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, @@ -1516,6 +1521,7 @@ "global.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, @@ -1551,6 +1557,7 @@ "us.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.875e-06, "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, @@ -1586,6 +1593,7 @@ "eu.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.875e-06, "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, @@ -1621,6 +1629,43 @@ "au.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, + "cache_read_input_token_cost": 5.5e-07, + "input_cost_per_token": 5.5e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.75e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true + }, + "jp.anthropic.claude-opus-4-8": { + "bedrock_converse_supports_strict_tools": false, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "cache_creation_input_token_cost": 6.875e-06, "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, @@ -1703,6 +1748,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1737,6 +1783,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1771,6 +1818,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1805,6 +1853,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1839,6 +1888,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -1873,6 +1923,7 @@ "search_context_size_medium": 0.01 }, "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -45231,6 +45282,17 @@ }, "fallback_generalizations": { "rules": [ + { + "name": "bedrock-anthropic-claude-mid-conversation-system", + "pattern": "anthropic\\.claude-[a-z]+-(?:4[-._](?:[89]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d)(?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)", + "description": "Bedrock Invoke ids for Claude 4.8 or higher: anthropic.claude- with minor 4.8 through 4.99, any 5.x or later major-minor, or a bare 5+ major, which also admits new families such as fable. These models accept mid-conversation role system messages in place (verified live on Opus 4.8, Sonnet 5 and Fable 5), so unmapped future Bedrock Claudes keep the cache-preserving in-place handling instead of the hoist-all default. Listed first so bare-id provider inference, which takes the first pattern hit, resolves these Bedrock ids to bedrock; model-info resolution skips provider-mismatched rules either way.", + "extends": "anthropic-claude", + "model_info": { + "litellm_provider": "bedrock", + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true + } + }, { "name": "anthropic-claude-adaptive-thinking", "pattern": "(?:opus|sonnet|haiku)[-._](?:4[-._](?:[6-9]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d{1,})[-._]\\d{1,2}(?!\\d))", diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 82f2604d492..4c4c8fe735b 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -1,9 +1,9 @@ """Shared fixtures for all live e2e suites under tests/e2e/. -Design rule: skip on environment, fail on behavior. Live tests (marked `e2e`) -skip when no proxy answers; once a request reaches the proxy, behavior is -asserted. Pure unit coverage of the harness itself carries no `e2e` marker and -runs regardless of whether a proxy is up. +Design rule: hard failures only. Live tests (marked `e2e`) fail when no proxy +answers or when credentials/env are missing; they never skip. Pure unit coverage +of the harness itself carries no `e2e` marker and runs regardless of whether a +proxy is up. Lifecycle: the `resources` fixture maps the init -> run -> teardown contract (lifecycle.E2ECase) onto pytest - setup is init(), the test body is run(), and @@ -40,7 +40,7 @@ def pytest_configure(config: pytest.Config) -> None: def _liveness_reason(label: str, base_url: str) -> str | None: - """None if `base_url` answers its liveness probe, else a skip reason.""" + """None if `base_url` answers its liveness probe, else a failure reason.""" try: resp = requests.get(f"{base_url}/health/liveliness", timeout=5) except requests.RequestException as exc: @@ -51,10 +51,10 @@ def _liveness_reason(label: str, base_url: str) -> str | None: @functools.lru_cache(maxsize=1) -def _proxy_skip_reason() -> str | None: - """Probe the proxy once per session. None if it answers, else a skip reason. In - a split deployment the management/admin control plane is a separate service, so - require it too (when it differs) - else its tests would fail rather than skip.""" +def _proxy_fail_reason() -> str | None: + """Probe the proxy once per session. None if it answers, else a failure reason. + In a split deployment the management/admin control plane is a separate service, + so require it too when it differs.""" reason = _liveness_reason("proxy", PROXY_BASE_URL) if reason is not None: return reason @@ -64,19 +64,19 @@ def _proxy_skip_reason() -> str | None: def pytest_runtest_setup(item: pytest.Item) -> None: - """Skip `e2e`-marked tests unless a proxy answers its liveness probe. Unmarked - tests (unit coverage of the harness) don't touch the proxy, so they run even - when none is up.""" + """Hard-fail `e2e`-marked tests unless a proxy answers its liveness probe. + Unmarked tests (unit coverage of the harness) don't touch the proxy, so they + run even when none is up. Never skip for a missing proxy.""" if item.get_closest_marker("e2e") is None: return - reason = _proxy_skip_reason() + reason = _proxy_fail_reason() if reason is not None: - pytest.skip(reason) + pytest.fail(reason) def pytest_runtest_call(item: pytest.Item) -> None: - """Mark that an e2e test body actually ran (not skipped at setup). Skipped - sessions never reach this hook, so the session-finish cleanup can use it as a + """Mark that an e2e test body actually ran (setup passed). Sessions that fail + setup never reach this hook, so the session-finish cleanup can use it as a guard before truncating the spend-log DB. Tests under `tests/e2e/` without the `e2e` marker (pure unit coverage for the harness itself) never hit the proxy, so they must not arm the destructive DB truncate.""" @@ -87,12 +87,12 @@ def pytest_runtest_call(item: pytest.Item) -> None: def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None: """Once the whole e2e session is done (all suites), truncate the spend logs so - the DB doesn't accumulate test rows. Skipped sessions (no live proxy, no test - actually executed) leave the DB alone so a `DATABASE_URL` pointing at a shared - instance is never wiped without an e2e run. Best-effort: a cleanup failure (no - DB reachable) must not fail the run. The spend_tracking dir goes on sys.path - only for this import and is removed after, so a broader `pytest tests/` run is - not left with a mutated path.""" + the DB doesn't accumulate test rows. Sessions where no e2e test body ran leave + the DB alone so a `DATABASE_URL` pointing at a shared instance is never wiped + without an e2e run. Best-effort: a cleanup failure (no DB reachable) must not + fail the run. The spend_tracking dir goes on sys.path only for this import and + is removed after, so a broader `pytest tests/` run is not left with a mutated + path.""" if not session.stash.get(_E2E_TEST_RAN, False): return spend_dir = str(Path(__file__).parent / "spend_tracking") @@ -102,7 +102,7 @@ def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None: reset_spend_logs() except Exception as exc: # noqa: BLE001 - cleanup is best-effort - print(f"spend-log cleanup skipped: {exc}") + print(f"spend-log cleanup best-effort failed: {exc}") finally: if spend_dir in sys.path: sys.path.remove(spend_dir) @@ -112,7 +112,7 @@ def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None: remediate(session) except Exception as exc: # noqa: BLE001 - remediation is best-effort - print(f"devin remediation skipped: {exc}") + print(f"devin remediation best-effort failed: {exc}") @pytest.fixture diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index 7b8d3045d3b..32005faed4b 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -107,13 +107,15 @@ class ProbeResult(BaseModel): class StreamingResponse(BaseModel): """Raw outcome for calls whose body is provider-native or streamed: status, the - x-litellm-call-id header (== SpendLogs.request_id), the content-type (which - tells streaming `text/event-stream` from non-streaming `application/json`), and - the body. Used by passthrough and streaming, where one validated JSON model - does not fit.""" + x-litellm-call-id header, the x-litellm-response-cost header (StandardLogging + response_cost), the content-type (which tells streaming `text/event-stream` from + non-streaming `application/json`), and the body. SpendLogs.request_id is the + completion body id, not call_id. Used by passthrough and streaming, where one + validated JSON model does not fit.""" status_code: int call_id: str | None = None # x-litellm-call-id header + response_cost: float | None = None # x-litellm-response-cost header content_type: str | None = None body: str chunks: int = 0 # streamed events (0 for non-streaming) @@ -260,13 +262,25 @@ def probe( return ProbeResult(status_code=resp.status_code, body=resp.text) +def _parse_response_cost(resp: requests.Response) -> float | None: + raw = _hdr(resp, "x-litellm-response-cost") + if raw is None or raw == "": + return None + try: + return float(raw) + except ValueError: + return None + + def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingResponse: call_id = _hdr(resp, "x-litellm-call-id") + response_cost = _parse_response_cost(resp) content_type = _hdr(resp, "content-type") if not stream or not (200 <= resp.status_code < 300): return StreamingResponse( status_code=resp.status_code, call_id=call_id, + response_cost=response_cost, content_type=content_type, body=resp.text, ) @@ -275,6 +289,7 @@ def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingRespon return StreamingResponse( status_code=resp.status_code, call_id=call_id, + response_cost=response_cost, content_type=content_type, body="", chunks=chunks, diff --git a/tests/e2e/logging/conftest.py b/tests/e2e/logging/conftest.py index 40c19aefca7..567355f23ac 100644 --- a/tests/e2e/logging/conftest.py +++ b/tests/e2e/logging/conftest.py @@ -1,37 +1,43 @@ -"""Fixtures for the Datadog logging suite. +"""Fixtures for the logging e2e suite. -These tests drive the Datadog batch-send path (#25663) directly against the real -Datadog logs intake with synthetic events - no LLM calls, no proxy, no log -read-back - so they need only the shipping credentials DD_API_KEY + DD_SITE -(DD_SERVICE is an optional tag). No Datadog Application key is required, and they -skip when the shipping credentials are absent from the environment. +Missing proxy, provider keys, or integration credentials are hard failures. +Never pytest.skip from this suite for environment gaps. """ +from __future__ import annotations + import os import pytest -from logging_client import LoggingClient, build_logging_client +from logging_client import LangfuseCreds, LoggingClient, build_logging_client, load_langfuse_creds def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line( "markers", - "covers: registry cell a test covers, e.g. logging.datadog.success.writes_object", + "covers: registry cell a test covers, e.g. logging.langfuse.success.logs_spend", ) @pytest.fixture(scope="session") def client() -> LoggingClient: """The logging suite's client: holds the shared Gateway so `resources` / - `scoped_key` clean up keys, and adds `/metrics` scraping.""" + `scoped_key` clean up keys and teams, and adds `/metrics` scraping plus + Langfuse read-back.""" return build_logging_client() @pytest.fixture def datadog_creds() -> None: - """Gate the suite on the Datadog shipping credentials. The DataDogLogger is built - inside each async test, not here, because its __init__ schedules a periodic-flush - task via asyncio.create_task and so needs a running event loop.""" + """Require Datadog shipping credentials. Hard-fail when absent; never skip.""" if not (os.getenv("DD_API_KEY") and os.getenv("DD_SITE")): - pytest.skip("set DD_API_KEY and DD_SITE to run the Datadog logging suite") + pytest.fail( + "Datadog e2e requires DD_API_KEY and DD_SITE; missing credentials is a hard failure, not a skip" + ) + + +@pytest.fixture(scope="session") +def langfuse_creds() -> LangfuseCreds: + """Require real Langfuse cloud credentials for team callback + trace poll.""" + return load_langfuse_creds() diff --git a/tests/e2e/logging/logging_client.py b/tests/e2e/logging/logging_client.py index a3213fbdb00..06b219fc6e2 100644 --- a/tests/e2e/logging/logging_client.py +++ b/tests/e2e/logging/logging_client.py @@ -1,48 +1,571 @@ -"""Client for the logging e2e suite: drive traffic and scrape the proxy's -Prometheus ``/metrics`` endpoint. +"""Client for the logging e2e suite: team/key/org-scoped Langfuse OTEL callbacks, +chat (including tools), Prometheus scrape, and Langfuse observation read-back. -Holds the shared Gateway so the ``resources`` fixture cleans up keys it creates. -``/metrics`` is exposed as plaintext (not a typed JSON body), so scraping goes -through ``transport.probe`` and returns the raw exposition text for a Prometheus -parser to read. +Holds the shared Gateway so the ``resources`` fixture cleans up keys, teams, +users, orgs, and models it creates. External Langfuse reads go through +``e2e_http`` (the only module allowed to call ``requests.*``). + +Uses the ``langfuse_otel`` callback (OTLP to ``{host}/api/public/otel``), not +the classic ``langfuse`` SDK callback. OTEL generations land as name +``litellm_request``; correlate by unique prompt marker and ``user_api_key_alias`` +in metadata. Spend is on ``calculatedTotalCost`` (StandardLogging response_cost). """ from __future__ import annotations +import base64 +import json +import os +import time from dataclasses import dataclass +from typing import Literal +import pytest +from pydantic import BaseModel, ConfigDict, Field + +from e2e_config import POLL_INTERVAL, POLL_TIMEOUT from e2e_gateway import Gateway, build_gateway -from e2e_http import NoBody, unwrap -from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody +from e2e_http import ( + URL, + AuthHeaders, + NoBody, + StreamingResponse, + Success, + get, + unwrap, +) +from models import ( + ChatBody, + ChatMessage, + ChatResponse, + ChatTool, + ChatToolFunction, + KeyGenerateBody, + KeyLoggingCallback, + KeyLoggingCallbackVars, + KeyMetadata, + LiteLLMParamsBody, + OrgDeleteBody, + OrgNewBody, + OrgNewResponse, + SpendLogRow, + TeamDeleteBody, + TeamNewBody, + TeamNewResponse, + UserDeleteBody, + UserNewBody, + UserNewResponse, +) + +# Deliberately invalid *upstream provider* key for failure-path tests. +# Not a LiteLLM virtual key; OpenAI must reject it after the proxy accepts the call. +INVALID_UPSTREAM_API_KEY = "sk-upstream-invalid-for-langfuse-e2e-only" + +WEATHER_TOOL = ChatTool( + type="function", + function=ChatToolFunction( + name="get_weather", + description="Get the current weather for a city", + parameters={ + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + ), +) + + +class TeamCallbackBody(BaseModel): + callback_name: Literal["langfuse_otel", "langfuse", "langsmith", "gcs"] + callback_type: Literal["success", "failure", "success_and_failure"] + callback_vars: dict[str, str] + + +class TeamCallbackResponse(BaseModel): + model_config = ConfigDict(extra="ignore") + + status: str + + +class GuardrailLitellmParams(BaseModel): + guardrail: str + mode: str + default_on: bool = False + rules: list[dict[str, object]] | None = None + default_action: str | None = None + on_disallowed_action: str | None = None + + +class GuardrailSpec(BaseModel): + guardrail_name: str + litellm_params: GuardrailLitellmParams + + +class CreateGuardrailBody(BaseModel): + guardrail: GuardrailSpec + + +class CreateGuardrailResponse(BaseModel): + model_config = ConfigDict(extra="ignore") + + guardrail_id: str | None = None + guardrail_name: str | None = None + + +class LangfuseObservation(BaseModel): + model_config = ConfigDict(extra="ignore", populate_by_name=True) + + id: str + trace_id: str | None = Field(default=None, alias="traceId") + name: str | None = None + type: str | None = None + calculated_total_cost: float | None = Field(default=None, alias="calculatedTotalCost") + level: str | None = None + input: object | None = None + output: object | None = None + metadata: object | None = None + usage: object | None = None + usage_details: object | None = Field(default=None, alias="usageDetails") + model: str | None = None + + +class LangfuseObservationList(BaseModel): + model_config = ConfigDict(extra="ignore") + + data: list[LangfuseObservation] = [] + + +class LangfuseListParams(BaseModel): + model_config = ConfigDict(populate_by_name=True) + + limit: int = 100 + trace_id: str | None = Field(default=None, alias="traceId") + name: str | None = None + from_start_time: str | None = Field(default=None, alias="fromStartTime") + + +@dataclass(frozen=True, slots=True) +class LangfuseCreds: + public_key: str + secret_key: str + host: str + + @property + def auth_headers(self) -> AuthHeaders: + token = base64.b64encode(f"{self.public_key}:{self.secret_key}".encode()).decode() + return AuthHeaders(authorization=f"Basic {token}") + + def callback_vars(self) -> dict[str, str]: + return { + "langfuse_public_key": self.public_key, + "langfuse_secret_key": self.secret_key, + "langfuse_host": self.host, + } + + def key_logging_metadata(self) -> KeyMetadata: + return KeyMetadata( + logging=[ + KeyLoggingCallback( + callback_name="langfuse_otel", + callback_type="success_and_failure", + callback_vars=KeyLoggingCallbackVars( + langfuse_public_key=self.public_key, + langfuse_secret_key=self.secret_key, + langfuse_host=self.host, + ), + ) + ] + ) + + +def load_langfuse_creds() -> LangfuseCreds: + public_key = os.getenv("LANGFUSE_PUBLIC_KEY") + secret_key = os.getenv("LANGFUSE_SECRET_KEY") + host = (os.getenv("LANGFUSE_BASE_URL") or os.getenv("LANGFUSE_HOST") or "").rstrip("/") + if not (public_key and secret_key and host): + pytest.fail( + "Langfuse e2e requires LANGFUSE_PUBLIC_KEY, LANGFUSE_SECRET_KEY, and " + "LANGFUSE_BASE_URL (or LANGFUSE_HOST); missing credentials is a hard failure, not a skip" + ) + return LangfuseCreds(public_key=public_key, secret_key=secret_key, host=host) + + +def observation_spend(obs: LangfuseObservation) -> float | None: + """Langfuse calculatedTotalCost is populated from StandardLogging response_cost.""" + return obs.calculated_total_cost + + +def costs_agree(expected: float, actual: float, *, rel_tol: float = 0.05) -> bool: + """Costs agree within 5% relative (or 1e-9 absolute for near-zero).""" + return abs(expected - actual) <= max(1e-9, abs(expected) * rel_tol) + + +def completion_response_id(body: str) -> str | None: + """SpendLogs.request_id is the chat completion body id, not x-litellm-call-id.""" + if not body or body == "": + return None + try: + parsed = json.loads(body) + except json.JSONDecodeError: + return None + if not isinstance(parsed, dict): + return None + raw = parsed.get("id") + return raw if isinstance(raw, str) and raw else None + + +def _matches_run(obs: LangfuseObservation, *, key_alias: str, prompt_marker: str) -> bool: + """Match a Langfuse generation for this run. + + langfuse_otel names generations ``litellm_request`` (not ``litellm:{alias}``). + Prefer the unique prompt marker in input; fall back to key alias in metadata + (user_api_key_alias) or the classic SDK generation name. + """ + if prompt_marker and prompt_marker in json.dumps(obs.input, default=str): + return True + meta_blob = json.dumps(obs.metadata, default=str) if obs.metadata is not None else "" + if key_alias and key_alias in meta_blob: + return True + if obs.name == f"litellm:{key_alias}": + return True + return False + + +def observation_mentions_tool(obs: LangfuseObservation, tool_name: str) -> bool: + blob = json.dumps( + {"input": obs.input, "output": obs.output, "metadata": obs.metadata}, + default=str, + ) + return tool_name in blob + + +def observation_has_guardrail(obs: LangfuseObservation, *, guardrail_name: str) -> bool: + blob = json.dumps(obs.metadata, default=str) if obs.metadata is not None else "" + if guardrail_name in blob or "guardrail" in blob.lower(): + return True + if obs.name is not None and "guardrail" in obs.name.lower(): + return True + return False @dataclass(frozen=True, slots=True) class LoggingClient: gateway: Gateway - def key_with_alias(self, alias: str, *, models: list[str]) -> str: + def key_with_alias( + self, + alias: str, + *, + models: list[str], + team_id: str | None = None, + user_id: str | None = None, + organization_id: str | None = None, + metadata: KeyMetadata | None = None, + ) -> str: return self.gateway.generate_key( - KeyGenerateBody(key_alias=alias, models=models, user_id=f"e2e-{alias}") + KeyGenerateBody( + key_alias=alias, + models=models, + user_id=user_id or f"e2e-{alias}", + team_id=team_id, + organization_id=organization_id, + metadata=metadata, + ) ) def delete_key(self, key: str) -> None: self.gateway.delete_key(key) + def create_team( + self, + alias: str, + *, + models: list[str], + organization_id: str | None = None, + ) -> str: + return unwrap( + self.gateway.transport.post( + "/team/new", + headers=self.gateway.transport.master, + json=TeamNewBody( + team_alias=alias, + models=models, + organization_id=organization_id, + ), + response_type=TeamNewResponse, + ) + ).team_id + + def delete_team(self, team_id: str) -> None: + _ = self.gateway.transport.post( + "/team/delete", + headers=self.gateway.transport.master, + json=TeamDeleteBody(team_ids=[team_id]), + response_type=NoBody, + ) + + def create_user(self, *, user_email: str, user_id: str | None = None) -> str: + return unwrap( + self.gateway.transport.post( + "/user/new", + headers=self.gateway.transport.master, + json=UserNewBody( + user_email=user_email, + user_role="internal_user", + user_id=user_id, + ), + response_type=UserNewResponse, + ) + ).user_id + + def delete_user(self, user_id: str) -> None: + _ = self.gateway.transport.post( + "/user/delete", + headers=self.gateway.transport.master, + json=UserDeleteBody(user_ids=[user_id]), + response_type=NoBody, + ) + + def create_org(self, alias: str, *, models: list[str]) -> str: + return unwrap( + self.gateway.transport.post( + "/organization/new", + headers=self.gateway.transport.master, + json=OrgNewBody(organization_alias=alias, models=models), + response_type=OrgNewResponse, + ) + ).organization_id + + def delete_org(self, organization_id: str) -> None: + _ = self.gateway.transport.delete( + "/organization/delete", + headers=self.gateway.transport.master, + json=OrgDeleteBody(organization_ids=[organization_id]), + response_type=NoBody, + ) + + def add_team_langfuse_callback( + self, + team_id: str, + creds: LangfuseCreds, + *, + callback_type: Literal["success", "failure", "success_and_failure"] = "success_and_failure", + ) -> None: + response = unwrap( + self.gateway.transport.post( + f"/team/{team_id}/callback", + headers=self.gateway.transport.master, + json=TeamCallbackBody( + callback_name="langfuse_otel", + callback_type=callback_type, + callback_vars=creds.callback_vars(), + ), + response_type=TeamCallbackResponse, + ) + ) + assert response.status == "success", ( + f"POST /team/{team_id}/callback must return status=success; got {response.status!r}" + ) + + def create_tool_permission_guardrail(self, name: str, *, allowed_tool: str) -> str: + """Register a tool_permission guardrail that allows one tool and denies the rest.""" + response = unwrap( + self.gateway.transport.post( + "/guardrails", + headers=self.gateway.transport.master, + json=CreateGuardrailBody( + guardrail=GuardrailSpec( + guardrail_name=name, + litellm_params=GuardrailLitellmParams( + guardrail="tool_permission", + mode="post_call", + default_on=False, + default_action="deny", + on_disallowed_action="block", + rules=[ + { + "id": "allow-named-tool", + "tool_name": allowed_tool, + "decision": "allow", + } + ], + ), + ) + ), + response_type=CreateGuardrailResponse, + ) + ) + guardrail_id = response.guardrail_id + assert guardrail_id, f"create guardrail returned no id: {response!r}" + return guardrail_id + + def delete_guardrail(self, guardrail_id: str) -> None: + _ = self.gateway.transport.delete( + f"/guardrails/{guardrail_id}", + headers=self.gateway.transport.master, + json=NoBody(), + response_type=NoBody, + ) + + def create_model(self, model_name: str, litellm_params: LiteLLMParamsBody) -> str: + return self.gateway.create_model(model_name, litellm_params) + + def delete_model(self, model_id: str) -> None: + self.gateway.delete_model(model_id) + def chat(self, key: str, model: str, text: str) -> ChatResponse: return unwrap( self.gateway.chat( key, ChatBody( model=model, - messages=[ChatMessage(role="user", content=text)], - max_tokens=64, + messages=[ChatMessage(role="user", content=text)], + max_tokens=64, ), ) ) + def chat_raw( + self, + key: str, + model: str, + text: str, + *, + stream: bool = False, + tools: list[ChatTool] | None = None, + tool_choice: str | None = None, + guardrails: list[str] | None = None, + max_tokens: int = 64, + ) -> StreamingResponse: + body = ChatBody( + model=model, + messages=[ChatMessage(role="user", content=text)], + max_tokens=max_tokens, + stream=stream, + tools=tools, + tool_choice=tool_choice, + guardrails=guardrails, + ) + if stream: + return self.gateway.chat_stream(key, body) + return self.gateway.transport.send( + "/chat/completions", + headers=self.gateway.transport.bearer(key), + json=body, + ) + def scrape_metrics(self) -> str: return self.gateway.probe("/metrics", params=NoBody()).body + def poll_proxy_spend_for_key( + self, + key: str, + *, + response_id: str | None = None, + require_positive_spend: bool = True, + ) -> SpendLogRow | None: + """Poll /spend/logs by virtual key. + + When ``response_id`` is set, only that SpendLogs.request_id may match. + When unset, any positive-spend row for the key is accepted. Never falls + back to an unmatched row; missing match returns None. + """ + + def _matches(row: SpendLogRow) -> bool: + if response_id is not None and row.request_id != response_id: + return False + if require_positive_spend and not (row.spend is not None and row.spend > 0): + return False + return True + + rows = self.gateway.poll_logs_for_key( + key, min_rows=1, predicate=lambda rs: any(_matches(r) for r in rs) + ) + for row in rows: + if _matches(row): + return row + return None + + def list_langfuse_observations( + self, + creds: LangfuseCreds, + *, + trace_id: str | None = None, + name: str | None = None, + from_start_time: str | None = None, + ) -> list[LangfuseObservation]: + result = get( + URL(f"{creds.host}/api/public/observations"), + headers=creds.auth_headers, + params=LangfuseListParams( + limit=100, + trace_id=trace_id, + name=name, + from_start_time=from_start_time, + ), + response_type=LangfuseObservationList, + timeout=30.0, + ) + match result: + case Success(data=page): + return page.data + case _: + return [] + + def find_langfuse_observation( + self, + creds: LangfuseCreds, + *, + key_alias: str, + prompt_marker: str, + ) -> LangfuseObservation | None: + # langfuse_otel generations are named litellm_request; classic SDK used + # litellm:{key_alias}. Search both, then a recent unfiltered page. + for name in ("litellm_request", f"litellm:{key_alias}"): + for obs in self.list_langfuse_observations(creds, name=name): + if _matches_run(obs, key_alias=key_alias, prompt_marker=prompt_marker): + return obs + for obs in self.list_langfuse_observations(creds): + if _matches_run(obs, key_alias=key_alias, prompt_marker=prompt_marker): + return obs + return None + + def poll_langfuse_observation( + self, + creds: LangfuseCreds, + *, + key_alias: str, + prompt_marker: str, + require_positive_cost: bool = False, + ) -> LangfuseObservation | None: + deadline = time.monotonic() + POLL_TIMEOUT + last: LangfuseObservation | None = None + while time.monotonic() < deadline: + last = self.find_langfuse_observation( + creds, key_alias=key_alias, prompt_marker=prompt_marker + ) + if last is not None: + cost = observation_spend(last) + if not require_positive_cost or (cost is not None and cost > 0): + return last + time.sleep(POLL_INTERVAL) + return last + + def poll_langfuse_trace_observations( + self, + creds: LangfuseCreds, + *, + key_alias: str, + prompt_marker: str, + ) -> list[LangfuseObservation]: + """Generation plus any sibling/child observations (guardrail spans, etc.).""" + gen = self.poll_langfuse_observation( + creds, key_alias=key_alias, prompt_marker=prompt_marker + ) + if gen is None or not gen.trace_id: + return [] if gen is None else [gen] + return self.list_langfuse_observations(creds, trace_id=gen.trace_id) or [gen] + def build_logging_client() -> LoggingClient: return LoggingClient(gateway=build_gateway()) diff --git a/tests/e2e/logging/test_langfuse_e2e.py b/tests/e2e/logging/test_langfuse_e2e.py new file mode 100644 index 00000000000..d014b5d8291 --- /dev/null +++ b/tests/e2e/logging/test_langfuse_e2e.py @@ -0,0 +1,534 @@ +"""Live e2e: Langfuse OTEL logs_spend for registry cells in logging.yaml P0. + +Registry cells: +- logging.langfuse.success.logs_spend (exercised_on chat_completions, messages, embeddings) +- logging.langfuse.failure.logs_spend (exercised_on chat_completions, messages) +- logging.langfuse.stream.logs_spend (exercised_on chat_completions, messages) + +Integration under test is ``langfuse_otel`` (OTLP to Langfuse), not the classic +``langfuse`` SDK callback. StandardLoggingPayload.response_cost is the spend +source of truth. Generations are named ``litellm_request``; correlate by unique +prompt marker and user_api_key_alias in metadata. + +Dynamic credentials by product surface: +- team: POST /team/{id}/callback with callback_name=langfuse_otel +- user/key: key metadata.logging with callback_name=langfuse_otel +- org: organization + team under it + team callback (no org-level callback API) + +Extra success paths assert tool calls and applied guardrails land on the trace. +""" + +from __future__ import annotations + +import json + +import pytest + +from e2e_config import unique_marker +from e2e_http import StreamingResponse, require_successful_call +from lifecycle import ResourceManager +from logging_client import ( + INVALID_UPSTREAM_API_KEY, + WEATHER_TOOL, + LangfuseCreds, + LoggingClient, + completion_response_id, + costs_agree, + observation_has_guardrail, + observation_mentions_tool, + observation_spend, +) +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + +DRIVER_MODEL = "gemini-2.5-flash" +FAIL_BACKEND = "openai/gpt-4o-mini" + + +def _json_blob(value: object) -> str: + return json.dumps(value, default=str) + + +def _assert_logs_spend( + client: LoggingClient, + *, + key: str, + outcome: StreamingResponse, + obs_cost: float | None, + scope: str, + require_positive: bool = True, +) -> None: + """logs_spend: Langfuse cost matches StandardLogging response_cost and proxy spend. + + Non-stream responses expose response_cost on x-litellm-response-cost. Streaming + sends headers before final cost is known, so stream paths rely on /spend/logs. + """ + if not require_positive: + assert obs_cost is not None, ( + f"{scope}: failure path must still track spend (0 is fine); cost={obs_cost!r}" + ) + return + + assert obs_cost is not None and obs_cost > 0, ( + f"{scope}: Langfuse must log positive spend; calculatedTotalCost={obs_cost!r}" + ) + # Stream responses send headers before final cost is known, so the cost header + # is often absent; non-stream must always expose x-litellm-response-cost. + if not outcome.is_streaming: + assert outcome.response_cost is not None and outcome.response_cost > 0, ( + f"{scope}: proxy must return positive x-litellm-response-cost; " + f"got {outcome.response_cost!r}" + ) + assert costs_agree(outcome.response_cost, obs_cost), ( + f"{scope}: Langfuse cost {obs_cost!r} disagrees with " + f"x-litellm-response-cost {outcome.response_cost!r}" + ) + elif outcome.response_cost is not None and outcome.response_cost > 0: + assert costs_agree(outcome.response_cost, obs_cost), ( + f"{scope}: Langfuse cost {obs_cost!r} disagrees with " + f"x-litellm-response-cost {outcome.response_cost!r}" + ) + spend_row = client.poll_proxy_spend_for_key( + key, + response_id=completion_response_id(outcome.body), + require_positive_spend=True, + ) + assert spend_row is not None and spend_row.spend is not None and spend_row.spend > 0, ( + f"{scope}: proxy /spend/logs never produced a positive spend row for key" + ) + assert costs_agree(spend_row.spend, obs_cost), ( + f"{scope}: Langfuse cost {obs_cost!r} disagrees with proxy spend " + f"{spend_row.spend!r} (request_id={spend_row.request_id!r})" + ) + + +class TestLangfuseTeamLogging: + """Team-scoped callback via POST /team/{id}/callback.""" + + def _team_key( + self, + client: LoggingClient, + resources: ResourceManager, + creds: LangfuseCreds, + *, + models: list[str], + organization_id: str | None = None, + ) -> tuple[str, str, str]: + marker = unique_marker() + key_alias = f"e2e-lf-team-key-{marker}" + team_id = client.create_team( + f"e2e-lf-team-{marker}", + models=models, + organization_id=organization_id, + ) + resources.defer(lambda: client.delete_team(team_id)) + client.add_team_langfuse_callback(team_id, creds) + key = client.key_with_alias(key_alias, models=models, team_id=team_id) + resources.defer(lambda: client.delete_key(key)) + return team_id, key, key_alias + + @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) + def test_success_logs_spend( + self, + client: LoggingClient, + resources: ResourceManager, + langfuse_creds: LangfuseCreds, + ) -> None: + _, key, key_alias = self._team_key( + client, resources, langfuse_creds, models=[DRIVER_MODEL] + ) + prompt_marker = unique_marker() + outcome = client.chat_raw( + key, DRIVER_MODEL, f"reply with one word only {prompt_marker}" + ) + require_successful_call(outcome) + + obs = client.poll_langfuse_observation( + langfuse_creds, + key_alias=key_alias, + prompt_marker=prompt_marker, + require_positive_cost=True, + ) + assert obs is not None, ( + f"team scope: Langfuse never received generation for key_alias={key_alias!r}" + ) + _assert_logs_spend( + client, + key=key, + outcome=outcome, + obs_cost=observation_spend(obs), + scope="team-success", + ) + + @pytest.mark.covers("logging.langfuse.failure.logs_spend", exercised_on=["chat_completions"]) + def test_failure_logs_spend( + self, + client: LoggingClient, + resources: ResourceManager, + langfuse_creds: LangfuseCreds, + ) -> None: + """Provider-auth failure still ships a Langfuse observation with spend tracked. + + Uses a throwaway deployment whose upstream OpenAI key is + INVALID_UPSTREAM_API_KEY (not a LiteLLM virtual key). + """ + prompt_marker = unique_marker() + model_name = f"e2e-lf-fail-{prompt_marker}" + model_id = client.create_model( + model_name, + LiteLLMParamsBody(model=FAIL_BACKEND, api_key=INVALID_UPSTREAM_API_KEY), + ) + resources.defer(lambda: client.delete_model(model_id)) + + _, key, key_alias = self._team_key( + client, resources, langfuse_creds, models=[model_name] + ) + outcome = client.chat_raw(key, model_name, f"this must fail {prompt_marker}") + assert not outcome.ok, ( + f"expected upstream provider failure for {INVALID_UPSTREAM_API_KEY!r}, " + f"got {outcome.status_code}: {outcome.body[:200]}" + ) + + obs = client.poll_langfuse_observation( + langfuse_creds, + key_alias=key_alias, + prompt_marker=prompt_marker, + require_positive_cost=False, + ) + assert obs is not None, ( + f"team failure path: Langfuse never received generation for key_alias={key_alias!r}" + ) + _assert_logs_spend( + client, + key=key, + outcome=outcome, + obs_cost=observation_spend(obs), + scope="team-failure", + require_positive=False, + ) + + @pytest.mark.covers("logging.langfuse.stream.logs_spend", exercised_on=["chat_completions"]) + def test_stream_logs_spend( + self, + client: LoggingClient, + resources: ResourceManager, + langfuse_creds: LangfuseCreds, + ) -> None: + _, key, key_alias = self._team_key( + client, resources, langfuse_creds, models=[DRIVER_MODEL] + ) + prompt_marker = unique_marker() + outcome = client.chat_raw( + key, DRIVER_MODEL, f"reply with one word only {prompt_marker}", stream=True + ) + require_successful_call(outcome) + assert outcome.is_streaming + assert outcome.chunks > 0 + + obs = client.poll_langfuse_observation( + langfuse_creds, + key_alias=key_alias, + prompt_marker=prompt_marker, + require_positive_cost=True, + ) + assert obs is not None + # Streamed body is elided; correlate cost via header + key spend row. + _assert_logs_spend( + client, + key=key, + outcome=outcome, + obs_cost=observation_spend(obs), + scope="team-stream", + ) + + @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) + def test_tool_calls_logged_with_cost( + self, + client: LoggingClient, + resources: ResourceManager, + langfuse_creds: LangfuseCreds, + ) -> None: + _, key, key_alias = self._team_key( + client, resources, langfuse_creds, models=[DRIVER_MODEL] + ) + prompt_marker = unique_marker() + outcome = client.chat_raw( + key, + DRIVER_MODEL, + f"Use get_weather for Paris. marker={prompt_marker}", + tools=[WEATHER_TOOL], + tool_choice="required", + max_tokens=128, + ) + require_successful_call(outcome) + assert "get_weather" in outcome.body or "tool_calls" in outcome.body, ( + f"gateway response must include a tool call; body={outcome.body[:300]}" + ) + + obs = client.poll_langfuse_observation( + langfuse_creds, + key_alias=key_alias, + prompt_marker=prompt_marker, + require_positive_cost=True, + ) + assert obs is not None + assert observation_mentions_tool(obs, "get_weather"), ( + f"Langfuse generation must record the tool; name={obs.name!r} " + f"input={str(obs.input)[:200]} output={str(obs.output)[:200]}" + ) + _assert_logs_spend( + client, + key=key, + outcome=outcome, + obs_cost=observation_spend(obs), + scope="team-tools", + ) + + @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) + def test_tool_permission_guardrail_logged( + self, + client: LoggingClient, + resources: ResourceManager, + langfuse_creds: LangfuseCreds, + ) -> None: + """tool_permission post_call guardrail must appear on the Langfuse trace + (StandardLogging guardrail_information -> Langfuse guardrail span).""" + marker = unique_marker() + guardrail_name = f"e2e-lf-tool-perm-{marker}" + guardrail_id = client.create_tool_permission_guardrail( + guardrail_name, allowed_tool="get_weather" + ) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + _, key, key_alias = self._team_key( + client, resources, langfuse_creds, models=[DRIVER_MODEL] + ) + prompt_marker = unique_marker() + outcome = client.chat_raw( + key, + DRIVER_MODEL, + f"Use get_weather for Berlin. marker={prompt_marker}", + tools=[WEATHER_TOOL], + tool_choice="required", + guardrails=[guardrail_name], + max_tokens=128, + ) + require_successful_call(outcome) + + observations = client.poll_langfuse_trace_observations( + langfuse_creds, key_alias=key_alias, prompt_marker=prompt_marker + ) + assert observations, ( + f"team+guardrail: no Langfuse observations for key_alias={key_alias!r}" + ) + gen = next( + ( + o + for o in observations + if prompt_marker in _json_blob(o.input) + or key_alias in _json_blob(o.metadata) + or o.name in (f"litellm:{key_alias}", "litellm_request") + ), + observations[0], + ) + _assert_logs_spend( + client, + key=key, + outcome=outcome, + obs_cost=observation_spend(gen), + scope="team-guardrail", + ) + assert any( + observation_has_guardrail(o, guardrail_name=guardrail_name) + or (o.name is not None and "guardrail" in o.name.lower()) + for o in observations + ), ( + f"Langfuse trace must include applied guardrail {guardrail_name!r}; " + f"observation names={[o.name for o in observations]}" + ) + + +class TestLangfuseUserKeyLogging: + """User-owned key with metadata.logging (key-level dynamic Langfuse credentials). + + Product surface: key metadata.logging on /key/generate, not a separate + /user/.../callback route. The key is bound to a real /user/new user_id. + """ + + def _user_key( + self, + client: LoggingClient, + resources: ResourceManager, + creds: LangfuseCreds, + *, + models: list[str], + ) -> tuple[str, str, str]: + marker = unique_marker() + key_alias = f"e2e-lf-user-key-{marker}" + user_id = client.create_user( + user_email=f"e2e-lf-user-{marker}@example.com", + user_id=f"e2e-lf-user-{marker}", + ) + resources.defer(lambda: client.delete_user(user_id)) + key = client.key_with_alias( + key_alias, + models=models, + user_id=user_id, + metadata=creds.key_logging_metadata(), + ) + resources.defer(lambda: client.delete_key(key)) + return user_id, key, key_alias + + @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) + def test_success_logs_spend( + self, + client: LoggingClient, + resources: ResourceManager, + langfuse_creds: LangfuseCreds, + ) -> None: + user_id, key, key_alias = self._user_key( + client, resources, langfuse_creds, models=[DRIVER_MODEL] + ) + prompt_marker = unique_marker() + outcome = client.chat_raw( + key, DRIVER_MODEL, f"reply with one word only {prompt_marker}" + ) + require_successful_call(outcome) + + obs = client.poll_langfuse_observation( + langfuse_creds, + key_alias=key_alias, + prompt_marker=prompt_marker, + require_positive_cost=True, + ) + assert obs is not None, ( + f"user/key scope: Langfuse never received generation for key_alias={key_alias!r}" + ) + meta_blob = _json_blob(obs.metadata) + assert user_id in meta_blob or key_alias in (obs.name or ""), ( + f"user/key scope should attribute the user or key; metadata={meta_blob[:300]}" + ) + _assert_logs_spend( + client, + key=key, + outcome=outcome, + obs_cost=observation_spend(obs), + scope="user-key", + ) + + @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) + def test_tool_calls_logged_with_cost( + self, + client: LoggingClient, + resources: ResourceManager, + langfuse_creds: LangfuseCreds, + ) -> None: + _, key, key_alias = self._user_key( + client, resources, langfuse_creds, models=[DRIVER_MODEL] + ) + prompt_marker = unique_marker() + outcome = client.chat_raw( + key, + DRIVER_MODEL, + f"Use get_weather for Tokyo. marker={prompt_marker}", + tools=[WEATHER_TOOL], + tool_choice="required", + max_tokens=128, + ) + require_successful_call(outcome) + + obs = client.poll_langfuse_observation( + langfuse_creds, + key_alias=key_alias, + prompt_marker=prompt_marker, + require_positive_cost=True, + ) + assert obs is not None + assert observation_mentions_tool(obs, "get_weather"), ( + f"user/key tool path: tool missing from Langfuse; output={str(obs.output)[:200]}" + ) + _assert_logs_spend( + client, + key=key, + outcome=outcome, + obs_cost=observation_spend(obs), + scope="user-key-tools", + ) + + +class TestLangfuseOrgScopedLogging: + """Org-scoped run: organization + team under it + team Langfuse callback. + + There is no /organization/.../callback today; logging attaches at the team + (or key) under the org. This class proves org-linked team keys still deliver + accurate Langfuse spend and team attribution (StandardLogging metadata + user_api_key_team_id / user_api_key_org_id). + """ + + def _org_team_key( + self, + client: LoggingClient, + resources: ResourceManager, + creds: LangfuseCreds, + *, + models: list[str], + ) -> tuple[str, str, str, str]: + marker = unique_marker() + key_alias = f"e2e-lf-org-key-{marker}" + org_id = client.create_org(f"e2e-lf-org-{marker}", models=models) + resources.defer(lambda: client.delete_org(org_id)) + team_id = client.create_team( + f"e2e-lf-org-team-{marker}", + models=models, + organization_id=org_id, + ) + resources.defer(lambda: client.delete_team(team_id)) + client.add_team_langfuse_callback(team_id, creds) + key = client.key_with_alias( + key_alias, + models=models, + team_id=team_id, + organization_id=org_id, + ) + resources.defer(lambda: client.delete_key(key)) + return org_id, team_id, key, key_alias + + @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) + def test_success_logs_spend_with_team_attribution( + self, + client: LoggingClient, + resources: ResourceManager, + langfuse_creds: LangfuseCreds, + ) -> None: + org_id, team_id, key, key_alias = self._org_team_key( + client, resources, langfuse_creds, models=[DRIVER_MODEL] + ) + prompt_marker = unique_marker() + outcome = client.chat_raw( + key, DRIVER_MODEL, f"reply with one word only {prompt_marker}" + ) + require_successful_call(outcome) + + obs = client.poll_langfuse_observation( + langfuse_creds, + key_alias=key_alias, + prompt_marker=prompt_marker, + require_positive_cost=True, + ) + assert obs is not None, ( + f"org scope: Langfuse never received generation for key_alias={key_alias!r}" + ) + meta_blob = _json_blob(obs.metadata) + assert team_id in meta_blob, ( + f"org-scoped team key must stamp team_id on Langfuse metadata; " + f"team_id={team_id!r} metadata={meta_blob[:400]}" + ) + _ = org_id + _assert_logs_spend( + client, + key=key, + outcome=outcome, + obs_cost=observation_spend(obs), + scope="org-team", + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index ab2835d87c4..e79d19215bf 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -23,6 +23,22 @@ class BudgetWindow(BaseModel): max_budget: float +class KeyLoggingCallbackVars(BaseModel): + langfuse_public_key: str | None = None + langfuse_secret_key: str | None = None + langfuse_host: str | None = None + + +class KeyLoggingCallback(BaseModel): + callback_name: str + callback_type: str = "success_and_failure" + callback_vars: KeyLoggingCallbackVars + + +class KeyMetadata(BaseModel): + logging: list[KeyLoggingCallback] | None = None + + class KeyGenerateBody(BaseModel): models: list[str] = [] duration: str | None = None @@ -31,6 +47,7 @@ class KeyGenerateBody(BaseModel): budget_duration: str | None = None user_id: str | None = None team_id: str | None = None + organization_id: str | None = None budget_id: str | None = None key_alias: str | None = None model_max_budget: dict[str, ModelBudgetEntry] | None = None @@ -39,6 +56,7 @@ class KeyGenerateBody(BaseModel): tpm_limit: int | None = None rpm_limit: int | None = None allowed_routes: list[str] | None = None + metadata: KeyMetadata | None = None class KeyGenerateResponse(BaseModel): @@ -105,6 +123,17 @@ class ThinkingParam(BaseModel): budget_tokens: int | None = None +class ChatToolFunction(BaseModel): + name: str + description: str | None = None + parameters: dict[str, object] | None = None + + +class ChatTool(BaseModel): + type: str = "function" + function: ChatToolFunction + + class ChatBody(BaseModel): model: str messages: list[ChatMessage] @@ -115,6 +144,9 @@ class ChatBody(BaseModel): reasoning_effort: str | None = None thinking: ThinkingParam | None = None service_tier: str | None = None + tools: list[ChatTool] | None = None + tool_choice: str | None = None + guardrails: list[str] | None = None class AnthropicMessagesBody(BaseModel): @@ -459,6 +491,7 @@ class TeamNewBody(BaseModel): team_alias: str models: list[str] = [] team_id: str | None = None + organization_id: str | None = None class TeamNewResponse(BaseModel): diff --git a/tests/llm_responses_api_testing/test_openai_responses_api.py b/tests/llm_responses_api_testing/test_openai_responses_api.py index ea8b8fa886c..bd1517dbffb 100644 --- a/tests/llm_responses_api_testing/test_openai_responses_api.py +++ b/tests/llm_responses_api_testing/test_openai_responses_api.py @@ -1628,23 +1628,27 @@ async def test_openai_responses_api_token_limit_error(): """ Relevant issue: https://github.com/BerriAI/litellm/issues/15785 - - When this fails you'll see: - "pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent" - in the console. + Parsing the in-stream ErrorEvent must not raise + "pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent". + The iterator now surfaces the event as litellm.APIError with status 400 + (invalid_request_error is a non-retriable client error, so no + MidStreamFallbackError wrapping) carrying the provider's message. """ litellm._turn_on_debug() # Generate text with >400k tokens to trigger token limit error oversized_text = "This is a test sentence. " * 50000 # ~400k tokens - # This will raise ValidationError instead of showing the real error response = await litellm.aresponses( model="gpt-5-mini", input=oversized_text, stream=True ) - async for event in response: - print(event) # Never reaches here - ValidationError is raised + with pytest.raises(litellm.APIError) as exc_info: + async for event in response: + print(event) + + assert exc_info.value.status_code == 400 + assert "exceeds the context window" in str(exc_info.value) async def test_openai_streaming_logging(): diff --git a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py index 25bf79cd575..2fb7bdfceb5 100644 --- a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py +++ b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py @@ -266,3 +266,220 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator(): ) assert out is wrapped mock_wrap.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_aresponses_fallback_on_in_stream_error_event(): + """A retriable in-stream error event (429) must trigger the router's mid-stream + fallback path: the wrapper catches MidStreamFallbackError raised by the source + iterator and yields the fallback stream instead of surfacing the error.""" + import json + from unittest.mock import Mock + + import litellm + from litellm.exceptions import MidStreamFallbackError + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig + from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator + from litellm.types.llms.openai import ErrorEvent, ErrorEventError + + router = _make_router() + + error_payload = { + "type": "error", + "error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "rate limited"}, + } + sse_bytes = f"data: {json.dumps(error_payload)}\n\n".encode() + + async def mock_aiter_bytes(): + yield sse_bytes + + mock_response = Mock() + mock_response.headers = {} + mock_response.aiter_bytes = mock_aiter_bytes + mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj) + mock_logging_obj.model_call_details = {"litellm_params": {}} + mock_logging_obj.completion_start_time = None + mock_config = Mock(spec=BaseResponsesAPIConfig) + mock_config.transform_streaming_response.return_value = ErrorEvent( + type=ResponsesAPIStreamEvents.ERROR, + sequence_number=0, + error=ErrorEventError(type="tokens", code="rate_limit_exceeded", message="rate limited"), + ) + + source = ResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-5", + responses_api_provider_config=mock_config, + logging_obj=mock_logging_obj, + custom_llm_provider="openai", + ) + + fallback_event = _make_completed_event(1, 1, 2) + + class _FallbackStream: + def __init__(self) -> None: + self._done = False + + def __aiter__(self): + return self + + async def __anext__(self): + if self._done: + raise StopAsyncIteration + self._done = True + return fallback_event + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(return_value=_FallbackStream()), + ) as mock_fallback: + wrapped = await router._aresponses_streaming_iterator( + response=source, + initial_kwargs={"model": "primary", "input": "original question"}, + ) + collected = [ev async for ev in wrapped] + + assert collected == [fallback_event] + mock_fallback.assert_awaited_once() + raised = mock_fallback.await_args.kwargs["e"] + assert isinstance(raised, MidStreamFallbackError) + assert raised.status_code == 429 + assert isinstance(raised.original_exception, litellm.APIError) + assert raised.original_exception.status_code == 429 + assert mock_fallback.await_args.kwargs["kwargs"]["input"] == "original question" + + +@pytest.mark.asyncio +async def test_aresponses_fallback_uses_continuation_input_after_partial_content(): + """When output text was already streamed before the error, the fallback re-entry + must carry a continuation input with the partial assistant text instead of + retrying the original input from scratch (which would duplicate streamed content).""" + import json + from unittest.mock import Mock + + from litellm.exceptions import MidStreamFallbackError + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig + from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator + from litellm.types.llms.openai import ErrorEvent, ErrorEventError + + router = _make_router() + + events = [ + {"type": "response.output_text.delta", "delta": "partial answer"}, + {"type": "error", "error": {"type": "server_error", "code": "internal_error", "message": "boom"}}, + ] + sse_payload = b"".join(f"data: {json.dumps(event)}\n\n".encode() for event in events) + + async def mock_aiter_bytes(): + yield sse_payload + + mock_response = Mock() + mock_response.headers = {} + mock_response.aiter_bytes = mock_aiter_bytes + mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj) + mock_logging_obj.model_call_details = {"litellm_params": {}} + mock_logging_obj.completion_start_time = None + mock_config = Mock(spec=BaseResponsesAPIConfig) + + def transform(model, parsed_chunk, logging_obj): + if parsed_chunk.get("type") == "error": + return ErrorEvent( + type=ResponsesAPIStreamEvents.ERROR, + sequence_number=0, + error=ErrorEventError(**parsed_chunk["error"]), + ) + delta_event = Mock() + delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA + delta_event.delta = parsed_chunk["delta"] + return delta_event + + mock_config.transform_streaming_response.side_effect = transform + + source = ResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-5", + responses_api_provider_config=mock_config, + logging_obj=mock_logging_obj, + custom_llm_provider="openai", + ) + + fallback_event = _make_completed_event(1, 1, 2) + + class _FallbackStream: + def __init__(self) -> None: + self._done = False + + def __aiter__(self): + return self + + async def __anext__(self): + if self._done: + raise StopAsyncIteration + self._done = True + return fallback_event + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(return_value=_FallbackStream()), + ) as mock_fallback: + wrapped = await router._aresponses_streaming_iterator( + response=source, + initial_kwargs={"model": "primary", "input": "original question"}, + ) + collected = [ev async for ev in wrapped] + + assert collected[-1] == fallback_event + raised = mock_fallback.await_args.kwargs["e"] + assert isinstance(raised, MidStreamFallbackError) + assert raised.is_pre_first_chunk is False + assert raised.generated_content == "partial answer" + continuation = mock_fallback.await_args.kwargs["kwargs"]["input"] + assert isinstance(continuation, list) + assert continuation[0]["content"][0]["text"] == "original question" + assert continuation[-2]["role"] == "developer" + assert continuation[-1]["role"] == "assistant" + assert continuation[-1]["content"][0]["text"] == "partial answer" + + +@pytest.mark.asyncio +async def test_aresponses_client_error_event_skips_fallback(): + """A 400-mapped in-stream error (raised as APIError, not MidStreamFallbackError) + must surface to the caller without invoking the router's fallback path.""" + import litellm + + router = _make_router() + + class _ClientErrorSource: + completed_response = None + + def __aiter__(self): + return self + + async def __anext__(self): + raise litellm.APIError( + status_code=400, + message="bad request", + llm_provider="openai", + model="gpt-5", + ) + + wrapped = await router._aresponses_streaming_iterator( + response=_ClientErrorSource(), + initial_kwargs={"model": "primary"}, + ) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(), + ) as mock_fallback: + with pytest.raises(litellm.APIError) as exc_info: + async for _ in wrapped: + pass + + assert exc_info.value.status_code == 400 + mock_fallback.assert_not_awaited() diff --git a/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py b/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py index d1c7a4032fb..f645379a4f4 100644 --- a/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py +++ b/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py @@ -1,3 +1,4 @@ +import asyncio from unittest.mock import AsyncMock, Mock, patch import httpx @@ -6,16 +7,20 @@ from httpx import Request, Response from litellm.integrations.datadog.datadog import DataDogLogger from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError -from litellm.types.integrations.datadog import DD_MAX_BATCH_SIZE, DatadogPayload +from litellm.types.integrations.datadog import ( + DD_MAX_BATCH_SIZE, + DD_MAX_PAYLOAD_SIZE_BYTES, + DatadogPayload, +) -def _payloads(n): +def _payloads(n, message=None): return [ DatadogPayload( ddsource="litellm", ddtags="env:test", hostname="host", - message=f'{{"event": {i}}}', + message=f"{message}{i}" if message else f'{{"event": {i}}}', service="svc", status="info", ) @@ -177,6 +182,87 @@ async def test_413_returned_response_also_splits(datadog_env): assert logger.log_queue == [] +def _make_recording_send(sent_batches, delivered): + async def _send(data): + sent_batches.append(list(data)) + delivered.extend(data) + return Response( + 202, request=Request("POST", "https://example.com"), text="Accepted" + ) + + return _send + + +@pytest.mark.asyncio +async def test_oversized_payload_splits_before_any_send(datadog_env): + """Regression for LIT-4325: a batch above Datadog's uncompressed payload limit is + split proactively, so the intake never has to reject it with a 413.""" + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + + with patch("asyncio.create_task"): + logger = DataDogLogger() + + events = _payloads(3, message="x" * 3_000_000) + logger.log_queue = list(events) + sent_batches: list = [] + delivered: list = [] + logger.async_send_compressed_data = AsyncMock( + side_effect=_make_recording_send(sent_batches, delivered) + ) + + await logger.async_send_batch() + + assert delivered == events + assert len(sent_batches) == 3 + assert all( + len(safe_dumps(batch).encode("utf-8")) <= DD_MAX_PAYLOAD_SIZE_BYTES + for batch in sent_batches + ) + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_batch_over_max_event_count_splits_before_any_send(datadog_env): + """Datadog caps a payload at 1000 events; a queue that grew past that (e.g. after + re-queues) must be sent in count-compliant chunks.""" + with patch("asyncio.create_task"): + logger = DataDogLogger() + + events = _payloads(DD_MAX_BATCH_SIZE + 1) + logger.log_queue = list(events) + sent_batches: list = [] + delivered: list = [] + logger.async_send_compressed_data = AsyncMock( + side_effect=_make_recording_send(sent_batches, delivered) + ) + + await logger.async_send_batch() + + assert delivered == events + assert all(len(batch) <= DD_MAX_BATCH_SIZE for batch in sent_batches) + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_single_event_above_payload_cap_is_still_sent(datadog_env): + """A lone event over the byte cap cannot be split further; it must be sent once + (Datadog decides), never looped on.""" + with patch("asyncio.create_task"): + logger = DataDogLogger() + + logger.log_queue = _payloads(1, message="x" * (DD_MAX_PAYLOAD_SIZE_BYTES + 1)) + sent_batches: list = [] + delivered: list = [] + send = AsyncMock(side_effect=_make_recording_send(sent_batches, delivered)) + logger.async_send_compressed_data = send + + await asyncio.wait_for(logger.async_send_batch(), timeout=10) + + assert send.await_count == 1 + assert len(delivered) == 1 + assert logger.log_queue == [] + + @pytest.mark.asyncio async def test_partial_delivery_then_transient_error_requeues_only_undelivered( datadog_env, diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metric_guard.py b/tests/test_litellm/integrations/test_prometheus_budget_metric_guard.py new file mode 100644 index 00000000000..ef844c80d4d --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_budget_metric_guard.py @@ -0,0 +1,359 @@ +""" +Unit tests for the NoOpMetric guard in _increment_remaining_budget_metrics +and the per-entity guards in _set_*_budget_metrics_after_api_request. + +Regression tests that the specific bug can never happen again: +when budget gauges are excluded from prometheus_metrics_config (and therefore +created as NoOpMetric instances), the DB/cache lookup helpers must not be called. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from prometheus_client import REGISTRY + +import litellm +from litellm.integrations.prometheus import PrometheusLogger +from litellm.types.integrations.prometheus import NoOpMetric + +_BUDGET_EXCLUDED_CONFIG = [ + { + "group": "core-only", + "metrics": [ + "litellm_requests_metric", + "litellm_total_tokens_metric", + ], + } +] + + +@pytest.fixture(autouse=True) +def cleanup_prometheus_registry(): + old_config = litellm.prometheus_metrics_config + for collector in list(REGISTRY._collector_to_names.keys()): + try: + REGISTRY.unregister(collector) + except Exception: + pass + yield + litellm.prometheus_metrics_config = old_config + for collector in list(REGISTRY._collector_to_names.keys()): + try: + REGISTRY.unregister(collector) + except Exception: + pass + + +def make_logger_with_budget_metrics_disabled() -> PrometheusLogger: + litellm.prometheus_metrics_config = _BUDGET_EXCLUDED_CONFIG + return PrometheusLogger() + + +def make_logger_with_all_metrics_enabled() -> PrometheusLogger: + litellm.prometheus_metrics_config = None + return PrometheusLogger() + + +COMMON_KWARGS = dict( + user_api_team="team-123", + user_api_team_alias="my-team", + user_api_key="hashed-key", + user_api_key_alias="my-key", + litellm_params={"metadata": {}}, + response_cost=0.001, + user_id="user-1", + user_api_key_org_id="org-1", +) + + +class TestBudgetGaugesAreNoopWhenExcluded: + def test_team_gauge_is_noop(self): + logger = make_logger_with_budget_metrics_disabled() + assert isinstance(logger.litellm_remaining_team_budget_metric, NoOpMetric) + + def test_api_key_gauge_is_noop(self): + logger = make_logger_with_budget_metrics_disabled() + assert isinstance(logger.litellm_remaining_api_key_budget_metric, NoOpMetric) + + def test_user_gauge_is_noop(self): + logger = make_logger_with_budget_metrics_disabled() + assert isinstance(logger.litellm_remaining_user_budget_metric, NoOpMetric) + + def test_org_gauge_is_noop(self): + logger = make_logger_with_budget_metrics_disabled() + assert isinstance(logger.litellm_remaining_org_budget_metric, NoOpMetric) + + def test_gauges_are_real_when_all_metrics_enabled(self): + logger = make_logger_with_all_metrics_enabled() + assert not isinstance(logger.litellm_remaining_team_budget_metric, NoOpMetric) + assert not isinstance(logger.litellm_remaining_api_key_budget_metric, NoOpMetric) + assert not isinstance(logger.litellm_remaining_user_budget_metric, NoOpMetric) + assert not isinstance(logger.litellm_remaining_org_budget_metric, NoOpMetric) + + +class TestTopLevelGuard: + @pytest.mark.asyncio + async def test_no_db_lookups_when_all_budget_gauges_are_noop(self): + """Regression: _increment_remaining_budget_metrics must return early + without any I/O when all four budget gauges are NoOpMetric.""" + logger = make_logger_with_budget_metrics_disabled() + + assemble_team = AsyncMock(return_value=MagicMock()) + assemble_key = AsyncMock(return_value=MagicMock()) + assemble_user = AsyncMock(return_value=MagicMock()) + + with ( + patch.object(logger, "_assemble_team_object", assemble_team), + patch.object(logger, "_assemble_key_object", assemble_key), + patch.object(logger, "_assemble_user_object", assemble_user), + ): + await logger._increment_remaining_budget_metrics(**COMMON_KWARGS) + + assemble_team.assert_not_called() + assemble_key.assert_not_called() + assemble_user.assert_not_called() + + @pytest.mark.asyncio + async def test_db_lookups_run_when_budget_gauges_are_real(self): + """When budget gauges are real Prometheus metrics, the assemble helpers + must be called so I/O proceeds normally.""" + logger = make_logger_with_all_metrics_enabled() + + assemble_team = AsyncMock( + return_value=MagicMock( + team_id="team-123", + team_alias="my-team", + spend=0.001, + max_budget=None, + budget_reset_at=None, + ) + ) + assemble_key = AsyncMock( + return_value=MagicMock( + token="hashed-key", + key_alias="my-key", + spend=0.001, + max_budget=None, + budget_reset_at=None, + ) + ) + assemble_user = AsyncMock( + return_value=MagicMock( + user_id="user-1", + spend=0.001, + max_budget=None, + budget_reset_at=None, + user_email=None, + user_alias=None, + ) + ) + + with ( + patch.object(logger, "_assemble_team_object", assemble_team), + patch.object(logger, "_assemble_key_object", assemble_key), + patch.object(logger, "_assemble_user_object", assemble_user), + patch.object(logger, "_set_team_budget_metrics", MagicMock()), + patch.object(logger, "_set_key_budget_metrics", MagicMock()), + patch.object(logger, "_set_user_budget_metrics", MagicMock()), + patch.object(logger, "_set_org_budget_metrics_after_api_request", AsyncMock()), + ): + await logger._increment_remaining_budget_metrics(**COMMON_KWARGS) + + assemble_team.assert_called_once() + assemble_key.assert_called_once() + assemble_user.assert_called_once() + + +class TestPerEntityGuards: + @pytest.mark.asyncio + async def test_team_guard_skips_lookup_when_team_gauge_is_noop(self): + """Per-entity guard: team assemble helper is not called when team gauge is NoOp, + even when key and user gauges are real.""" + logger = make_logger_with_all_metrics_enabled() + logger.litellm_remaining_team_budget_metric = NoOpMetric() + + assemble_team = AsyncMock(return_value=MagicMock()) + assemble_key = AsyncMock( + return_value=MagicMock( + token="hashed-key", + key_alias="my-key", + spend=0.001, + max_budget=None, + budget_reset_at=None, + ) + ) + assemble_user = AsyncMock( + return_value=MagicMock( + user_id="user-1", + spend=0.001, + max_budget=None, + budget_reset_at=None, + user_email=None, + user_alias=None, + ) + ) + + with ( + patch.object(logger, "_assemble_team_object", assemble_team), + patch.object(logger, "_assemble_key_object", assemble_key), + patch.object(logger, "_assemble_user_object", assemble_user), + patch.object(logger, "_set_team_budget_metrics", MagicMock()), + patch.object(logger, "_set_key_budget_metrics", MagicMock()), + patch.object(logger, "_set_user_budget_metrics", MagicMock()), + patch.object(logger, "_set_org_budget_metrics_after_api_request", AsyncMock()), + ): + await logger._increment_remaining_budget_metrics(**COMMON_KWARGS) + + assemble_team.assert_not_called() + assemble_key.assert_called_once() + assemble_user.assert_called_once() + + @pytest.mark.asyncio + async def test_key_guard_skips_lookup_when_key_gauge_is_noop(self): + """Per-entity guard: key assemble helper is not called when key gauge is NoOp, + even when team and user gauges are real.""" + logger = make_logger_with_all_metrics_enabled() + logger.litellm_remaining_api_key_budget_metric = NoOpMetric() + + assemble_team = AsyncMock( + return_value=MagicMock( + team_id="team-123", + team_alias="my-team", + spend=0.001, + max_budget=None, + budget_reset_at=None, + ) + ) + assemble_key = AsyncMock(return_value=MagicMock()) + assemble_user = AsyncMock( + return_value=MagicMock( + user_id="user-1", + spend=0.001, + max_budget=None, + budget_reset_at=None, + user_email=None, + user_alias=None, + ) + ) + + with ( + patch.object(logger, "_assemble_team_object", assemble_team), + patch.object(logger, "_assemble_key_object", assemble_key), + patch.object(logger, "_assemble_user_object", assemble_user), + patch.object(logger, "_set_team_budget_metrics", MagicMock()), + patch.object(logger, "_set_key_budget_metrics", MagicMock()), + patch.object(logger, "_set_user_budget_metrics", MagicMock()), + patch.object(logger, "_set_org_budget_metrics_after_api_request", AsyncMock()), + ): + await logger._increment_remaining_budget_metrics(**COMMON_KWARGS) + + assemble_key.assert_not_called() + assemble_team.assert_called_once() + assemble_user.assert_called_once() + + @pytest.mark.asyncio + async def test_user_guard_skips_lookup_when_user_gauge_is_noop(self): + """Per-entity guard: user assemble helper is not called when user gauge is NoOp, + even when team and key gauges are real.""" + logger = make_logger_with_all_metrics_enabled() + logger.litellm_remaining_user_budget_metric = NoOpMetric() + + assemble_team = AsyncMock( + return_value=MagicMock( + team_id="team-123", + team_alias="my-team", + spend=0.001, + max_budget=None, + budget_reset_at=None, + ) + ) + assemble_key = AsyncMock( + return_value=MagicMock( + token="hashed-key", + key_alias="my-key", + spend=0.001, + max_budget=None, + budget_reset_at=None, + ) + ) + assemble_user = AsyncMock(return_value=MagicMock()) + + with ( + patch.object(logger, "_assemble_team_object", assemble_team), + patch.object(logger, "_assemble_key_object", assemble_key), + patch.object(logger, "_assemble_user_object", assemble_user), + patch.object(logger, "_set_team_budget_metrics", MagicMock()), + patch.object(logger, "_set_key_budget_metrics", MagicMock()), + patch.object(logger, "_set_user_budget_metrics", MagicMock()), + patch.object(logger, "_set_org_budget_metrics_after_api_request", AsyncMock()), + ): + await logger._increment_remaining_budget_metrics(**COMMON_KWARGS) + + assemble_user.assert_not_called() + assemble_team.assert_called_once() + assemble_key.assert_called_once() + + @pytest.mark.asyncio + async def test_set_team_budget_metrics_directly_skips_when_gauge_is_noop(self): + """_set_team_budget_metrics_after_api_request returns early when team gauge is NoOp.""" + logger = make_logger_with_budget_metrics_disabled() + assemble_team = AsyncMock(return_value=MagicMock()) + + with patch.object(logger, "_assemble_team_object", assemble_team): + await logger._set_team_budget_metrics_after_api_request( + user_api_team="team-123", + user_api_team_alias="my-team", + team_spend=0.5, + team_max_budget=10.0, + response_cost=0.001, + ) + + assemble_team.assert_not_called() + + @pytest.mark.asyncio + async def test_set_api_key_budget_metrics_directly_skips_when_gauge_is_noop(self): + """_set_api_key_budget_metrics_after_api_request returns early when key gauge is NoOp.""" + logger = make_logger_with_budget_metrics_disabled() + assemble_key = AsyncMock(return_value=MagicMock()) + + with patch.object(logger, "_assemble_key_object", assemble_key): + await logger._set_api_key_budget_metrics_after_api_request( + user_api_key="hashed-key", + user_api_key_alias="my-key", + response_cost=0.001, + key_max_budget=10.0, + key_spend=0.5, + ) + + assemble_key.assert_not_called() + + @pytest.mark.asyncio + async def test_set_user_budget_metrics_directly_skips_when_gauge_is_noop(self): + """_set_user_budget_metrics_after_api_request returns early when user gauge is NoOp.""" + logger = make_logger_with_budget_metrics_disabled() + assemble_user = AsyncMock(return_value=MagicMock()) + + with patch.object(logger, "_assemble_user_object", assemble_user): + await logger._set_user_budget_metrics_after_api_request( + user_id="user-1", + user_spend=0.5, + user_max_budget=10.0, + response_cost=0.001, + ) + + assemble_user.assert_not_called() + + @pytest.mark.asyncio + async def test_set_org_budget_metrics_directly_skips_when_gauge_is_noop(self): + """_set_org_budget_metrics_after_api_request returns early when org gauge is NoOp. + The guard fires before any import of auth_checks, so prisma_client is never touched.""" + logger = make_logger_with_budget_metrics_disabled() + + set_org_metrics = MagicMock() + with patch.object(logger, "_set_org_budget_metrics", set_org_metrics): + await logger._set_org_budget_metrics_after_api_request( + org_id="org-1", + response_cost=0.001, + ) + + set_org_metrics.assert_not_called() diff --git a/tests/test_litellm/litellm_core_utils/test_fallback_generalizations.py b/tests/test_litellm/litellm_core_utils/test_fallback_generalizations.py index 410014958b6..cf74bed9c15 100644 --- a/tests/test_litellm/litellm_core_utils/test_fallback_generalizations.py +++ b/tests/test_litellm/litellm_core_utils/test_fallback_generalizations.py @@ -16,6 +16,7 @@ sys.path.insert(0, os.path.abspath("../../..")) import litellm from litellm.litellm_core_utils.fallback_generalizations import ( get_fallback_generalization_rules, + match_all_fallback_generalizations, match_fallback_generalization, set_fallback_generalizations, ) @@ -57,6 +58,48 @@ def test_match_returns_model_info_of_first_matching_rule(restore_generalizations assert matched["tag"] == "first" +def test_match_all_returns_every_matching_rule_in_order(restore_generalizations): + restore_generalizations( + [ + { + "name": "first", + "pattern": r"^acme-", + "model_info": {"litellm_provider": "openai", "tag": "first"}, + }, + { + "name": "second", + "pattern": r"^acme-pro-", + "model_info": {"litellm_provider": "anthropic", "tag": "second"}, + }, + ] + ) + assert [m["tag"] for m in match_all_fallback_generalizations("acme-pro-1")] == ["first", "second"] + assert match_all_fallback_generalizations("gpt-4o") == [] + + +def test_provider_scoped_rule_is_skipped_for_other_providers(restore_generalizations): + """Model-info resolution must fall through a provider-mismatched earlier rule to a + later applicable one, instead of discarding the model name at the first pattern hit.""" + restore_generalizations( + [ + { + "name": "bedrock-scoped", + "pattern": r"^acme-", + "model_info": {"litellm_provider": "bedrock", "supports_vision": False}, + }, + { + "name": "openai-scoped", + "pattern": r"^acme-", + "model_info": {"litellm_provider": "openai", "mode": "chat", "supports_vision": True}, + }, + ] + ) + litellm.get_model_info.cache_clear() + info = litellm.get_model_info("acme-fast-1", custom_llm_provider="openai") + assert info["litellm_provider"] == "openai" + assert info["supports_vision"] is True + + def test_match_is_case_insensitive(restore_generalizations): restore_generalizations( [{"name": "r", "pattern": r"^claude-opus", "model_info": {"ok": True}}] @@ -298,3 +341,47 @@ def test_shipped_adaptive_rule_gates_on_version_not_pricing(shipped_cost_map): assert non_adaptive not in litellm.model_cost assert AnthropicModelInfo._is_adaptive_thinking_model(adaptive) is True assert AnthropicModelInfo._is_adaptive_thinking_model(non_adaptive) is False + + +def test_shipped_bedrock_rule_resolves_unmapped_future_claude_for_bedrock(shipped_cost_map): + """An unmapped Bedrock Claude >= 4.8 resolves via the bedrock-scoped + ``bedrock-anthropic-claude-mid-conversation-system`` rule even when the lookup + carries ``custom_llm_provider="bedrock"``, which the provider check uses to drop + the anthropic-scoped rules. It inherits base capabilities, gains both + version-gated flags, and stays unpriced.""" + model = "us.anthropic.claude-opus-4-9" + assert model not in litellm.model_cost + info = litellm.get_model_info(model, custom_llm_provider="bedrock") + assert info["litellm_provider"] == "bedrock" + assert info["supports_mid_conversation_system"] is True + assert info["supports_adaptive_thinking"] is True + assert info["supports_function_calling"] is True + assert not info.get("input_cost_per_token") + + +def test_shipped_bedrock_mid_conversation_rule_gates_on_version_and_naming(shipped_cost_map): + """The bedrock rule only claims Bedrock-style ids at 4.8+, bare 5+ majors and + new families included; pre-4.8 Bedrock ids and native ids never gain the flag, + and the rule outranks the anthropic-scoped ones for Bedrock ids because it is + listed first.""" + for flagged in ( + "us.anthropic.claude-opus-4-8", + "jp.anthropic.claude-opus-4-8", + "anthropic.claude-sonnet-5", + "us.anthropic.claude-fable-5", + "anthropic.claude-sonnet-5-20260101-v1:0", + ): + matched = match_fallback_generalization(flagged) + assert matched is not None, flagged + assert matched["litellm_provider"] == "bedrock", flagged + assert matched["supports_mid_conversation_system"] is True, flagged + for unflagged in ( + "us.anthropic.claude-opus-4-7", + "us.anthropic.claude-sonnet-4-6", + "us.anthropic.claude-haiku-4-5-20251001-v1:0", + "anthropic.claude-3-5-sonnet-20240620-v1:0", + "claude-opus-4-9", + "claude-sonnet-5", + ): + matched = match_fallback_generalization(unflagged) + assert matched is None or not matched.get("supports_mid_conversation_system"), unflagged diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py new file mode 100644 index 00000000000..06d3effcfbb --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py @@ -0,0 +1,189 @@ +import pytest + +from litellm.constants import ( + DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, + DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, + DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET, +) +from litellm.llms.anthropic.common_utils import AnthropicError +from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, +) + + +def _claude_code_payload(effort="medium", max_tokens=8192, **output_config_extra): + """The exact adaptive-thinking shape Claude Code (claude-cli) sends.""" + output_config = {"effort": effort, **output_config_extra} + return { + "max_tokens": max_tokens, + "thinking": {"type": "adaptive"}, + "output_config": output_config, + } + + +def _transform(model, params, litellm_params=None): + return AnthropicMessagesConfig().transform_anthropic_messages_request( + model=model, + messages=[{"role": "user", "content": "Hello"}], + anthropic_messages_optional_request_params=dict(params), + litellm_params=litellm_params or {}, + headers={}, + ) + + +def test_effort_translated_to_legacy_thinking_for_haiku_4_5(): + """Core regression: Claude Code sends adaptive thinking + effort to Haiku 4.5 + (thinking-capable, pre-4.6). Effort must be translated to legacy extended + thinking rather than forwarded raw (which Anthropic rejects with "This model + does not support the effort parameter").""" + result = _transform("claude-haiku-4-5", _claude_code_payload(effort="medium")) + + assert result["thinking"] == { + "type": "enabled", + "budget_tokens": DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, + } + assert "output_config" not in result + + +def test_effort_high_maps_to_high_budget_for_sonnet_4_5(): + result = _transform("claude-sonnet-4-5", _claude_code_payload(effort="high")) + + assert result["thinking"] == { + "type": "enabled", + "budget_tokens": DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, + } + assert "output_config" not in result + + +def test_adaptive_effort_passes_through_untouched_for_4_6(): + """4.6+ natively supports the adaptive interface, so it must not be rewritten.""" + result = _transform("claude-sonnet-4-6", _claude_code_payload(effort="high")) + + assert result["thinking"] == {"type": "adaptive"} + assert result["output_config"] == {"effort": "high"} + + +def test_thinking_and_effort_dropped_for_non_reasoning_model(): + """A model with no reasoning support cannot take thinking or effort, so both are + silently dropped (no drop_params required) so the request still succeeds.""" + result = _transform("claude-3-5-haiku-latest", _claude_code_payload(effort="medium")) + + assert "thinking" not in result + assert "output_config" not in result + + +def test_residual_output_config_preserved_after_effort_translation(): + """output_config may carry `format` (structured outputs) alongside effort. Only + the consumed effort key is removed; the residual is left for provider subclasses + (bedrock/vertex) to handle, and effort is translated to legacy thinking.""" + result = _transform( + "claude-haiku-4-5", + _claude_code_payload(effort="medium", format={"type": "json_schema"}), + ) + + assert result["thinking"] == { + "type": "enabled", + "budget_tokens": DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, + } + assert result["output_config"] == {"format": {"type": "json_schema"}} + + +def test_opus_4_5_keeps_effort_but_drops_adaptive_thinking(): + """Regression: Opus 4.5 advertises supports_output_config (accepts + output_config.effort) but is NOT adaptive, so thinking:{type:adaptive} is + rejected by Anthropic. The effort must be kept and only the adaptive thinking + block dropped, rather than early-returning and forwarding adaptive thinking raw.""" + result = _transform("claude-opus-4-5", _claude_code_payload(effort="medium")) + + assert result["output_config"] == {"effort": "medium"} + assert "thinking" not in result + + +def test_opus_4_5_preserves_native_effort_without_adaptive_thinking(): + """A caller sending output_config.effort alone (no adaptive thinking) to Opus 4.5 + must pass through untouched, since the model supports it natively.""" + result = AnthropicMessagesConfig().transform_anthropic_messages_request( + model="claude-opus-4-5", + messages=[{"role": "user", "content": "Hello"}], + anthropic_messages_optional_request_params={ + "max_tokens": 8192, + "output_config": {"effort": "high"}, + }, + litellm_params={}, + headers={}, + ) + + assert result["output_config"] == {"effort": "high"} + assert "thinking" not in result + + +def test_opus_4_5_unsupported_effort_level_translated_to_legacy_thinking(): + """Opus 4.5 accepts output_config.effort but only levels low/medium/high; + Claude Code defaults to xhigh on newer models, and forwarding that level raw + would be rejected with "effort='xhigh' is not supported by this model". An + unsupported level must fall through to the legacy translation (budget-based + thinking, effort stripped) instead of being preserved.""" + result = _transform("claude-opus-4-5", _claude_code_payload(effort="xhigh", max_tokens=64000)) + + assert result["thinking"] == { + "type": "enabled", + "budget_tokens": DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET, + } + assert "output_config" not in result + + +def test_opus_4_5_effort_only_unsupported_level_left_for_provider_normalization(): + """An effort-only request (no adaptive thinking) must pass through untouched even + when the level exceeds what the model supports: provider subclasses own their + level normalization (bedrock clamps xhigh to the model's ceiling after this base + transform runs), so consuming the effort here breaks that contract.""" + result = _transform( + "claude-opus-4-5", + {"max_tokens": 4096, "output_config": {"effort": "xhigh"}}, + ) + + assert result["output_config"] == {"effort": "xhigh"} + assert "thinking" not in result + + +def test_budget_capped_below_max_tokens(): + """Adaptive thinking carries no budget, so the translated legacy budget must be + capped below max_tokens (Anthropic requires max_tokens > budget_tokens). A + high-effort budget (4096) with max_tokens=3000 must be capped to 2999.""" + result = _transform("claude-haiku-4-5", _claude_code_payload(effort="high", max_tokens=3000)) + + assert result["thinking"] == {"type": "enabled", "budget_tokens": 2999} + + +def test_thinking_dropped_when_max_tokens_too_small_for_min_budget(): + """When max_tokens can't fit even the minimum thinking budget, thinking is + silently dropped so the request still succeeds rather than being rejected.""" + result = _transform("claude-haiku-4-5", _claude_code_payload(effort="medium", max_tokens=512)) + + assert "thinking" not in result + assert "output_config" not in result + + +def test_unrecognized_effort_raises_clean_400(): + """An unrecognized effort value (e.g. a future Anthropic tier) must surface as a + clean AnthropicError 400, matching _translate_reasoning_effort_to_anthropic, + rather than leaking litellm's internal BadRequestError.""" + with pytest.raises(AnthropicError) as exc_info: + _transform("claude-haiku-4-5", _claude_code_payload(effort="turbo")) + + assert exc_info.value.status_code == 400 + + +def test_non_adaptive_request_without_effort_is_untouched(): + """A non-adaptive model receiving a request with no adaptive interface (no + effort, no adaptive thinking) must pass through untouched.""" + result = AnthropicMessagesConfig().transform_anthropic_messages_request( + model="claude-haiku-4-5", + messages=[{"role": "user", "content": "Hello"}], + anthropic_messages_optional_request_params={"max_tokens": 1024}, + litellm_params={}, + headers={}, + ) + + assert "thinking" not in result + assert "output_config" not in result diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index a2fbb68bbb9..f7a234c8ae8 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -772,7 +772,6 @@ class TestProxyOAuthHeaderForwarding: self, ): """OAuth Authorization header IS forwarded when x-litellm-api-key was used for proxy auth.""" - from unittest.mock import patch from starlette.datastructures import Headers @@ -1724,3 +1723,48 @@ class TestClaudeOpus48AdaptiveThinking: from litellm.llms.anthropic.common_utils import AnthropicModelInfo assert AnthropicModelInfo._is_adaptive_thinking_model(model) is False + + +class TestDefaultSuffixAdaptiveThinking: + """@default-suffixed Vertex AI model names (e.g. vertex_ai/claude-opus-4-8@default) + must resolve as adaptive thinking. Before the fix, _model_map_lookup_candidates + never stripped the @default suffix, so the lookup fell through to the bare + model name without @default, which may or may not have the flag, and for + provider-prefixed forms the lookup always missed (issue #31760).""" + + @pytest.mark.parametrize( + "model", + [ + "vertex_ai/claude-opus-4-8@default", + "vertex_ai/claude-sonnet-4-6@default", + "vertex_ai/claude-opus-4-7@default", + "vertex_ai/claude-opus-4-6@default", + "vertex_ai/claude-fable-5@default", + ], + ) + def test_default_suffix_models_are_adaptive_thinking( + self, local_model_cost_map, model: str + ) -> None: + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + assert AnthropicModelInfo._is_adaptive_thinking_model(model) is True, ( + f"{model} not classified as adaptive thinking. " + "Check _model_map_lookup_candidates strips @default suffix." + ) + + @pytest.mark.parametrize( + "model,expected_bare", + [ + ("vertex_ai/claude-opus-4-8@default", "claude-opus-4-8"), + ("vertex_ai/claude-sonnet-4-6@default", "claude-sonnet-4-6"), + ], + ) + def test_lookup_candidates_include_bare_name( + self, model: str, expected_bare: str + ) -> None: + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + candidates = AnthropicModelInfo._model_map_lookup_candidates(model) + assert expected_bare in candidates, ( + f"Expected '{expected_bare}' in candidates for '{model}', got: {candidates}" + ) diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index ffe45cd5e79..7d104ff1f2b 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -1974,14 +1974,24 @@ def test_bedrock_invoke_transform_merges_list_content_system_role_into_system(): ] -def test_bedrock_invoke_transform_keeps_mid_conversation_system_role_in_place(): +@pytest.mark.parametrize( + "model", + [ + "anthropic.claude-opus-4-8", + "jp.anthropic.claude-opus-4-8", + "us.anthropic.claude-sonnet-5", + "us.anthropic.claude-fable-5", + ], +) +def test_bedrock_invoke_transform_keeps_mid_conversation_system_role_in_place(local_model_cost_map, model): """Regression test for the Bedrock prompt-cache collapse: hoisting a mid-conversation ``role: "system"`` message (e.g. Claude Code's ``mid-conversation-system-2026-04-07`` reminders) into the top-level ``system`` field mutates the cache prefix and invalidates the cached message - history, so such entries must be forwarded in place. Invoke only rejects a - system entry at ``messages.0``. Billing-header blocks must still be stripped - from the top-level ``system`` field even when nothing is hoisted.""" + history, so on models flagged ``supports_mid_conversation_system`` (Claude + 4.8+, which Invoke accepts the role on) such entries must be forwarded + in place. Billing-header blocks must still be stripped from the top-level + ``system`` field even when nothing is hoisted.""" from litellm.types.router import GenericLiteLLMParams cfg = AmazonAnthropicClaudeMessagesConfig() @@ -1993,7 +2003,7 @@ def test_bedrock_invoke_transform_keeps_mid_conversation_system_role_in_place(): ] result = cfg.transform_anthropic_messages_request( - model="anthropic.claude-opus-4-8", + model=model, messages=copy.deepcopy(messages), anthropic_messages_optional_request_params={ "max_tokens": 256, @@ -2013,10 +2023,11 @@ def test_bedrock_invoke_transform_keeps_mid_conversation_system_role_in_place(): ] -def test_bedrock_invoke_transform_hoists_only_leading_system_run(): - """Only the leading run of ``role: "system"`` messages is hoisted into the - top-level ``system`` field; a later system entry keeps its position in - ``messages`` so the serialized prefix stays stable across turns.""" +def test_bedrock_invoke_transform_hoists_only_leading_system_run(local_model_cost_map): + """On models flagged ``supports_mid_conversation_system``, only the leading + run of ``role: "system"`` messages is hoisted into the top-level ``system`` + field; a later system entry keeps its position in ``messages`` so the + serialized prefix stays stable across turns.""" from litellm.types.router import GenericLiteLLMParams cfg = AmazonAnthropicClaudeMessagesConfig() @@ -2047,6 +2058,133 @@ def test_bedrock_invoke_transform_hoists_only_leading_system_run(): ] +def test_bedrock_invoke_transform_hoists_mid_conversation_system_for_older_claude(local_model_cost_map): + """Regression test for Claude Code 400s on pre-Opus-4.8 Bedrock models: + Invoke rejects ``role: "system"`` in every position on Opus 4.7, Sonnet 4.6, + Haiku 4.5, etc. ("role 'system' is not supported on this model"), so on + models without ``supports_mid_conversation_system`` every system entry must + be hoisted into the top-level ``system`` field, mid-conversation ones + included.""" + from litellm.types.router import GenericLiteLLMParams + + cfg = AmazonAnthropicClaudeMessagesConfig() + messages = [ + {"role": "user", "content": "read the file"}, + {"role": "system", "content": "[Truncated: PARTIAL view of big1.txt]"}, + {"role": "assistant", "content": "reading"}, + {"role": "user", "content": "continue"}, + ] + + result = cfg.transform_anthropic_messages_request( + model="us.anthropic.claude-opus-4-7", + messages=copy.deepcopy(messages), + anthropic_messages_optional_request_params={ + "max_tokens": 256, + "stream": False, + "system": [{"type": "text", "text": "Base."}], + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert result["messages"] == [ + {"role": "user", "content": "read the file"}, + {"role": "assistant", "content": "reading"}, + {"role": "user", "content": "continue"}, + ] + assert result["system"] == [ + {"type": "text", "text": "Base."}, + {"type": "text", "text": "[Truncated: PARTIAL view of big1.txt]"}, + ] + + +def test_bedrock_invoke_transform_hoists_all_system_for_unmapped_model(local_model_cost_map): + """A model with no cost-map entry and no fallback-generalization rule gets + the hoist-everything behavior: the safe default is a mutated cache prefix, + never a provider 400 from forwarding a role the model may not accept.""" + from litellm.types.router import GenericLiteLLMParams + + cfg = AmazonAnthropicClaudeMessagesConfig() + messages = [ + {"role": "user", "content": "hi"}, + {"role": "system", "content": "mid-conversation reminder"}, + {"role": "assistant", "content": "hello"}, + {"role": "user", "content": "continue"}, + ] + + result = cfg.transform_anthropic_messages_request( + model="us.anthropic.claude-opus-3-9", + messages=copy.deepcopy(messages), + anthropic_messages_optional_request_params={"max_tokens": 256, "stream": False}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert result["messages"] == [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "hello"}, + {"role": "user", "content": "continue"}, + ] + assert result["system"] == [{"type": "text", "text": "mid-conversation reminder"}] + + +def test_bedrock_invoke_transform_keeps_system_in_place_for_unmapped_future_claude(local_model_cost_map): + """An unmapped Bedrock Claude at 4.8 or higher resolves through the + ``bedrock-anthropic-claude-mid-conversation-system`` fallback rule, so a + future model that has not landed in the cost map yet keeps the + cache-preserving in-place behavior instead of falling back to hoist-all.""" + from litellm.types.router import GenericLiteLLMParams + + cfg = AmazonAnthropicClaudeMessagesConfig() + messages = [ + {"role": "user", "content": "hi"}, + {"role": "system", "content": "mid-conversation reminder"}, + {"role": "assistant", "content": "hello"}, + {"role": "user", "content": "continue"}, + ] + + result = cfg.transform_anthropic_messages_request( + model="us.anthropic.claude-opus-4-9", + messages=copy.deepcopy(messages), + anthropic_messages_optional_request_params={"max_tokens": 256, "stream": False}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert result["messages"] == messages + assert "system" not in result + + +def test_bedrock_claude_4_8_plus_cost_map_entries_carry_mid_conversation_system_flag(): + """Exact cost-map hits resolve before fallback-generalization rules, so a + mapped Bedrock Claude 4.8+ entry without ``supports_mid_conversation_system`` + silently loses the cache-preserving in-place handling that the + ``bedrock-anthropic-claude-mid-conversation-system`` rule grants unmapped + ids. Every mapped entry the rule's own pattern matches must carry the flag + explicitly.""" + import re + + import litellm + + cost_map_path = os.path.join(os.path.dirname(litellm.__file__), "model_prices_and_context_window_backup.json") + with open(cost_map_path) as f: + cost_map = json.load(f) + rules = cost_map["fallback_generalizations"]["rules"] + pattern = re.compile( + next(r["pattern"] for r in rules if r["name"] == "bedrock-anthropic-claude-mid-conversation-system"), + re.IGNORECASE, + ) + missing = [ + key + for key, info in cost_map.items() + if isinstance(info, dict) + and str(info.get("litellm_provider", "")).startswith("bedrock") + and pattern.search(key) + and info.get("supports_mid_conversation_system") is not True + ] + assert missing == [] + + def test_as_system_content_blocks_handles_each_shape(): """``_as_system_content_blocks`` normalizes every system shape: ``None`` -> empty, a string -> a single text block, a list -> a shallow copy, and any other value diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 9cf60ea5f91..d8e3f495ced 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -1959,6 +1959,110 @@ class TestUsageTransformation: assert response_usage.output_tokens_details.text_tokens == 50 assert response_usage.output_tokens_details.image_tokens == 100 + def test_reasoning_tokens_not_forced_to_zero_when_absent(self): + # Regression: previously the else branch wrote reasoning_tokens=0 even when + # completion_tokens_details had no reasoning (reasoning_tokens=None). That caused + # the proxy to always report reasoning_tokens=0 for non-thinking responses. + usage = Usage( + prompt_tokens=10, + completion_tokens=50, + total_tokens=60, + completion_tokens_details=CompletionTokensDetailsWrapper( + text_tokens=50, + # reasoning_tokens intentionally absent -> None + ), + ) + + chat_completion_response = ModelResponse( + id="test-response-id", + created=1234567890, + model="claude-haiku-4-5", + object="chat.completion", + usage=usage, + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message(content="Hello!", role="assistant"), + ) + ], + ) + + response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + chat_completion_response=chat_completion_response + ) + + assert response_usage.output_tokens_details is not None + assert response_usage.output_tokens_details.reasoning_tokens is None + + def test_reasoning_tokens_preserved_when_thinking_occurred(self): + # Regression: reasoning_tokens must survive the chat->responses translation + # when the provider actually did thinking. + usage = Usage( + prompt_tokens=100, + completion_tokens=612, + total_tokens=712, + completion_tokens_details=CompletionTokensDetailsWrapper( + reasoning_tokens=512, + text_tokens=100, + ), + ) + + chat_completion_response = ModelResponse( + id="test-response-id", + created=1234567890, + model="claude-haiku-4-5", + object="chat.completion", + usage=usage, + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message(content="Hello!", role="assistant"), + ) + ], + ) + + response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + chat_completion_response=chat_completion_response + ) + + assert response_usage.output_tokens_details is not None + assert response_usage.output_tokens_details.reasoning_tokens == 512 + + def test_reasoning_tokens_explicit_zero_preserved(self): + usage = Usage( + prompt_tokens=10, + completion_tokens=50, + total_tokens=60, + completion_tokens_details=CompletionTokensDetailsWrapper( + reasoning_tokens=0, + text_tokens=50, + ), + ) + + chat_completion_response = ModelResponse( + id="test-response-id", + created=1234567890, + model="gpt-5.6", + object="chat.completion", + usage=usage, + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message(content="Hello!", role="assistant"), + ) + ], + ) + + response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + chat_completion_response=chat_completion_response + ) + + assert response_usage.output_tokens_details is not None + assert response_usage.output_tokens_details.reasoning_tokens == 0 + class TestStreamingIDConsistency: """Test cases for consistent IDs across streaming events (issue #14962)""" diff --git a/tests/test_litellm/responses/test_streaming_iterator_error_events.py b/tests/test_litellm/responses/test_streaming_iterator_error_events.py new file mode 100644 index 00000000000..1a2dcd0fcb7 --- /dev/null +++ b/tests/test_litellm/responses/test_streaming_iterator_error_events.py @@ -0,0 +1,357 @@ +""" +Regression: in-stream error events (type="error", type="response.failed") must +raise instead of being returned as benign chunks, mirroring chat streaming +semantics (_handle_stream_fallback_error): non-retriable 4xx (except 429) +raise litellm.APIError directly; 429 and 5xx are wrapped in +MidStreamFallbackError so the Router's mid-stream fallback machinery fires. + +Status mapping must consider both the OpenAI error `type` (e.g. +"invalid_request_error") and `code` (e.g. "invalid_prompt", +"rate_limit_exceeded") fields — previously only `code` was read, so +type-classified client errors fell through to 500. + +Also covers: ErrorEventError.param must accept dict payloads without raising a +Pydantic ValidationError (previously typed as Optional[str]). +""" + +import json +import os +import sys +from unittest.mock import Mock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +import litellm +from litellm.exceptions import MidStreamFallbackError +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig +from litellm.responses.streaming_iterator import ( + BaseResponsesAPIStreamingIterator, + ResponsesAPIStreamingIterator, + SyncResponsesAPIStreamingIterator, +) +from litellm.types.llms.openai import ( + ErrorEvent, + ErrorEventError, + ResponseAPIUsage, + ResponsesAPIStreamEvents, +) + + +def _make_iterator() -> BaseResponsesAPIStreamingIterator: + mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.model_call_details = {"litellm_params": {}} + mock_config = Mock(spec=BaseResponsesAPIConfig) + mock_response = Mock() + mock_response.headers = {} + return BaseResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-5", + responses_api_provider_config=mock_config, + logging_obj=mock_logging_obj, + custom_llm_provider="openai", + ) + + +def _make_error_chunk(error_type: str, code: str, message: str = "err") -> ErrorEvent: + error_obj = ErrorEventError(type=error_type, code=code, message=message) + return ErrorEvent(type=ResponsesAPIStreamEvents.ERROR, sequence_number=0, error=error_obj) + + +def test_maybe_raise_for_error_event_wraps_unknown_error_in_mid_stream_fallback(): + iterator = _make_iterator() + chunk = _make_error_chunk("server_error", "internal_error", "something went wrong") + with pytest.raises(MidStreamFallbackError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert exc_info.value.status_code == 500 + assert isinstance(exc_info.value.original_exception, litellm.APIError) + assert exc_info.value.original_exception.status_code == 500 + + +def test_maybe_raise_for_error_event_maps_rate_limit_code_to_429_mid_stream_fallback(): + """429 is retriable: it must be wrapped so the Router can fall back, carrying the mapped APIError.""" + iterator = _make_iterator() + chunk = _make_error_chunk("tokens", "rate_limit_exceeded", "Too many requests") + with pytest.raises(MidStreamFallbackError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert exc_info.value.status_code == 429 + assert exc_info.value.generated_content == "" + assert exc_info.value.is_pre_first_chunk is True + assert isinstance(exc_info.value.original_exception, litellm.APIError) + assert exc_info.value.original_exception.status_code == 429 + + +def test_maybe_raise_for_error_event_maps_invalid_request_type_to_400(): + """Client errors classified via the `type` field must raise APIError directly (no fallback).""" + iterator = _make_iterator() + chunk = _make_error_chunk("invalid_request_error", "invalid_prompt", "bad request") + with pytest.raises(litellm.APIError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert exc_info.value.status_code == 400 + assert not isinstance(exc_info.value, MidStreamFallbackError) + + +def test_maybe_raise_for_error_event_maps_context_length_code_to_400(): + """Client errors classified via the `code` field alone must still map to 400.""" + iterator = _make_iterator() + chunk = Mock() + chunk.type = "error" + chunk.error = {"code": "context_length_exceeded", "message": "too long"} + with pytest.raises(litellm.APIError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert exc_info.value.status_code == 400 + assert not isinstance(exc_info.value, MidStreamFallbackError) + + +def test_maybe_raise_for_error_event_maps_insufficient_quota_to_429(): + """OpenAI returns HTTP 429 for insufficient_quota; it must not map to 400 even though its type + is invalid_request_error-adjacent, and it must be wrapped for fallback.""" + iterator = _make_iterator() + chunk = _make_error_chunk("invalid_request_error", "insufficient_quota", "You exceeded your current quota") + with pytest.raises(MidStreamFallbackError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert exc_info.value.status_code == 429 + + +def test_maybe_raise_for_error_event_passes_through_normal_chunk(): + iterator = _make_iterator() + chunk = Mock() + chunk.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA + iterator._maybe_raise_for_error_event(chunk) # must not raise + + +def test_error_event_error_param_accepts_dict(): + error_obj = ErrorEventError( + type="invalid_request_error", + code="context_length_exceeded", + message="too long", + param={"field": "messages", "index": 0}, + ) + assert isinstance(error_obj.param, dict) + + +def _make_async_iterator_with_events(events: list) -> ResponsesAPIStreamingIterator: + sse_payload = b"".join(f"data: {json.dumps(event)}\n\n".encode() for event in events) + + async def mock_aiter_bytes(): + yield sse_payload + + mock_response = Mock() + mock_response.headers = {} + mock_response.aiter_bytes = mock_aiter_bytes + mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.model_call_details = {"litellm_params": {}} + mock_logging_obj.completion_start_time = None + mock_config = Mock(spec=BaseResponsesAPIConfig) + + def transform(model, parsed_chunk, logging_obj): + if parsed_chunk.get("type") == "error": + return ErrorEvent( + type=ResponsesAPIStreamEvents.ERROR, + sequence_number=0, + error=ErrorEventError(**parsed_chunk["error"]), + ) + delta_event = Mock() + delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA + delta_event.delta = parsed_chunk.get("delta", "") + return delta_event + + mock_config.transform_streaming_response.side_effect = transform + + return ResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-5", + responses_api_provider_config=mock_config, + logging_obj=mock_logging_obj, + custom_llm_provider="openai", + ) + + +@pytest.mark.asyncio +async def test_async_iterator_raises_mid_stream_fallback_on_rate_limit_error_event(): + iterator = _make_async_iterator_with_events( + [ + { + "type": "error", + "error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "rate limited"}, + } + ] + ) + + with pytest.raises(MidStreamFallbackError) as exc_info: + async for _ in iterator: + pass + assert exc_info.value.status_code == 429 + assert exc_info.value.is_pre_first_chunk is True + assert exc_info.value.generated_content == "" + assert isinstance(exc_info.value.original_exception, litellm.APIError) + assert exc_info.value.original_exception.status_code == 429 + + +@pytest.mark.asyncio +async def test_async_iterator_error_after_first_chunk_carries_generated_content(): + """An error after streamed output must expose the accumulated text so the router's + fallback can build a continuation input instead of restarting from scratch.""" + iterator = _make_async_iterator_with_events( + [ + {"type": "response.output_text.delta", "delta": "hello "}, + {"type": "response.output_text.delta", "delta": "world"}, + { + "type": "error", + "error": {"type": "server_error", "code": "internal_error", "message": "boom"}, + }, + ] + ) + + chunks = [] + with pytest.raises(MidStreamFallbackError) as exc_info: + async for chunk in iterator: + chunks.append(chunk) + assert len(chunks) == 2 + assert exc_info.value.status_code == 500 + assert exc_info.value.is_pre_first_chunk is False + assert exc_info.value.generated_content == "hello world" + + +def test_maybe_raise_for_response_failed_event_with_dict_error(): + """response.failed chunks carry a dict error on .response.error; covers dict branch.""" + iterator = _make_iterator() + mock_response_obj = Mock() + mock_response_obj.error = {"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"} + chunk = Mock() + chunk.type = "response.failed" + chunk.response = mock_response_obj + with pytest.raises(MidStreamFallbackError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert exc_info.value.status_code == 429 + + +def test_maybe_raise_for_error_event_null_error_obj(): + """error chunk with no error field: message and code default; wrapped as 500.""" + iterator = _make_iterator() + chunk = Mock() + chunk.type = "error" + chunk.error = None + with pytest.raises(MidStreamFallbackError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert exc_info.value.status_code == 500 + assert "Response API in-stream error" in str(exc_info.value) + + +def _make_failed_chunk(error: dict, usage: ResponseAPIUsage | None = None) -> Mock: + mock_response_obj = Mock() + mock_response_obj.error = error + mock_response_obj.usage = usage + chunk = Mock() + chunk.type = "response.failed" + chunk.response = mock_response_obj + return chunk + + +def test_handle_logging_failed_response_maps_rate_limit_to_429(): + """The exception logged to failure handlers must carry the mapped status, not a hardcoded 500.""" + iterator = _make_iterator() + iterator.completed_response = _make_failed_chunk( + {"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"} + ) + with ( + patch("litellm.responses.streaming_iterator.run_async_function") as mock_run_async, + patch("litellm.responses.streaming_iterator.executor"), + ): + iterator._handle_logging_failed_response() + logged_exception = mock_run_async.call_args.kwargs["exception"] + assert isinstance(logged_exception, litellm.APIError) + assert logged_exception.status_code == 429 + assert "throttled" in str(logged_exception) + + +def test_handle_logging_failed_response_maps_type_field_to_400(): + """Status derivation for failed-response logging must also read the error `type` field.""" + iterator = _make_iterator() + iterator.completed_response = _make_failed_chunk( + {"type": "invalid_request_error", "code": "invalid_prompt", "message": "bad prompt"} + ) + with ( + patch("litellm.responses.streaming_iterator.run_async_function") as mock_run_async, + patch("litellm.responses.streaming_iterator.executor"), + ): + iterator._handle_logging_failed_response() + logged_exception = mock_run_async.call_args.kwargs["exception"] + assert isinstance(logged_exception, litellm.APIError) + assert logged_exception.status_code == 400 + + +def test_handle_logging_failed_response_records_usage_and_cost(): + """Usage on a response.failed event must reach failure spend accounting via combined_usage_object.""" + iterator = _make_iterator() + usage = ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15) + chunk = _make_failed_chunk( + {"type": "server_error", "code": "server_error", "message": "boom"}, + usage=usage, + ) + iterator.completed_response = chunk + iterator.logging_obj._response_cost_calculator.return_value = 0.0042 + with ( + patch("litellm.responses.streaming_iterator.run_async_function"), + patch("litellm.responses.streaming_iterator.executor"), + ): + iterator._handle_logging_failed_response() + combined_usage = iterator.logging_obj.model_call_details["combined_usage_object"] + assert isinstance(combined_usage, litellm.Usage) + assert combined_usage.prompt_tokens == 10 + assert combined_usage.completion_tokens == 5 + assert combined_usage.total_tokens == 15 + assert iterator.logging_obj.model_call_details["response_cost"] == 0.0042 + iterator.logging_obj._response_cost_calculator.assert_called_once_with(result=chunk.response) + + +def test_handle_logging_failed_response_without_usage_skips_recording(): + iterator = _make_iterator() + iterator.completed_response = _make_failed_chunk( + {"type": "server_error", "code": "server_error", "message": "boom"} + ) + with ( + patch("litellm.responses.streaming_iterator.run_async_function"), + patch("litellm.responses.streaming_iterator.executor"), + ): + iterator._handle_logging_failed_response() + assert "combined_usage_object" not in iterator.logging_obj.model_call_details + iterator.logging_obj._response_cost_calculator.assert_not_called() + + +def test_sync_iterator_raises_mid_stream_fallback_on_rate_limit_error_event(): + """SyncResponsesAPIStreamingIterator must wrap retriable error events for fallback.""" + error_payload = { + "type": "error", + "error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"}, + } + sse_bytes = f"data: {json.dumps(error_payload)}\n\n".encode() + + mock_response = Mock() + mock_response.headers = {} + mock_response.iter_bytes.return_value = iter([sse_bytes]) + mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.model_call_details = {"litellm_params": {}} + mock_logging_obj.completion_start_time = None + mock_config = Mock(spec=BaseResponsesAPIConfig) + + error_obj = ErrorEventError(type="tokens", code="rate_limit_exceeded", message="throttled") + mock_config.transform_streaming_response.return_value = ErrorEvent( + type=ResponsesAPIStreamEvents.ERROR, sequence_number=0, error=error_obj + ) + + iterator = SyncResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-5", + responses_api_provider_config=mock_config, + logging_obj=mock_logging_obj, + custom_llm_provider="openai", + ) + + with pytest.raises(MidStreamFallbackError) as exc_info: + for _ in iterator: + pass + assert exc_info.value.status_code == 429 + assert isinstance(exc_info.value.original_exception, litellm.APIError) diff --git a/tests/test_litellm/test_gpt_5_6_model_metadata.py b/tests/test_litellm/test_gpt_5_6_model_metadata.py index af9bb117f78..5a7b621d521 100644 --- a/tests/test_litellm/test_gpt_5_6_model_metadata.py +++ b/tests/test_litellm/test_gpt_5_6_model_metadata.py @@ -122,6 +122,8 @@ def test_azure_gpt_5_6_regional_model_info(model): assert info is not None, f"{model} not found in model_prices_and_context_window.json" assert info["litellm_provider"] == "azure" + assert info["mode"] == "chat" + input_cost, output_cost, cache_read_cost, _ = STANDARD_PRICING[_tier_key(model)] assert info["input_cost_per_token"] == pytest.approx(input_cost * 1.1) @@ -132,6 +134,13 @@ def test_azure_gpt_5_6_regional_model_info(model): assert info["input_cost_per_token_priority"] == pytest.approx(input_cost * 2.75) assert info["output_cost_per_token_priority"] == pytest.approx(output_cost * 2.75) + assert info["max_input_tokens"] == 1050000 + assert info["max_output_tokens"] == 128000 + assert info["supports_reasoning"] is True + + _, provider, _, _ = get_llm_provider(model=model) + assert provider == "azure" + def test_gpt_5_6_backup_matches_main(): """Ensure the bundled model cost map stays in sync with the canonical file.""" diff --git a/tests/test_litellm/test_responses_api_bridge_non_stream.py b/tests/test_litellm/test_responses_api_bridge_non_stream.py index 8905293d6b6..25a3bc2dbba 100644 --- a/tests/test_litellm/test_responses_api_bridge_non_stream.py +++ b/tests/test_litellm/test_responses_api_bridge_non_stream.py @@ -269,29 +269,30 @@ def test_transform_usage_with_zero_values(): """ Test transformation when token details are explicitly set to 0. - This ensures 0 values are preserved and not treated as None. + cached_tokens=0 is preserved (cache was available; nothing was cached). + reasoning_tokens=0 is preserved the same way: an explicit provider-reported + zero passes through, while an absent value (None) is omitted. """ completion_response = create_mock_completion_response( model="gpt-4", prompt_tokens=100, completion_tokens=50, total_tokens=150, - cached_tokens=0, # Explicitly 0 - reasoning_tokens=0, # Explicitly 0 + cached_tokens=0, # Explicitly 0 — preserved + reasoning_tokens=0, # Explicitly 0 — preserved ) responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( completion_response ) - # Should preserve 0 values assert responses_usage.input_tokens_details is not None assert responses_usage.input_tokens_details.cached_tokens == 0 assert responses_usage.output_tokens_details is not None assert responses_usage.output_tokens_details.reasoning_tokens == 0 - print("✓ Transformation preserves explicit 0 values") + print("✓ Transformation preserves explicit reasoning_tokens=0 and omits absent values") def test_input_tokens_details_requires_cached_tokens(): @@ -315,25 +316,23 @@ def test_input_tokens_details_requires_cached_tokens(): print("✓ InputTokensDetails correctly defaults cached_tokens to 0") -def test_output_tokens_details_requires_reasoning_tokens(): +def test_output_tokens_details_reasoning_tokens(): """ - Test that OutputTokensDetails has reasoning_tokens as an int with default value 0. + Test OutputTokensDetails.reasoning_tokens field semantics. - This ensures backward compatibility while making the field non-optional. + reasoning_tokens is Optional[int] = None: present only when reasoning actually occurred. """ - # Should work with reasoning_tokens=0 - details1 = OutputTokensDetails(reasoning_tokens=0) - assert details1.reasoning_tokens == 0 + details_explicit_zero = OutputTokensDetails(reasoning_tokens=0) + assert details_explicit_zero.reasoning_tokens == 0 - # Should work with reasoning_tokens=100 - details2 = OutputTokensDetails(reasoning_tokens=100) - assert details2.reasoning_tokens == 100 + details_positive = OutputTokensDetails(reasoning_tokens=100) + assert details_positive.reasoning_tokens == 100 - # Should work without reasoning_tokens (defaults to 0) - details3 = OutputTokensDetails() - assert details3.reasoning_tokens == 0 + # Default is None — absence means reasoning did not occur (or was not tracked) + details_default = OutputTokensDetails() + assert details_default.reasoning_tokens is None - print("✓ OutputTokensDetails correctly defaults reasoning_tokens to 0") + print("✓ OutputTokensDetails.reasoning_tokens defaults to None") def test_all_providers_transformation_scenarios(): @@ -419,7 +418,7 @@ if __name__ == "__main__": test_transform_usage_with_both_token_details() test_transform_usage_with_zero_values() test_input_tokens_details_requires_cached_tokens() - test_output_tokens_details_requires_reasoning_tokens() + test_output_tokens_details_reasoning_tokens() test_all_providers_transformation_scenarios() print("\n" + "=" * 60) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index d739f9c116a..26ae6ff7f4c 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -855,6 +855,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "supports_xhigh_reasoning_effort": {"type": "boolean"}, "supports_max_reasoning_effort": {"type": "boolean"}, "supports_adaptive_thinking": {"type": "boolean"}, + "supports_mid_conversation_system": {"type": "boolean"}, "supports_sampling_params": {"type": "boolean"}, "supports_output_config": {"type": "boolean"}, "supports_speed": {"type": "boolean"}, @@ -1091,8 +1092,8 @@ def test_get_model_info_bedrock_regional_inference_profile_pricing(local_model_c def test_get_model_info_bedrock_regional_profile_without_entry_falls_back_to_base(local_model_cost_map): """A regional profile with no dedicated cost-map entry must still resolve to its region-stripped base entry.""" - assert "jp.anthropic.claude-opus-4-8" not in litellm.model_cost - info = litellm.get_model_info(model="bedrock/jp.anthropic.claude-opus-4-8") + assert "apac.anthropic.claude-opus-4-8" not in litellm.model_cost + info = litellm.get_model_info(model="bedrock/apac.anthropic.claude-opus-4-8") assert info["key"] == "anthropic.claude-opus-4-8" diff --git a/ui/litellm-dashboard/eslint-metrics.json b/ui/litellm-dashboard/eslint-metrics.json index fe6f9f5b482..e99cf762c5d 100644 --- a/ui/litellm-dashboard/eslint-metrics.json +++ b/ui/litellm-dashboard/eslint-metrics.json @@ -1,7 +1,7 @@ { - "@typescript-eslint/no-explicit-any": 1976, - "complexity": 130, - "local/no-large-inline-object-arg": 509, + "@typescript-eslint/no-explicit-any": 1969, + "complexity": 129, + "local/no-large-inline-object-arg": 501, "local/no-long-condition-chain": 234, "max-depth": 59, "no-console": 16 diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 49fdb38c47b..648503abb27 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -28,6 +28,7 @@ "moment": "2.30.1", "next": "16.2.6", "openai": "4.104.0", + "openapi-fetch": "^0.17.0", "papaparse": "5.5.3", "react": "18.3.1", "react-copy-to-clipboard": "5.1.1", @@ -10534,6 +10535,15 @@ "integrity": "sha512-JlCMO+ehdEIKqlFxk6IfVoAUVmgz7cU7zD/h9XZ0qzeosSHmUJVOzSQvvYSYWXkFXC+IfLKSIffhv0sVZup6pA==", "license": "MIT" }, + "node_modules/openapi-fetch": { + "version": "0.17.0", + "resolved": "https://registry.npmjs.org/openapi-fetch/-/openapi-fetch-0.17.0.tgz", + "integrity": "sha512-PsbZR1wAPcG91eEthKhN+Zn92FMHxv+/faECIwjXdxfTODGSGegYv0sc1Olz+HYPvKOuoXfp+0pA2XVt2cI0Ig==", + "license": "MIT", + "dependencies": { + "openapi-typescript-helpers": "^0.1.0" + } + }, "node_modules/openapi-typescript": { "version": "7.13.0", "resolved": "https://registry.npmjs.org/openapi-typescript/-/openapi-typescript-7.13.0.tgz", @@ -10555,6 +10565,12 @@ "typescript": "^5.x" } }, + "node_modules/openapi-typescript-helpers": { + "version": "0.1.0", + "resolved": "https://registry.npmjs.org/openapi-typescript-helpers/-/openapi-typescript-helpers-0.1.0.tgz", + "integrity": "sha512-OKTGPthhivLw/fHz6c3OPtg72vi86qaMlqbJuVJ23qOvQ+53uw1n7HdmkJFibloF7QEjDrDkzJiOJuockM/ljw==", + "license": "MIT" + }, "node_modules/openapi-typescript/node_modules/supports-color": { "version": "10.2.2", "resolved": "https://registry.npmjs.org/supports-color/-/supports-color-10.2.2.tgz", diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index ab01ec70cea..e35620c4996 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -44,6 +44,7 @@ "moment": "2.30.1", "next": "16.2.6", "openai": "4.104.0", + "openapi-fetch": "^0.17.0", "papaparse": "5.5.3", "react": "18.3.1", "react-copy-to-clipboard": "5.1.1", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.test.ts index 716d6f75399..1e614b709e2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.test.ts @@ -1,334 +1,104 @@ -import { allEndUsersCall } from "@/components/networking"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { renderHook, waitFor } from "@testing-library/react"; import React, { ReactNode } from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import type { Customer, CustomersResponse } from "./useCustomers"; -import { useCustomers } from "./useCustomers"; +import { useCustomers, type EndUser } from "./useCustomers"; -// Mock the networking function -vi.mock("@/components/networking", () => ({ - allEndUsersCall: vi.fn(), +const mockGet = vi.fn(); +vi.mock("@/lib/http/api", () => ({ + fetchClient: { GET: (...args: unknown[]) => mockGet(...args) }, })); -// Mock useAuthorized hook - we can override this in individual tests const mockUseAuthorized = vi.fn(); vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => mockUseAuthorized(), })); -// Import actual roles instead of mocking them - -// Mock data -const mockCustomers: Customer[] = [ - { - user_id: "customer-1", - alias: "Test Customer 1", - spend: 150.5, - blocked: false, - allowed_model_region: "us-east-1", - default_model: "gpt-3.5-turbo", - budget_id: "budget-1", - litellm_budget_table: { - budget_id: "budget-1", - max_budget: 1000, - soft_budget: 800, - max_parallel_requests: 10, - tpm_limit: 1000, - rpm_limit: 100, - model_max_budget: { "gpt-4": 500 }, - budget_duration: "monthly", - budget_reset_at: "2024-02-01T00:00:00Z", - created_at: "2024-01-01T00:00:00Z", - created_by: "admin-1", - updated_at: "2024-01-01T00:00:00Z", - updated_by: "admin-1", - }, - }, - { - user_id: "customer-2", - alias: null, - spend: 0, - blocked: true, - allowed_model_region: null, - default_model: null, - budget_id: null, - litellm_budget_table: null, - }, +const mockCustomers: EndUser[] = [ + { user_id: "customer-1", alias: "Test Customer 1", spend: 150.5, blocked: false }, + { user_id: "customer-2", alias: null, spend: 0, blocked: true }, ]; -const mockCustomersResponse: CustomersResponse = mockCustomers; +const authorized = { + accessToken: "test-access-token", + userRole: "Admin", + userId: "test-user-id", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, +}; describe("useCustomers", () => { let queryClient: QueryClient; beforeEach(() => { - queryClient = new QueryClient({ - defaultOptions: { - queries: { - retry: false, - }, - }, - }); - - // Reset all mocks + queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); vi.clearAllMocks(); - - // Set default mock for useAuthorized (enabled state) - mockUseAuthorized.mockReturnValue({ - accessToken: "test-access-token", - userRole: "Admin", - userId: "test-user-id", - token: "test-token", - userEmail: "test@example.com", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, - }); + mockUseAuthorized.mockReturnValue(authorized); }); const wrapper = ({ children }: { children: ReactNode }) => React.createElement(QueryClientProvider, { client: queryClient }, children); - it("should return customers data when query is successful", async () => { - // Mock successful API call - (allEndUsersCall as any).mockResolvedValue(mockCustomersResponse); + it("fetches /customer/list and returns the typed list on success", async () => { + mockGet.mockResolvedValue({ data: mockCustomers }); const { result } = renderHook(() => useCustomers(), { wrapper }); - // Initially loading expect(result.current.isLoading).toBe(true); - expect(result.current.data).toBeUndefined(); - // Wait for success await waitFor(() => { - expect(result.current.isLoading).toBe(false); expect(result.current.isSuccess).toBe(true); }); - expect(result.current.data).toEqual(mockCustomersResponse); - expect(result.current.error).toBeNull(); - expect(allEndUsersCall).toHaveBeenCalledWith("test-access-token"); - expect(allEndUsersCall).toHaveBeenCalledTimes(1); + expect(result.current.data).toEqual(mockCustomers); + expect(mockGet).toHaveBeenCalledWith("/customer/list"); + expect(mockGet).toHaveBeenCalledTimes(1); }); - it("should handle error when allEndUsersCall fails", async () => { - const errorMessage = "Failed to fetch customers"; - const testError = new Error(errorMessage); - - // Mock failed API call - (allEndUsersCall as any).mockRejectedValue(testError); + it("surfaces an error when the request rejects", async () => { + const testError = new Error("Failed to fetch customers"); + mockGet.mockRejectedValue(testError); const { result } = renderHook(() => useCustomers(), { wrapper }); - // Initially loading - expect(result.current.isLoading).toBe(true); - - // Wait for error await waitFor(() => { - expect(result.current.isLoading).toBe(false); expect(result.current.isError).toBe(true); }); expect(result.current.error).toEqual(testError); expect(result.current.data).toBeUndefined(); - expect(allEndUsersCall).toHaveBeenCalledWith("test-access-token"); - expect(allEndUsersCall).toHaveBeenCalledTimes(1); }); - it("should not execute query when accessToken is missing", async () => { - // Mock missing accessToken - mockUseAuthorized.mockReturnValue({ - accessToken: null, - userRole: "Admin", - userId: "test-user-id", - token: null, - userEmail: "test@example.com", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, - }); + it("falls back to an empty list when the response has no body", async () => { + mockGet.mockResolvedValue({ data: undefined }); const { result } = renderHook(() => useCustomers(), { wrapper }); - // Query should not execute - expect(result.current.isLoading).toBe(false); - expect(result.current.data).toBeUndefined(); - expect(result.current.isFetched).toBe(false); - - // API should not be called - expect(allEndUsersCall).not.toHaveBeenCalled(); - }); - - it("should not execute query when userRole is not an admin role", async () => { - // Mock non-admin userRole - mockUseAuthorized.mockReturnValue({ - accessToken: "test-access-token", - userRole: "member", // Not in all_admin_roles - userId: "test-user-id", - token: "test-token", - userEmail: "test@example.com", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, - }); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - // Query should not execute - expect(result.current.isLoading).toBe(false); - expect(result.current.data).toBeUndefined(); - expect(result.current.isFetched).toBe(false); - - // API should not be called - expect(allEndUsersCall).not.toHaveBeenCalled(); - }); - - it("should not execute query when userRole is null", async () => { - // Mock null userRole - mockUseAuthorized.mockReturnValue({ - accessToken: "test-access-token", - userRole: null, - userId: "test-user-id", - token: "test-token", - userEmail: "test@example.com", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, - }); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - // Query should not execute - expect(result.current.isLoading).toBe(false); - expect(result.current.data).toBeUndefined(); - expect(result.current.isFetched).toBe(false); - - // API should not be called - expect(allEndUsersCall).not.toHaveBeenCalled(); - }); - - it("should not execute query when userRole is empty string", async () => { - // Mock empty string userRole - mockUseAuthorized.mockReturnValue({ - accessToken: "test-access-token", - userRole: "", - userId: "test-user-id", - token: "test-token", - userEmail: "test@example.com", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, - }); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - // Query should not execute - expect(result.current.isLoading).toBe(false); - expect(result.current.data).toBeUndefined(); - expect(result.current.isFetched).toBe(false); - - // API should not be called - expect(allEndUsersCall).not.toHaveBeenCalled(); - }); - - it("should not execute query when both accessToken and userRole are missing", async () => { - // Mock both auth values missing - mockUseAuthorized.mockReturnValue({ - accessToken: null, - userRole: null, - userId: "test-user-id", - token: null, - userEmail: "test@example.com", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, - }); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - // Query should not execute - expect(result.current.isLoading).toBe(false); - expect(result.current.data).toBeUndefined(); - expect(result.current.isFetched).toBe(false); - - // API should not be called - expect(allEndUsersCall).not.toHaveBeenCalled(); - }); - - it("should execute query when accessToken is present and userRole is Admin", async () => { - // Mock successful API call - (allEndUsersCall as any).mockResolvedValue(mockCustomersResponse); - - // Ensure auth values are set (already done in beforeEach) - const { result } = renderHook(() => useCustomers(), { wrapper }); - - // Wait for query to execute await waitFor(() => { - expect(result.current.isLoading).toBe(false); - }); - - expect(allEndUsersCall).toHaveBeenCalledWith("test-access-token"); - expect(allEndUsersCall).toHaveBeenCalledTimes(1); - }); - - it("should execute query when accessToken is present and userRole is proxy_admin", async () => { - // Mock successful API call - (allEndUsersCall as any).mockResolvedValue(mockCustomersResponse); - - // Mock proxy_admin role - mockUseAuthorized.mockReturnValue({ - accessToken: "test-access-token", - userRole: "proxy_admin", - userId: "test-user-id", - token: "test-token", - userEmail: "test@example.com", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, - }); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - // Wait for query to execute - await waitFor(() => { - expect(result.current.isLoading).toBe(false); - }); - - expect(allEndUsersCall).toHaveBeenCalledWith("test-access-token"); - expect(allEndUsersCall).toHaveBeenCalledTimes(1); - }); - - it("should return empty customers array when API returns empty data", async () => { - // Mock API returning empty customers array - (allEndUsersCall as any).mockResolvedValue([]); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - // Wait for success - await waitFor(() => { - expect(result.current.isLoading).toBe(false); expect(result.current.isSuccess).toBe(true); }); expect(result.current.data).toEqual([]); - expect(allEndUsersCall).toHaveBeenCalledWith("test-access-token"); }); - it("should handle network timeout error", async () => { - const timeoutError = new Error("Network timeout"); - - // Mock network timeout - (allEndUsersCall as any).mockRejectedValue(timeoutError); + it("does not fetch when the access token is missing", () => { + mockUseAuthorized.mockReturnValue({ ...authorized, accessToken: null, token: null }); const { result } = renderHook(() => useCustomers(), { wrapper }); - // Wait for error - await waitFor(() => { - expect(result.current.isError).toBe(true); - }); + expect(result.current.isFetched).toBe(false); + expect(mockGet).not.toHaveBeenCalled(); + }); - expect(result.current.error).toEqual(timeoutError); - expect(result.current.data).toBeUndefined(); + it("does not fetch when the user is not an admin", () => { + mockUseAuthorized.mockReturnValue({ ...authorized, userRole: "member" }); + + const { result } = renderHook(() => useCustomers(), { wrapper }); + + expect(result.current.isFetched).toBe(false); + expect(mockGet).not.toHaveBeenCalled(); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.ts index d9f3e7cbb36..25e2e3f5e90 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.ts @@ -1,42 +1,19 @@ -import { allEndUsersCall } from "@/components/networking"; import { useQuery } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; +import { fetchClient } from "@/lib/http/api"; import { all_admin_roles } from "@/utils/roles"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import type { components } from "@/lib/http/schema"; + +export type EndUser = components["schemas"]["CustomerResponse"]; + const customersKeys = createQueryKeys("customers"); -export interface Customer { - user_id: string; - alias?: string | null; - spend: number; - blocked: boolean; - allowed_model_region?: string | null; - default_model?: string | null; - budget_id?: string | null; - litellm_budget_table?: { - budget_id: string; - max_budget?: number | null; - soft_budget?: number | null; - max_parallel_requests?: number | null; - tpm_limit?: number | null; - rpm_limit?: number | null; - model_max_budget?: Record | null; - budget_duration?: string | null; - budget_reset_at?: string | null; - created_at: string; - created_by: string; - updated_at: string; - updated_by: string; - } | null; -} - -export type CustomersResponse = Customer[]; - export const useCustomers = () => { const { accessToken, userRole } = useAuthorized(); - return useQuery({ + return useQuery({ queryKey: customersKeys.list({}), - queryFn: async () => await allEndUsersCall(accessToken!), + queryFn: async () => (await fetchClient.GET("/customer/list")).data ?? [], enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!), }); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx index b668cdee0bd..9afa07251c2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx @@ -1,9 +1,16 @@ import React, { useState, useEffect, useCallback, useMemo } from "react"; import { Button, Collapse, Drawer, Empty, Spin, Tooltip, Typography } from "antd"; import { ReloadOutlined } from "@ant-design/icons"; -import type { ColumnDef } from "@tanstack/react-table"; +import type { ColumnDef, ColumnFiltersState } from "@tanstack/react-table"; import { proxyBaseUrl } from "@/components/networking"; -import { DataTable } from "@/components/shared/DataTable"; +import { + DataTable, + DataTableFilterDrawer, + DataTableFilterField, + DataTableToolbar, +} from "@/components/shared/DataTable"; +import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; const { Text } = Typography; @@ -59,6 +66,15 @@ const STATUS_DOT: Record = { failed: "#ef4444", }; +const RUN_STATUS_OPTIONS: RunStatus[] = ["pending", "running", "paused", "completed", "failed"]; +const STATUS_LABELS: Record = { + pending: "Pending", + running: "Running", + paused: "Paused", + completed: "Completed", + failed: "Failed", +}; + const EVENT_COLOR: Record = { "step.started": { bar: "#f0fdf4", border: "#86efac", text: "#16a34a" }, "step.failed": { bar: "#fef2f2", border: "#fca5a5", text: "#dc2626" }, @@ -482,6 +498,9 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { const [messages, setMessages] = useState([]); const [loadingDetail, setLoadingDetail] = useState(false); const [drawerOpen, setDrawerOpen] = useState(false); + const [columnFilters, setColumnFilters] = useState([]); + const [globalFilter, setGlobalFilter] = useState(""); + const [filtersOpen, setFiltersOpen] = useState(false); const fetchRuns = useCallback(async () => { if (!accessToken) return; @@ -547,7 +566,9 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { () => [ { id: "run", + accessorFn: (row) => `${runTitle(row)} ${row.run_id}`, header: "Run", + meta: { title: "Run", skeleton: "twoLine" }, cell: ({ row }) => { const run = row.original; return ( @@ -564,13 +585,18 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { { accessorKey: "workflow_type", header: "Type", + meta: { title: "Type" }, + filterFn: "includesString", cell: ({ row }) => ( {row.original.workflow_type} ), }, { id: "status", + accessorKey: "status", header: "Status", + meta: { title: "Status" }, + filterFn: "equalsString", cell: ({ row }) => { const run = row.original; return ( @@ -586,6 +612,7 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { { accessorKey: "created_at", header: "Created", + meta: { title: "Created" }, cell: ({ row }) => {timeAgo(row.original.created_at)}, }, ], @@ -603,28 +630,11 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { }} > {/* page header */} -
-
-
Workflow Runs
-
- Durable state tracking for agents and automated workflows -
+
+
Workflow Runs
+
+ Durable state tracking for agents and automated workflows
-
= ({ accessToken }) => { } paginationMode="client" pageSizeOptions={[50, 100]} + filterMode="client" + columnFilters={columnFilters} + onColumnFiltersChange={setColumnFilters} + globalFilter={globalFilter} + onGlobalFilterChange={setGlobalFilter} onRowClick={fetchRunDetail} size="compact" + toolbar={(table) => ( + <> + setFiltersOpen(true)} + /> + + {({ get, set }) => ( + <> + + + + + set("workflow_type", event.target.value)} + placeholder="Filter by type…" + /> + + + )} + + + )} /> {/* detail drawer */} diff --git a/ui/litellm-dashboard/src/components/DashboardHeader.test.tsx b/ui/litellm-dashboard/src/components/DashboardHeader.test.tsx new file mode 100644 index 00000000000..aa83e70c56d --- /dev/null +++ b/ui/litellm-dashboard/src/components/DashboardHeader.test.tsx @@ -0,0 +1,57 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; +import { act, fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { DashboardHeader } from "./DashboardHeader"; + +const { mockUsePluginMode, mockUseUISettings, state } = vi.hoisted(() => { + const state = { + plugins: [] as { name: string; display_name: string; url: string }[], + enableChatUI: false, + }; + return { + state, + mockUsePluginMode: vi.fn(() => ({ mode: "ai-gateway", setMode: vi.fn(), plugins: state.plugins })), + mockUseUISettings: vi.fn(() => ({ data: { values: { enable_chat_ui: state.enableChatUI } } })), + }; +}); + +vi.mock("@/contexts/PluginModeContext", () => ({ usePluginMode: mockUsePluginMode })); +vi.mock("@/app/(dashboard)/hooks/uiSettings/useUISettings", () => ({ useUISettings: mockUseUISettings })); +vi.mock("next/navigation", () => ({ usePathname: () => "/ui/" })); +vi.mock("@/utils/migratedPages", () => ({ migratedHref: (seg: string) => `/ui/${seg}` })); +vi.mock("@/hooks/useWorker", () => ({ useWorker: () => ({ isControlPlane: false, selectedWorker: null }) })); +vi.mock("@/app/(dashboard)/hooks/useDisableShowPrompts", () => ({ useDisableShowPrompts: () => false })); +vi.mock("@/components/Navbar/BlogDropdown/BlogDropdown", () => ({ BlogDropdown: () => null })); +vi.mock("@/components/Navbar/CommunityEngagementButtons/CommunityEngagementButtons", () => ({ + CommunityEngagementButtons: () => null, +})); +vi.mock("@/components/Navbar/NotificationsBell/NotificationsBell", () => ({ NotificationsBell: () => null })); +vi.mock("@/components/Navbar/WorkerDropdown/WorkerDropdown", () => ({ default: () => null })); + +describe("DashboardHeader breadcrumb", () => { + afterEach(() => { + state.plugins = []; + state.enableChatUI = false; + }); + + it("roots the breadcrumb in the AI Gateway selector (with a Chat option) and drops the static section crumb when the selector is available", async () => { + state.enableChatUI = true; + render(); + + expect(screen.getByText("Logs")).toBeInTheDocument(); + expect(screen.queryByText("Observability")).not.toBeInTheDocument(); + + const selector = screen.getByRole("button", { name: /AI Gateway/i }); + act(() => { + fireEvent.click(selector); + }); + await waitFor(() => expect(screen.getByText("Chat")).toBeInTheDocument()); + }); + + it("keeps the AI Gateway selector at the root even when there is nothing to switch to (discovery)", () => { + render(); + + expect(screen.getByRole("button", { name: /AI Gateway/i })).toBeInTheDocument(); + expect(screen.getByText("Logs")).toBeInTheDocument(); + expect(screen.queryByText("Observability")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/DashboardHeader.tsx b/ui/litellm-dashboard/src/components/DashboardHeader.tsx index 4f285cefed9..8b928f8e6fe 100644 --- a/ui/litellm-dashboard/src/components/DashboardHeader.tsx +++ b/ui/litellm-dashboard/src/components/DashboardHeader.tsx @@ -27,7 +27,7 @@ interface DashboardHeaderProps { // Top bar for the dashboard shell. Sits only over the content column (the brand // lives in the sidebar header); mirrors the design's breadcrumb-left / tools-right layout. export function DashboardHeader({ page }: DashboardHeaderProps) { - const { section, title } = getBreadcrumb(page); + const { title } = getBreadcrumb(page); const { isControlPlane, selectedWorker } = useWorker(); const showWorkerSwitch = isControlPlane && selectedWorker !== null; const hideCommunityLinks = useDisableShowPrompts(); @@ -44,12 +44,10 @@ export function DashboardHeader({ page }: DashboardHeaderProps) {
- {section && ( - <> - {section} - - - )} + + + + {title} @@ -76,8 +74,6 @@ export function DashboardHeader({ page }: DashboardHeaderProps) { {!hideCommunityLinks && } - -
); diff --git a/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.test.tsx b/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.test.tsx index bdbf74cdf7d..d9a28b2ab8c 100644 --- a/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.test.tsx +++ b/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.test.tsx @@ -49,9 +49,23 @@ describe("ViewSwitcher", () => { state.setMode.mockClear(); }); - it("renders nothing with no plugins, chat disabled, and a non-admin user", () => { - const { container } = render(); - expect(container.firstChild).toBeNull(); + it("still renders the selector with a disabled Chat hint when there are no plugins and chat is off", async () => { + render(); + + const button = screen.getByRole("button"); + expect(button).toHaveTextContent("AI Gateway"); + + act(() => { + fireEvent.click(button); + }); + await waitFor(() => expect(screen.getByText("Chat")).toBeInTheDocument()); + expect(screen.getByText(/Admins can enable in Settings/i)).toBeInTheDocument(); + + act(() => { + fireEvent.click(screen.getByText("Chat")); + }); + expect(assignSpy).not.toHaveBeenCalled(); + expect(state.setMode).not.toHaveBeenCalled(); }); it("labels the button from the active plugin and lists AI Gateway + each plugin", async () => { @@ -95,7 +109,7 @@ describe("ViewSwitcher", () => { fireEvent.click(screen.getByRole("button")); }); await waitFor(() => expect(screen.getByText("Chat")).toBeInTheDocument()); - expect(screen.queryByText(/Enable in Admin Settings/i)).not.toBeInTheDocument(); + expect(screen.queryByText(/Admins can enable in Settings/i)).not.toBeInTheDocument(); act(() => { fireEvent.click(screen.getByText("Chat")); @@ -121,7 +135,7 @@ describe("ViewSwitcher", () => { expect(assignSpy).toHaveBeenCalledWith("/ui/"); }); - it("hides the Chat entry from everyone when disabled", async () => { + it("shows Chat as a disabled, non-navigating entry with an admin hint when disabled", async () => { state.enableChatUI = false; state.plugins = [{ name: "obs", display_name: "Observability", url: "http://localhost:9000" }]; render(); @@ -130,6 +144,12 @@ describe("ViewSwitcher", () => { fireEvent.click(screen.getByRole("button")); }); await waitFor(() => expect(screen.getByText("Observability")).toBeInTheDocument()); - expect(screen.queryByText("Chat")).not.toBeInTheDocument(); + expect(screen.getByText("Chat")).toBeInTheDocument(); + expect(screen.getByText(/Admins can enable in Settings/i)).toBeInTheDocument(); + + act(() => { + fireEvent.click(screen.getByText("Chat")); + }); + expect(assignSpy).not.toHaveBeenCalled(); }); }); diff --git a/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx b/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx index 834966aca35..09a4538ae18 100644 --- a/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx +++ b/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx @@ -1,7 +1,8 @@ import React from "react"; import { usePathname } from "next/navigation"; import { Dropdown } from "antd"; -import { AppstoreOutlined, CheckOutlined, DownOutlined } from "@ant-design/icons"; +import { AppstoreOutlined, CheckOutlined } from "@ant-design/icons"; +import { ChevronsUpDown } from "lucide-react"; import type { MenuProps } from "antd"; import { usePluginMode } from "@/contexts/PluginModeContext"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; @@ -17,8 +18,6 @@ export default function ViewSwitcher() { const chatEnabled = Boolean(uiSettings?.values?.enable_chat_ui); - if (plugins.length === 0 && !chatEnabled) return null; - const chatHref = migratedHref(CHAT); const normalizedPathname = (pathname ?? "").replace(/\/+$/, ""); const isChatRoute = chatEnabled && (normalizedPathname === chatHref || normalizedPathname.startsWith(`${chatHref}/`)); @@ -30,6 +29,29 @@ export default function ViewSwitcher() { ...plugins.map((p) => ({ key: p.name, label: p.display_name })), ]; + const chatItem = chatEnabled + ? { + key: CHAT, + label: ( +
+ Chat + {isChatRoute && } +
+ ), + } + : { + key: CHAT, + disabled: true, + label: ( +
+ Chat + + Admins can enable in Settings + +
+ ), + }; + const items: MenuProps["items"] = [ ...modeEntries.map((e) => ({ key: e.key, @@ -40,19 +62,7 @@ export default function ViewSwitcher() {
), })), - ...(chatEnabled - ? [ - { - key: CHAT, - label: ( -
- Chat - {isChatRoute && } -
- ), - }, - ] - : []), + chatItem, ]; const onClick: MenuProps["onClick"] = ({ key }) => { @@ -72,11 +82,13 @@ export default function ViewSwitcher() { ); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/DcrBridgeToggle.tsx b/ui/litellm-dashboard/src/components/mcp_tools/DcrBridgeToggle.tsx new file mode 100644 index 00000000000..f1b642293c0 --- /dev/null +++ b/ui/litellm-dashboard/src/components/mcp_tools/DcrBridgeToggle.tsx @@ -0,0 +1,40 @@ +import React from "react"; +import { Form, Switch, Tooltip } from "antd"; +import { InfoCircleOutlined } from "@ant-design/icons"; +import { isClientForwardedTokenMode } from "./types"; + +/** + * DCR-bridge toggle for the client-forwarded token modes (true_passthrough / + * oauth_delegate); self-gates to those two auth types and renders nothing + * otherwise. When on, OAuth-only clients like Claude Desktop can register and + * sign in through the gateway; when off, the gateway relays the upstream + * server's own OAuth metadata instead. `initialChecked` seeds the antd + * Form.Item `initialValue` (not the Switch's DOM defaultChecked): the create + * form defaults it on, the edit form seeds it from the stored value. + */ +export default function DcrBridgeToggle({ + authType, + initialChecked, +}: { + authType?: string | null; + initialChecked?: boolean; +}) { + if (!isClientForwardedTokenMode(authType)) return null; + return ( + + Gateway-hosted sign-in (DCR bridge) + + + + + } + name="dcr_bridge" + valuePropName="checked" + initialValue={initialChecked} + > + + + ); +} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.test.tsx new file mode 100644 index 00000000000..0a09ef3f856 --- /dev/null +++ b/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.test.tsx @@ -0,0 +1,54 @@ +import React from "react"; +import { describe, it, expect } from "vitest"; +import { render, screen } from "@testing-library/react"; +import { Form } from "antd"; +import PassthroughAuthorizeSection from "./PassthroughAuthorizeSection"; + +const WithForm: React.FC<{ children: React.ReactNode }> = ({ children }) => { + const [form] = Form.useForm(); + return
{children}
; +}; + +const noopFlow = { startOAuthFlow: () => {}, status: "idle", error: null, tokenResponse: null }; + +describe("PassthroughAuthorizeSection credential-class-aware copy", () => { + it("shows keep-existing copy when the credential class is unchanged (true_passthrough <-> oauth_delegate)", () => { + render( + + + , + ); + expect(screen.getByPlaceholderText("Leave blank to keep the currently saved app (if any)")).toBeInTheDocument(); + expect(screen.getByPlaceholderText("Leave blank to keep the currently saved secret (if any)")).toBeInTheDocument(); + }); + + it("shows the discard warning copy when switching from a different class (oauth2 -> true_passthrough)", () => { + render( + + + , + ); + expect(screen.getByPlaceholderText("Leave blank to use dynamic client registration")).toBeInTheDocument(); + expect(screen.getByPlaceholderText("Leave blank for public clients / PKCE")).toBeInTheDocument(); + expect(screen.getByText(/Switching the auth type discards the previously saved app/)).toBeInTheDocument(); + }); + + it("shows the keep+warn banner when the upstream may no longer match", () => { + render( + + + , + ); + expect(screen.getByText(/registered for the previous upstream/)).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.tsx b/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.tsx index af81f2713ae..0ed4ee555d1 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.tsx @@ -1,6 +1,7 @@ import React from "react"; -import { Button, Form, Input } from "antd"; -import { isClientForwardedTokenMode } from "./types"; +import { Button, Checkbox, Form, Input } from "antd"; +import DcrBridgeToggle from "./DcrBridgeToggle"; +import { credentialAuthClass, isClientForwardedTokenMode } from "./types"; interface PassthroughOAuthFlow { startOAuthFlow: () => void | Promise; @@ -11,20 +12,42 @@ interface PassthroughOAuthFlow { /** * Browser-only Authorize & Fetch for the client-forwarded token modes - * (true_passthrough / oauth_delegate). LiteLLM never stores upstream - * credentials for these modes, so the token obtained here lives in this - * browser session only: it is forwarded per-server for the tools preview and - * allowlist configuration, and is never written to the server row or the - * per-user credential store. The optional client credentials cover IdPs - * without dynamic client registration (e.g. a pre-registered Slack app) and - * ride the temporary authorize session only. + * (true_passthrough / oauth_delegate). Tokens are never stored: the token + * obtained here lives in this browser session only, forwarded per-server for + * the tools preview and allowlist configuration, and is never written to the + * server row or the per-user credential store. The optional OAuth client + * credentials cover IdPs without dynamic client registration (e.g. a + * pre-registered Slack app); unlike the token they ARE saved onto the server + * as declared config, so internal users' Authorize relays through the org's + * app instead of dead-ending on upstreams that cannot mint clients. + * + * Blank fields follow the same convention as the M2M credential fields. On + * create they mean "no app configured" (dynamic client registration). On edit + * they mean "keep existing" ONLY when the credential class is unchanged: the + * backend merges a partial update within the client-forwarded class, so a + * true_passthrough <-> oauth_delegate switch keeps the stored app, but a switch + * from a different class (e.g. oauth2) replaces it, so blanks then mean "no + * app". Removing a stored app is an explicit checkbox (edit only) that writes + * an explicit-null credential. */ export default function PassthroughAuthorizeSection({ authType, oauthFlow, + dcrBridgeInitialChecked, + isEditing = false, + savedAuthType, + removeStoredApp = false, + onRemoveStoredAppChange, + appMayNotMatchUpstream = false, }: { authType?: string | null; oauthFlow: PassthroughOAuthFlow; + dcrBridgeInitialChecked?: boolean; + isEditing?: boolean; + savedAuthType?: string | null; + removeStoredApp?: boolean; + onRemoveStoredAppChange?: (remove: boolean) => void; + appMayNotMatchUpstream?: boolean; }) { if (!isClientForwardedTokenMode(authType)) return null; const authorizeButtonLabels: Record = { @@ -32,32 +55,61 @@ export default function PassthroughAuthorizeSection({ exchanging: "Exchanging authorization code...", }; const authorizeButtonLabel = authorizeButtonLabels[oauthFlow.status] ?? "Authorize & Fetch Tools (browser-only)"; + // On edit, "keep existing" only holds when the stored credential class is unchanged; a cross-class + // switch (e.g. oauth2 -> true_passthrough) replaces credentials, so blanks then mean "no app". + const classUnchanged = isEditing && credentialAuthClass(savedAuthType) === credentialAuthClass(authType); + const clientIdPlaceholder = classUnchanged + ? "Leave blank to keep the currently saved app (if any)" + : "Leave blank to use dynamic client registration"; + const clientSecretPlaceholder = classUnchanged + ? "Leave blank to keep the currently saved secret (if any)" + : "Leave blank for public clients / PKCE"; + const clientIdExtra = classUnchanged + ? "Set this to make everyone authorize through a specific app; required for upstreams without dynamic client registration (e.g. a pre-registered Slack app)." + : "Switching the auth type discards the previously saved app; enter a client ID here or leave blank to use dynamic client registration."; return (

- Callers bring their own upstream token for this auth type, so LiteLLM stores no upstream credentials. To preview - tools and configure the tool allowlist, authorize against the upstream here: the token stays in this browser - session only and is never saved to LiteLLM. + Callers bring their own upstream token for this auth type, so LiteLLM never stores tokens. To preview tools and + configure the tool allowlist, authorize against the upstream here: the token stays in this browser session only + and is never saved to LiteLLM. An OAuth app configured below IS saved with the server, so internal users who + authorize from the Tools page go through it.

+ {appMayNotMatchUpstream && ( +

+ You changed the upstream URL or endpoints; the OAuth app entered here was registered for the previous upstream + and may not be valid. Update the client ID, or clear it to use dynamic client registration. +

+ )} OAuth Client ID (optional, not saved)} + label={OAuth Client ID (optional)} name={["credentials", "client_id"]} - extra="Only needed when the upstream does not support dynamic client registration (e.g. a pre-registered Slack app). Used for this browser authorization only." + extra={clientIdExtra} > OAuth Client Secret (optional, not saved)} + label={OAuth Client Secret (optional)} name={["credentials", "client_secret"]} > + + {isEditing && onRemoveStoredAppChange && ( + onRemoveStoredAppChange(e.target.checked)}> + + Remove the saved OAuth app on save (the server goes back to dynamic client registration) + + + )}
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx index 02374a3ffa4..32fe439a316 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx @@ -31,6 +31,7 @@ const oauthHook = vi.hoisted(() => ({ | ((token: Record | null, registeredClient?: { clientId?: string; clientSecret?: string }) => void) | null, getCredentials: null as (() => Record | undefined) | null, + getTemporaryPayload: null as (() => Record | null) | null, })); vi.mock("@/hooks/useMcpOAuthFlow", () => ({ useMcpOAuthFlow: (opts: { @@ -39,9 +40,11 @@ vi.mock("@/hooks/useMcpOAuthFlow", () => ({ registeredClient?: { clientId?: string; clientSecret?: string }, ) => void; getCredentials?: () => Record | undefined; + getTemporaryPayload?: () => Record | null; }) => { oauthHook.onTokenReceived = opts.onTokenReceived; oauthHook.getCredentials = opts.getCredentials ?? null; + oauthHook.getTemporaryPayload = opts.getTemporaryPayload ?? null; return { startOAuthFlow: vi.fn(), status: "idle", @@ -202,8 +205,8 @@ describe("CreateMCPServer", () => { await waitFor(() => { expect(screen.getByRole("button", { name: "Authorize & Fetch Tools (browser-only)" })).toBeInTheDocument(); }); - expect(screen.getByText("OAuth Client ID (optional, not saved)")).toBeInTheDocument(); - expect(screen.getByText("OAuth Client Secret (optional, not saved)")).toBeInTheDocument(); + expect(screen.getByText("OAuth Client ID (optional)")).toBeInTheDocument(); + expect(screen.getByText("OAuth Client Secret (optional)")).toBeInTheDocument(); }, ); @@ -436,6 +439,434 @@ describe("CreateMCPServer", () => { ); }); + it.each([ + ["true_passthrough", "True Passthrough (no LiteLLM auth)"], + ["oauth_delegate", "OAuth Delegate (client-supplied upstream token)"], + ])( + "persists admin-entered OAuth app credentials on create for %s while the token stays browser-held", + async (_authType, optionLabel) => { + oauthHook.tokenResponse = { access_token: "upstream-tok", token_type: "Bearer" }; + await selectHttpTransport(); + + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "CF_App_Server"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + + await selectAntOption("Authentication", optionLabel); + + // Admin declares the org's pre-registered upstream app; unlike the browser-authorized + // token, this is config and must survive onto the server row so internal users' + // Tools-page Authorize relays through it (required for non-DCR upstreams like Slack). + await user.type( + screen.getByPlaceholderText("Leave blank to use dynamic client registration"), + "org-app-client-id", + ); + await user.type(screen.getByPlaceholderText("Leave blank for public clients / PKCE"), "org-app-secret"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!({ access_token: "upstream-tok", token_type: "Bearer" }, undefined); + }); + + const createdServer = { + server_id: "new-cf-app-server", + server_name: "CF_App_Server", + alias: "CF_App_Server", + url: "https://example.com/mcp", + transport: "http", + auth_type: _authType, + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + vi.mocked(networking.createMCPServer).mockResolvedValue(createdServer); + + const submitButton = screen.getByRole("button", { name: "Add MCP Server" }); + await act(async () => { + fireEvent.click(submitButton); + }); + + await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + + // The declared app persists; the browser-authorized token still appears nowhere in the + // payload and no per-user DB credential is written. + expect(payload.credentials).toEqual({ + client_id: "org-app-client-id", + client_secret: "org-app-secret", + }); + expect(JSON.stringify(payload)).not.toContain("upstream-tok"); + expect(networking.storeMCPOAuthUserCredential).not.toHaveBeenCalled(); + expect(setToken).toHaveBeenCalledWith( + "new-cf-app-server", + expect.objectContaining({ access_token: "upstream-tok" }), + undefined, + ); + }, + ); + + it("preserves admin-entered app credentials when the URL changes after authorize for true_passthrough", async () => { + await selectHttpTransport(); + + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "CF_Keep_Server"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + + await user.type( + screen.getByPlaceholderText("Leave blank to use dynamic client registration"), + "org-app-client-id", + ); + await user.type(screen.getByPlaceholderText("Leave blank for public clients / PKCE"), "org-app-secret"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!({ access_token: "upstream-tok", token_type: "Bearer" }, undefined); + }); + + // Editing the URL after authorize invalidates the held token (identity change), but the + // declared app is config, not minted material: it must survive the invalidation instead of + // being silently reset, or the server would persist without the configured app. + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), { + target: { value: "https://other.example.com/mcp" }, + }); + }); + + const keptAppServer = { + server_id: "kept-app-server", + server_name: "CF_Keep_Server", + alias: "CF_Keep_Server", + url: "https://other.example.com/mcp", + transport: "http", + auth_type: "true_passthrough", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + vi.mocked(networking.createMCPServer).mockResolvedValue(keptAppServer); + + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Add MCP Server" })); + }); + + await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + expect(payload.url).toBe("https://other.example.com/mcp"); + expect(payload.credentials).toEqual({ + client_id: "org-app-client-id", + client_secret: "org-app-secret", + }); + expect(JSON.stringify(payload)).not.toContain("upstream-tok"); + }); + + it("wipes oauth2-minted credentials when the auth type switches to a client-forwarded mode", async () => { + await selectHttpTransport(); + + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "Switch_Server"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + + await selectAntOption("Authentication", "OAuth"); + + // The oauth2 onTokenReceived branch writes the fetched token AND the DCR client into + // form.credentials; both are minted for the oauth2 identity. + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!( + { access_token: "oauth2-minted-tok", refresh_token: "oauth2-minted-refresh", token_type: "Bearer" }, + { clientId: "dcr-minted-client", clientSecret: "dcr-minted-secret" }, + ); + }); + + // Switching into a client-forwarded mode changes the identity with auth_type in the changed + // values, so the preserve carve-out must NOT apply: the minted material would otherwise ride + // into a mode that now persists credentials onto the server row. + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + + const switchedServer = { + server_id: "switched-server", + server_name: "Switch_Server", + alias: "Switch_Server", + url: "https://example.com/mcp", + transport: "http", + auth_type: "true_passthrough", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + vi.mocked(networking.createMCPServer).mockResolvedValue(switchedServer); + + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Add MCP Server" })); + }); + + await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + expect(payload.credentials).toBeUndefined(); + expect(JSON.stringify(payload)).not.toContain("dcr-minted-client"); + expect(JSON.stringify(payload)).not.toContain("oauth2-minted-tok"); + }); + + it("keeps the DCR-minted client out of form.credentials but reuses it via getCredentials", async () => { + await selectHttpTransport(); + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "DCR_Server"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + await selectAntOption("Authentication", "OAuth"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!( + { access_token: "oauth2-tok", token_type: "Bearer" }, + { clientId: "dcr-client", clientSecret: "dcr-secret" }, + ); + }); + + // The DCR client must NOT be in the form store (or it could be collected as a CF server's app), + // but getCredentials merges it so a re-authorize reuses the registered client instead of re-DCRing. + expect(oauthHook.getCredentials?.()?.client_id).toBe("dcr-client"); + // getTemporaryPayload must mirror getCredentials for oauth2, or a re-authorize's temp session omits + // the registered client and useMcpOAuthFlow re-registers instead of reusing it. + expect(oauthHook.getTemporaryPayload?.()?.credentials).toMatchObject({ client_id: "dcr-client" }); + }); + + it("clears the DCR ref and the upstream warning when the modal closes so nothing leaks to the next session", async () => { + const { rerender } = render(); + await selectAntOption("Transport Type", "Streamable HTTP"); + await waitFor(() => expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument()); + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "Leak_Server"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + await selectAntOption("Authentication", "OAuth"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!( + { access_token: "oauth2-tok", token_type: "Bearer" }, + { clientId: "leak-client", clientSecret: "leak-secret" }, + ); + }); + // Ref is held while the modal is open. + expect(oauthHook.getCredentials?.()?.client_id).toBe("leak-client"); + + // A parent dismiss (isModalVisible -> false) that does not route through Cancel/Create must still + // clear the DCR ref, or the next server's oauth2 submit would carry this server's registered client. + await act(async () => { + rerender(); + }); + + expect(oauthHook.getCredentials?.()?.client_id).toBeUndefined(); + expect(oauthHook.getTemporaryPayload?.()?.credentials ?? {}).not.toMatchObject({ client_id: "leak-client" }); + }); + + it("persists the DCR client on an oauth2 submit via the ref", async () => { + await selectHttpTransport(); + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "DCR_Submit_Server"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + await selectAntOption("Authentication", "OAuth"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!( + { access_token: "oauth2-tok", token_type: "Bearer" }, + { clientId: "dcr-client", clientSecret: "dcr-secret" }, + ); + }); + + const dcrSubmitServer = { + server_id: "dcr-submit", + server_name: "DCR_Submit_Server", + alias: "DCR_Submit_Server", + url: "https://example.com/mcp", + transport: "http", + auth_type: "oauth2", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + vi.mocked(networking.createMCPServer).mockResolvedValue(dcrSubmitServer); + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Add MCP Server" })); + }); + + await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + expect(payload.credentials.client_id).toBe("dcr-client"); + expect(payload.credentials.client_secret).toBe("dcr-secret"); + }); + + // These two tests drive multiple antd auth-type switches; use single-shot fireEvent.change for the + // text fields (not per-keystroke userEvent.type) and a wider timeout so they do not flake under CI + // resource contention. The behavior under test is the credential preserve across the switches. + const fillText = (el: HTMLElement, value: string) => fireEvent.change(el, { target: { value } }); + + it("preserves the typed app across a switch between the two client-forwarded modes", async () => { + await selectHttpTransport(); + fillText(getServerNameInput(), "CF_Switch_Keep"); + fillText(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + fillText(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), "app-id"); + fillText(screen.getByPlaceholderText("Leave blank for public clients / PKCE"), "app-secret"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!({ access_token: "cf-tok", token_type: "Bearer" }, undefined); + }); + + await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)"); + + const switched = { + server_id: "cf-switch-keep", + server_name: "CF_Switch_Keep", + alias: "CF_Switch_Keep", + url: "https://example.com/mcp", + transport: "http", + auth_type: "oauth_delegate", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + vi.mocked(networking.createMCPServer).mockResolvedValue(switched); + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Add MCP Server" })); + }); + + await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + expect(payload.credentials).toEqual({ client_id: "app-id", client_secret: "app-secret" }); + }, 60_000); + + it("preserves the typed app across a client-forwarded -> oauth2 -> client-forwarded round trip", async () => { + await selectHttpTransport(); + fillText(getServerNameInput(), "CF_Round"); + fillText(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + fillText(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), "app-id"); + fillText(screen.getByPlaceholderText("Leave blank for public clients / PKCE"), "app-secret"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!({ access_token: "cf-tok", token_type: "Bearer" }, undefined); + }); + + await selectAntOption("Authentication", "OAuth"); + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + + const cfRoundServer = { + server_id: "cf-round", + server_name: "CF_Round", + alias: "CF_Round", + url: "https://example.com/mcp", + transport: "http", + auth_type: "true_passthrough", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + vi.mocked(networking.createMCPServer).mockResolvedValue(cfRoundServer); + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Add MCP Server" })); + }); + + await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + expect(payload.credentials).toEqual({ client_id: "app-id", client_secret: "app-secret" }); + }, 60_000); + + it("keeps the typed app but warns when the URL changes after a client-forwarded authorize", async () => { + await selectHttpTransport(); + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "CF_Warn"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + await user.type(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), "app-id"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!({ access_token: "cf-tok", token_type: "Bearer" }, undefined); + }); + + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), { + target: { value: "https://other.example.com/mcp" }, + }); + }); + + // Keep + warn: the app stays in the field, and a non-blocking warning appears. + expect(screen.getByText(/OAuth app entered here was registered for the previous upstream/)).toBeInTheDocument(); + }); + + it("keeps client_secret when only client_id is edited after a client-forwarded authorize", async () => { + await selectHttpTransport(); + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "CF_Keystroke"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + await user.type(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), "app-id"); + await user.type(screen.getByPlaceholderText("Leave blank for public clients / PKCE"), "app-secret"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + await act(async () => { + oauthHook.onTokenReceived!({ access_token: "cf-tok", token_type: "Bearer" }, undefined); + }); + + // Editing only client_id fires an invalidation whose changedValues carries only the client_id + // sub-field; the preserve + deep-merge re-apply must keep client_secret from being dropped. + await user.type(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), "2"); + + const cfKeystrokeServer = { + server_id: "cf-keystroke", + server_name: "CF_Keystroke", + alias: "CF_Keystroke", + url: "https://example.com/mcp", + transport: "http", + auth_type: "true_passthrough", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + vi.mocked(networking.createMCPServer).mockResolvedValue(cfKeystrokeServer); + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Add MCP Server" })); + }); + + await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + expect(payload.credentials).toEqual({ client_id: "app-id2", client_secret: "app-secret" }); + }); + + it("replaces the token set on re-authorize instead of leaving stale siblings", async () => { + await selectHttpTransport(); + const user = userEvent.setup({ delay: null }); + await user.type(getServerNameInput(), "Reauth_Server"); + await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp"); + await selectAntOption("Authentication", "OAuth"); + + await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy()); + const firstToken = { access_token: "T1", refresh_token: "R1", scope: "read", token_type: "Bearer" }; + await act(async () => { + oauthHook.onTokenReceived!(firstToken, undefined); + }); + await act(async () => { + oauthHook.onTokenReceived!({ access_token: "T2", token_type: "Bearer" }, undefined); + }); + + const creds = oauthHook.getCredentials?.() ?? {}; + expect(creds.access_token).toBe("T2"); + expect(creds.refresh_token).toBeUndefined(); + expect(creds.scope).toBeUndefined(); + }); + it("should not show auth value field when None auth type is selected", async () => { await selectHttpTransport(); @@ -1509,3 +1940,181 @@ describe("CreateMCPServer oauth2_flow persistence", () => { expect(payload.oauth2_flow).toBeUndefined(); }); }); + +describe("CreateMCPServer dcr_bridge toggle", () => { + beforeEach(() => { + vi.clearAllMocks(); + oauthHook.tokenResponse = null; + oauthHook.onTokenReceived = null; + }); + + const createdServer = { + server_id: "new-cf-server", + server_name: "CF_Server", + alias: "CF_Server", + url: "https://example.com/mcp", + transport: "http", + auth_type: "true_passthrough", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }; + + const getDcrToggle = () => document.getElementById("dcr_bridge"); + + async function setupHttpServerForm() { + render(); + await selectAntOption("Transport Type", "Streamable HTTP"); + await waitFor(() => { + expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument(); + }); + await act(async () => { + fireEvent.change(getServerNameInput(), { target: { value: "CF_Server" } }); + }); + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), { + target: { value: "https://example.com/mcp" }, + }); + }); + } + + async function submitCreate() { + const submitButton = screen.getByRole("button", { name: "Add MCP Server" }); + await act(async () => { + fireEvent.click(submitButton); + }); + await waitFor(() => { + expect(networking.createMCPServer).toHaveBeenCalledTimes(1); + }); + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + return payload; + } + + it.each([["True Passthrough (no LiteLLM auth)"], ["OAuth Delegate (client-supplied upstream token)"]])( + "renders the toggle default-checked when %s is selected", + async (optionLabel) => { + await setupHttpServerForm(); + + await selectAntOption("Authentication", optionLabel); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + expect(screen.getByText("Gateway-hosted sign-in (DCR bridge)")).toBeInTheDocument(); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "true"); + }, + ); + + it.each([["None"], ["API Key"], ["OAuth"]])("does not render the toggle for %s", async (optionLabel) => { + await setupHttpServerForm(); + + await selectAntOption("Authentication", optionLabel); + + await waitFor(() => { + expect(screen.queryByText("Gateway-hosted sign-in (DCR bridge)")).not.toBeInTheDocument(); + }); + expect(getDcrToggle()).not.toBeInTheDocument(); + }); + + it("renders the toggle between the OAuth client fields and the Authorize button", async () => { + await setupHttpServerForm(); + + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + const toggle = getDcrToggle() as HTMLElement; + const secretInput = screen.getByPlaceholderText("Leave blank for public clients / PKCE"); + const authorizeButton = screen.getByRole("button", { name: "Authorize & Fetch Tools (browser-only)" }); + expect(secretInput.compareDocumentPosition(toggle) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy(); + expect(toggle.compareDocumentPosition(authorizeButton) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy(); + }); + + it.each([ + ["true_passthrough", "True Passthrough (no LiteLLM auth)"], + ["oauth_delegate", "OAuth Delegate (client-supplied upstream token)"], + ])("sends dcr_bridge: true by default on create for %s", async (authType, optionLabel) => { + vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: authType }); + await setupHttpServerForm(); + await selectAntOption("Authentication", optionLabel); + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + + const payload = await submitCreate(); + expect(payload.dcr_bridge).toBe(true); + }); + + it("sends an explicit dcr_bridge: false when the toggle is unchecked", async () => { + vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: "oauth_delegate" }); + await setupHttpServerForm(); + await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)"); + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + + await act(async () => { + fireEvent.click(getDcrToggle()!); + }); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "false"); + + const payload = await submitCreate(); + expect(payload.dcr_bridge).toBe(false); + }); + + it.each([ + ["none", "None"], + ["api_key", "API Key"], + ["oauth2", "OAuth"], + ])("forces an explicit dcr_bridge: false for %s", async (authType, optionLabel) => { + vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: authType }); + await setupHttpServerForm(); + await selectAntOption("Authentication", optionLabel); + + const payload = await submitCreate(); + expect(payload.dcr_bridge).toBe(false); + }); + + it("forces dcr_bridge: false when the auth type is switched away after toggling", async () => { + vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: "none" }); + await setupHttpServerForm(); + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + await act(async () => { + fireEvent.click(getDcrToggle()!); + }); + + await selectAntOption("Authentication", "None"); + await waitFor(() => { + expect(getDcrToggle()).not.toBeInTheDocument(); + }); + + const payload = await submitCreate(); + expect(payload.dcr_bridge).toBe(false); + }); + + it("preserves the toggle value when switching between the two client-forwarded modes", async () => { + vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: "oauth_delegate" }); + await setupHttpServerForm(); + await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "true"); + + // The Form.Item is mounted in both client-forwarded modes, so switching between them keeps the + // live toggle value rather than forcing it back to the default or to false. + await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)"); + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "true"); + + const payload = await submitCreate(); + expect(payload.dcr_bridge).toBe(true); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index eb48fd02474..b9a6f5d229d 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -18,6 +18,8 @@ import { getOAuthAuthorizationIdentity, CLEARED_ON_INVALIDATION, isHeldOAuthTokenStale, + preservedDeclaredAppCredentials, + withoutMintedTokenCredentials, } from "./types"; import OAuthFormFields from "./OAuthFormFields"; import TruePassthroughWarning from "./TruePassthroughWarning"; @@ -60,6 +62,8 @@ const AUTH_TYPES_REQUIRING_CREDENTIALS = [ AUTH_TYPE.OAUTH2, AUTH_TYPE.OAUTH2_TOKEN_EXCHANGE, AUTH_TYPE.AWS_SIGV4, + AUTH_TYPE.TRUE_PASSTHROUGH, + AUTH_TYPE.OAUTH_DELEGATE, ]; const CREATE_OAUTH_UI_STATE_KEY = "litellm-mcp-oauth-create-state"; @@ -106,6 +110,14 @@ const CreateMCPServer: React.FC = ({ // was fetched; undefined when no valid token is held. If any mint-relevant field diverges from this, // the held token is stale and is discarded so the admin must re-authorize. const [authorizedIdentity, setAuthorizedIdentity] = useState(undefined); + // The DCR-minted OAuth client from an interactive (oauth2) Authorize. Held OUT of form.credentials so + // it can never be collected as a client-forwarded server's declared app; injected into the payload + // only on an oauth2 submit (where persisting the registered client is correct), and cleared on any + // invalidation or modal close. An abandoned authorize leaves it null, which is the desired asymmetry. + const dcrClientRef = React.useRef<{ client_id: string; client_secret?: string } | null>(null); + // Set when the upstream identity (url/endpoints) changed while a declared app is present, so the + // section can warn that the saved app may not match the new upstream (the app is kept, not wiped). + const [appMayNotMatchUpstream, setAppMayNotMatchUpstream] = useState(false); // Single hook call shared by MCPConnectionStatus and MCPToolConfiguration to avoid duplicate requests. const { @@ -147,6 +159,9 @@ const CreateMCPServer: React.FC = ({ searchValue, aliasManuallyEdited, logoUrl, + // Persist the identity so invalidation stays armed across the OAuth redirect round trip: a + // post-restore url/mode edit must still discard the held token instead of silently keeping it. + authorizedIdentity, }; setSecureItem(CREATE_OAUTH_UI_STATE_KEY, JSON.stringify(uiState)); } catch (err) { @@ -162,7 +177,12 @@ const CreateMCPServer: React.FC = ({ reset: resetOAuthFlow, } = useMcpOAuthFlow({ accessToken, - getCredentials: () => form.getFieldValue("credentials"), + // Merge the ref-held DCR client so a re-authorize reuses the registered client instead of + // re-registering; the form store itself never holds the DCR client (see onTokenReceived). + getCredentials: () => ({ + ...((form.getFieldValue("credentials") as Record | undefined) ?? {}), + ...(dcrClientRef.current ?? {}), + }), getTemporaryPayload: () => { const values = form.getFieldsValue(true); const transport = values.transport || transportType; @@ -184,7 +204,12 @@ const CreateMCPServer: React.FC = ({ url, transport: transport === TRANSPORT.OPENAPI ? "http" : transport, auth_type: isClientForwardedTokenMode(values.auth_type) ? values.auth_type : AUTH_TYPE.OAUTH2, - credentials: values.credentials, + // Mirror getCredentials: merge the ref-held DCR client for oauth2 so a re-authorize reuses the + // registered client (useMcpOAuthFlow keys reuse off credentials.client_id) instead of re-DCRing; + // the client-forwarded modes carry only the declared app. + credentials: isClientForwardedTokenMode(values.auth_type) + ? preservedDeclaredAppCredentials(values.credentials) + : { ...((values.credentials as Record | undefined) ?? {}), ...(dcrClientRef.current ?? {}) }, authorization_url: values.authorization_url, token_url: values.token_url, registration_url: values.registration_url, @@ -209,23 +234,36 @@ const CreateMCPServer: React.FC = ({ // edit form's onTokenReceived early return. setAuthorizedIdentity(getOAuthAuthorizationIdentity(form.getFieldsValue(true))); NotificationsManager.success( - "Token held for this browser session. Tools can now be previewed and configured; nothing will be saved to LiteLLM.", + "Token held for this browser session. Tools can now be previewed and configured; the token is not saved to LiteLLM.", ); return; } - const credentials = { + // The DCR-minted client is held in a ref, NOT written into form.credentials, so it can never be + // collected as a client-forwarded server's declared app; it is injected into the payload only on + // an oauth2 submit. An admin-typed client already lives in form.credentials and is left untouched. + dcrClientRef.current = registeredClient?.clientId + ? { + client_id: registeredClient.clientId, + ...(registeredClient.clientSecret && { client_secret: registeredClient.clientSecret }), + } + : null; + + const current = (form.getFieldValue("credentials") as Record | undefined) ?? {}; + const nextCredentials = { + ...(preservedDeclaredAppCredentials(current) ?? {}), + ...(current.scopes !== undefined && { scopes: current.scopes }), access_token: token.access_token, ...(token.refresh_token && { refresh_token: token.refresh_token }), ...(token.expires_in && { expires_in: token.expires_in }), ...(token.scope && { scope: token.scope }), - ...(registeredClient?.clientId && { client_id: registeredClient.clientId }), - ...(registeredClient?.clientSecret && { client_secret: registeredClient.clientSecret }), }; - - form.setFieldsValue({ credentials }); - // Capture the identity AFTER writing the DCR'd credentials so the held token is not spuriously - // invalidated by its own credential write. + // Path-replace (not deep-merge) so a re-authorize with fewer token fields does not leave stale + // siblings from the previous token behind; the admin-typed client keys and scopes are carried + // explicitly above. + form.setFieldValue("credentials", nextCredentials); + // Capture the identity AFTER writing the token so the held token is not spuriously invalidated by + // its own credential write. setAuthorizedIdentity(getOAuthAuthorizationIdentity(form.getFieldsValue(true))); NotificationsManager.success( @@ -246,7 +284,17 @@ const CreateMCPServer: React.FC = ({ clearTools(); resetOAuthFlow(); setAuthorizedIdentity(undefined); + dcrClientRef.current = null; + // Capture the admin-typed app before resetFields destroys it, then re-apply it: the app is + // upstream-scoped config, not minted material, so it survives every invalidation (the token is + // what gets discarded). Token-shaped keys are excluded by the helper's key filter. + const keptAppCredentials = preservedDeclaredAppCredentials(form.getFieldValue("credentials")); form.resetFields([...CLEARED_ON_INVALIDATION]); + if (keptAppCredentials) { + form.setFieldsValue({ credentials: keptAppCredentials }); + } + // Re-apply the in-flight edit last; rc-field-form deep-merges nested objects, so a changed + // credentials sub-field composes with the preserved sibling instead of replacing the object. const preserved = Object.fromEntries( CLEARED_ON_INVALIDATION.filter((key) => key in changedValues).map((key) => [key, changedValues[key]]), ); @@ -274,7 +322,18 @@ const CreateMCPServer: React.FC = ({ setTransportType(restoredTransport); } if (parsed.formValues) { - setPendingRestoredValues({ values: parsed.formValues, transport: restoredTransport }); + // Assign the cleaned credentials (strip minted token material so a stale token never rehydrates); + // the declared app the admin typed is kept. Create has no server-side stored app to merge. + const restoredValues = { + ...parsed.formValues, + credentials: withoutMintedTokenCredentials(parsed.formValues.credentials), + }; + setPendingRestoredValues({ values: restoredValues, transport: restoredTransport }); + } + if (typeof parsed.authorizedIdentity === "string") { + // Re-arm invalidation: without this the remounted form has authorizedIdentity=undefined, so a + // post-restore mode/url edit would never fire the stale-token discard. + setAuthorizedIdentity(parsed.authorizedIdentity); } if (parsed.costConfig) { setCostConfig(parsed.costConfig); @@ -380,6 +439,7 @@ const CreateMCPServer: React.FC = ({ available_on_public_internet: availableOnPublicInternetRaw, delegate_auth_to_upstream: delegateAuthToUpstreamRaw, oauth_passthrough: oauthPassthroughRaw, + dcr_bridge: dcrBridgeRaw, token_validation_json: rawTokenValidationJson, ...restValues } = values; @@ -486,6 +546,11 @@ const CreateMCPServer: React.FC = ({ available_on_public_internet: Boolean(availableOnPublicInternetRaw), delegate_auth_to_upstream: Boolean(delegateAuthToUpstreamRaw), oauth_passthrough: Boolean(oauthPassthroughRaw), + // ``dcr_bridge`` is only meaningful for the client-forwarded token + // modes (true_passthrough / oauth_delegate) and defaults on when the + // toggle is shown; force false for any other auth type so a stale + // ``true`` is never persisted. Mirrors the sibling flags above. + dcr_bridge: isClientForwardedTokenMode(restValues.auth_type) ? Boolean(dcrBridgeRaw ?? true) : false, ...(restValues.auth_type === AUTH_TYPE.OAUTH2 ? { oauth2_flow: @@ -500,8 +565,20 @@ const CreateMCPServer: React.FC = ({ const includeCredentials = restValues.auth_type && AUTH_TYPES_REQUIRING_CREDENTIALS.includes(restValues.auth_type); - if (includeCredentials && credentialsPayload && Object.keys(credentialsPayload).length > 0) { - payload.credentials = credentialsPayload; + // Client-forwarded rows persist ONLY the declared app; strip any token material that lingered in + // the form (e.g. from a prior oauth2 authorize on the same session) so it can never reach the row. + const submitCredentials = isClientForwardedTokenMode(restValues.auth_type) + ? preservedDeclaredAppCredentials(credentialsPayload) + : credentialsPayload; + + if (includeCredentials && submitCredentials && Object.keys(submitCredentials).length > 0) { + payload.credentials = submitCredentials; + } + + // An interactive (oauth2) create persists its DCR-minted client from the ref (kept out of the + // form store); reuse a re-authorize's registered client instead of re-registering. + if (restValues.auth_type === AUTH_TYPE.OAUTH2 && dcrClientRef.current) { + payload.credentials = { ...(payload.credentials ?? {}), ...dcrClientRef.current }; } if (accessToken != null) { @@ -576,6 +653,9 @@ const CreateMCPServer: React.FC = ({ setHasToolAllowlistInteraction(false); setAliasManuallyEdited(false); setLogoUrl(undefined); + setAuthorizedIdentity(undefined); + dcrClientRef.current = null; + setAppMayNotMatchUpstream(false); setModalVisible(false); }; @@ -655,6 +735,8 @@ const CreateMCPServer: React.FC = ({ clearTools(); resetOAuthFlow(); setAuthorizedIdentity(undefined); + dcrClientRef.current = null; + setAppMayNotMatchUpstream(false); } }, [isModalVisible, form, clearTools, resetOAuthFlow]); @@ -663,12 +745,31 @@ const CreateMCPServer: React.FC = ({ const handleFormValuesChange = (changedValues: Record, allValues: Record) => { // Any change to a mint-relevant field (url, auth_type, oauth_flow_type, client creds/scopes, or the // authorization/token/registration endpoints — see getOAuthAuthorizationIdentity) makes a held token - // stale, so discard it and force a fresh authorize. When that happens, formValues must be rebuilt - // from the form's post-reset state, not the pre-reset allValues snapshot: the snapshot still holds - // the discarded token in credentials, and useTestMCPConnection reads formValues for tool preview. - if (isHeldOAuthTokenStale(allValues, authorizedIdentity)) { + // stale, so discard it and force a fresh authorize. The stale check reads getFieldsValue(true): the + // onValuesChange allValues argument holds only MOUNTED paths, so an unmounted identity field (e.g. + // an oauth_flow_type initialValue while in a client-forwarded mode) would compare as changed on + // every keystroke and churn the held token. When a clear happens, formValues is rebuilt from the + // form's post-reset state (not the pre-reset snapshot, which still holds the discarded token). + // Editing the client fields is the admin managing/acknowledging the app, so it always dismisses + // the "may not match upstream" warning regardless of the stale-token branch below. + // Editing the client fields is the admin managing/acknowledging the app, so it dismisses the "may + // not match upstream" warning. Otherwise a url/endpoint change while a declared app is present keeps + // the app but flags that it may not match the new upstream (the "keep + warn" behavior). This is + // independent of the held-token stale check below so it fires even without an authorize this session. + if ("credentials" in changedValues) { + setAppMayNotMatchUpstream(false); + } else { + const upstreamChanged = ["url", "spec_path", "authorization_url", "token_url", "registration_url"].some( + (key) => key in changedValues, + ); + const hasDeclaredApp = preservedDeclaredAppCredentials(form.getFieldValue("credentials")) !== undefined; + if (upstreamChanged && hasDeclaredApp) { + setAppMayNotMatchUpstream(true); + } + } + if (isHeldOAuthTokenStale(form.getFieldsValue(true), authorizedIdentity)) { clearHeldOAuthToken(changedValues); - setFormValues({ ...form.getFieldsValue(true), ...changedValues }); + setFormValues(form.getFieldsValue(true)); return; } setFormValues(allValues); @@ -989,12 +1090,14 @@ const CreateMCPServer: React.FC = ({ {shouldShowAuthValueField && ( diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx index adb3e161da5..5b8d0aac8c0 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx @@ -1,7 +1,9 @@ import React from "react"; import { describe, it, expect, vi, beforeEach } from "vitest"; import { render, screen, waitFor, fireEvent, act } from "@testing-library/react"; -import MCPServerEdit from "./mcp_server_edit"; +import userEvent from "@testing-library/user-event"; +import MCPServerEdit, { EDIT_OAUTH_UI_STATE_KEY } from "./mcp_server_edit"; +import { setSecureItem } from "@/utils/secureStorage"; import * as networking from "../networking"; import NotificationsManager from "../molecules/notifications_manager"; import { selectAntOption } from "./testUtils"; @@ -1377,6 +1379,258 @@ describe("MCPServerEdit (OAuth token persistence on save)", () => { }, ); + it.each([["true_passthrough"], ["oauth_delegate"]])( + "persists admin-entered OAuth app credentials in the update payload for the %s mode", + async (authType) => { + mockOauth.tokenResponse = { access_token: "cf-tok", expires_in: 1800, token_type: "bearer" }; + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: authType, + }); + + render( + , + ); + + const user = userEvent.setup({ delay: null }); + await user.type( + screen.getByPlaceholderText("Leave blank to keep the currently saved app (if any)"), + "org-app-client-id", + ); + await user.type( + screen.getByPlaceholderText("Leave blank to keep the currently saved secret (if any)"), + "org-app-secret", + ); + + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: "Save Changes" })[0]); + }); + + await waitFor(() => expect(networking.updateMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + + // The declared app is config and persists onto the row; the browser-held token still never + // reaches the payload or the per-user credential store. + expect(payload.credentials).toMatchObject({ + client_id: "org-app-client-id", + client_secret: "org-app-secret", + }); + expect(JSON.stringify(payload)).not.toContain("cf-tok"); + expect(networking.storeMCPOAuthUserCredential).not.toHaveBeenCalled(); + }, + ); + + it.each([["true_passthrough"], ["oauth_delegate"]])( + "preserves admin-entered app credentials when the URL changes after authorize for the %s mode", + async (authType) => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: authType, + }); + + render( + , + ); + + const user = userEvent.setup({ delay: null }); + await user.type( + screen.getByPlaceholderText("Leave blank to keep the currently saved app (if any)"), + "org-app-client-id", + ); + await user.type( + screen.getByPlaceholderText("Leave blank to keep the currently saved secret (if any)"), + "org-app-secret", + ); + + act(() => { + mockOauth.onTokenReceived?.({ access_token: "cf-tok", token_type: "bearer" }); + }); + + // The URL edit invalidates the held browser token (removeToken fires), but the declared app + // is config and must survive the invalidation into the update payload. + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), { + target: { value: "https://other.example.com/mcp" }, + }); + }); + expect(mockRemoveToken).toHaveBeenCalledWith("oauth_server_1", "user-1"); + + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: "Save Changes" })[0]); + }); + + await waitFor(() => expect(networking.updateMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.url).toBe("https://other.example.com/mcp"); + expect(payload.credentials).toMatchObject({ + client_id: "org-app-client-id", + client_secret: "org-app-secret", + }); + expect(JSON.stringify(payload)).not.toContain("cf-tok"); + }, + ); + + it("sends an explicit-null credential write when removing the saved app for true_passthrough", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: "true_passthrough", + }); + + render( + , + ); + + // Blank fields keep the stored app (the backend merges partial credential updates), so the + // edit form states that convention and removal is an explicit checkbox that saves nulls. + expect(screen.getByPlaceholderText("Leave blank to keep the currently saved app (if any)")).toBeInTheDocument(); + + fireEvent.click( + screen.getByRole("checkbox", { + name: /Remove the saved OAuth app on save/, + }), + ); + + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: "Save Changes" })[0]); + }); + + await waitFor(() => expect(networking.updateMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.credentials).toEqual({ client_id: null, client_secret: null }); + }); + + it("warns that the saved app may not match after a URL change on a client-forwarded server", async () => { + render( + , + ); + + // No warning until the upstream changes. + expect(screen.queryByText(/registered for the previous upstream/)).not.toBeInTheDocument(); + + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), { + target: { value: "https://different.example.com/mcp" }, + }); + }); + + // Keep + warn parity with the create form: the stored app is kept, and the banner appears. + expect(screen.getByText(/registered for the previous upstream/)).toBeInTheDocument(); + }); + + it("preserves a stored client_id on OAuth-resume restore even when the saved snapshot is token-only", async () => { + // Post-redirect restore: the sessionStorage snapshot carries only a minted token (no client keys), + // while the loaded server has a stored client_id. The restore must merge the server's declared app + // under the snapshot before stripping tokens, so the stored client_id is never cleared to blank. + setSecureItem( + EDIT_OAUTH_UI_STATE_KEY, + JSON.stringify({ + serverId: "oauth_server_1", + formValues: { auth_type: "true_passthrough", credentials: { access_token: "leftover-token" } }, + }), + ); + + render( + , + ); + + const clientIdField = await screen.findByPlaceholderText("Leave blank to keep the currently saved app (if any)"); + await waitFor(() => expect((clientIdField as HTMLInputElement).value).toBe("stored-client")); + // The leftover minted token must not have rehydrated anywhere. + expect(document.body.innerHTML).not.toContain("leftover-token"); + }); + + it("resets the remove-app checkbox on a server switch so it never deletes the next server's stored app", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: "true_passthrough", + }); + + const { rerender } = render( + , + ); + + // Check "remove saved app" on server A. + fireEvent.click(screen.getByRole("checkbox", { name: /Remove the saved OAuth app on save/ })); + expect( + (screen.getByRole("checkbox", { name: /Remove the saved OAuth app on save/ }) as HTMLInputElement).checked, + ).toBe(true); + + // Switch the panel to server B without unmounting. + rerender( + , + ); + + // The checkbox must have reset, so saving server B does not send the explicit-null delete write. + expect( + (screen.getByRole("checkbox", { name: /Remove the saved OAuth app on save/ }) as HTMLInputElement).checked, + ).toBe(false); + + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: "Save Changes" })[0]); + }); + + await waitFor(() => expect(networking.updateMCPServer).toHaveBeenCalledTimes(1)); + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.credentials).not.toEqual({ client_id: null, client_secret: null }); + }); + it("forwards a newly authorized browser-held token for tool loading before the form is saved", async () => { // Regression: fetchTools keyed the browser-held decision off the saved mcpServer.auth_type, so // after switching the form to true_passthrough and authorizing, the fresh token was not sent as @@ -1737,3 +1991,168 @@ describe("MCPServerEdit (max concurrent requests)", () => { expect(payload.max_concurrent_requests).toBeNull(); }); }); + +describe("MCPServerEdit (dcr_bridge toggle)", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockOauth.tokenResponse = null; + }); + + const getDcrToggle = () => document.getElementById("dcr_bridge"); + + function renderEdit(server: Record) { + render( + , + ); + } + + async function saveAndGetPayload() { + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + return payload; + } + + it.each([["true_passthrough"], ["oauth_delegate"]])("renders the toggle for a %s server", async (authType) => { + renderEdit({ auth_type: authType }); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + expect(screen.getByText("Gateway-hosted sign-in (DCR bridge)")).toBeInTheDocument(); + }); + + it.each([["oauth2"], ["api_key"], ["none"]])("does not render the toggle for an %s server", async (authType) => { + renderEdit({ auth_type: authType }); + + await waitFor(() => { + expect(screen.getAllByRole("button", { name: "Save Changes" }).length).toBeGreaterThan(0); + }); + expect(screen.queryByText("Gateway-hosted sign-in (DCR bridge)")).not.toBeInTheDocument(); + expect(getDcrToggle()).not.toBeInTheDocument(); + }); + + it("renders the toggle between the OAuth client fields and the Authorize button", async () => { + renderEdit({ auth_type: "true_passthrough" }); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + const toggle = getDcrToggle() as HTMLElement; + const secretInput = screen.getByPlaceholderText("Leave blank to keep the currently saved secret (if any)"); + const authorizeButton = screen.getByRole("button", { name: "Authorize & Fetch Tools (browser-only)" }); + expect(secretInput.compareDocumentPosition(toggle) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy(); + expect(toggle.compareDocumentPosition(authorizeButton) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy(); + }); + + it("initializes unchecked from a null stored value and saves an explicit false", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: "true_passthrough", + }); + renderEdit({ auth_type: "true_passthrough", dcr_bridge: null }); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "false"); + + const payload = await saveAndGetPayload(); + expect(payload.dcr_bridge).toBe(false); + }); + + it("initializes checked from a stored true and saves an explicit true", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: "oauth_delegate", + dcr_bridge: true, + }); + renderEdit({ auth_type: "oauth_delegate", dcr_bridge: true }); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "true"); + + const payload = await saveAndGetPayload(); + expect(payload.dcr_bridge).toBe(true); + }); + + it("saves an explicit false after the admin unchecks a stored true", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: "true_passthrough", + dcr_bridge: false, + }); + renderEdit({ auth_type: "true_passthrough", dcr_bridge: true }); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + await act(async () => { + fireEvent.click(getDcrToggle()!); + }); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "false"); + + const payload = await saveAndGetPayload(); + expect(payload.dcr_bridge).toBe(false); + }); + + it("forces dcr_bridge: false when the auth type is switched away", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: "api_key", + }); + renderEdit({ auth_type: "true_passthrough", dcr_bridge: true }); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + + await selectAntOption("Authentication", "API Key"); + await waitFor(() => { + expect(getDcrToggle()).not.toBeInTheDocument(); + }); + + // Mirrors the sibling delegate_auth_to_upstream / oauth_passthrough force-false: a stale true is + // never left behind to silently re-activate if the mode is switched back. + const payload = await saveAndGetPayload(); + expect(payload.dcr_bridge).toBe(false); + }); + + it("preserves the toggle value when switching between the two client-forwarded modes", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: "oauth_delegate", + dcr_bridge: true, + }); + renderEdit({ auth_type: "true_passthrough", dcr_bridge: true }); + + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "true"); + + // The Form.Item stays mounted across the two client-forwarded modes, so the live toggle value is + // preserved rather than forced false by the switch. + await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)"); + await waitFor(() => { + expect(getDcrToggle()).toBeInTheDocument(); + }); + expect(getDcrToggle()).toHaveAttribute("aria-checked", "true"); + + const payload = await saveAndGetPayload(); + expect(payload.dcr_bridge).toBe(true); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index 7446c96c40e..6709eb02c65 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -8,6 +8,8 @@ import { getOAuthAuthorizationIdentity, CLEARED_ON_INVALIDATION, isHeldOAuthTokenStale, + preservedDeclaredAppCredentials, + withoutMintedTokenCredentials, OAUTH_FLOW, MCP_OAUTH2_FLOW_M2M, MCP_OAUTH2_FLOW_INTERACTIVE, @@ -56,6 +58,8 @@ const AUTH_TYPES_REQUIRING_CREDENTIALS = [ AUTH_TYPE.OAUTH2, AUTH_TYPE.OAUTH2_TOKEN_EXCHANGE, AUTH_TYPE.AWS_SIGV4, + AUTH_TYPE.TRUE_PASSTHROUGH, + AUTH_TYPE.OAUTH_DELEGATE, ]; export const EDIT_OAUTH_UI_STATE_KEY = "litellm-mcp-oauth-edit-state"; @@ -74,6 +78,10 @@ const MCPServerEdit: React.FC = ({ const [toolsError, setToolsError] = useState(null); const [searchValue, setSearchValue] = useState(""); const [aliasManuallyEdited, setAliasManuallyEdited] = useState(false); + const [removeStoredApp, setRemoveStoredApp] = useState(false); + // Set when the upstream identity (url/endpoints) changed while a declared app is present, so the + // section warns that the saved app may not match the new upstream (the app is kept, not wiped). + const [appMayNotMatchUpstream, setAppMayNotMatchUpstream] = useState(false); const [allowedTools, setAllowedTools] = useState([]); const [hasToolAllowlistInteraction, setHasToolAllowlistInteraction] = useState(false); const [toolNameToDisplayName, setToolNameToDisplayName] = useState>({}); @@ -179,7 +187,9 @@ const MCPServerEdit: React.FC = ({ url, transport, auth_type: isClientForwardedTokenMode(values.auth_type) ? values.auth_type : AUTH_TYPE.OAUTH2, - credentials: values.credentials, + credentials: isClientForwardedTokenMode(values.auth_type) + ? preservedDeclaredAppCredentials(values.credentials) + : values.credentials, mcp_access_groups: values.mcp_access_groups || mcpServer.mcp_access_groups, static_headers: staticHeaders, command: values.command, @@ -202,19 +212,23 @@ const MCPServerEdit: React.FC = ({ }; setToken(mcpServer.server_id, browserHeldToken, userID); NotificationsManager.success( - "Token held for this browser session. Tools can now be loaded and configured; nothing was saved to LiteLLM.", + "Token held for this browser session. Tools can now be loaded and configured; the token is not saved to LiteLLM.", ); return; } - const credentials = { + const current = (form.getFieldValue("credentials") as Record | undefined) ?? {}; + const nextCredentials = { + ...(preservedDeclaredAppCredentials(current) ?? {}), + ...(current.scopes !== undefined && { scopes: current.scopes }), access_token: token.access_token, ...(token.refresh_token && { refresh_token: token.refresh_token }), ...(token.expires_in && { expires_in: token.expires_in }), ...(token.scope && { scope: token.scope }), }; - - form.setFieldsValue({ credentials }); + // Path-replace (not deep-merge) so a re-authorize with fewer token fields does not leave stale + // siblings behind; the admin-typed client keys and scopes are carried explicitly above. + form.setFieldValue("credentials", nextCredentials); // Re-capture after writing credentials so the token is not invalidated by its own credential write. authorizedIdentityRef.current = getOAuthAuthorizationIdentity(form.getFieldsValue(true)); @@ -276,6 +290,7 @@ const MCPServerEdit: React.FC = ({ env_vars: initialEnvVars, extra_headers: mcpServer.extra_headers || [], oauth_flow_type: oauth2FlowToFormValue(mcpServer.oauth2_flow), + dcr_bridge: Boolean(mcpServer.dcr_bridge), token_validation_json: mcpServer.token_validation ? JSON.stringify(mcpServer.token_validation, null, 2) : undefined, @@ -295,6 +310,11 @@ const MCPServerEdit: React.FC = ({ } syncedServerIdRef.current = mcpServer.server_id; form.setFieldsValue(initialValues); + // Reset per-server OAuth UI state so it never carries across a server switch without an unmount: a + // stale removeStoredApp would send an explicit-null credential write that deletes the new server's + // stored app, and a stale warning would show on a server whose upstream did not change. + setAppMayNotMatchUpstream(false); + setRemoveStoredApp(false); }, [mcpServer.server_id, initialValues, form]); // Initialize cost config from existing server data @@ -332,8 +352,24 @@ const MCPServerEdit: React.FC = ({ return; } if (parsed.formValues) { - setPendingRestoredValues({ ...mcpServer, ...parsed.formValues }); + // Rebuild credentials from the declared app in EITHER the loaded server or the saved snapshot, + // then strip minted token material. Merging the two (server under snapshot) before stripping is + // what guarantees a token-only snapshot never clears a stored client_id/client_secret: the + // server's declared app survives and only the token keys drop. Assigning the cleaned result (not + // spreading the raw snapshot) also ensures a stale token can never rehydrate into the form. + const restoredCredentials = withoutMintedTokenCredentials({ + ...(mcpServer.credentials ?? {}), + ...((parsed.formValues.credentials as Record | undefined) ?? {}), + }); + const restoredValues = { + ...mcpServer, + ...parsed.formValues, + credentials: restoredCredentials, + }; + setPendingRestoredValues(restoredValues); } + // The ref is re-armed by onTokenReceived when the redirect completes the code exchange, so there + // is no separate restore-side re-arm here (writing a ref inside an effect is disallowed). if (parsed.costConfig) { setCostConfig(parsed.costConfig); } @@ -407,7 +443,13 @@ const MCPServerEdit: React.FC = ({ } setTools([]); resetOAuthFlow(); + // The admin-typed app is upstream-scoped config, not minted material, so it survives every + // invalidation; only the held token is discarded. Token-shaped keys are excluded by the filter. + const keptAppCredentials = preservedDeclaredAppCredentials(form.getFieldValue("credentials")); form.resetFields([...CLEARED_ON_INVALIDATION]); + if (keptAppCredentials) { + form.setFieldsValue({ credentials: keptAppCredentials }); + } const preserved = Object.fromEntries( CLEARED_ON_INVALIDATION.filter((key) => key in changedValues).map((key) => [key, changedValues[key]]), ); @@ -417,6 +459,21 @@ const MCPServerEdit: React.FC = ({ }; const handleFormValuesChange = (changedValues: Record) => { + // Editing the client fields dismisses the "may not match upstream" warning; otherwise a url/endpoint + // change while a declared app is present keeps the app but flags that it may not match the new + // upstream (the "keep + warn" behavior). Mirrors the create form; independent of the held-token + // stale check so it fires even without an authorize this session (the stored app is for the old url). + if ("credentials" in changedValues) { + setAppMayNotMatchUpstream(false); + } else { + const upstreamChanged = ["url", "spec_path", "authorization_url", "token_url", "registration_url"].some( + (key) => key in changedValues, + ); + const hasDeclaredApp = preservedDeclaredAppCredentials(form.getFieldValue("credentials")) !== undefined; + if (upstreamChanged && hasDeclaredApp) { + setAppMayNotMatchUpstream(true); + } + } if (isHeldOAuthTokenStale(form.getFieldsValue(true), authorizedIdentityRef.current)) { clearHeldOAuthToken(changedValues); } @@ -627,6 +684,7 @@ const MCPServerEdit: React.FC = ({ available_on_public_internet: availableOnPublicInternetRaw, delegate_auth_to_upstream: delegateAuthToUpstreamRaw, oauth_passthrough: oauthPassthroughRaw, + dcr_bridge: dcrBridgeRaw, token_validation_json: rawTokenValidationJson, ...restValues } = values; @@ -837,6 +895,15 @@ const MCPServerEdit: React.FC = ({ ? Boolean(oauthPassthroughRaw ?? mcpServer.oauth_passthrough) : false; })(), + // ``dcr_bridge`` is only meaningful for the client-forwarded token + // modes (true_passthrough / oauth_delegate). The Form.Item is + // conditionally rendered so the value drops out of the form on + // auth_type change; force false for any other configuration to avoid + // persisting a stale ``true`` that would silently re-activate if the + // mode is later switched back. + dcr_bridge: isClientForwardedTokenMode(restValues.auth_type) + ? Boolean(dcrBridgeRaw ?? mcpServer.dcr_bridge) + : false, ...(restValues.auth_type === AUTH_TYPE.OAUTH2 && restValues.oauth_flow_type ? { oauth2_flow: @@ -850,8 +917,22 @@ const MCPServerEdit: React.FC = ({ const includeCredentials = restValues.auth_type && AUTH_TYPES_REQUIRING_CREDENTIALS.includes(restValues.auth_type); - if (includeCredentials && credentialsPayload && Object.keys(credentialsPayload).length > 0) { - payload.credentials = credentialsPayload; + // Client-forwarded rows persist ONLY the declared app; strip any token material lingering in the + // form (e.g. from a prior oauth2 authorize this session) so it can never reach the row. + const submitCredentials = isClientForwardedTokenMode(restValues.auth_type) + ? preservedDeclaredAppCredentials(credentialsPayload) + : credentialsPayload; + + if (includeCredentials && submitCredentials && Object.keys(submitCredentials).length > 0) { + payload.credentials = submitCredentials; + } + + // Explicit removal of a saved app for the client-forwarded modes, applied AFTER the filter so it + // always wins. Blank fields are the keep-existing convention (the backend merges partial + // credential updates), so removal must be an explicit-null write: encrypt skips nulls and the + // merge overrides the stored keys, returning the server to dynamic client registration. + if (removeStoredApp && isClientForwardedTokenMode(restValues.auth_type)) { + payload.credentials = { client_id: null, client_secret: null }; } const updated = await updateMCPServer(accessToken, payload); @@ -895,6 +976,7 @@ const MCPServerEdit: React.FC = ({ } NotificationsManager.success("MCP Server updated successfully"); + setAppMayNotMatchUpstream(false); onSuccess(updated); } catch (error: any) { NotificationsManager.fromBackend("Failed to update MCP Server" + (error?.message ? `: ${error.message}` : "")); @@ -1040,6 +1122,11 @@ const MCPServerEdit: React.FC = ({ error: oauthError, tokenResponse: oauthTokenResponse, }} + isEditing + savedAuthType={mcpServer.auth_type} + removeStoredApp={removeStoredApp} + onRemoveStoredAppChange={setRemoveStoredApp} + appMayNotMatchUpstream={appMayNotMatchUpstream} /> )} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.test.tsx index c6faca1fb51..7ce803d583c 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.test.tsx @@ -10,6 +10,9 @@ import { getOAuthAuthorizationIdentity, isHeldOAuthTokenStale, oauth2FlowToFormValue, + preservedDeclaredAppCredentials, + withoutMintedTokenCredentials, + credentialAuthClass, } from "./types"; describe("getOAuthAuthorizationIdentity", () => { @@ -180,3 +183,51 @@ describe("oauth2FlowToFormValue", () => { expect(oauth2FlowToFormValue(undefined)).toBeUndefined(); }); }); + +describe("preservedDeclaredAppCredentials", () => { + it("keeps only non-empty string declared-app keys and never token-shaped keys", () => { + expect(preservedDeclaredAppCredentials(undefined)).toBeUndefined(); + expect(preservedDeclaredAppCredentials({})).toBeUndefined(); + expect(preservedDeclaredAppCredentials({ client_id: 123 })).toBeUndefined(); + expect(preservedDeclaredAppCredentials({ client_id: "" })).toBeUndefined(); + expect(preservedDeclaredAppCredentials({ client_id: "a", access_token: "t", scopes: ["s"] })).toEqual({ + client_id: "a", + }); + expect(preservedDeclaredAppCredentials({ client_secret: "s" })).toEqual({ client_secret: "s" }); + expect(preservedDeclaredAppCredentials({ client_id: "a", client_secret: "b", refresh_token: "r" })).toEqual({ + client_id: "a", + client_secret: "b", + }); + }); +}); + +describe("withoutMintedTokenCredentials", () => { + it("drops token keys and keeps the declared app and other config", () => { + expect(withoutMintedTokenCredentials(undefined)).toBeUndefined(); + const mixed = { + client_id: "a", + client_secret: "b", + access_token: "t", + refresh_token: "r", + expires_in: 3600, + scope: "read", + scopes: ["read"], + }; + expect(withoutMintedTokenCredentials(mixed)).toEqual({ client_id: "a", client_secret: "b", scopes: ["read"] }); + }); + + it("returns undefined (not {}) when only minted keys are present, so a restore never blanks the fields", () => { + expect(withoutMintedTokenCredentials({ access_token: "t", refresh_token: "r", expires_in: 3600 })).toBeUndefined(); + // A declared client is always kept, so a stored client_id can never be overwritten with empty. + expect(withoutMintedTokenCredentials({ client_id: "x", access_token: "t" })).toEqual({ client_id: "x" }); + }); +}); + +describe("credentialAuthClass", () => { + it("collapses the client-forwarded modes to one class and leaves others distinct", () => { + expect(credentialAuthClass(AUTH_TYPE.TRUE_PASSTHROUGH)).toBe("client_forwarded"); + expect(credentialAuthClass(AUTH_TYPE.OAUTH_DELEGATE)).toBe("client_forwarded"); + expect(credentialAuthClass(AUTH_TYPE.OAUTH2)).toBe(AUTH_TYPE.OAUTH2); + expect(credentialAuthClass(null)).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 3eba8b30968..04766aad7b4 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -96,6 +96,54 @@ export const getOAuthAuthorizationIdentity = (values: Record): // edit forms so what gets wiped cannot drift. export const CLEARED_ON_INVALIDATION = ["credentials"] as const; +// The declared-app filter over form.credentials. It is a pure key filter with no mode/transition +// guard because the surrounding code establishes that a client_id/client_secret in form.credentials +// is ALWAYS admin-typed in every reachable state: the create form holds the DCR-minted client in a +// ref and never writes it into the form store, the edit form's onTokenReceived never writes client +// keys, and the invalidation reset clears the whole object atomically. So preserving the string +// client keys across any invalidation (URL/endpoint edit, true_passthrough<->oauth_delegate switch, +// or a round trip through another mode) is always legitimate, while the output key filter excludes +// token-shaped keys so a preserve can never carry minted material through. Shared by both forms. +const DECLARED_APP_CREDENTIAL_KEYS = ["client_id", "client_secret"] as const; + +// Minted token material the oauth2 authorize path writes beside the app keys; stripped from restored +// snapshots and from any credentials that transit to the temp-session preview so a stale token never +// reaches the backend or a client-forwarded server row. +export const MINTED_TOKEN_CREDENTIAL_KEYS = ["access_token", "refresh_token", "expires_in", "scope"] as const; + +export const preservedDeclaredAppCredentials = ( + credentials: Record | null | undefined, +): Record | undefined => { + if (!credentials) return undefined; + const kept = Object.fromEntries( + DECLARED_APP_CREDENTIAL_KEYS.filter((key) => typeof credentials[key] === "string" && credentials[key] !== "").map( + (key) => [key, credentials[key] as string], + ), + ); + return Object.keys(kept).length > 0 ? kept : undefined; +}; + +// Drop minted token keys, keeping everything else (the declared app plus any non-token config). +export const withoutMintedTokenCredentials = ( + credentials: Record | null | undefined, +): Record | undefined => { + if (!credentials) return undefined; + const kept = Object.fromEntries( + Object.entries(credentials).filter(([key]) => !(MINTED_TOKEN_CREDENTIAL_KEYS as readonly string[]).includes(key)), + ); + // Return undefined (not {}) when only minted keys were present, so a restore spreads `credentials: + // undefined` (the fields keep their placeholder / keep-existing state) rather than blanking them. + return Object.keys(kept).length > 0 ? kept : undefined; +}; + +// The client-forwarded modes share one credential class (same declared app, same authorize relay), so +// a switch between them must NOT be treated as an app change. Mirrors the backend _credential_auth_class +// in db.py; kept in sync so the UI's keep-existing copy and the backend's merge cannot disagree. +export const credentialAuthClass = (authType: string | null | undefined): string | null => { + if (authType === AUTH_TYPE.TRUE_PASSTHROUGH || authType === AUTH_TYPE.OAUTH_DELEGATE) return "client_forwarded"; + return authType ?? null; +}; + // True when a token was authorized in this session (authorizedIdentity recorded at mint time) and the // form's current identity no longer matches it. Every invalidation decision in both forms goes through // this single check: onValuesChange for user edits, and an explicit recheck after any programmatic @@ -319,7 +367,10 @@ export interface MCPServer { available_on_public_internet?: boolean; delegate_auth_to_upstream?: boolean; oauth_passthrough?: boolean; + dcr_bridge?: boolean | null; max_concurrent_requests?: number | null; + /** Redacted to null in server responses; present when constructing a server locally. */ + credentials?: Record | null; /** Stdio-only fields (present when transport === 'stdio') */ command?: string | null; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 781799286e9..da6bb079876 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -22,7 +22,8 @@ export const getCallbackConfigsCall = async (accessToken: string) => { * Helper file for calls being made to proxy */ import MessageManager from "@/components/molecules/message_manager"; -import { clearTokenCookies, storeLoginToken } from "@/utils/cookieUtils"; +import { clearTokenCookies, getCookie, storeLoginToken } from "@/utils/cookieUtils"; +import { decodeToken } from "@/utils/jwtUtils"; import { TagNewRequest, TagUpdateRequest, TagListResponse, TagInfoResponse } from "./tag_management/types"; import { Team } from "./key_team_helpers/key_list"; import { EmailEventSettingsResponse, EmailEventSettingsUpdateRequest } from "./email_events/types"; @@ -38,6 +39,12 @@ import type { import { MCP_TOOLS_PREVIEW_FORBIDDEN_MESSAGE } from "./mcp_tools/constants"; import { createApiClient, deriveErrorMessage } from "@/lib/http/client"; import { resolveApiBase } from "@/lib/http/resolveApiBase"; +import { + registerAuthHeaderNameGetter, + registerAuthTokenGetter, + registerBaseUrlGetter, + registerErrorHandler, +} from "@/lib/http/runtime"; import { serverRootPath, setServerRootPath } from "@/lib/serverRootPath"; export { serverRootPath }; @@ -371,6 +378,11 @@ const apiClient = createApiClient({ onError: handleError, }); +registerBaseUrlGetter(getProxyBaseUrl); +registerAuthHeaderNameGetter(getGlobalLitellmHeaderName); +registerAuthTokenGetter(() => decodeToken(getCookie("token"))?.key ?? null); +registerErrorHandler(handleError); + export const makeModelGroupPublic = async (accessToken: string, modelGroups: string[]) => { const url = proxyBaseUrl ? `${proxyBaseUrl}/model_group/make_public` : `/model_group/make_public`; const response = await fetch(url, { diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx index df42c156975..ef0d842ad3e 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx @@ -29,6 +29,16 @@ const nameCellColumns: ColumnDef[] = [ }, ]; +const filterableColumns: ColumnDef[] = [ + { + accessorKey: "name", + header: "Name", + meta: { title: "Name" }, + filterFn: (row, columnId, value) => row.getValue(columnId) === value, + cell: ({ row }) => {row.original.name}, + }, +]; + const headerCycleColumns: ColumnDef[] = [ { accessorKey: "name", @@ -195,10 +205,110 @@ describe("DataTable pagination", () => { }); }); +describe("DataTable filtering", () => { + it("client mode filters rows by columnFilters", () => { + const { rerender } = render( + , + ); + expect(names()).toEqual(["Charlie", "Alice", "Bob"]); + + rerender( + , + ); + expect(names()).toEqual(["Alice"]); + }); + + it("client global filter matches substrings across columns", () => { + const { rerender } = render( + , + ); + expect(names()).toEqual(["Charlie", "Alice", "Bob"]); + + rerender( + , + ); + expect(names()).toEqual(["Alice"]); + }); + + it("server mode never filters locally even when columnFilters is set", () => { + render( + , + ); + expect(names()).toEqual(["Charlie", "Alice", "Bob"]); + }); + + it("throws when server filtering is missing required props", () => { + const spy = vi.spyOn(console, "error").mockImplementation(() => {}); + expect(() => render()).toThrow( + /filterMode='server'/, + ); + spy.mockRestore(); + }); +}); + +describe("DataTable loading", () => { + it("renders skeleton rows while loading and real rows once loaded", () => { + const { rerender } = render(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByTestId("name-cell")).toBeNull(); + + rerender(); + expect(screen.queryAllByTestId("skeleton-row")).toHaveLength(0); + expect(names()).toEqual(["Charlie", "Alice", "Bob"]); + }); + + it("varies skeleton shape and width per column instead of one fixed bar", () => { + const columns: ColumnDef[] = [ + { accessorKey: "name", header: "Name", meta: { skeleton: "twoLine" }, cell: () => null }, + { accessorKey: "email", header: "Email", cell: () => null }, + ]; + render(); + + const firstRow = screen.getAllByTestId("skeleton-row").at(0); + expect(firstRow).toBeDefined(); + const bars = Array.from(firstRow?.querySelectorAll('[data-slot="skeleton"]') ?? []); + + // twoLine column contributes a main + sub bar (2); the text column contributes 1 + expect(bars).toHaveLength(3); + // per-column widths differ instead of every cell sharing one fixed width + expect(new Set(bars.map((bar) => bar.className)).size).toBeGreaterThan(1); + }); +}); + describe("DataTable column visibility", () => { it("hides a column when toggled off in the view-options menu", async () => { const user = userEvent.setup(); - render( + const { container } = render( { />, ); - expect(screen.getByText("Email")).toBeInTheDocument(); + expect(container.querySelector('th[data-header-id="email"]')).not.toBeNull(); await user.click(screen.getByTestId("view-options-trigger")); await user.click(await screen.findByTestId("view-option-email")); - await waitFor(() => expect(screen.queryByText("Email")).not.toBeInTheDocument()); + await waitFor(() => expect(container.querySelector('th[data-header-id="email"]')).toBeNull()); await user.click(screen.getByTestId("view-option-email")); - await waitFor(() => expect(screen.getByText("Email")).toBeInTheDocument()); + await waitFor(() => expect(container.querySelector('th[data-header-id="email"]')).not.toBeNull()); }); it("omits columns that opt out of hiding from the menu", async () => { diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx index 2e95ee170fa..758ca5a597b 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx @@ -4,12 +4,14 @@ import { type Cell, type Column, type ColumnDef, + type ColumnFiltersState, type ColumnPinningState, type ColumnSizingState, type ExpandedState, flexRender, getCoreRowModel, getExpandedRowModel, + getFilteredRowModel, getPaginationRowModel, getSortedRowModel, type Header, @@ -21,9 +23,11 @@ import { useReactTable, type VisibilityState, } from "@tanstack/react-table"; +import { SearchX } from "lucide-react"; import * as React from "react"; import { Fragment, useState } from "react"; +import { Skeleton } from "@/components/ui/skeleton"; import { Table as TableRoot, TableBody, @@ -37,7 +41,7 @@ import { cn } from "@/lib/cva.config"; import "./columnMeta"; import { DataTablePagination, DEFAULT_PAGE_SIZE_OPTIONS } from "./DataTablePagination"; -import type { ColumnPinnedSide, DataTableProps, DataTableSize, PaginationMode, SortingMode } from "./types"; +import type { ColumnPinnedSide, DataTableProps, DataTableSize, FilterMode, PaginationMode, SortingMode } from "./types"; const INTERACTIVE_SELECTOR = "button, a, input, select, textarea, [role=checkbox], [data-row-click-exempt]"; @@ -60,14 +64,22 @@ export function validateDataTableConfig( props.pagination === undefined || props.onPaginationChange === undefined || props.rowCount === undefined; const serverPaginationIncomplete = props.paginationMode === "server" && serverPaginationPropsMissing; + const serverFilteringIncomplete = + props.filterMode === "server" && (props.columnFilters === undefined || props.onColumnFiltersChange === undefined); + const bothSortingSources = props.defaultSorting !== undefined && props.sorting !== undefined; + const bothFilterSources = props.defaultColumnFilters !== undefined && props.columnFilters !== undefined; return [ serverSortingIncomplete ? "sortingMode='server' requires both `sorting` and `onSortingChange`." : null, serverPaginationIncomplete ? "paginationMode='server' requires `pagination`, `onPaginationChange`, and `rowCount`." : null, + serverFilteringIncomplete ? "filterMode='server' requires both `columnFilters` and `onColumnFiltersChange`." : null, bothSortingSources ? "Provide either `defaultSorting` (uncontrolled) or `sorting` (controlled), not both." : null, + bothFilterSources + ? "Provide either `defaultColumnFilters` (uncontrolled) or `columnFilters` (controlled), not both." + : null, ].filter((message): message is string => message !== null); } @@ -93,9 +105,11 @@ function derivePinning(columns: ColumnDef[]): Colu function buildRowModels( sortingMode: SortingMode, paginationMode: PaginationMode, + filterMode: FilterMode, getRowCanExpand: ((row: Row) => boolean) | undefined, ): Partial> { return { + ...(filterMode === "client" ? { getFilteredRowModel: getFilteredRowModel() } : {}), ...(sortingMode === "client" ? { getSortedRowModel: getSortedRowModel() } : {}), ...(paginationMode === "client" ? { getPaginationRowModel: getPaginationRowModel() } : {}), ...(getRowCanExpand !== undefined ? { getRowCanExpand, getExpandedRowModel: getExpandedRowModel() } : {}), @@ -307,6 +321,65 @@ function MessageRow({ colSpan, children }: { colSpan: number; children: React.Re ); } +function DefaultEmptyState() { + return ( +
+
+ +
+
No results
+
No rows match your search or filters.
+
+ ); +} + +const SKELETON_WIDTHS = ["w-[58%]", "w-[44%]", "w-[70%]", "w-[50%]", "w-[64%]", "w-[48%]"] as const; + +function SkeletonCell({ column, index }: { column: Column | undefined; index: number }) { + const meta = column?.columnDef.meta; + const width = SKELETON_WIDTHS[index % SKELETON_WIDTHS.length]; + if (meta?.skeleton === "twoLine") { + return ( +
+ + +
+ ); + } + return ; +} + +function SkeletonRows({ + rowCount, + columns, + size, + message, +}: { + rowCount: number; + columns: readonly Column[]; + size: DataTableSize; + message?: string; +}) { + const rowKeys = Array.from({ length: Math.max(rowCount, 1) }, (_, index) => index); + const cells = columns.length > 0 ? columns : [undefined]; + return ( + + {rowKeys.map((rowKey) => ( + + {cells.map((column, columnKey) => ( + + + {rowKey === 0 && columnKey === 0 && message !== undefined ? ( + {message} + ) : null} + + ))} + + ))} + + ); +} + function useControllable( controlled: T | undefined, controlledOnChange: OnChangeFn | undefined, @@ -334,6 +407,12 @@ function useDataTableInstance(props: DataTablePro onPaginationChange, rowCount, pageSizeOptions = DEFAULT_PAGE_SIZE_OPTIONS, + filterMode = "none", + columnFilters, + onColumnFiltersChange, + defaultColumnFilters, + globalFilter, + onGlobalFilterChange, enableColumnResizing = false, columnResizeMode = "onEnd", defaultColumnVisibility, @@ -348,6 +427,12 @@ function useDataTableInstance(props: DataTablePro pageIndex: 0, pageSize: pageSizeOptions[0] ?? 25, }); + const filterState = useControllable( + columnFilters, + onColumnFiltersChange, + defaultColumnFilters ?? [], + ); + const globalFilterState = useControllable(globalFilter, onGlobalFilterChange, ""); const expandedState = useControllable(expanded, onExpandedChange, {}); const [columnVisibility, setColumnVisibility] = useState(defaultColumnVisibility ?? {}); const [columnSizing, setColumnSizing] = useState({}); @@ -360,6 +445,8 @@ function useDataTableInstance(props: DataTablePro state: { sorting: sortingState.value, pagination: paginationState.value, + columnFilters: filterState.value, + globalFilter: globalFilterState.value, expanded: expandedState.value, columnVisibility, columnSizing, @@ -367,16 +454,19 @@ function useDataTableInstance(props: DataTablePro initialState: { columnPinning }, manualSorting: sortingMode === "server", manualPagination: paginationMode === "server", + manualFiltering: filterMode === "server", enableSortingRemoval, enableColumnResizing, columnResizeMode, onSortingChange: sortingState.onChange, onPaginationChange: paginationState.onChange, + onColumnFiltersChange: filterState.onChange, + onGlobalFilterChange: globalFilterState.onChange, onExpandedChange: expandedState.onChange, onColumnVisibilityChange: setColumnVisibility, onColumnSizingChange: setColumnSizing, getCoreRowModel: getCoreRowModel(), - ...buildRowModels(sortingMode, paginationMode, expansionGuard), + ...buildRowModels(sortingMode, paginationMode, filterMode, expansionGuard), ...(getRowId !== undefined ? { getRowId } : {}), ...(paginationMode === "server" && rowCount !== undefined ? { rowCount } : {}), }; @@ -397,7 +487,8 @@ export function DataTable(props: DataTableProps(props: DataTableProps { if (isLoading) { - return {loadingMessage}; + return ( + + ); } if (rows.length === 0) { - return {noDataMessage}; + return {noDataMessage ?? }; } return rows.map((row) => ( (props: DataTableProps - {toolbar !== undefined &&
{toolbar(table)}
} -
- - - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - - ))} - - ))} - - {renderBody()} - {footer !== undefined && {footer(table)}} - +
+ {toolbar !== undefined &&
{toolbar(table)}
} +
+ + + {table.getHeaderGroups().map((headerGroup) => ( + + {headerGroup.headers.map((header) => ( + + ))} + + ))} + + {renderBody()} + {footer !== undefined && {footer(table)}} + +
+ {paginationNode !== null &&
{paginationNode}
}
- {renderPagination()}
); } diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.test.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.test.tsx new file mode 100644 index 00000000000..0770c17cba6 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.test.tsx @@ -0,0 +1,98 @@ +import type { ColumnDef, ColumnFiltersState } from "@tanstack/react-table"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { useState } from "react"; +import { describe, expect, it } from "vitest"; + +import { DataTable } from "./DataTable"; +import { DataTableFilterDrawer } from "./DataTableFilterDrawer"; +import { DataTableToolbar } from "./DataTableToolbar"; + +interface Person { + id: string; + name: string; +} + +const DATA: Person[] = [ + { id: "a", name: "Alice" }, + { id: "b", name: "Bob" }, + { id: "c", name: "Carol" }, +]; + +const columns: ColumnDef[] = [ + { + accessorKey: "name", + header: "Name", + meta: { title: "Name" }, + filterFn: (row, columnId, value) => row.getValue(columnId) === value, + cell: ({ row }) => {row.original.name}, + }, +]; + +const names = (): (string | null)[] => screen.getAllByTestId("name-cell").map((el) => el.textContent); + +function Harness({ initialFilters }: { initialFilters?: ColumnFiltersState }) { + const [open, setOpen] = useState(false); + return ( + ( + <> + setOpen(true)} /> + + {({ get, set }) => ( + set("name", event.target.value)} + /> + )} + + + )} + /> + ); +} + +describe("DataTableFilterDrawer", () => { + it("stages edits and only commits them to the table on Apply", async () => { + const user = userEvent.setup(); + render(); + expect(names()).toEqual(["Alice", "Bob", "Carol"]); + + await user.click(screen.getByTestId("datatable-filters-trigger")); + await user.type(await screen.findByTestId("draft-name"), "Bob"); + + expect(names()).toEqual(["Alice", "Bob", "Carol"]); + expect(screen.queryByTestId("filter-chip-name")).toBeNull(); + + await user.click(screen.getByTestId("filter-drawer-apply")); + expect(names()).toEqual(["Bob"]); + expect(screen.getByTestId("filter-chip-name")).toHaveTextContent("Bob"); + }); + + it("seeds the draft from committed filters when opened", async () => { + const user = userEvent.setup(); + render(); + expect(names()).toEqual(["Bob"]); + + await user.click(screen.getByTestId("datatable-filters-trigger")); + expect(await screen.findByTestId("draft-name")).toHaveValue("Bob"); + }); + + it("reset clears the committed filters and the draft", async () => { + const user = userEvent.setup(); + render(); + + await user.click(screen.getByTestId("datatable-filters-trigger")); + await user.click(await screen.findByTestId("filter-drawer-reset")); + + expect(names()).toEqual(["Alice", "Bob", "Carol"]); + expect(screen.queryByTestId("filter-chip-name")).toBeNull(); + expect(screen.getByTestId("draft-name")).toHaveValue(""); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.tsx new file mode 100644 index 00000000000..8aaec3f13a0 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableFilterDrawer.tsx @@ -0,0 +1,108 @@ +"use client"; + +import type { ColumnFiltersState, Table } from "@tanstack/react-table"; +import * as React from "react"; + +import { Button } from "@/components/ui/button"; +import { Label } from "@/components/ui/label"; +import { Sheet, SheetContent, SheetDescription, SheetFooter, SheetHeader, SheetTitle } from "@/components/ui/sheet"; + +export interface FilterDraft { + get: (columnId: string) => unknown; + set: (columnId: string, value: unknown) => void; +} + +interface DataTableFilterDrawerProps { + table: Table; + open: boolean; + onOpenChange: (open: boolean) => void; + title?: string; + description?: React.ReactNode; + applyLabel?: string; + resetLabel?: string; + children: (draft: FilterDraft) => React.ReactNode; +} + +function isEmpty(value: unknown): boolean { + if (Array.isArray(value)) { + return value.length === 0; + } + return value === undefined || value === null || value === ""; +} + +function toDraft(filters: ColumnFiltersState): Record { + return Object.fromEntries(filters.map((filter) => [filter.id, filter.value])); +} + +function toFilters(draft: Record): ColumnFiltersState { + return Object.entries(draft) + .filter(([, value]) => !isEmpty(value)) + .map(([id, value]) => ({ id, value })); +} + +export function DataTableFilterDrawer({ + table, + open, + onOpenChange, + title = "Filters", + description, + applyLabel = "Apply Filters", + resetLabel = "Reset", + children, +}: DataTableFilterDrawerProps) { + const [draft, setDraft] = React.useState>(() => toDraft(table.getState().columnFilters)); + const [wasOpen, setWasOpen] = React.useState(open); + + if (open !== wasOpen) { + setWasOpen(open); + if (open) { + setDraft(toDraft(table.getState().columnFilters)); + } + } + + const helpers: FilterDraft = { + get: (columnId) => draft[columnId], + set: (columnId, value) => setDraft((previous) => ({ ...previous, [columnId]: value })), + }; + + const apply = () => { + table.setColumnFilters(toFilters(draft)); + onOpenChange(false); + }; + + const reset = () => { + setDraft({}); + table.setColumnFilters([]); + }; + + return ( + + + + {title} + {description !== undefined && {description}} + +
+ {children(helpers)} +
+ + + + +
+
+ ); +} + +export function DataTableFilterField({ label, children }: { label: string; children: React.ReactNode }) { + return ( +
+ + {children} +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTablePagination.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTablePagination.tsx index 5a30b12f27f..7802465ef74 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTablePagination.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTablePagination.tsx @@ -37,7 +37,7 @@ export function DataTablePagination({ const lastPage = Math.max(pageCount - 1, 0); return ( -
+
Rows per page onSearchChange(event.target.value)} + placeholder={searchPlaceholder} + className="h-8 w-56 pl-8" + data-testid="datatable-search" + /> +
)} - {onToggleFilters !== undefined && ( - + {filters.map((filter) => ( + + {labelFor(filter.id)}: + {valueFor(filter.id, filter.value)} + + + ))} + {filters.length > 0 && ( + + )} +
+
+ {children} + {onRefresh !== undefined && ( + + )} + {showViewOptions && } + {onOpenFilters !== undefined && ( + )} - {showReset && }
- {children !== undefined &&
{children}
}
); } diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableViewOptions.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableViewOptions.tsx index ab56aafe7b5..f462481ef9c 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableViewOptions.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableViewOptions.tsx @@ -2,7 +2,7 @@ import { Menu } from "@base-ui/react/menu"; import type { Table } from "@tanstack/react-table"; -import { Check, SlidersHorizontal } from "lucide-react"; +import { Check, Columns3 } from "lucide-react"; import { Button } from "@/components/ui/button"; @@ -24,7 +24,7 @@ export function DataTableViewOptions({ table, label = "View", className } - + {label} } @@ -44,7 +44,8 @@ export function DataTableViewOptions({ table, label = "View", className } - {column.columnDef.meta?.title ?? column.id} + {column.columnDef.meta?.title ?? + (typeof column.columnDef.header === "string" ? column.columnDef.header : column.id)} ))} diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/columnMeta.ts b/ui/litellm-dashboard/src/components/shared/DataTable/columnMeta.ts index 46e72226038..0f14c277c6f 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/columnMeta.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/columnMeta.ts @@ -1,6 +1,6 @@ import type { RowData } from "@tanstack/react-table"; -import type { ColumnPinnedSide } from "./types"; +import type { ColumnPinnedSide, DataTableSkeletonShape } from "./types"; declare module "@tanstack/react-table" { interface ColumnMeta { @@ -9,5 +9,6 @@ declare module "@tanstack/react-table" { headerClassName?: string; title?: string; pinned?: ColumnPinnedSide; + skeleton?: DataTableSkeletonShape; } } diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/index.ts b/ui/litellm-dashboard/src/components/shared/DataTable/index.ts index 49a4430bbee..c4218f6051a 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/index.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/index.ts @@ -1,6 +1,7 @@ import "./columnMeta"; export { DataTable, DataTableConfigError, validateDataTableConfig } from "./DataTable"; +export { DataTableFilterDrawer, DataTableFilterField, type FilterDraft } from "./DataTableFilterDrawer"; export { DataTablePagination, DEFAULT_PAGE_SIZE_OPTIONS } from "./DataTablePagination"; export { DataTableToolbar } from "./DataTableToolbar"; export { DataTableViewOptions } from "./DataTableViewOptions"; @@ -11,6 +12,7 @@ export type { ColumnResizeMode, DataTableProps, DataTableSize, + FilterMode, PaginationMode, SortingMode, } from "./types"; diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts index 8fa6f21c4d3..f5130b4c823 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts @@ -1,5 +1,6 @@ import type { ColumnDef, + ColumnFiltersState, ExpandedState, OnChangeFn, PaginationState, @@ -13,9 +14,11 @@ import type * as React from "react"; export type SortingMode = "none" | "client" | "server"; export type PaginationMode = "none" | "client" | "server"; +export type FilterMode = "none" | "client" | "server"; export type ColumnResizeMode = "onEnd" | "onChange"; export type DataTableSize = "compact" | "default"; export type ColumnPinnedSide = "left" | "right"; +export type DataTableSkeletonShape = "text" | "twoLine"; export interface DataTableProps { data: TData[]; @@ -24,6 +27,7 @@ export interface DataTableProps { isLoading?: boolean; loadingMessage?: string; + skeletonRowCount?: number; noDataMessage?: React.ReactNode; sortingMode?: SortingMode; @@ -38,6 +42,14 @@ export interface DataTableProps { rowCount?: number; pageSizeOptions?: number[]; + filterMode?: FilterMode; + columnFilters?: ColumnFiltersState; + onColumnFiltersChange?: OnChangeFn; + defaultColumnFilters?: ColumnFiltersState; + + globalFilter?: string; + onGlobalFilterChange?: OnChangeFn; + enableColumnResizing?: boolean; columnResizeMode?: ColumnResizeMode; defaultColumnVisibility?: VisibilityState; diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index 4ccfb891417..f0e27607ffb 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -583,7 +583,7 @@ describe("TeamInfoView", () => { await waitFor(() => { expect(screen.getByRole("button", { name: "Filters" })).toBeInTheDocument(); }); - expect(screen.getByRole("button", { name: "Reset Filters" })).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Columns" })).toBeInTheDocument(); expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-1 of 1"); expect(screen.getByTestId("pagination-prev")).toBeInTheDocument(); expect(screen.getByTestId("pagination-next")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx index fd81aaaad99..1c9b14d0c9d 100644 --- a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx @@ -4,8 +4,6 @@ import { beforeEach, describe, expect, it, vi, MockedFunction } from "vitest"; import { renderWithProviders } from "../../../tests/test-utils"; import { TeamVirtualKeysTable } from "./TeamVirtualKeysTable"; import { KeysResponse, useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; -import { fetchTeamFilterOptions } from "../key_team_helpers/filter_helpers"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { KeyResponse } from "../key_team_helpers/key_list"; import { Organization } from "../networking"; @@ -13,18 +11,6 @@ vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({ useKeys: vi.fn(), })); -vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ - default: vi.fn(), -})); - -vi.mock("../key_team_helpers/filter_helpers", () => ({ - fetchTeamFilterOptions: vi.fn().mockResolvedValue({ - keyAliases: [], - organizationIds: [], - userIds: [], - }), -})); - vi.mock("../key_team_helpers/fetch_available_models_team_key", () => ({ getModelDisplayName: vi.fn((model: string) => model), })); @@ -38,8 +24,12 @@ vi.mock("../templates/key_info_view", () => ({ )), })); +// Resolve the debounced search synchronously so typed input lands in the useKeys query within the test tick. +vi.mock("@tanstack/react-pacer/debouncer", () => ({ + useDebouncedValue: (value: unknown) => [value, { cancel: vi.fn(), flush: vi.fn() }], +})); + const mockUseKeys = useKeys as MockedFunction; -const mockUseAuthorized = useAuthorized as MockedFunction; const createMockKey = (overrides: Partial = {}): KeyResponse => ({ @@ -85,7 +75,6 @@ describe("TeamVirtualKeysTable", () => { beforeEach(() => { vi.clearAllMocks(); - mockUseAuthorized.mockReturnValue({ accessToken: "test-token" } as any); mockUseKeys.mockReturnValue({ data: { keys: [], total_count: 0, current_page: 1, total_pages: 1 } as KeysResponse, isPending: false, @@ -262,30 +251,48 @@ describe("TeamVirtualKeysTable", () => { await waitFor(() => expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.anything())); }); - it("resets the sort order to the default when filters are reset", async () => { + it("maps the User ID drawer filter to a server-side useKeys query and clears it", async () => { const user = userEvent.setup(); - const result = { + mockUseKeys.mockReturnValue({ data: { keys: [createMockKey()], total_count: 1, current_page: 1, total_pages: 1 }, isPending: false, isFetching: false, refetch: vi.fn(), - } as unknown as ReturnType; - mockUseKeys.mockReturnValue(result); + } as unknown as ReturnType); renderWithProviders(); - await user.click(await screen.findByTestId("sort-header-created_at")); + await user.click(await screen.findByTestId("datatable-filters-trigger")); + const drawerBody = await screen.findByTestId("filter-drawer-body"); + const userInput = drawerBody.querySelector("input") as HTMLElement; + await user.type(userInput, "user-42"); + await user.click(screen.getByTestId("filter-drawer-apply")); + await waitFor(() => - expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ sortOrder: "asc" })), + expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ userID: "user-42" })), ); - await user.click(screen.getByRole("button", { name: "Reset Filters" })); + await user.click(screen.getByTestId("datatable-clear-filters")); await waitFor(() => - expect(mockUseKeys).toHaveBeenLastCalledWith( - 1, - 50, - expect.objectContaining({ sortBy: "created_at", sortOrder: "desc" }), - ), + expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ userID: undefined })), + ); + }); + + it("maps the search box to a server-side key-alias query", async () => { + const user = userEvent.setup(); + mockUseKeys.mockReturnValue({ + data: { keys: [createMockKey()], total_count: 1, current_page: 1, total_pages: 1 }, + isPending: false, + isFetching: false, + refetch: vi.fn(), + } as unknown as ReturnType); + + renderWithProviders(); + + await user.type(await screen.findByTestId("datatable-search"), "check-002"); + + await waitFor(() => + expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ selectedKeyAlias: "check-002" })), ); }); @@ -304,7 +311,7 @@ describe("TeamVirtualKeysTable", () => { }); }); - it("should show No keys found when keys array is empty", async () => { + it("should show the empty state when keys array is empty", async () => { mockUseKeys.mockReturnValue({ data: { keys: [], total_count: 0, current_page: 1, total_pages: 1 } as KeysResponse, isPending: false, @@ -315,26 +322,7 @@ describe("TeamVirtualKeysTable", () => { renderWithProviders(); await waitFor(() => { - expect(screen.getByText("No keys found")).toBeInTheDocument(); - }); - }); - - it("should fetch team-scoped filter options for Key Alias, Organization ID, and User ID", async () => { - const mockFetchTeamFilterOptions = vi.mocked(fetchTeamFilterOptions); - mockFetchTeamFilterOptions.mockResolvedValue({ - keyAliases: ["alice_key_team1", "charlie_key_team1"], - organizationIds: ["org-123"], - userIds: [ - { id: "user-1", email: "alice@example.com" }, - { id: "user-2", email: "charlie@example.com" }, - ], - }); - - // Use unique teamId to avoid cache hit from previous tests (refetchOnMount: false) - renderWithProviders(); - - await waitFor(() => { - expect(mockFetchTeamFilterOptions).toHaveBeenCalledWith("test-token", "team-filter-options-test"); + expect(screen.getByText("No rows match your search or filters.")).toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx index 73f524e13f8..207e6f2ccfe 100644 --- a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx @@ -1,21 +1,25 @@ "use client"; import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; import { DateCell, IdCell, MoneyCell } from "@/components/shared/table_cells"; -import { DataTable, DataTablePagination, DataTableSortHeader } from "@/components/shared/DataTable"; +import { + DataTable, + DataTableFilterDrawer, + DataTableFilterField, + DataTableSortHeader, + DataTableToolbar, +} from "@/components/shared/DataTable"; +import { Input } from "@/components/ui/input"; import { ChevronDownIcon, ChevronRightIcon } from "@heroicons/react/outline"; -import { ColumnDef, PaginationState, SortingState } from "@tanstack/react-table"; +import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; +import { ColumnDef, ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; import { Badge, Icon, Text } from "@tremor/react"; import { Popover, Tooltip, Typography } from "antd"; import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; import React, { useCallback, useEffect, useMemo, useState } from "react"; import { getModelDisplayName } from "../key_team_helpers/fetch_available_models_team_key"; import { KeyResponse, Team } from "../key_team_helpers/key_list"; -import FilterComponent, { FilterOption } from "../molecules/filter"; import { Organization } from "../networking"; import KeyInfoView from "../templates/key_info_view"; -import { useQuery } from "@tanstack/react-query"; -import { fetchTeamFilterOptions } from "../key_team_helpers/filter_helpers"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; interface TeamVirtualKeysTableProps { teamId: string; @@ -30,18 +34,29 @@ interface TeamVirtualKeysTableProps { const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }]; export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVirtualKeysTableProps) { - const { accessToken } = useAuthorized(); const [selectedKey, setSelectedKey] = useState(null); const [sorting, setSorting] = useState(DEFAULT_SORTING); const [tablePagination, setTablePagination] = useState({ pageIndex: 0, pageSize: 50, }); - const [filters, setFilters] = useState>({ - "Organization ID": "", - "Key Alias": "", - "User ID": "", - }); + const [columnFilters, setColumnFilters] = useState([]); + const [filtersOpen, setFiltersOpen] = useState(false); + const [searchInput, setSearchInput] = useState(""); + const [searchQuery] = useDebouncedValue(searchInput, { wait: 300 }); + + const handleSearchChange = useCallback((value: string) => { + setSearchInput(value); + setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); + }, []); + + const getFilterValue = useCallback( + (columnId: string): string | undefined => { + const entry = columnFilters.find((filter) => filter.id === columnId); + return typeof entry?.value === "string" && entry.value.trim() ? entry.value.trim() : undefined; + }, + [columnFilters], + ); const sortBy = sorting.length > 0 ? sorting[0].id : "created_at"; const sortOrder = sorting.length > 0 ? (sorting[0].desc ? "desc" : "asc") : "desc"; @@ -56,9 +71,8 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi refetch, } = useKeys(pageIndex + 1, pageSize, { teamID: teamId, - organizationID: filters["Organization ID"]?.trim() || undefined, - selectedKeyAlias: filters["Key Alias"]?.trim() || undefined, - userID: filters["User ID"]?.trim() || undefined, + selectedKeyAlias: searchQuery.trim() || undefined, + userID: getFilterValue("user_id"), sortBy: sortBy || undefined, sortOrder: sortOrder || undefined, expand: "user", @@ -95,18 +109,6 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi [teamId, teamAlias, organization], ); - const teamFilterOptionsQuery = useQuery({ - queryKey: ["teamFilterOptions", teamId, accessToken], - queryFn: async () => fetchTeamFilterOptions(accessToken, teamId), - enabled: !!accessToken && !!teamId, - staleTime: 30000, // 30 seconds - align with useKeys - }); - const teamFilterOptions = teamFilterOptionsQuery.data || { - keyAliases: [], - organizationIds: [], - userIds: [], - }; - const handleStorageChange = useCallback(() => { refetch?.(); }, [refetch]); @@ -116,76 +118,17 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi return () => window.removeEventListener("storage", handleStorageChange); }, [handleStorageChange]); - const handleFilterChange = useCallback((newFilters: Record) => { - setFilters((prev) => ({ - ...prev, - "Organization ID": newFilters["Organization ID"] ?? prev["Organization ID"], - "Key Alias": newFilters["Key Alias"] ?? prev["Key Alias"], - "User ID": newFilters["User ID"] ?? prev["User ID"], - })); + const handleColumnFiltersChange = useCallback>((updaterOrValue) => { + setColumnFilters(updaterOrValue); setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); }, []); - const handleFilterReset = useCallback(() => { - setFilters({ - "Organization ID": "", - "Key Alias": "", - "User ID": "", - }); - setSorting(DEFAULT_SORTING); - setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); - }, []); - - const filterOptions: FilterOption[] = useMemo( - () => [ - { - name: "Organization ID", - label: "Organization ID", - isSearchable: true, - searchFn: async (searchText: string) => { - const { organizationIds } = teamFilterOptions; - if (!organizationIds.length) return []; - const lower = searchText.toLowerCase(); - const filtered = lower ? organizationIds.filter((id) => id.toLowerCase().includes(lower)) : organizationIds; - return filtered.map((id) => ({ label: id, value: id })); - }, - }, - { - name: "Key Alias", - label: "Key Alias", - isSearchable: true, - searchFn: async (searchText: string) => { - const { keyAliases } = teamFilterOptions; - const lower = searchText.toLowerCase(); - const filtered = lower ? keyAliases.filter((alias) => alias.toLowerCase().includes(lower)) : keyAliases; - return filtered.map((alias) => ({ label: alias, value: alias })); - }, - }, - { - name: "User ID", - label: "User ID", - isSearchable: true, - searchFn: async (searchText: string) => { - const { userIds } = teamFilterOptions; - const lower = searchText.toLowerCase(); - const filtered = lower - ? userIds.filter((u) => u.id.toLowerCase().includes(lower) || u.email.toLowerCase().includes(lower)) - : userIds; - return filtered.map((u) => ({ - label: u.email ? `${u.id} (${u.email})` : u.id, - value: u.id, - })); - }, - }, - ], - [teamFilterOptions], - ); - const columns: ColumnDef[] = useMemo( () => [ { id: "token", accessorKey: "token", + meta: { title: "Key ID" }, header: ({ column }) => , size: 120, enableSorting: true, @@ -196,6 +139,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi { id: "key_alias", accessorKey: "key_alias", + meta: { title: "Key Alias" }, header: ({ column }) => , size: 150, enableSorting: true, @@ -268,6 +212,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi { id: "created_at", accessorKey: "created_at", + meta: { title: "Created At" }, header: ({ column }) => , size: 120, enableSorting: true, @@ -335,6 +280,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi { id: "updated_at", accessorKey: "updated_at", + meta: { title: "Updated At" }, header: ({ column }) => , size: 120, enableSorting: true, @@ -359,6 +305,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi { id: "spend", accessorKey: "spend", + meta: { title: "Spend (USD)" }, header: ({ column }) => , size: 100, enableSorting: true, @@ -367,6 +314,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi { id: "max_budget", accessorKey: "max_budget", + meta: { title: "Budget (USD)" }, header: ({ column }) => , size: 110, enableSorting: true, @@ -503,27 +451,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi onDelete={refetch} /> ) : ( -
-
- -
- -
- setTablePagination((prev) => ({ ...prev, pageIndex: nextPage }))} - onPageSizeChange={(nextSize) => setTablePagination({ pageIndex: 0, pageSize: nextSize })} - isLoading={isLoading || isFetching} - /> -
- +
null} + filterMode="server" + columnFilters={columnFilters} + onColumnFiltersChange={handleColumnFiltersChange} enableColumnResizing columnResizeMode="onChange" isLoading={isLoading || isFetching} loadingMessage="Loading keys..." - noDataMessage="No keys found" maxBodyHeight="75vh" size="compact" + toolbar={(table) => ( + <> + refetch?.()} + isRefreshing={isFetching} + onOpenFilters={() => setFiltersOpen(true)} + filterLabels={{ user_id: "User ID" }} + /> + + {({ get, set }) => ( + + set("user_id", event.target.value)} + placeholder="Filter by user ID…" + /> + + )} + + + )} />
)} diff --git a/ui/litellm-dashboard/src/components/ui/sheet.tsx b/ui/litellm-dashboard/src/components/ui/sheet.tsx new file mode 100644 index 00000000000..b619c927bbb --- /dev/null +++ b/ui/litellm-dashboard/src/components/ui/sheet.tsx @@ -0,0 +1,100 @@ +"use client"; + +import * as React from "react"; +import { Dialog as SheetPrimitive } from "@base-ui/react/dialog"; + +import { cn } from "@/lib/cva.config"; +import { Button } from "@/components/ui/button"; +import { XIcon } from "lucide-react"; + +function Sheet({ ...props }: SheetPrimitive.Root.Props) { + return ; +} + +function SheetTrigger({ ...props }: SheetPrimitive.Trigger.Props) { + return ; +} + +function SheetClose({ ...props }: SheetPrimitive.Close.Props) { + return ; +} + +function SheetPortal({ ...props }: SheetPrimitive.Portal.Props) { + return ; +} + +function SheetOverlay({ className, ...props }: SheetPrimitive.Backdrop.Props) { + return ( + + ); +} + +function SheetContent({ + className, + children, + side = "right", + showCloseButton = true, + ...props +}: SheetPrimitive.Popup.Props & { + side?: "top" | "right" | "bottom" | "left"; + showCloseButton?: boolean; +}) { + return ( + + + + {children} + {showCloseButton && ( + } + > + + Close + + )} + + + ); +} + +function SheetHeader({ className, ...props }: React.ComponentProps<"div">) { + return
; +} + +function SheetFooter({ className, ...props }: React.ComponentProps<"div">) { + return
; +} + +function SheetTitle({ className, ...props }: SheetPrimitive.Title.Props) { + return ( + + ); +} + +function SheetDescription({ className, ...props }: SheetPrimitive.Description.Props) { + return ( + + ); +} + +export { Sheet, SheetTrigger, SheetClose, SheetContent, SheetHeader, SheetFooter, SheetTitle, SheetDescription }; diff --git a/ui/litellm-dashboard/src/lib/http/api.test.ts b/ui/litellm-dashboard/src/lib/http/api.test.ts new file mode 100644 index 00000000000..fba2b7f77dc --- /dev/null +++ b/ui/litellm-dashboard/src/lib/http/api.test.ts @@ -0,0 +1,104 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { fetchClient } from "./api"; +import { + registerAuthHeaderNameGetter, + registerAuthTokenGetter, + registerBaseUrlGetter, + registerErrorHandler, +} from "./runtime"; + +const jsonResponse = (status: number, body: unknown): Response => + new Response(JSON.stringify(body), { status, headers: { "Content-Type": "application/json" } }); + +const capturingFetch = (response: Response) => { + const requests: Request[] = []; + const fetch = vi.fn(async (request: Request) => { + requests.push(request); + return response; + }); + return { fetch, requests }; +}; + +describe("typed api client middleware", () => { + beforeEach(() => { + registerBaseUrlGetter(() => ""); + registerAuthHeaderNameGetter(() => "Authorization"); + registerErrorHandler(() => {}); + registerAuthTokenGetter(() => null); + }); + + afterEach(() => { + vi.restoreAllMocks(); + }); + + it("injects the bearer token under the registered auth header name", async () => { + registerAuthTokenGetter(() => "sk-test"); + registerAuthHeaderNameGetter(() => "x-litellm-key"); + const { fetch, requests } = capturingFetch(jsonResponse(200, { data: [] })); + + await fetchClient.GET("/model_group/info", { fetch }); + + expect(requests[0].headers.get("x-litellm-key")).toBe("Bearer sk-test"); + expect(requests[0].headers.get("Authorization")).toBeNull(); + }); + + it("omits the auth header when no token is set", async () => { + const { fetch, requests } = capturingFetch(jsonResponse(200, { data: [] })); + + await fetchClient.GET("/model_group/info", { fetch }); + + expect(requests[0].headers.get("Authorization")).toBeNull(); + }); + + it("rebases the request onto the registered base url, preserving path and query", async () => { + registerBaseUrlGetter(() => "https://proxy.example.com/"); + const { fetch, requests } = capturingFetch(jsonResponse(200, { data: [] })); + + await fetchClient.GET("/model_group/info", { fetch, params: { query: { model_group: "gpt-4o" } } }); + + const url = new URL(requests[0].url); + expect(url.origin).toBe("https://proxy.example.com"); + expect(url.pathname).toBe("/model_group/info"); + expect(url.searchParams.get("model_group")).toBe("gpt-4o"); + }); + + it("maps a non-2xx response to an ApiError carrying status and the derived message", async () => { + const { fetch } = capturingFetch(jsonResponse(403, { error: { message: "no access" } })); + + await expect(fetchClient.GET("/model_group/info", { fetch })).rejects.toMatchObject({ + name: "ApiError", + status: 403, + message: "no access", + }); + }); + + it("returns the parsed body on a successful response", async () => { + const body = { data: [{ model_group: "gpt-4o" }] }; + const { fetch } = capturingFetch(jsonResponse(200, body)); + + const { data, error } = await fetchClient.GET("/model_group/info", { fetch }); + + expect(error).toBeUndefined(); + expect(data).toEqual(body); + }); + + it("reports the derived message to the registered error handler on a non-2xx response", async () => { + const onError = vi.fn(); + registerErrorHandler(onError); + const { fetch } = capturingFetch(jsonResponse(401, { error: { message: "Authentication Error - Expired Key" } })); + + await expect(fetchClient.GET("/model_group/info", { fetch })).rejects.toBeInstanceOf(Error); + + expect(onError).toHaveBeenCalledWith("Authentication Error - Expired Key"); + }); + + it("does not call the error handler on a successful response", async () => { + const onError = vi.fn(); + registerErrorHandler(onError); + const { fetch } = capturingFetch(jsonResponse(200, { data: [] })); + + await fetchClient.GET("/model_group/info", { fetch }); + + expect(onError).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/lib/http/api.ts b/ui/litellm-dashboard/src/lib/http/api.ts new file mode 100644 index 00000000000..2e8e1d36c5b --- /dev/null +++ b/ui/litellm-dashboard/src/lib/http/api.ts @@ -0,0 +1,48 @@ +import createFetchClient, { type Middleware } from "openapi-fetch"; +import type { paths } from "./schema"; +import { ApiError, deriveErrorMessage } from "./client"; +import { getAuthHeaderName, getAuthToken, getRequestBaseUrl, reportError } from "./runtime"; + +const rebaseUrl = (requestUrl: string, base: string): string => { + const { pathname, search } = new URL(requestUrl); + return `${base.replace(/\/+$/, "")}${pathname}${search}`; +}; + +const middleware: Middleware = { + onRequest({ request }) { + const base = getRequestBaseUrl(); + const next = new Request(base ? rebaseUrl(request.url, base) : request.url, request); + const token = getAuthToken(); + if (token) { + next.headers.set(getAuthHeaderName(), `Bearer ${token}`); + } + return next; + }, + async onResponse({ response }) { + if (response.ok) return response; + const raw = await response.clone().text(); + let body: unknown = raw; + let message: string; + try { + body = JSON.parse(raw); + message = deriveErrorMessage(body); + } catch { + message = raw || `HTTP ${response.status}`; + } + reportError(message); + throw new ApiError(message, response.status, body); + }, +}; + +/** + * The typed, schema-bound HTTP client. Use it inside TanStack Query hooks + * (`fetchClient.GET("/path", { params })`) and for imperative calls; path + * params, query params, and request bodies are inferred from schema.d.ts. + * + * The creation-time base is the current origin so request URLs are absolute; the + * middleware rebases each call onto the runtime base when one is registered (a + * split-origin proxy or worker URL), injects the auth header, and maps non-2xx + * responses to ApiError so query functions can just read `.data`. + */ +export const fetchClient = createFetchClient({ baseUrl: globalThis.location?.origin ?? "" }); +fetchClient.use(middleware); diff --git a/ui/litellm-dashboard/src/lib/http/runtime.test.ts b/ui/litellm-dashboard/src/lib/http/runtime.test.ts new file mode 100644 index 00000000000..a475cc2bb10 --- /dev/null +++ b/ui/litellm-dashboard/src/lib/http/runtime.test.ts @@ -0,0 +1,22 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; +import { getAuthHeaderName, getRequestBaseUrl } from "./runtime"; + +describe("runtime request config defaults", () => { + afterEach(() => { + vi.unstubAllEnvs(); + }); + + it("resolves the default base URL from NEXT_PUBLIC_BASE_URL before a getter is registered", () => { + vi.stubEnv("NEXT_PUBLIC_BASE_URL", "https://proxy.example.com/"); + expect(getRequestBaseUrl()).toBe("https://proxy.example.com"); + }); + + it("defaults the base URL to same-origin when NEXT_PUBLIC_BASE_URL is unset", () => { + vi.stubEnv("NEXT_PUBLIC_BASE_URL", ""); + expect(getRequestBaseUrl()).toBe(""); + }); + + it("defaults the auth header name to Authorization", () => { + expect(getAuthHeaderName()).toBe("Authorization"); + }); +}); diff --git a/ui/litellm-dashboard/src/lib/http/runtime.ts b/ui/litellm-dashboard/src/lib/http/runtime.ts new file mode 100644 index 00000000000..dd105255fb1 --- /dev/null +++ b/ui/litellm-dashboard/src/lib/http/runtime.ts @@ -0,0 +1,46 @@ +import { resolveApiBase } from "./resolveApiBase"; + +/** + * Runtime request config the typed client reads on every call. The values are + * mutable at runtime (base URL can switch to a worker origin; the auth header + * name and token come from the logged-in session), and they are owned outside + * this module: networking.tsx registers the base URL / header-name / token + * getters and the error handler. The token getter reads the session cookie, the + * same source useAuthorized decodes, so the client's token and the gate that + * enables a query cannot diverge. Keeping the seam here (not importing from the + * component tree) lets api.ts stay in lib/http without a layering inversion. + * + * The base URL default resolves from NEXT_PUBLIC_BASE_URL so a request still + * hits the right origin if it fires before networking registers its fuller + * getter (which additionally folds in the server root path from the live UI + * config). The auth header name has no build-time source, so it defaults to + * "Authorization" until the session's JWT supplies a custom one. + */ + +type Getter = () => T; + +let baseUrlGetter: Getter = () => resolveApiBase({ explicitBase: process.env.NEXT_PUBLIC_BASE_URL }); +let authHeaderNameGetter: Getter = () => "Authorization"; +let authTokenGetter: Getter = () => null; +let errorHandler: (message: string) => void = () => {}; + +export const registerBaseUrlGetter = (getter: Getter): void => { + baseUrlGetter = getter; +}; + +export const registerAuthHeaderNameGetter = (getter: Getter): void => { + authHeaderNameGetter = getter; +}; + +export const registerAuthTokenGetter = (getter: Getter): void => { + authTokenGetter = getter; +}; + +export const registerErrorHandler = (handler: (message: string) => void): void => { + errorHandler = handler; +}; + +export const getRequestBaseUrl = (): string => baseUrlGetter(); +export const getAuthHeaderName = (): string => authHeaderNameGetter(); +export const getAuthToken = (): string | null => authTokenGetter(); +export const reportError = (message: string): void => errorHandler(message);