mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
4486d9fe0c
74 changed files with 6169 additions and 799 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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))",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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))",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
534
tests/e2e/logging/test_langfuse_e2e.py
Normal file
534
tests/e2e/logging/test_langfuse_e2e.py
Normal 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",
|
||||
)
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
16
ui/litellm-dashboard/package-lock.json
generated
16
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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!),
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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 */}
|
||||
|
|
|
|||
57
ui/litellm-dashboard/src/components/DashboardHeader.test.tsx
Normal file
57
ui/litellm-dashboard/src/components/DashboardHeader.test.tsx
Normal 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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 && (
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
/>
|
||||
</>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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, {
|
||||
|
|
|
|||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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("");
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
100
ui/litellm-dashboard/src/components/ui/sheet.tsx
Normal file
100
ui/litellm-dashboard/src/components/ui/sheet.tsx
Normal 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 };
|
||||
104
ui/litellm-dashboard/src/lib/http/api.test.ts
Normal file
104
ui/litellm-dashboard/src/lib/http/api.test.ts
Normal 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();
|
||||
});
|
||||
});
|
||||
48
ui/litellm-dashboard/src/lib/http/api.ts
Normal file
48
ui/litellm-dashboard/src/lib/http/api.ts
Normal 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);
|
||||
22
ui/litellm-dashboard/src/lib/http/runtime.test.ts
Normal file
22
ui/litellm-dashboard/src/lib/http/runtime.test.ts
Normal 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");
|
||||
});
|
||||
});
|
||||
46
ui/litellm-dashboard/src/lib/http/runtime.ts
Normal file
46
ui/litellm-dashboard/src/lib/http/runtime.ts
Normal 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);
|
||||
Loading…
Add table
Reference in a new issue