Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_complexity_router_keyword_tiers

# Conflicts:
#	ui/litellm-dashboard/eslint-metrics.json
This commit is contained in:
mateo-berri 2026-07-11 18:22:44 +00:00
commit 4486d9fe0c
No known key found for this signature in database
74 changed files with 6169 additions and 799 deletions

View file

@ -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

View file

@ -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,

View file

@ -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)

View file

@ -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)))

View file

@ -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)

View file

@ -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 ###
#########################################################

View file

@ -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-<family> 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))",

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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):

View file

@ -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):

View file

@ -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]

View file

@ -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),

View file

@ -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-<family> 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))",

View file

@ -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

View file

@ -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="<streamed>",
chunks=chunks,

View file

@ -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()

View file

@ -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 == "<streamed>":
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())

View file

@ -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",
)

View file

@ -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):

View file

@ -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():

View file

@ -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()

View file

@ -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,

View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -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}"
)

View file

@ -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

View file

@ -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)"""

View file

@ -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)

View file

@ -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."""

View file

@ -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)

View file

@ -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"

View file

@ -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

View file

@ -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",

View file

@ -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",

View file

@ -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();
});
});

View file

@ -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<string, unknown> | 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<CustomersResponse>({
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!),
});
};

View file

@ -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<RunStatus, string> = {
failed: "#ef4444",
};
const RUN_STATUS_OPTIONS: RunStatus[] = ["pending", "running", "paused", "completed", "failed"];
const STATUS_LABELS: Record<RunStatus, string> = {
pending: "Pending",
running: "Running",
paused: "Paused",
completed: "Completed",
failed: "Failed",
};
const EVENT_COLOR: Record<string, { bar: string; border: string; text: string }> = {
"step.started": { bar: "#f0fdf4", border: "#86efac", text: "#16a34a" },
"step.failed": { bar: "#fef2f2", border: "#fca5a5", text: "#dc2626" },
@ -482,6 +498,9 @@ const WorkflowRuns: React.FC<WorkflowRunsProps> = ({ accessToken }) => {
const [messages, setMessages] = useState<WorkflowRunMessage[]>([]);
const [loadingDetail, setLoadingDetail] = useState(false);
const [drawerOpen, setDrawerOpen] = useState(false);
const [columnFilters, setColumnFilters] = useState<ColumnFiltersState>([]);
const [globalFilter, setGlobalFilter] = useState("");
const [filtersOpen, setFiltersOpen] = useState(false);
const fetchRuns = useCallback(async () => {
if (!accessToken) return;
@ -547,7 +566,9 @@ const WorkflowRuns: React.FC<WorkflowRunsProps> = ({ 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<WorkflowRunsProps> = ({ accessToken }) => {
{
accessorKey: "workflow_type",
header: "Type",
meta: { title: "Type" },
filterFn: "includesString",
cell: ({ row }) => (
<span style={{ fontFamily: "monospace", fontSize: 12, color: "#71717a" }}>{row.original.workflow_type}</span>
),
},
{
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<WorkflowRunsProps> = ({ accessToken }) => {
{
accessorKey: "created_at",
header: "Created",
meta: { title: "Created" },
cell: ({ row }) => <span style={{ fontSize: 12, color: "#a1a1aa" }}>{timeAgo(row.original.created_at)}</span>,
},
],
@ -603,28 +630,11 @@ const WorkflowRuns: React.FC<WorkflowRunsProps> = ({ accessToken }) => {
}}
>
{/* page header */}
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "space-between",
marginBottom: 20,
}}
>
<div>
<div style={{ fontSize: 18, fontWeight: 600, color: "#18181b" }}>Workflow Runs</div>
<div style={{ fontSize: 13, color: "#71717a", marginTop: 2 }}>
Durable state tracking for agents and automated workflows
</div>
<div style={{ marginBottom: 20 }}>
<div style={{ fontSize: 18, fontWeight: 600, color: "#18181b" }}>Workflow Runs</div>
<div style={{ fontSize: 13, color: "#71717a", marginTop: 2 }}>
Durable state tracking for agents and automated workflows
</div>
<Button
icon={<ReloadOutlined />}
onClick={fetchRuns}
loading={loadingRuns}
style={{ color: "#71717a", borderColor: "#e4e4e7" }}
>
Refresh
</Button>
</div>
<DataTable
@ -641,8 +651,64 @@ const WorkflowRuns: React.FC<WorkflowRunsProps> = ({ accessToken }) => {
}
paginationMode="client"
pageSizeOptions={[50, 100]}
filterMode="client"
columnFilters={columnFilters}
onColumnFiltersChange={setColumnFilters}
globalFilter={globalFilter}
onGlobalFilterChange={setGlobalFilter}
onRowClick={fetchRunDetail}
size="compact"
toolbar={(table) => (
<>
<DataTableToolbar
table={table}
searchValue={globalFilter}
onSearchChange={setGlobalFilter}
searchPlaceholder="Search runs…"
onRefresh={fetchRuns}
isRefreshing={loadingRuns}
onOpenFilters={() => setFiltersOpen(true)}
/>
<DataTableFilterDrawer
table={table}
open={filtersOpen}
onOpenChange={setFiltersOpen}
title="Filters"
description="Narrow down workflow runs"
>
{({ get, set }) => (
<>
<DataTableFilterField label="Status">
<Select
items={STATUS_LABELS}
value={(get("status") as string) || null}
onValueChange={(value: string | null) => set("status", value ?? "")}
>
<SelectTrigger className="w-full">
<SelectValue placeholder="All statuses" />
</SelectTrigger>
<SelectContent>
<SelectItem value={null}>All statuses</SelectItem>
{RUN_STATUS_OPTIONS.map((status) => (
<SelectItem key={status} value={status}>
{STATUS_LABELS[status]}
</SelectItem>
))}
</SelectContent>
</Select>
</DataTableFilterField>
<DataTableFilterField label="Type">
<Input
value={(get("workflow_type") as string) ?? ""}
onChange={(event) => set("workflow_type", event.target.value)}
placeholder="Filter by type…"
/>
</DataTableFilterField>
</>
)}
</DataTableFilterDrawer>
</>
)}
/>
{/* detail drawer */}

View file

@ -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(<DashboardHeader page="logs" />);
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(<DashboardHeader page="logs" />);
expect(screen.getByRole("button", { name: /AI Gateway/i })).toBeInTheDocument();
expect(screen.getByText("Logs")).toBeInTheDocument();
expect(screen.queryByText("Observability")).not.toBeInTheDocument();
});
});

View file

@ -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) {
<header className="flex h-14 flex-none items-center justify-between gap-4 border-b border-border bg-background px-4">
<Breadcrumb className="min-w-0">
<BreadcrumbList className="flex-nowrap">
{section && (
<>
<BreadcrumbItem className="whitespace-nowrap">{section}</BreadcrumbItem>
<BreadcrumbSeparator />
</>
)}
<BreadcrumbItem className="flex-none">
<ViewSwitcher />
</BreadcrumbItem>
<BreadcrumbSeparator />
<BreadcrumbItem className="min-w-0">
<BreadcrumbPage className="truncate">{title}</BreadcrumbPage>
</BreadcrumbItem>
@ -76,8 +74,6 @@ export function DashboardHeader({ page }: DashboardHeaderProps) {
{!hideCommunityLinks && <CommunityEngagementButtons />}
<Separator orientation="vertical" className="mx-1.5 h-5" />
<NotificationsBell />
<Separator orientation="vertical" className="mx-1.5 h-5" />
<ViewSwitcher />
</div>
</header>
);

View file

@ -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(<ViewSwitcher />);
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(<ViewSwitcher />);
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(<ViewSwitcher />);
@ -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();
});
});

View file

@ -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: (
<div className="flex items-center justify-between gap-6 py-0.5">
<span className="font-medium">Chat</span>
{isChatRoute && <CheckOutlined className="text-blue-600" />}
</div>
),
}
: {
key: CHAT,
disabled: true,
label: (
<div className="flex max-w-[220px] flex-col py-0.5">
<span className="font-medium">Chat</span>
<span className="whitespace-normal text-xs leading-snug text-muted-foreground">
Admins can enable in Settings
</span>
</div>
),
};
const items: MenuProps["items"] = [
...modeEntries.map((e) => ({
key: e.key,
@ -40,19 +62,7 @@ export default function ViewSwitcher() {
</div>
),
})),
...(chatEnabled
? [
{
key: CHAT,
label: (
<div className="flex items-center justify-between gap-6 py-0.5">
<span className="font-medium">Chat</span>
{isChatRoute && <CheckOutlined className="text-blue-600" />}
</div>
),
},
]
: []),
chatItem,
];
const onClick: MenuProps["onClick"] = ({ key }) => {
@ -72,11 +82,13 @@ export default function ViewSwitcher() {
<Dropdown menu={{ items, onClick, selectedKeys: [isChatRoute ? CHAT : mode] }} trigger={["click"]}>
<button
type="button"
className="flex items-center gap-2 rounded-md border border-gray-200 px-2.5 py-1.5 text-sm font-medium text-gray-700 transition-colors hover:bg-gray-50"
className="flex h-8 max-w-[220px] items-center gap-1.5 rounded-md border border-border bg-background pl-1.5 pr-2 text-sm font-medium text-foreground transition-colors hover:bg-accent"
>
<AppstoreOutlined className="text-gray-500" />
<span>{activeLabel}</span>
<DownOutlined className="text-[10px] text-gray-400" />
<span className="flex size-5 flex-none items-center justify-center rounded bg-muted text-muted-foreground">
<AppstoreOutlined className="text-[13px]" />
</span>
<span className="truncate">{activeLabel}</span>
<ChevronsUpDown className="size-3.5 flex-none text-muted-foreground" />
</button>
</Dropdown>
);

View file

@ -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 (
<Form.Item
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Gateway-hosted sign-in (DCR bridge)
<Tooltip title="Lets OAuth-only clients like Claude Desktop register and sign in through the gateway. Turn off to relay the upstream server's own OAuth metadata instead (for clients pre-registered with the upstream IdP).">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
</span>
}
name="dcr_bridge"
valuePropName="checked"
initialValue={initialChecked}
>
<Switch />
</Form.Item>
);
}

View file

@ -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 <Form form={form}>{children}</Form>;
};
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(
<WithForm>
<PassthroughAuthorizeSection
authType="oauth_delegate"
oauthFlow={noopFlow}
isEditing
savedAuthType="true_passthrough"
/>
</WithForm>,
);
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(
<WithForm>
<PassthroughAuthorizeSection
authType="true_passthrough"
oauthFlow={noopFlow}
isEditing
savedAuthType="oauth2"
/>
</WithForm>,
);
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(
<WithForm>
<PassthroughAuthorizeSection authType="true_passthrough" oauthFlow={noopFlow} appMayNotMatchUpstream />
</WithForm>,
);
expect(screen.getByText(/registered for the previous upstream/)).toBeInTheDocument();
});
});

View file

@ -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<void>;
@ -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<string, string> = {
@ -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 (
<div className="rounded-lg border border-dashed border-gray-300 p-4 space-y-2 mb-4">
<p className="text-sm text-gray-600">
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.
</p>
{appMayNotMatchUpstream && (
<p className="text-sm text-amber-600">
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.
</p>
)}
<Form.Item
label={<span className="text-sm font-medium text-gray-700">OAuth Client ID (optional, not saved)</span>}
label={<span className="text-sm font-medium text-gray-700">OAuth Client ID (optional)</span>}
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}
>
<Input.Password
placeholder="Leave blank to use dynamic client registration"
placeholder={clientIdPlaceholder}
disabled={removeStoredApp}
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
/>
</Form.Item>
<Form.Item
label={<span className="text-sm font-medium text-gray-700">OAuth Client Secret (optional, not saved)</span>}
label={<span className="text-sm font-medium text-gray-700">OAuth Client Secret (optional)</span>}
name={["credentials", "client_secret"]}
>
<Input.Password
placeholder="Leave blank for public clients / PKCE"
placeholder={clientSecretPlaceholder}
disabled={removeStoredApp}
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
/>
</Form.Item>
<DcrBridgeToggle authType={authType} initialChecked={dcrBridgeInitialChecked} />
{isEditing && onRemoveStoredAppChange && (
<Checkbox checked={removeStoredApp} onChange={(e) => onRemoveStoredAppChange(e.target.checked)}>
<span className="text-sm text-gray-700">
Remove the saved OAuth app on save (the server goes back to dynamic client registration)
</span>
</Checkbox>
)}
<Button
onClick={oauthFlow.startOAuthFlow}
disabled={oauthFlow.status === "authorizing" || oauthFlow.status === "exchanging"}
@ -67,7 +119,8 @@ export default function PassthroughAuthorizeSection({
{oauthFlow.error && <p className="text-sm text-red-500">{oauthFlow.error}</p>}
{oauthFlow.status === "success" && oauthFlow.tokenResponse?.access_token && (
<p className="text-sm text-green-600">
Token held for this browser session. Tools can now be previewed and configured; nothing was saved to LiteLLM.
Token held for this browser session. Tools can now be previewed and configured; the token was not saved to
LiteLLM.
</p>
)}
</div>

View file

@ -31,6 +31,7 @@ const oauthHook = vi.hoisted(() => ({
| ((token: Record<string, unknown> | null, registeredClient?: { clientId?: string; clientSecret?: string }) => void)
| null,
getCredentials: null as (() => Record<string, unknown> | undefined) | null,
getTemporaryPayload: null as (() => Record<string, unknown> | null) | null,
}));
vi.mock("@/hooks/useMcpOAuthFlow", () => ({
useMcpOAuthFlow: (opts: {
@ -39,9 +40,11 @@ vi.mock("@/hooks/useMcpOAuthFlow", () => ({
registeredClient?: { clientId?: string; clientSecret?: string },
) => void;
getCredentials?: () => Record<string, unknown> | undefined;
getTemporaryPayload?: () => Record<string, unknown> | 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(<CreateMCPServer {...defaultProps} />);
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(<CreateMCPServer {...defaultProps} isModalVisible={false} />);
});
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(<CreateMCPServer {...defaultProps} />);
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);
});
});

View file

@ -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<CreateMCPServerProps> = ({
// 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<string | undefined>(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<CreateMCPServerProps> = ({
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<CreateMCPServerProps> = ({
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<string, unknown> | undefined) ?? {}),
...(dcrClientRef.current ?? {}),
}),
getTemporaryPayload: () => {
const values = form.getFieldsValue(true);
const transport = values.transport || transportType;
@ -184,7 +204,12 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
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<string, unknown> | 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<CreateMCPServerProps> = ({
// 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<string, unknown> | 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<CreateMCPServerProps> = ({
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<CreateMCPServerProps> = ({
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<CreateMCPServerProps> = ({
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<CreateMCPServerProps> = ({
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<CreateMCPServerProps> = ({
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<CreateMCPServerProps> = ({
setHasToolAllowlistInteraction(false);
setAliasManuallyEdited(false);
setLogoUrl(undefined);
setAuthorizedIdentity(undefined);
dcrClientRef.current = null;
setAppMayNotMatchUpstream(false);
setModalVisible(false);
};
@ -655,6 +735,8 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
clearTools();
resetOAuthFlow();
setAuthorizedIdentity(undefined);
dcrClientRef.current = null;
setAppMayNotMatchUpstream(false);
}
}, [isModalVisible, form, clearTools, resetOAuthFlow]);
@ -663,12 +745,31 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
const handleFormValuesChange = (changedValues: Record<string, unknown>, allValues: Record<string, unknown>) => {
// 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<CreateMCPServerProps> = ({
<PassthroughAuthorizeSection
authType={authType}
dcrBridgeInitialChecked
oauthFlow={{
startOAuthFlow,
status: oauthStatus,
error: oauthError,
tokenResponse: oauthTokenResponse,
}}
appMayNotMatchUpstream={appMayNotMatchUpstream}
/>
{shouldShowAuthValueField && (

View file

@ -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(
<MCPServerEdit
mcpServer={{ ...interactiveOAuthServer, auth_type: authType }}
accessToken="access-token"
userID="user-1"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
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(
<MCPServerEdit
mcpServer={{ ...interactiveOAuthServer, auth_type: authType }}
accessToken="access-token"
userID="user-1"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
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(
<MCPServerEdit
mcpServer={{ ...interactiveOAuthServer, auth_type: "true_passthrough" }}
accessToken="access-token"
userID="user-1"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
// 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(
<MCPServerEdit
mcpServer={{
...interactiveOAuthServer,
auth_type: "true_passthrough",
credentials: { client_id: "stored-client" },
}}
accessToken="access-token"
userID="user-1"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
// 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(
<MCPServerEdit
mcpServer={{
...interactiveOAuthServer,
auth_type: "true_passthrough",
credentials: { client_id: "stored-client" },
}}
accessToken="access-token"
userID="user-1"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
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(
<MCPServerEdit
mcpServer={{ ...interactiveOAuthServer, server_id: "server-A", auth_type: "true_passthrough" }}
accessToken="access-token"
userID="user-1"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
// 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(
<MCPServerEdit
mcpServer={{ ...interactiveOAuthServer, server_id: "server-B", auth_type: "true_passthrough" }}
accessToken="access-token"
userID="user-1"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
// 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<string, unknown>) {
render(
<MCPServerEdit
mcpServer={{ ...interactiveOAuthServer, ...server }}
accessToken="access-token"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
}
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);
});
});

View file

@ -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<MCPServerEditProps> = ({
const [toolsError, setToolsError] = useState<string | null>(null);
const [searchValue, setSearchValue] = useState<string>("");
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<string[]>([]);
const [hasToolAllowlistInteraction, setHasToolAllowlistInteraction] = useState(false);
const [toolNameToDisplayName, setToolNameToDisplayName] = useState<Record<string, string>>({});
@ -179,7 +187,9 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
};
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<string, unknown> | 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<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
}
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<MCPServerEditProps> = ({
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<string, unknown> | 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<MCPServerEditProps> = ({
}
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<MCPServerEditProps> = ({
};
const handleFormValuesChange = (changedValues: Record<string, unknown>) => {
// 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<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
? 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<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
}
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<MCPServerEditProps> = ({
error: oauthError,
tokenResponse: oauthTokenResponse,
}}
isEditing
savedAuthType={mcpServer.auth_type}
removeStoredApp={removeStoredApp}
onRemoveStoredAppChange={setRemoveStoredApp}
appMayNotMatchUpstream={appMayNotMatchUpstream}
/>
</>
)}

View file

@ -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();
});
});

View file

@ -96,6 +96,54 @@ export const getOAuthAuthorizationIdentity = (values: Record<string, unknown>):
// 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<string, unknown> | null | undefined,
): Record<string, string> | 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<string, unknown> | null | undefined,
): Record<string, unknown> | 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<string, unknown> | null;
/** Stdio-only fields (present when transport === 'stdio') */
command?: string | null;

View file

@ -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, {

View file

@ -29,6 +29,16 @@ const nameCellColumns: ColumnDef<Person, unknown>[] = [
},
];
const filterableColumns: ColumnDef<Person, unknown>[] = [
{
accessorKey: "name",
header: "Name",
meta: { title: "Name" },
filterFn: (row, columnId, value) => row.getValue<string>(columnId) === value,
cell: ({ row }) => <span data-testid="name-cell">{row.original.name}</span>,
},
];
const headerCycleColumns: ColumnDef<Person, unknown>[] = [
{
accessorKey: "name",
@ -195,10 +205,110 @@ describe("DataTable pagination", () => {
});
});
describe("DataTable filtering", () => {
it("client mode filters rows by columnFilters", () => {
const { rerender } = render(
<DataTable
data={CHARLIE_ALICE_BOB}
columns={filterableColumns}
filterMode="client"
columnFilters={[]}
onColumnFiltersChange={vi.fn()}
/>,
);
expect(names()).toEqual(["Charlie", "Alice", "Bob"]);
rerender(
<DataTable
data={CHARLIE_ALICE_BOB}
columns={filterableColumns}
filterMode="client"
columnFilters={[{ id: "name", value: "Alice" }]}
onColumnFiltersChange={vi.fn()}
/>,
);
expect(names()).toEqual(["Alice"]);
});
it("client global filter matches substrings across columns", () => {
const { rerender } = render(
<DataTable
data={CHARLIE_ALICE_BOB}
columns={nameEmailColumns}
filterMode="client"
globalFilter=""
onGlobalFilterChange={vi.fn()}
/>,
);
expect(names()).toEqual(["Charlie", "Alice", "Bob"]);
rerender(
<DataTable
data={CHARLIE_ALICE_BOB}
columns={nameEmailColumns}
filterMode="client"
globalFilter="ali"
onGlobalFilterChange={vi.fn()}
/>,
);
expect(names()).toEqual(["Alice"]);
});
it("server mode never filters locally even when columnFilters is set", () => {
render(
<DataTable
data={CHARLIE_ALICE_BOB}
columns={filterableColumns}
filterMode="server"
columnFilters={[{ id: "name", value: "Alice" }]}
onColumnFiltersChange={vi.fn()}
/>,
);
expect(names()).toEqual(["Charlie", "Alice", "Bob"]);
});
it("throws when server filtering is missing required props", () => {
const spy = vi.spyOn(console, "error").mockImplementation(() => {});
expect(() => render(<DataTable data={[]} columns={filterableColumns} filterMode="server" />)).toThrow(
/filterMode='server'/,
);
spy.mockRestore();
});
});
describe("DataTable loading", () => {
it("renders skeleton rows while loading and real rows once loaded", () => {
const { rerender } = render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} isLoading />);
expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0);
expect(screen.queryByTestId("name-cell")).toBeNull();
rerender(<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} />);
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<Person, unknown>[] = [
{ accessorKey: "name", header: "Name", meta: { skeleton: "twoLine" }, cell: () => null },
{ accessorKey: "email", header: "Email", cell: () => null },
];
render(<DataTable data={CHARLIE_ALICE_BOB} columns={columns} isLoading />);
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(
<DataTable
data={CHARLIE_ALICE_BOB}
columns={nameEmailColumns}
@ -206,13 +316,13 @@ describe("DataTable column visibility", () => {
/>,
);
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 () => {

View file

@ -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<TData extends RowData, TValue>(
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<TData, TValue>(columns: ColumnDef<TData, TValue>[]): Colu
function buildRowModels<TData>(
sortingMode: SortingMode,
paginationMode: PaginationMode,
filterMode: FilterMode,
getRowCanExpand: ((row: Row<TData>) => boolean) | undefined,
): Partial<TableOptions<TData>> {
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 (
<div className="flex flex-col items-center gap-1 py-6">
<div className="mb-1 flex size-10 items-center justify-center rounded-lg bg-muted">
<SearchX className="size-5 text-muted-foreground" />
</div>
<div className="text-sm font-medium text-foreground">No results</div>
<div className="text-sm text-muted-foreground">No rows match your search or filters.</div>
</div>
);
}
const SKELETON_WIDTHS = ["w-[58%]", "w-[44%]", "w-[70%]", "w-[50%]", "w-[64%]", "w-[48%]"] as const;
function SkeletonCell<TData>({ column, index }: { column: Column<TData, unknown> | undefined; index: number }) {
const meta = column?.columnDef.meta;
const width = SKELETON_WIDTHS[index % SKELETON_WIDTHS.length];
if (meta?.skeleton === "twoLine") {
return (
<div className="flex flex-col gap-2">
<Skeleton className={cn("h-3.5", width)} />
<Skeleton className="h-2.5 w-2/5 opacity-65" />
</div>
);
}
return <Skeleton className={cn("h-3.5", width, meta?.numeric ? "ml-auto" : "")} />;
}
function SkeletonRows<TData>({
rowCount,
columns,
size,
message,
}: {
rowCount: number;
columns: readonly Column<TData, unknown>[];
size: DataTableSize;
message?: string;
}) {
const rowKeys = Array.from({ length: Math.max(rowCount, 1) }, (_, index) => index);
const cells = columns.length > 0 ? columns : [undefined];
return (
<Fragment>
{rowKeys.map((rowKey) => (
<TableRow key={`skeleton-${rowKey}`} className="hover:bg-transparent" data-testid="skeleton-row">
{cells.map((column, columnKey) => (
<TableCell key={column?.id ?? columnKey} className={size === "compact" ? "px-2 py-1" : ""}>
<SkeletonCell column={column} index={columnKey} />
{rowKey === 0 && columnKey === 0 && message !== undefined ? (
<span className="sr-only">{message}</span>
) : null}
</TableCell>
))}
</TableRow>
))}
</Fragment>
);
}
function useControllable<T>(
controlled: T | undefined,
controlledOnChange: OnChangeFn<T> | undefined,
@ -334,6 +407,12 @@ function useDataTableInstance<TData extends RowData, TValue>(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<TData extends RowData, TValue>(props: DataTablePro
pageIndex: 0,
pageSize: pageSizeOptions[0] ?? 25,
});
const filterState = useControllable<ColumnFiltersState>(
columnFilters,
onColumnFiltersChange,
defaultColumnFilters ?? [],
);
const globalFilterState = useControllable<string>(globalFilter, onGlobalFilterChange, "");
const expandedState = useControllable<ExpandedState>(expanded, onExpandedChange, {});
const [columnVisibility, setColumnVisibility] = useState<VisibilityState>(defaultColumnVisibility ?? {});
const [columnSizing, setColumnSizing] = useState<ColumnSizingState>({});
@ -360,6 +445,8 @@ function useDataTableInstance<TData extends RowData, TValue>(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<TData extends RowData, TValue>(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<TData extends RowData, TValue>(props: DataTableProps<T
const {
isLoading = false,
loadingMessage = "Loading…",
noDataMessage = "No results",
skeletonRowCount = 8,
noDataMessage,
paginationMode = "none",
rowCount,
pageSizeOptions = DEFAULT_PAGE_SIZE_OPTIONS,
@ -443,10 +534,17 @@ export function DataTable<TData extends RowData, TValue>(props: DataTableProps<T
const renderBody = (): React.ReactNode => {
if (isLoading) {
return <MessageRow colSpan={visibleColumnCount}>{loadingMessage}</MessageRow>;
return (
<SkeletonRows
rowCount={skeletonRowCount}
columns={table.getVisibleLeafColumns()}
size={size}
message={loadingMessage}
/>
);
}
if (rows.length === 0) {
return <MessageRow colSpan={visibleColumnCount}>{noDataMessage}</MessageRow>;
return <MessageRow colSpan={visibleColumnCount}>{noDataMessage ?? <DefaultEmptyState />}</MessageRow>;
}
return rows.map((row) => (
<DataTableBodyRow
@ -462,34 +560,38 @@ export function DataTable<TData extends RowData, TValue>(props: DataTableProps<T
));
};
const paginationNode = renderPagination();
return (
<div className="w-full">
{toolbar !== undefined && <div className="w-full">{toolbar(table)}</div>}
<div
className={cn("rounded-lg border border-border", stickyHeader ? "overflow-auto" : "overflow-x-auto")}
style={stickyHeader ? { maxHeight: maxBodyHeight } : undefined}
>
<TableRoot className={enableColumnResizing ? "table-fixed" : ""} style={tableStyle}>
<TableHeader className={stickyHeader ? "sticky top-0 z-20" : ""}>
{table.getHeaderGroups().map((headerGroup) => (
<TableRow key={headerGroup.id} className="bg-muted/50 hover:bg-muted/50">
{headerGroup.headers.map((header) => (
<DataTableHeadCell
key={header.id}
header={header}
size={size}
stickyHeader={stickyHeader}
enableColumnResizing={enableColumnResizing}
/>
))}
</TableRow>
))}
</TableHeader>
<TableBody>{renderBody()}</TableBody>
{footer !== undefined && <TableFooter>{footer(table)}</TableFooter>}
</TableRoot>
<div className="overflow-hidden rounded-lg border border-border">
{toolbar !== undefined && <div className="border-b border-border px-4 py-3">{toolbar(table)}</div>}
<div
className={stickyHeader ? "overflow-auto" : "overflow-x-auto"}
style={stickyHeader ? { maxHeight: maxBodyHeight } : undefined}
>
<TableRoot className={enableColumnResizing ? "table-fixed" : ""} style={tableStyle}>
<TableHeader className={stickyHeader ? "sticky top-0 z-20" : ""}>
{table.getHeaderGroups().map((headerGroup) => (
<TableRow key={headerGroup.id} className="bg-muted/50 hover:bg-muted/50">
{headerGroup.headers.map((header) => (
<DataTableHeadCell
key={header.id}
header={header}
size={size}
stickyHeader={stickyHeader}
enableColumnResizing={enableColumnResizing}
/>
))}
</TableRow>
))}
</TableHeader>
<TableBody>{renderBody()}</TableBody>
{footer !== undefined && <TableFooter>{footer(table)}</TableFooter>}
</TableRoot>
</div>
{paginationNode !== null && <div className="border-t border-border">{paginationNode}</div>}
</div>
{renderPagination()}
</div>
);
}

View file

@ -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<Person, unknown>[] = [
{
accessorKey: "name",
header: "Name",
meta: { title: "Name" },
filterFn: (row, columnId, value) => row.getValue<string>(columnId) === value,
cell: ({ row }) => <span data-testid="name-cell">{row.original.name}</span>,
},
];
const names = (): (string | null)[] => screen.getAllByTestId("name-cell").map((el) => el.textContent);
function Harness({ initialFilters }: { initialFilters?: ColumnFiltersState }) {
const [open, setOpen] = useState(false);
return (
<DataTable
data={DATA}
columns={columns}
filterMode="client"
defaultColumnFilters={initialFilters}
toolbar={(table) => (
<>
<DataTableToolbar table={table} onOpenFilters={() => setOpen(true)} />
<DataTableFilterDrawer table={table} open={open} onOpenChange={setOpen} title="Filters">
{({ get, set }) => (
<input
aria-label="name filter"
data-testid="draft-name"
value={(get("name") as string | undefined) ?? ""}
onChange={(event) => set("name", event.target.value)}
/>
)}
</DataTableFilterDrawer>
</>
)}
/>
);
}
describe("DataTableFilterDrawer", () => {
it("stages edits and only commits them to the table on Apply", async () => {
const user = userEvent.setup();
render(<Harness />);
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(<Harness initialFilters={[{ id: "name", value: "Bob" }]} />);
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(<Harness initialFilters={[{ id: "name", value: "Bob" }]} />);
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("");
});
});

View file

@ -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<TData> {
table: Table<TData>;
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<string, unknown> {
return Object.fromEntries(filters.map((filter) => [filter.id, filter.value]));
}
function toFilters(draft: Record<string, unknown>): ColumnFiltersState {
return Object.entries(draft)
.filter(([, value]) => !isEmpty(value))
.map(([id, value]) => ({ id, value }));
}
export function DataTableFilterDrawer<TData>({
table,
open,
onOpenChange,
title = "Filters",
description,
applyLabel = "Apply Filters",
resetLabel = "Reset",
children,
}: DataTableFilterDrawerProps<TData>) {
const [draft, setDraft] = React.useState<Record<string, unknown>>(() => 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 (
<Sheet open={open} onOpenChange={onOpenChange}>
<SheetContent side="right">
<SheetHeader>
<SheetTitle>{title}</SheetTitle>
{description !== undefined && <SheetDescription>{description}</SheetDescription>}
</SheetHeader>
<div className="flex min-h-0 flex-1 flex-col gap-4 overflow-y-auto p-4" data-testid="filter-drawer-body">
{children(helpers)}
</div>
<SheetFooter className="flex-row">
<Button variant="outline" className="flex-1" onClick={reset} data-testid="filter-drawer-reset">
{resetLabel}
</Button>
<Button className="flex-1" onClick={apply} data-testid="filter-drawer-apply">
{applyLabel}
</Button>
</SheetFooter>
</SheetContent>
</Sheet>
);
}
export function DataTableFilterField({ label, children }: { label: string; children: React.ReactNode }) {
return (
<div className="flex flex-col gap-1.5">
<Label>{label}</Label>
{children}
</div>
);
}

View file

@ -37,7 +37,7 @@ export function DataTablePagination({
const lastPage = Math.max(pageCount - 1, 0);
return (
<div className={cn("flex flex-wrap items-center justify-between gap-4 px-2 py-2", className)}>
<div className={cn("flex flex-wrap items-center justify-between gap-4 px-4 py-2.5", className)}>
<div className="flex items-center gap-2 text-sm text-muted-foreground">
<span>Rows per page</span>
<Select
@ -65,6 +65,9 @@ export function DataTablePagination({
<span data-testid="pagination-range" className="text-sm text-muted-foreground tabular-nums">
{rowCount === 0 ? "No results" : `Showing ${start}-${end} of ${rowCount}`}
</span>
<span data-testid="pagination-page" className="text-sm text-muted-foreground tabular-nums">
Page {page + 1} of {Math.max(pageCount, 1)}
</span>
<div className="flex items-center gap-1">
<Button
variant="outline"

View file

@ -1,35 +1,106 @@
import type { ColumnDef } from "@tanstack/react-table";
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import type * as React from "react";
import { describe, expect, it, vi } from "vitest";
import { DataTable } from "./DataTable";
import { DataTableToolbar } from "./DataTableToolbar";
interface Person {
id: string;
name: string;
}
const DATA: Person[] = [
{ id: "a", name: "Alice" },
{ id: "b", name: "Bob" },
];
const columns: ColumnDef<Person, unknown>[] = [
{
accessorKey: "name",
header: "Name",
meta: { title: "Name" },
filterFn: (row, columnId, value) => row.getValue<string>(columnId) === value,
cell: ({ row }) => <span data-testid="name-cell">{row.original.name}</span>,
},
];
const names = (): (string | null)[] => screen.getAllByTestId("name-cell").map((el) => el.textContent);
function Harness({
onOpenFilters,
onRefresh,
children,
}: {
onOpenFilters?: () => void;
onRefresh?: () => void;
children?: React.ReactNode;
}) {
return (
<DataTable
data={DATA}
columns={columns}
filterMode="client"
defaultColumnFilters={[{ id: "name", value: "Alice" }]}
toolbar={(table) => (
<DataTableToolbar table={table} onOpenFilters={onOpenFilters} onRefresh={onRefresh}>
{children}
</DataTableToolbar>
)}
/>
);
}
describe("DataTableToolbar", () => {
it("renders a chip for each active filter with its label and value", () => {
render(<Harness />);
expect(names()).toEqual(["Alice"]);
const chip = screen.getByTestId("filter-chip-name");
expect(chip).toHaveTextContent("Name:");
expect(chip).toHaveTextContent("Alice");
});
it("removes a single filter when its chip remove button is clicked", async () => {
const user = userEvent.setup();
render(<Harness />);
await user.click(screen.getByTestId("filter-chip-remove-name"));
expect(screen.queryByTestId("filter-chip-name")).toBeNull();
expect(names()).toEqual(["Alice", "Bob"]);
});
it("clears every filter via Clear all", async () => {
const user = userEvent.setup();
render(<Harness />);
await user.click(screen.getByTestId("datatable-clear-filters"));
expect(screen.queryByTestId("filter-chip-name")).toBeNull();
expect(names()).toEqual(["Alice", "Bob"]);
});
it("shows the active filter count and fires onOpenFilters", async () => {
const user = userEvent.setup();
const onOpenFilters = vi.fn();
render(<Harness onOpenFilters={onOpenFilters} />);
expect(screen.getByTestId("datatable-filter-count")).toHaveTextContent("1");
await user.click(screen.getByTestId("datatable-filters-trigger"));
expect(onOpenFilters).toHaveBeenCalledTimes(1);
});
it("renders slotted action children", () => {
render(
<DataTableToolbar>
<Harness>
<button data-testid="toolbar-action">Action</button>
</DataTableToolbar>,
</Harness>,
);
expect(screen.getByTestId("toolbar-action")).toBeInTheDocument();
});
it("shows the reset button only when there are active filters", async () => {
it("fires onRefresh when the refresh button is clicked", async () => {
const user = userEvent.setup();
const onResetFilters = vi.fn();
const { rerender } = render(<DataTableToolbar onResetFilters={onResetFilters} hasActiveFilters={false} />);
expect(screen.queryByText("Reset Filters")).toBeNull();
rerender(<DataTableToolbar onResetFilters={onResetFilters} hasActiveFilters />);
await user.click(screen.getByText("Reset Filters"));
expect(onResetFilters).toHaveBeenCalledTimes(1);
});
it("wires the filters toggle button", async () => {
const user = userEvent.setup();
const onToggleFilters = vi.fn();
render(<DataTableToolbar onToggleFilters={onToggleFilters} />);
await user.click(screen.getByText("Filters"));
expect(onToggleFilters).toHaveBeenCalledTimes(1);
const onRefresh = vi.fn();
render(<Harness onRefresh={onRefresh} />);
await user.click(screen.getByTestId("datatable-refresh"));
expect(onRefresh).toHaveBeenCalledTimes(1);
});
});

View file

@ -1,55 +1,128 @@
"use client";
import { Search } from "lucide-react";
import type { Table } from "@tanstack/react-table";
import { RefreshCw, Search, SlidersHorizontal, X } from "lucide-react";
import type * as React from "react";
import { FilterInput } from "@/components/common_components/Filters/FilterInput";
import { FiltersButton } from "@/components/common_components/Filters/FiltersButton";
import { ResetFiltersButton } from "@/components/common_components/Filters/ResetFiltersButton";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { cn } from "@/lib/cva.config";
interface DataTableToolbarProps {
import { DataTableViewOptions } from "./DataTableViewOptions";
interface DataTableToolbarProps<TData> {
table: Table<TData>;
searchValue?: string;
onSearchChange?: (value: string) => void;
searchPlaceholder?: string;
filtersActive?: boolean;
hasActiveFilters?: boolean;
onToggleFilters?: () => void;
onResetFilters?: () => void;
onOpenFilters?: () => void;
onRefresh?: () => void;
isRefreshing?: boolean;
filterLabels?: Record<string, string>;
formatFilterValue?: (columnId: string, value: unknown) => string;
showViewOptions?: boolean;
children?: React.ReactNode;
className?: string;
}
export function DataTableToolbar({
function defaultFormatValue(value: unknown): string {
if (Array.isArray(value)) {
return value.join(", ");
}
return String(value);
}
export function DataTableToolbar<TData>({
table,
searchValue,
onSearchChange,
searchPlaceholder = "Search",
filtersActive = false,
hasActiveFilters = false,
onToggleFilters,
onResetFilters,
onOpenFilters,
onRefresh,
isRefreshing = false,
filterLabels,
formatFilterValue,
showViewOptions = true,
children,
className,
}: DataTableToolbarProps) {
const showReset = onResetFilters !== undefined && hasActiveFilters;
}: DataTableToolbarProps<TData>) {
const filters = table.getState().columnFilters;
const labelFor = (columnId: string): string =>
filterLabels?.[columnId] ?? table.getColumn(columnId)?.columnDef.meta?.title ?? columnId;
const valueFor = (columnId: string, value: unknown): string =>
formatFilterValue?.(columnId, value) ?? defaultFormatValue(value);
return (
<div className={cn("flex flex-wrap items-center justify-between gap-2 pb-3", className)}>
<div className="flex flex-wrap items-center gap-2">
<div className={cn("flex flex-wrap items-center justify-between gap-2", className)}>
<div className="flex flex-1 flex-wrap items-center gap-2">
{onSearchChange !== undefined && (
<FilterInput
value={searchValue ?? ""}
onChange={onSearchChange}
placeholder={searchPlaceholder}
icon={Search}
/>
<div className="relative">
<Search className="pointer-events-none absolute top-1/2 left-2.5 size-4 -translate-y-1/2 text-muted-foreground" />
<Input
value={searchValue ?? ""}
onChange={(event) => onSearchChange(event.target.value)}
placeholder={searchPlaceholder}
className="h-8 w-56 pl-8"
data-testid="datatable-search"
/>
</div>
)}
{onToggleFilters !== undefined && (
<FiltersButton onClick={onToggleFilters} active={filtersActive} hasActiveFilters={hasActiveFilters} />
{filters.map((filter) => (
<Badge key={filter.id} variant="outline" className="gap-1 py-1" data-testid={`filter-chip-${filter.id}`}>
<span className="text-muted-foreground">{labelFor(filter.id)}:</span>
{valueFor(filter.id, filter.value)}
<button
type="button"
aria-label={`Remove ${labelFor(filter.id)} filter`}
data-testid={`filter-chip-remove-${filter.id}`}
onClick={() => table.setColumnFilters((previous) => previous.filter((entry) => entry.id !== filter.id))}
className="ml-0.5 rounded-full text-muted-foreground hover:text-foreground"
>
<X className="size-3" />
</button>
</Badge>
))}
{filters.length > 0 && (
<Button
variant="ghost"
size="sm"
onClick={() => table.setColumnFilters([])}
data-testid="datatable-clear-filters"
>
Clear all
</Button>
)}
</div>
<div className="flex flex-wrap items-center gap-2">
{children}
{onRefresh !== undefined && (
<Button
variant="outline"
size="icon-sm"
onClick={onRefresh}
disabled={isRefreshing}
aria-label="Refresh"
title="Refresh"
data-testid="datatable-refresh"
>
<RefreshCw className={isRefreshing ? "animate-spin" : ""} />
</Button>
)}
{showViewOptions && <DataTableViewOptions table={table} label="Columns" />}
{onOpenFilters !== undefined && (
<Button variant="outline" size="sm" onClick={onOpenFilters} data-testid="datatable-filters-trigger">
<SlidersHorizontal />
Filters
{filters.length > 0 && (
<Badge className="ml-1 h-5 min-w-5 justify-center rounded-full px-1" data-testid="datatable-filter-count">
{filters.length}
</Badge>
)}
</Button>
)}
{showReset && <ResetFiltersButton onClick={onResetFilters} />}
</div>
{children !== undefined && <div className="flex flex-wrap items-center gap-2">{children}</div>}
</div>
);
}

View file

@ -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<TData>({ table, label = "View", className }
<Menu.Trigger
render={
<Button variant="outline" size="sm" className={className} data-testid="view-options-trigger">
<SlidersHorizontal />
<Columns3 />
{label}
</Button>
}
@ -44,7 +44,8 @@ export function DataTableViewOptions<TData>({ table, label = "View", className }
<Menu.CheckboxItemIndicator className="absolute left-2 flex size-4 items-center justify-center">
<Check className="size-3.5" />
</Menu.CheckboxItemIndicator>
{column.columnDef.meta?.title ?? column.id}
{column.columnDef.meta?.title ??
(typeof column.columnDef.header === "string" ? column.columnDef.header : column.id)}
</Menu.CheckboxItem>
))}
</Menu.Popup>

View file

@ -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<TData extends RowData, TValue> {
@ -9,5 +9,6 @@ declare module "@tanstack/react-table" {
headerClassName?: string;
title?: string;
pinned?: ColumnPinnedSide;
skeleton?: DataTableSkeletonShape;
}
}

View file

@ -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";

View file

@ -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<TData extends RowData, TValue> {
data: TData[];
@ -24,6 +27,7 @@ export interface DataTableProps<TData extends RowData, TValue> {
isLoading?: boolean;
loadingMessage?: string;
skeletonRowCount?: number;
noDataMessage?: React.ReactNode;
sortingMode?: SortingMode;
@ -38,6 +42,14 @@ export interface DataTableProps<TData extends RowData, TValue> {
rowCount?: number;
pageSizeOptions?: number[];
filterMode?: FilterMode;
columnFilters?: ColumnFiltersState;
onColumnFiltersChange?: OnChangeFn<ColumnFiltersState>;
defaultColumnFilters?: ColumnFiltersState;
globalFilter?: string;
onGlobalFilterChange?: OnChangeFn<string>;
enableColumnResizing?: boolean;
columnResizeMode?: ColumnResizeMode;
defaultColumnVisibility?: VisibilityState;

View file

@ -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();

View file

@ -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<typeof useKeys>;
const mockUseAuthorized = useAuthorized as MockedFunction<typeof useAuthorized>;
const createMockKey = (overrides: Partial<KeyResponse> = {}): 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<typeof useKeys>;
mockUseKeys.mockReturnValue(result);
} as unknown as ReturnType<typeof useKeys>);
renderWithProviders(<TeamVirtualKeysTable {...defaultProps} />);
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<typeof useKeys>);
renderWithProviders(<TeamVirtualKeysTable {...defaultProps} />);
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(<TeamVirtualKeysTable {...defaultProps} />);
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(<TeamVirtualKeysTable {...defaultProps} teamId="team-filter-options-test" />);
await waitFor(() => {
expect(mockFetchTeamFilterOptions).toHaveBeenCalledWith("test-token", "team-filter-options-test");
expect(screen.getByText("No rows match your search or filters.")).toBeInTheDocument();
});
});

View file

@ -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<KeyResponse | null>(null);
const [sorting, setSorting] = useState<SortingState>(DEFAULT_SORTING);
const [tablePagination, setTablePagination] = useState<PaginationState>({
pageIndex: 0,
pageSize: 50,
});
const [filters, setFilters] = useState<Record<string, string>>({
"Organization ID": "",
"Key Alias": "",
"User ID": "",
});
const [columnFilters, setColumnFilters] = useState<ColumnFiltersState>([]);
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<string, string>) => {
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<OnChangeFn<ColumnFiltersState>>((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<KeyResponse>[] = useMemo(
() => [
{
id: "token",
accessorKey: "token",
meta: { title: "Key ID" },
header: ({ column }) => <DataTableSortHeader column={column} title="Key ID" variant="header-cycle" />,
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 }) => <DataTableSortHeader column={column} title="Key Alias" variant="header-cycle" />,
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 }) => <DataTableSortHeader column={column} title="Created At" variant="header-cycle" />,
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 }) => <DataTableSortHeader column={column} title="Updated At" variant="header-cycle" />,
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 }) => <DataTableSortHeader column={column} title="Spend (USD)" variant="header-cycle" />,
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 }) => <DataTableSortHeader column={column} title="Budget (USD)" variant="header-cycle" />,
size: 110,
enableSorting: true,
@ -503,27 +451,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi
onDelete={refetch}
/>
) : (
<div className="border-b py-4 flex-1 overflow-hidden">
<div className="w-full mb-6">
<FilterComponent
options={filterOptions}
onApplyFilters={handleFilterChange}
initialValues={filters}
onResetFilters={handleFilterReset}
/>
</div>
<div className="w-full mb-4">
<DataTablePagination
page={pageIndex}
pageSize={pageSize}
rowCount={rowCount}
onPageChange={(nextPage) => setTablePagination((prev) => ({ ...prev, pageIndex: nextPage }))}
onPageSizeChange={(nextSize) => setTablePagination({ pageIndex: 0, pageSize: nextSize })}
isLoading={isLoading || isFetching}
/>
</div>
<div className="py-4 flex-1 overflow-hidden">
<DataTable
data={displayKeys}
columns={columns}
@ -534,14 +462,46 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi
pagination={tablePagination}
onPaginationChange={setTablePagination}
rowCount={rowCount}
paginationSlot={() => 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) => (
<>
<DataTableToolbar
table={table}
searchValue={searchInput}
onSearchChange={handleSearchChange}
searchPlaceholder="Search by key alias…"
onRefresh={() => refetch?.()}
isRefreshing={isFetching}
onOpenFilters={() => setFiltersOpen(true)}
filterLabels={{ user_id: "User ID" }}
/>
<DataTableFilterDrawer
table={table}
open={filtersOpen}
onOpenChange={setFiltersOpen}
title="Filters"
description={`Narrow down keys for ${teamAlias ?? "this team"}`}
>
{({ get, set }) => (
<DataTableFilterField label="User ID">
<Input
value={(get("user_id") as string) ?? ""}
onChange={(event) => set("user_id", event.target.value)}
placeholder="Filter by user ID…"
/>
</DataTableFilterField>
)}
</DataTableFilterDrawer>
</>
)}
/>
</div>
)}

View file

@ -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 <SheetPrimitive.Root data-slot="sheet" {...props} />;
}
function SheetTrigger({ ...props }: SheetPrimitive.Trigger.Props) {
return <SheetPrimitive.Trigger data-slot="sheet-trigger" {...props} />;
}
function SheetClose({ ...props }: SheetPrimitive.Close.Props) {
return <SheetPrimitive.Close data-slot="sheet-close" {...props} />;
}
function SheetPortal({ ...props }: SheetPrimitive.Portal.Props) {
return <SheetPrimitive.Portal data-slot="sheet-portal" {...props} />;
}
function SheetOverlay({ className, ...props }: SheetPrimitive.Backdrop.Props) {
return (
<SheetPrimitive.Backdrop
data-slot="sheet-overlay"
className={cn(
"fixed inset-0 z-50 bg-black/10 transition-opacity duration-150 data-ending-style:opacity-0 data-starting-style:opacity-0 supports-backdrop-filter:backdrop-blur-xs",
className,
)}
{...props}
/>
);
}
function SheetContent({
className,
children,
side = "right",
showCloseButton = true,
...props
}: SheetPrimitive.Popup.Props & {
side?: "top" | "right" | "bottom" | "left";
showCloseButton?: boolean;
}) {
return (
<SheetPortal>
<SheetOverlay />
<SheetPrimitive.Popup
data-slot="sheet-content"
data-side={side}
className={cn(
"fixed z-50 flex flex-col gap-4 bg-popover bg-clip-padding text-sm text-popover-foreground shadow-lg transition duration-200 ease-in-out data-ending-style:opacity-0 data-starting-style:opacity-0 data-[side=bottom]:inset-x-0 data-[side=bottom]:bottom-0 data-[side=bottom]:h-auto data-[side=bottom]:border-t data-[side=bottom]:data-ending-style:translate-y-[2.5rem] data-[side=bottom]:data-starting-style:translate-y-[2.5rem] data-[side=left]:inset-y-0 data-[side=left]:left-0 data-[side=left]:h-full data-[side=left]:w-3/4 data-[side=left]:border-r data-[side=left]:data-ending-style:translate-x-[-2.5rem] data-[side=left]:data-starting-style:translate-x-[-2.5rem] data-[side=right]:inset-y-0 data-[side=right]:right-0 data-[side=right]:h-full data-[side=right]:w-3/4 data-[side=right]:border-l data-[side=right]:data-ending-style:translate-x-[2.5rem] data-[side=right]:data-starting-style:translate-x-[2.5rem] data-[side=top]:inset-x-0 data-[side=top]:top-0 data-[side=top]:h-auto data-[side=top]:border-b data-[side=top]:data-ending-style:translate-y-[-2.5rem] data-[side=top]:data-starting-style:translate-y-[-2.5rem] data-[side=left]:sm:max-w-sm data-[side=right]:sm:max-w-sm",
className,
)}
{...props}
>
{children}
{showCloseButton && (
<SheetPrimitive.Close
data-slot="sheet-close"
render={<Button variant="ghost" className="absolute top-4 right-4" size="icon-sm" />}
>
<XIcon />
<span className="sr-only">Close</span>
</SheetPrimitive.Close>
)}
</SheetPrimitive.Popup>
</SheetPortal>
);
}
function SheetHeader({ className, ...props }: React.ComponentProps<"div">) {
return <div data-slot="sheet-header" className={cn("flex flex-col gap-1.5 p-4", className)} {...props} />;
}
function SheetFooter({ className, ...props }: React.ComponentProps<"div">) {
return <div data-slot="sheet-footer" className={cn("mt-auto flex flex-col gap-2 p-4", className)} {...props} />;
}
function SheetTitle({ className, ...props }: SheetPrimitive.Title.Props) {
return (
<SheetPrimitive.Title data-slot="sheet-title" className={cn("font-medium text-foreground", className)} {...props} />
);
}
function SheetDescription({ className, ...props }: SheetPrimitive.Description.Props) {
return (
<SheetPrimitive.Description
data-slot="sheet-description"
className={cn("text-sm text-muted-foreground", className)}
{...props}
/>
);
}
export { Sheet, SheetTrigger, SheetClose, SheetContent, SheetHeader, SheetFooter, SheetTitle, SheetDescription };

View file

@ -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();
});
});

View file

@ -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<paths>({ baseUrl: globalThis.location?.origin ?? "" });
fetchClient.use(middleware);

View file

@ -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");
});
});

View file

@ -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> = () => T;
let baseUrlGetter: Getter<string> = () => resolveApiBase({ explicitBase: process.env.NEXT_PUBLIC_BASE_URL });
let authHeaderNameGetter: Getter<string> = () => "Authorization";
let authTokenGetter: Getter<string | null> = () => null;
let errorHandler: (message: string) => void = () => {};
export const registerBaseUrlGetter = (getter: Getter<string>): void => {
baseUrlGetter = getter;
};
export const registerAuthHeaderNameGetter = (getter: Getter<string>): void => {
authHeaderNameGetter = getter;
};
export const registerAuthTokenGetter = (getter: Getter<string | null>): 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);