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

# Conflicts:
#	ui/litellm-dashboard/eslint-metrics.json
#	ui/litellm-dashboard/eslint-suppressions.json
This commit is contained in:
ryan-crabbe-berri 2026-07-08 14:54:02 -07:00
commit d5f9870612
108 changed files with 5541 additions and 440 deletions

View file

@ -1029,6 +1029,8 @@ jobs:
- *python312_image
working_directory: ~/project
resource_class: large
environment:
REQUEST_TIMEOUT: "180"
steps:
- checkout
@ -1058,7 +1060,8 @@ jobs:
-v -x \
--junitxml=test-results/junit.xml \
--durations=5 \
-n 8"
-n 8 \
--reruns 1 --only-rerun Timeout"
no_output_timeout: 15m
# Store test results

View file

@ -52,6 +52,7 @@ from litellm.integrations.otel.model.semconv import (
GenAIProvider,
JsonRpc,
LiteLLM,
LiteLLMError,
MCPMethod,
Metric,
Network,
@ -87,6 +88,7 @@ __all__ = [
"HTTP",
"JsonRpc",
"LiteLLM",
"LiteLLMError",
"MCP",
"MCPMethod",
"Metric",

View file

@ -16,9 +16,10 @@ from litellm.integrations.otel.model.payloads import (
MCPListToolsSpanData,
MCPToolCallSpanData,
ServiceSpanData,
SpanError,
)
from litellm.integrations.otel.plumbing.providers import to_otel_span_kind
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent, LiteLLMError
from litellm.integrations.otel.model.spans import (
SPAN_REGISTRY,
SpanRole,
@ -49,6 +50,27 @@ _NAME_BUILDERS: dict[SpanRole, Callable[..., str]] = {
_DEDUP_CACHE_MAX = 10_000
def _stamp_otel_error_attributes(span: Span, error_type: str, resolved_message: str) -> None:
"""Stamp the OTel-semconv error attributes (``error.type`` + ``error.message``).
``error_type`` and ``resolved_message`` are ``finish_span``'s already-computed
fallback chains, so the pair on the status, event, and attributes stays in
lockstep."""
span.set_attribute(Error.TYPE, error_type)
span.set_attribute(Error.MESSAGE, resolved_message)
def _stamp_litellm_error_attributes(span: Span, error: SpanError) -> None:
"""Stamp litellm-specific error detail attributes. Emitted only when the
corresponding field is populated so guardrail-shape errors carrying only a
message aren't polluted with empty detail keys."""
if error.code:
span.set_attribute(LiteLLMError.CODE, error.code)
if error.stack_trace:
span.set_attribute(LiteLLMError.STACK_TRACE, error.stack_trace)
if error.llm_provider:
span.set_attribute(LiteLLMError.LLM_PROVIDER, error.llm_provider)
class SpanEmitter:
def __init__(
self,
@ -190,12 +212,13 @@ class SpanEmitter:
if error and (error.error_type or error.message):
error_type = error.error_type or "error"
message = error.message or error.error_type or "error"
span.set_attribute(Error.TYPE, error_type)
_stamp_otel_error_attributes(span, error_type, message)
_stamp_litellm_error_attributes(span, error)
span.set_status(Status(StatusCode.ERROR, message))
# Carry the full message on the standard ``exception`` event so backends
# map it as full text under ``exception.message``. Setting it as a bare
# string attribute instead lets backends like Elasticsearch dynamic-map
# it to a ``keyword`` capped at 1024 chars, truncating the message.
# Also emit the semconv ``exception`` event so backends that
# dynamic-map unknown string span attrs to ``keyword`` (e.g.
# Elasticsearch with a 1024-char ``ignore_above``) still see the
# full untruncated message on the recognized event field.
span.add_event(
ExceptionEvent.NAME,
{ExceptionEvent.TYPE: error_type, ExceptionEvent.MESSAGE: message},

View file

@ -141,6 +141,9 @@ class LLMCost:
class SpanError:
error_type: str | None = None
message: str | None = None
code: str | None = None
stack_trace: str | None = None
llm_provider: str | None = None
@dataclass(frozen=True)
@ -571,6 +574,9 @@ def _parse_error(payload: "StandardLoggingPayload") -> SpanError | None:
return SpanError(
error_type=as_str(info.get("error_class")) or as_str(info.get("error_code")),
message=as_str(info.get("error_message")) or as_str(payload.get("error_str")),
code=as_str(info.get("error_code")),
stack_trace=as_str(info.get("traceback")),
llm_provider=as_str(info.get("llm_provider")),
)

View file

@ -144,7 +144,27 @@ class Client:
class Error:
"""OTel-defined error attribute keys, from the semconv ``error.*`` registry.
``MESSAGE`` is marked *Deprecated* upstream in favor of domain-specific
error message keys plus ``exception.message`` on the exception event, but
is still defined and stamped by litellm's v1 integration; keeping it here
for byte-for-byte parity."""
TYPE: Final = "error.type"
MESSAGE: Final = "error.message"
class LiteLLMError:
"""LiteLLM-specific error attribute keys. Emitted under the ``error.*``
namespace (not ``litellm.*``) for byte-for-byte compat with the v1
integration in ``opentelemetry.py``; consumers reading these keys on v1
spans read the same keys on v2 spans. OTel semconv does not define any of
these three, and per its extension rules a namespace may carry additional
vendor keys as long as they don't collide with defined names."""
CODE: Final = "error.code"
STACK_TRACE: Final = "error.stack_trace"
LLM_PROVIDER: Final = "error.llm_provider"
class ExceptionEvent:

View file

@ -3619,6 +3619,7 @@ class PrometheusLogger(CustomLogger):
hashed_token=user_api_key_dict.token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_cache_only=True,
)
if key_object:
user_api_key_dict.budget_reset_at = key_object.budget_reset_at

View file

@ -95,6 +95,17 @@ class HealthCheckHelpers:
"""
import litellm
logging_obj = filtered_model_params.get("litellm_logging_obj")
if logging_obj is not None:
api_base = filtered_model_params.get("api_base")
logging_obj.update_from_kwargs(
kwargs=filtered_model_params,
model=filtered_model_params.get("model"),
user=None,
optional_params={},
litellm_params={"api_base": api_base} if api_base else None,
)
if custom_llm_provider in LIST_BATCHES_SUPPORTED_PROVIDERS:
return await litellm.alist_batches(**filtered_model_params)
else:

View file

@ -4377,6 +4377,7 @@ class BedrockConverseMessagesProcessor:
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
message_block=cast(OpenAIMessageContentListBlock, element),
block_type="content_block",
model=model,
)
if _cache_point_block is not None:
_parts.append(_cache_point_block)
@ -4384,7 +4385,7 @@ class BedrockConverseMessagesProcessor:
elif message_block["content"] and isinstance(message_block["content"], str):
_part = BedrockContentBlock(text=messages[msg_i]["content"])
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
message_block, block_type="content_block"
message_block, block_type="content_block", model=model
)
user_content.append(_part)
if _cache_point_block is not None:
@ -4509,6 +4510,7 @@ class BedrockConverseMessagesProcessor:
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
message_block=cast(OpenAIMessageContentListBlock, element),
block_type="content_block",
model=model,
)
if _cache_point_block is not None:
assistants_parts.append(_cache_point_block)
@ -4520,7 +4522,7 @@ class BedrockConverseMessagesProcessor:
# If content is empty/whitespace, skip it (don't add a placeholder)
# Add cache point block for assistant string content
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
assistant_message_block, block_type="content_block"
assistant_message_block, block_type="content_block", model=model
)
if _cache_point_block is not None:
assistant_content.append(_cache_point_block)
@ -4745,6 +4747,7 @@ def _bedrock_converse_messages_pt(
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
message_block=cast(OpenAIMessageContentListBlock, element),
block_type="content_block",
model=model,
)
if _cache_point_block is not None:
_parts.append(_cache_point_block)
@ -4752,7 +4755,7 @@ def _bedrock_converse_messages_pt(
elif message_block["content"] and isinstance(message_block["content"], str):
_part = BedrockContentBlock(text=messages[msg_i]["content"])
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
message_block, block_type="content_block"
message_block, block_type="content_block", model=model
)
user_content.append(_part)
if _cache_point_block is not None:
@ -4882,6 +4885,7 @@ def _bedrock_converse_messages_pt(
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
message_block=cast(OpenAIMessageContentListBlock, element),
block_type="content_block",
model=model,
)
if _cache_point_block is not None:
assistants_parts.append(_cache_point_block)
@ -4892,7 +4896,7 @@ def _bedrock_converse_messages_pt(
assistant_content.append(BedrockContentBlock(text=_assistant_content))
# Add cache point block for assistant string content
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
assistant_message_block, block_type="content_block"
assistant_message_block, block_type="content_block", model=model
)
if _cache_point_block is not None:
assistant_content.append(_cache_point_block)

View file

@ -448,6 +448,12 @@ async def _store_per_user_token_server_side(
)
return # Don't warm Redis if DB write failed
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server.server_id)
# Warm the Redis cache so the first subsequent MCP call is a cache hit
ttl = _compute_per_user_token_ttl(server, expires_in)
await mcp_per_user_token_cache.set(
@ -575,8 +581,12 @@ async def exchange_token_with_server(
if mcp_server.token_url is None:
raise HTTPException(status_code=400, detail="MCP server token url is not set")
# The id and secret must come from the same source. When the server-side client_id wins,
# falling back to the caller's secret pairs the persisted client with a foreign secret; the
# register short-circuit hands clients a placeholder secret ("dummy"), so a re-auth against a
# persisted public PKCE client (no stored secret) would send that placeholder and the IdP 401s.
resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id
resolved_client_secret = mcp_server.client_secret if mcp_server.client_secret else client_secret
resolved_client_secret = mcp_server.client_secret if mcp_server.client_id else client_secret
try:
client_auth = build_token_endpoint_client_auth(
auth_method=mcp_server.token_endpoint_auth_method,

View file

@ -73,3 +73,18 @@ class MCPUpstreamAuthError(Exception):
detail=detail,
headers={"www-authenticate": challenge} if challenge else None,
)
class MCPToolResultError(Exception):
"""An MCP tool call completed with ``isError=True`` in its result.
Never raised on the wire path: streamable HTTP MCP correctly returns tool
failures as HTTP 200 with ``result.isError: true`` per the MCP spec. This
exception only drives the standard failure logging (``status="failure"``
payload, OTel ERROR span) for such results.
Lives here rather than ``utils.py`` deliberately: tests reload ``utils``
to re-read its env-derived constants, and a reload would fork this class
into two identities, breaking ``isinstance`` checks against instances
created before the reload.
"""

View file

@ -70,6 +70,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import
to_server_spec,
to_subject,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
InvalidatableOAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
)
@ -689,9 +692,16 @@ class MCPServerManager:
"""
return auth_type == MCPAuth.oauth2_token_exchange and not (token_exchange_endpoint or token_url)
def __init__(self, cred_provider: Optional[UpstreamCredentialProvider] = None):
def __init__(
self,
cred_provider: Optional[UpstreamCredentialProvider] = None,
per_user_oauth_token_store: Optional[InvalidatableOAuthTokenStore] = None,
):
self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore(
self.get_mcp_server_by_id
)
self._cred_provider = cred_provider or UpstreamCredentialProvider(
oauth_token_store=LazyPerUserOAuthTokenStore(self.get_mcp_server_by_id),
oauth_token_store=self._per_user_oauth_token_store,
token_exchanger=build_token_exchanger(),
)
self.registry: dict[str, MCPServer] = {}
@ -715,8 +725,10 @@ class MCPServerManager:
# Per-server outbound tool-call concurrency limiters, lazily created from
# each server's max_concurrent_requests. Keyed by server_id so the cap
# survives the registry atomic-swap on config reload; a missing key means
# the server has no configured limit.
self._server_call_semaphores: dict[str, asyncio.Semaphore] = {}
# the server has no configured limit. The limit is cached alongside the
# semaphore so an edited limit rebuilds it instead of keeping the old cap
# until restart.
self._server_call_semaphores: dict[str, tuple[int, asyncio.Semaphore]] = {}
self.tool_name_to_mcp_server_name_mapping: dict[str, str] = {}
"""
{
@ -3594,10 +3606,11 @@ class MCPServerManager:
limit = mcp_server.max_concurrent_requests
if limit is None or limit <= 0:
return None
semaphore = self._server_call_semaphores.get(mcp_server.server_id)
if semaphore is None:
semaphore = asyncio.Semaphore(limit)
self._server_call_semaphores[mcp_server.server_id] = semaphore
cached = self._server_call_semaphores.get(mcp_server.server_id)
if cached is not None and cached[0] == limit:
return cached[1]
semaphore = asyncio.Semaphore(limit)
self._server_call_semaphores[mcp_server.server_id] = (limit, semaphore)
return semaphore
@asynccontextmanager
@ -3919,6 +3932,19 @@ class MCPServerManager:
return False
return await self._cred_provider.has_user_token(to_subject(user_api_key_auth, None), spec)
async def invalidate_user_oauth_token_cache(self, user_id: str, server_id: str) -> None:
"""Drop the v2 chain's cached token for ``(user_id, server_id)`` after the credential row
changes (re-auth, revoke), so the next resolve reads the new row instead of serving the
replaced token until its cache TTL. Best-effort: a cache-drop failure is logged, never
raised, because the DB write already succeeded and the TTL remains the backstop.
"""
try:
await self._per_user_oauth_token_store.invalidate(user_id, server_id)
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
verbose_logger.warning(
"Failed to invalidate cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
)
async def _resolve_oauth2_headers_for_tool_call(
self,
mcp_server: MCPServer,

View file

@ -69,6 +69,17 @@ class OAuthTokenStore(Protocol):
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: ...
class InvalidatableOAuthTokenStore(OAuthTokenStore, Protocol):
"""An ``OAuthTokenStore`` whose cached entry for a ``(user, server)`` pair can be dropped.
The write side calls ``invalidate`` after a (re)authorization or revocation changes the
credential row, so reads stop serving the replaced token immediately instead of until its
cache TTL. ``CachedOAuthTokenStore`` (the top of the per-user chain) satisfies this.
"""
async def invalidate(self, user_id: str, server_id: str) -> None: ...
class TokenRefresher(Protocol):
"""Mints a fresh token from an expired one and persists it, returning the new token.

View file

@ -24,8 +24,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.dual_cache_toke
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
CachedOAuthTokenStore,
InvalidatableOAuthTokenStore,
OAuthToken,
OAuthTokenStore,
RefreshCoordinator,
RefreshingTokenStore,
TokenCacheBackend,
@ -51,7 +51,7 @@ if TYPE_CHECKING:
_DEFAULT_TTL_SECONDS = 300.0
ServerLookup = Callable[[str], "MCPServer | None"]
StoreBuilder = Callable[[ServerLookup], tuple[OAuthTokenStore, bool]]
StoreBuilder = Callable[[ServerLookup], tuple[InvalidatableOAuthTokenStore, bool]]
async def _read_credential(user_id: str, server_id: str) -> dict[str, object] | None:
@ -185,7 +185,7 @@ class LazyPerUserOAuthTokenStore:
self._server_lookup = server_lookup
self._store_builder = store_builder
self._redis_available = redis_available
self._store: OAuthTokenStore | None = None
self._store: InvalidatableOAuthTokenStore | None = None
self._uses_redis = False
self._fetch_lock = asyncio.Condition()
self._local_fetches = 0
@ -203,7 +203,26 @@ class LazyPerUserOAuthTokenStore:
if not uses_redis:
await self._finish_local_fetch()
async def _store_for_fetch(self) -> tuple[OAuthTokenStore, bool]:
async def invalidate(self, user_id: str, server_id: str) -> None:
"""Drop the chain's cached entry for ``(user_id, server_id)`` after the credential row
changes (re-auth, revoke). Builds the chain if no fetch has run yet, so a shared (Redis)
cache entry written by another worker is dropped too; the in-process case is then a no-op
on an empty cache.
"""
if self._uses_redis:
store = self._store
if store is not None:
await store.invalidate(user_id, server_id)
return
store, uses_redis = await self._store_for_fetch()
try:
await store.invalidate(user_id, server_id)
finally:
if not uses_redis:
await self._finish_local_fetch()
async def _store_for_fetch(self) -> tuple[InvalidatableOAuthTokenStore, bool]:
async with self._fetch_lock:
while (
self._store is not None and not self._uses_redis and self._redis_available() and self._local_fetches > 0

View file

@ -8,6 +8,7 @@ from typing import (
Dict,
List,
Literal,
Mapping,
Optional,
Set,
Tuple,
@ -78,7 +79,7 @@ if MCP_AVAILABLE:
MCPInfo,
MCPServer,
_apply_toolset_scope,
_fire_mcp_success_logging,
_fire_mcp_tool_call_logging,
_tool_name_matches,
execute_mcp_tool,
filter_tools_by_allowed_tools,
@ -86,23 +87,32 @@ if MCP_AVAILABLE:
########################################################
############ MCP Server REST API Routes #################
async def _safe_fire_mcp_success_logging(
async def _safe_fire_mcp_tool_call_logging(
logging_obj: Optional[Any],
result: Any,
start_time: datetime,
end_time: datetime,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
request_data: Optional[Mapping[str, object]] = None,
) -> None:
if logging_obj is None:
return
logging_results = await asyncio.gather(
_fire_mcp_success_logging(logging_obj, result, start_time, end_time),
_fire_mcp_tool_call_logging(
logging_obj,
result,
start_time,
end_time,
user_api_key_auth=user_api_key_auth,
request_data=request_data,
),
return_exceptions=True,
)
logging_error = logging_results[0]
if isinstance(logging_error, asyncio.CancelledError):
raise logging_error
if isinstance(logging_error, BaseException):
verbose_logger.warning("MCP tool success logging failed (continuing): %s", logging_error)
verbose_logger.warning("MCP tool call logging failed (continuing): %s", logging_error)
def _get_server_auth_header(
server,
@ -872,7 +882,14 @@ if MCP_AVAILABLE:
raw_headers=virtual_raw_headers,
litellm_logging_obj=virtual_logging_obj,
)
await _safe_fire_mcp_success_logging(virtual_logging_obj, result, _tool_start_time, datetime.now())
await _safe_fire_mcp_tool_call_logging(
virtual_logging_obj,
result,
_tool_start_time,
datetime.now(),
user_api_key_auth=user_api_key_dict,
request_data=data,
)
return result
# Validate required parameters early
@ -955,7 +972,14 @@ if MCP_AVAILABLE:
litellm_logging_obj=data.get("litellm_logging_obj"),
requested_server_id=canonical_server_id,
)
await _safe_fire_mcp_success_logging(logging_obj, result, _tool_start_time, datetime.now())
await _safe_fire_mcp_tool_call_logging(
logging_obj,
result,
_tool_start_time,
datetime.now(),
user_api_key_auth=user_api_key_dict,
request_data=data,
)
return result
except MCPMissingUserEnvVarsError as e:
verbose_logger.info(

View file

@ -20,6 +20,7 @@ from typing import (
Callable,
Dict,
List,
Mapping,
Optional,
Set,
Tuple,
@ -47,7 +48,10 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
get_request_base_url,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPToolResultError,
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
_mcp_gateway_initialize_instructions,
@ -60,6 +64,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
LITELLM_MCP_SERVER_VERSION,
MCPMissingUserEnvVarsError,
add_server_prefix_to_name,
extract_mcp_tool_result_error_message,
get_server_prefix,
iter_known_server_prefixes,
)
@ -716,7 +721,7 @@ if MCP_AVAILABLE:
if not (host_ctx and hasattr(host_ctx, "meta") and host_ctx.meta):
return None
host_token = getattr(host_ctx.meta, "progressToken", None)
if not (host_token and hasattr(host_ctx, "session") and host_ctx.session):
if host_token is None or not (hasattr(host_ctx, "session") and host_ctx.session):
return None
host_session = host_ctx.session
@ -732,7 +737,7 @@ if MCP_AVAILABLE:
except Exception as e:
verbose_logger.error(f"Failed to forward progress to Host: {e}")
verbose_logger.debug(f"Host progressToken captured: {host_token[:8]}...")
verbose_logger.debug(f"Host progressToken captured: {str(host_token)[:8]}...")
return forward_progress
async def _build_virtual_call_logging_obj(
@ -2743,12 +2748,40 @@ if MCP_AVAILABLE:
return response
async def _fire_mcp_success_logging(
_MCP_CREDENTIAL_REQUEST_FIELDS = frozenset(
{
"raw_headers",
"mcp_auth_header",
"mcp_server_auth_headers",
"oauth2_headers",
"user_api_key_auth",
}
)
async def _fire_mcp_tool_call_logging(
logging_obj: LiteLLMLoggingObj,
result: Any,
start_time: datetime,
end_time: datetime,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
request_data: Optional[Mapping[str, object]] = None,
) -> None:
"""Fire post-call logging for an executed MCP tool call.
A result with ``isError=True`` is logged as a failure (``status="failure"``
payload, so OTel marks the span ERROR) while the HTTP wire behavior stays
200 + ``isError: true`` per the MCP spec. The error check runs after
``async_post_mcp_tool_call_hook`` because guardrails may flip the result
to ``isError=True`` in that hook. Raised exceptions never reach here (the
``@client`` wrapper and ``call_mcp_tool``'s except path log those), so
this cannot double-log a failure.
``request_data`` may carry credential-bearing fields (the REST path puts
``raw_headers``, ``mcp_auth_header``, ``mcp_server_auth_headers``, and
``oauth2_headers`` at the top level of its data dict), so those are
stripped before the dict is handed to ``post_call_failure_hook``
callbacks.
"""
logging_obj.post_call(original_response=result)
await logging_obj.async_post_mcp_tool_call_hook(
kwargs=logging_obj.model_call_details,
@ -2757,7 +2790,31 @@ if MCP_AVAILABLE:
end_time=end_time,
)
logging_obj.call_type = CallTypes.call_mcp_tool.value
await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time)
error_message = extract_mcp_tool_result_error_message(result)
if error_message is None:
await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time)
return
logging_obj.has_run_logging(event_type="sync_success")
logging_obj.has_run_logging(event_type="async_success")
tool_error = MCPToolResultError(error_message)
logging_obj.failure_handler(tool_error, "", start_time, end_time)
await logging_obj.async_failure_handler(tool_error, "", start_time, end_time)
if user_api_key_auth is None:
return
from litellm.proxy.proxy_server import proxy_logging_obj
if proxy_logging_obj:
sanitized_request_data = {
key: value for key, value in (request_data or {}).items() if key not in _MCP_CREDENTIAL_REQUEST_FIELDS
}
await proxy_logging_obj.post_call_failure_hook(
request_data=sanitized_request_data,
original_exception=tool_error,
user_api_key_dict=user_api_key_auth,
route="/mcp/call_tool",
)
@client
async def call_mcp_tool(
@ -2833,7 +2890,14 @@ if MCP_AVAILABLE:
raise
if litellm_logging_obj:
await _fire_mcp_success_logging(litellm_logging_obj, response, start_time, datetime.now())
await _fire_mcp_tool_call_logging(
logging_obj=litellm_logging_obj,
result=response,
start_time=start_time,
end_time=datetime.now(),
user_api_key_auth=user_api_key_auth,
request_data=kwargs,
)
return response
async def mcp_get_prompt(

View file

@ -415,6 +415,25 @@ def validate_mcp_server_name(server_name: str, raise_http_exception: bool = Fals
raise Exception(error_message)
def extract_mcp_tool_result_error_message(result: object) -> Optional[str]:
"""The first text content of an ``isError=True`` tool result, or ``None``
when the result is not an error.
Accepts both ``mcp.types.CallToolResult`` objects and their dict
equivalents, duck-typed so the ``mcp`` package is not required.
"""
is_error: object = result.get("isError") if isinstance(result, Mapping) else getattr(result, "isError", None)
if is_error is not True:
return None
content: object = result.get("content") if isinstance(result, Mapping) else getattr(result, "content", None)
if isinstance(content, (list, tuple)):
for item in content:
text: object = item.get("text") if isinstance(item, Mapping) else getattr(item, "text", None)
if isinstance(text, str) and text:
return text
return "MCP tool call returned isError=true"
TOOL_DISPLAY_NAME_PATTERN = re.compile(r"^[a-zA-Z0-9_-]+$")

View file

@ -1045,6 +1045,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
model_rpm_limit: Optional[dict] = None
model_tpm_limit: Optional[dict] = None
mcp_rpm_limit: Optional[Dict[str, int]] = None
tag_rpm_limit: Optional[dict[str, int]] = None
guardrails: Optional[List[str]] = None
policies: Optional[List[str]] = None
prompts: Optional[List[str]] = None
@ -3869,6 +3870,7 @@ LiteLLM_ManagementEndpoint_MetadataFields = [
"model_rpm_limit",
"model_tpm_limit",
"mcp_rpm_limit",
"tag_rpm_limit",
"rpm_limit_type",
"tpm_limit_type",
"enforced_params",

View file

@ -1474,7 +1474,7 @@ def _should_check_db(key: str, last_db_access_time: LimitedSizeOrderedDict, db_c
elif last_db_access_time[key][0] is not None: # check db for non-null values (for refresh operations)
return True
elif last_db_access_time[key][0] is None:
if current_time - last_db_access_time[key] >= db_cache_expiry:
if current_time - last_db_access_time[key][1] >= db_cache_expiry:
return True
return False
@ -1649,6 +1649,12 @@ async def get_user_object(
include={"organization_memberships": True},
)
else:
if should_check_db:
_update_last_db_access_time(
key=db_access_time_key,
value=None,
last_db_access_time=last_db_access_time,
)
raise Exception
if response.organization_memberships is not None and len(response.organization_memberships) > 0:

View file

@ -975,6 +975,20 @@ def get_team_mcp_rpm_limit(
return None
def get_key_tag_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[dict[str, int]]:
"""
Get the per-request-tag rpm limit configured on a given api key.
The returned dict is keyed by request tag, so each tag/group tracked on
the key gets its own independent RPM counter.
"""
if user_api_key_dict.metadata:
return user_api_key_dict.metadata.get("tag_rpm_limit")
return None
def get_project_model_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, int]]:

View file

@ -31,8 +31,12 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata
from litellm.proxy.auth.auth_utils import (
get_key_tag_rpm_limit,
get_model_rate_limit_from_metadata,
)
from litellm.proxy.auth.budget_throttle import throttled_limit
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
from litellm.proxy.common_utils.proxy_rate_limit_error import (
ProxyRateLimitError,
map_v3_rate_limit_type,
@ -1300,6 +1304,43 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
)
def _add_tag_per_key_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
data: dict,
descriptors: list[RateLimitDescriptor],
) -> None:
"""
Add per-request-tag rpm limit descriptors for the API key.
Each tag carried on the request that has a configured limit gets its own
``{api_key}:{tag}`` counter, so a burst on one tag/group never consumes
another's budget. Tags without a configured limit fall through to the
key-level descriptor.
"""
if not user_api_key_dict.api_key:
return
tag_rpm_limit = get_key_tag_rpm_limit(user_api_key_dict) or {}
if not tag_rpm_limit:
return
for tag in dict.fromkeys(get_tags_from_request_body(data)):
rpm_limit = tag_rpm_limit.get(tag)
if rpm_limit is None:
continue
descriptors.append(
RateLimitDescriptor(
key="tag_per_key",
value=f"{user_api_key_dict.api_key}:{tag}",
rate_limit={
"requests_per_unit": rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
)
)
def _add_mcp_per_key_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -1645,6 +1686,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
descriptors=descriptors,
)
# Per-request-tag rate limits scoped to this key
self._add_tag_per_key_rate_limit_descriptor(
user_api_key_dict=user_api_key_dict,
data=data,
descriptors=descriptors,
)
# REST MCP calls pass the raw body through this hook before server
# resolution; only the later synthetic hook payload may carry this key.
if call_type == CallTypes.call_mcp_tool.value and "server_id" not in data:
@ -1961,6 +2009,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# Org Level Rate Limits
descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model))
# Only check rate limits if we have descriptors with actual limits
if descriptors:
# First pass: RPM and max_parallel_requests sliding-window check.

View file

@ -377,6 +377,7 @@ async def new_user(
- budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
- model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user.
- tag_rpm_limit: Optional[dict] - Per-request-tag rpm limit, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Enforced for keys only; values set on a user are stored but not enforced per user.
- model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
- agent_id: Optional[str] - The agent id associated with the user.
@ -1379,6 +1380,7 @@ async def user_update(
- budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
- model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user.
- tag_rpm_limit: Optional[dict] - Per-request-tag rpm limit, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Enforced for keys only; values set on a user are stored but not enforced per user.
- model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
- agent_id: Optional[str] - The agent id associated with the user.

View file

@ -111,6 +111,7 @@ from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
from litellm.router import Router
from litellm.secret_managers.base_secret_manager import raise_if_unsafe_secret_name
from litellm.secret_managers.main import get_secret
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequest,
@ -1496,6 +1497,7 @@ async def generate_key_fn(
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
- tag_rpm_limit: Optional[dict] - key-specific per-request-tag rpm limit, keyed by request tag. Example - {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; requests whose tag is absent fall back to the key-level rpm limit.
- tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
- rpm_limit_type: Optional[str] - Type of rpm limit. Options: "best_effort_throughput" (no error if we're overallocating rpm), "guaranteed_throughput" (raise an error if we're overallocating rpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
- allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request
@ -2514,6 +2516,7 @@ async def update_key_fn(
- rpm_limit: Optional[int] - Requests per minute limit
- model_rpm_limit: Optional[dict] - Model-specific RPM limits {"gpt-4": 100, "claude-v1": 200}
- mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200}
- tag_rpm_limit: Optional[dict] - Per-request-tag RPM limits, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; absent tags fall back to the key-level rpm limit.
- model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000}
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
@ -3551,6 +3554,7 @@ async def generate_key_helper_fn(
model_rpm_limit: Optional[dict] = None,
model_tpm_limit: Optional[dict] = None,
mcp_rpm_limit: Optional[dict] = None,
tag_rpm_limit: Optional[dict] = None,
guardrails: Optional[list] = None,
policies: Optional[list] = None,
prompts: Optional[list] = None,
@ -3624,6 +3628,9 @@ async def generate_key_helper_fn(
if mcp_rpm_limit is not None:
metadata = metadata or {}
metadata["mcp_rpm_limit"] = mcp_rpm_limit
if tag_rpm_limit is not None:
metadata = metadata or {}
metadata["tag_rpm_limit"] = tag_rpm_limit
if guardrails is not None:
metadata = metadata or {}
metadata["guardrails"] = guardrails
@ -6285,8 +6292,13 @@ def _validate_key_alias_format(key_alias: Optional[str]) -> None:
"""
Validate the format of the key_alias.
Gated behind ``litellm.enable_key_alias_format_validation`` (default **False**).
When disabled, no validation is performed so existing workflows are not broken.
A baseline validation always runs, regardless of
``litellm.enable_key_alias_format_validation``.
The remaining charset/length rules are gated behind
``litellm.enable_key_alias_format_validation`` (default **False**). When disabled,
only the baseline validation above is performed, so existing workflows are not
broken.
Rules (when enabled):
- None is OK (no alias).
@ -6294,10 +6306,20 @@ def _validate_key_alias_format(key_alias: Optional[str]) -> None:
- start/end with alphanumeric
- only allow a-zA-Z0-9_-/.@
"""
if not litellm.enable_key_alias_format_validation:
if key_alias is None:
return
if key_alias is None:
try:
raise_if_unsafe_secret_name(key_alias)
except ValueError:
raise ProxyException(
message="Invalid key_alias",
type=ProxyErrorTypes.bad_request_error,
param="key_alias",
code=400,
)
if not litellm.enable_key_alias_format_validation:
return
if not _KEY_ALIAS_PATTERN.match(key_alias):

View file

@ -1913,6 +1913,11 @@ if MCP_AVAILABLE:
expires_in=payload.expires_in,
scopes=payload.scopes,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id)
# Read back the persisted record so the response reflects the stored
# expires_at rather than recomputing it here (which could diverge by
# milliseconds or if the storage logic ever adds a grace period).
@ -1953,6 +1958,11 @@ if MCP_AVAILABLE:
await delete_user_credential(prisma_client, user_id, server_id)
except RecordNotFoundError:
pass # Already gone — treat as a successful delete
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id)
return MCPOAuthUserCredentialStatus(
server_id=server_id,
has_credential=False,

View file

@ -113,7 +113,7 @@ from litellm.proxy.utils import (
from litellm.repositories.table_repositories import SSOConfigRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
from litellm.secret_managers.main import get_secret_bool, str_to_bool
from litellm.secret_managers.main import get_secret_bool, get_secret_str, str_to_bool
from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
@ -2134,12 +2134,16 @@ async def cli_poll_key(
models=session_data.get("models", []),
)
user_db_obj = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
try:
user_db_obj = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except ValueError as e:
verbose_proxy_logger.debug(f"CLI poll: user lookup failed, proceeding without user budget: {e}")
user_db_obj = None
user_budget = user_db_obj.max_budget if user_db_obj is not None else None
team_budget: Optional[float] = None
@ -3733,8 +3737,7 @@ class MicrosoftSSOHandler:
Handles Microsoft SSO callback response and returns a CustomOpenID object
"""
graph_api_base_url = "https://graph.microsoft.com/v1.0"
graph_api_user_groups_endpoint = f"{graph_api_base_url}/me/memberOf"
DEFAULT_GRAPH_API_BASE_URL = "https://graph.microsoft.com/v1.0"
"""
Constants
@ -3744,6 +3747,19 @@ class MicrosoftSSOHandler:
# used for debugging to show the user groups litellm found from Graph API
GRAPH_API_RESPONSE_KEY = "graph_api_user_groups"
@staticmethod
def get_graph_api_base_url() -> str:
"""
Returns the Microsoft Graph API base URL, configurable via the
`MICROSOFT_GRAPH_ENDPOINT` env var so non-default clouds such as Azure
Government (GCC High) can point at `https://graph.microsoft.us/v1.0`
"""
return get_secret_str("MICROSOFT_GRAPH_ENDPOINT") or MicrosoftSSOHandler.DEFAULT_GRAPH_API_BASE_URL
@staticmethod
def get_graph_api_user_groups_endpoint() -> str:
return f"{MicrosoftSSOHandler.get_graph_api_base_url()}/me/memberOf"
@staticmethod
async def get_microsoft_callback_response(
request: Request,
@ -3920,7 +3936,7 @@ class MicrosoftSSOHandler:
# Fetch user membership from Microsoft Graph API
all_group_ids = []
next_link: Optional[str] = MicrosoftSSOHandler.graph_api_user_groups_endpoint
next_link: Optional[str] = MicrosoftSSOHandler.get_graph_api_user_groups_endpoint()
auth_headers = {"Authorization": f"Bearer {access_token}"}
page_count = 0
@ -4003,7 +4019,7 @@ class MicrosoftSSOHandler:
Users use Enterprise Applications to manage Groups and Users on Microsoft Entra ID
"""
base_url = "https://graph.microsoft.com/v1.0"
base_url = MicrosoftSSOHandler.get_graph_api_base_url()
# Endpoint to get app role assignments for the given service principal
endpoint = f"/servicePrincipals/{service_principal_id}/appRoleAssignedTo"
url = base_url + endpoint

View file

@ -7,7 +7,7 @@ import traceback
from base64 import b64encode
from datetime import datetime
from itertools import groupby
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union, cast
from typing import Any, AsyncGenerator, Dict, List, Mapping, Optional, Tuple, Union, cast
from urllib.parse import urlencode, urlparse
import httpx
@ -389,18 +389,24 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
forward_multipart: bool = False,
) -> httpx.Response:
"""
Handle non-streaming HTTP requests
Handle non-SSE HTTP requests
Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests
Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests.
GET and generic requests are sent with httpx stream semantics so the caller can
decide from the response headers whether to buffer the body (JSON, inspected for
logging/guardrails) or relay it to the client without materializing it in memory
(LIT-4009: large batch results files must not be buffered in proxy RSS).
"""
if request.method == "GET":
response = await async_client.request(
method=request.method,
url=url,
get_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
)
elif HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart:
return await async_client.send(get_request, stream=True)
if HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart:
# Forward multipart via make_multipart_http_request even when _parsed_body is
# non-empty (pass_through_request always injects litellm_logging_obj, etc.).
# forward_multipart is False when custom_body was supplied (JSON body despite
@ -412,16 +418,14 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
headers=headers,
requested_query_params=requested_query_params,
)
else:
# Generic httpx method
response = await async_client.request(
method=request.method,
url=url,
headers=headers,
params=requested_query_params,
json=_parsed_body,
)
return response
generic_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
json=_parsed_body,
)
return await async_client.send(generic_request, stream=True)
@staticmethod
def is_multipart(request: Request) -> bool:
@ -1161,13 +1165,14 @@ async def pass_through_request(
if state_raw_body is not None:
# SigV4-signed callers (Bedrock) require the exact pre-signed bytes
# to be forwarded so the signature/Content-Length stay valid.
response = await async_client.request(
method=request.method,
url=url,
raw_body_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
content=state_raw_body,
)
response = await async_client.send(raw_body_request, stream=True)
else:
response = await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler(
request=request,
@ -1223,6 +1228,40 @@ async def pass_through_request(
status_code=response.status_code,
)
if not _should_buffer_passthrough_response(response):
relay_custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=litellm_call_id,
model_id=None,
cache_key=None,
api_base=str(url._uri_reference),
)
relay_callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=_parsed_body or {},
user_api_key_dict=user_api_key_dict,
response=response,
request_headers=dict(request.headers),
)
if relay_callback_headers:
relay_custom_headers.update(relay_callback_headers)
return StreamingResponse(
_relay_passthrough_response_bytes(
response=response,
request_body=_parsed_body or {},
url_route=str(url),
start_time=start_time,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
success_handler_kwargs=kwargs,
),
status_code=response.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=response.headers,
custom_headers=relay_custom_headers,
),
)
content = await response.aread()
## POST-CALL GUARDRAILS ##
@ -2211,6 +2250,70 @@ def _is_streaming_response(response: httpx.Response) -> bool:
return False
def _should_buffer_passthrough_response(response: httpx.Response) -> bool:
"""
Decide from the response headers whether the body must be read into memory.
JSON bodies (and upstream errors) stay buffered: spend logging, guardrails and
managed-id rewriting inspect them, and they are small in practice. Everything
else (jsonl batch results, octet-stream files, ...) is relayed to the client
chunk by chunk so a large body is never resident in full (LIT-4009). A missing
content-type is buffered because the body cannot be classified.
"""
if response.status_code >= 400:
return True
media_type = response.headers.get("content-type", "").split(";")[0].strip().lower()
return media_type in ("", "application/json") or media_type.endswith("+json")
async def _relay_passthrough_response_bytes(
response: httpx.Response,
request_body: dict,
url_route: str,
start_time: datetime,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str],
success_handler_kwargs: dict,
) -> AsyncGenerator[bytes, None]:
"""
Yield upstream bytes to the client without accumulating them, then fire the
passthrough success handler with response_body=None (uninspected body). The
finally block also runs on client disconnect (GeneratorExit) so partial
downloads still produce a spend-log row, mirroring chunk_processor; a
disconnect additionally logs a warning with the number of bytes relayed so
partial deliveries are distinguishable from complete ones in proxy logs.
"""
bytes_relayed = 0
upstream_fully_relayed = False
try:
async for chunk in response.aiter_bytes():
bytes_relayed += len(chunk)
yield chunk
upstream_fully_relayed = True
finally:
if not upstream_fully_relayed:
verbose_proxy_logger.warning(
f"Passthrough stream for {url_route} ended before upstream body was fully relayed; "
f"{bytes_relayed} bytes were sent to the client"
)
await response.aclose()
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler(
httpx_response=response,
response_body=None,
url_route=url_route,
result="",
start_time=start_time,
end_time=datetime.now(),
logging_obj=logging_obj,
cache_hit=False,
request_body=request_body,
custom_llm_provider=custom_llm_provider,
**success_handler_kwargs,
)
)
def _extract_model_from_vertex_ai_setup(setup_response: dict) -> Optional[str]:
"""
Extract the model name from Vertex AI Live setup response.

View file

@ -34,6 +34,18 @@ from .llm_provider_handlers.vertex_passthrough_logging_handler import (
cohere_passthrough_logging_handler = CoherePassthroughLoggingHandler()
def _safe_response_text(httpx_response: httpx.Response) -> str:
"""
Streamed passthrough responses are relayed to the client without being read
into memory, so accessing .text on them raises ResponseNotRead. Their body is
intentionally uninspected; log an empty string instead of failing the row.
"""
try:
return httpx_response.text
except httpx.ResponseNotRead:
return ""
class PassThroughEndpointLogging:
def __init__(self):
self.TRACKED_VERTEX_ROUTES = [
@ -306,7 +318,9 @@ class PassThroughEndpointLogging:
]
kwargs = normalized_llm_passthrough_logging_payload["kwargs"]
if standard_logging_response_object is None:
standard_logging_response_object = StandardPassThroughResponseObject(response=httpx_response.text)
standard_logging_response_object = StandardPassThroughResponseObject(
response=_safe_response_text(httpx_response)
)
kwargs = self._set_cost_per_request(
logging_obj=logging_obj,

View file

@ -1117,45 +1117,6 @@ _OPENAPI_HTTP_METHODS = {
# `_SSO_SENSITIVE_FIELDS` / `_CACHE_SENSITIVE_FIELDS` constants in the SSO
# and cache endpoint files.
_ALERTING_SENSITIVE_VARS: Set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"}
_DB_LITELLM_PARAM_ENV_REF_KEYS = frozenset(
{
"api_key",
"client_secret",
"vertex_credentials",
"vertex_ai_credentials",
"aws_access_key_id",
"aws_secret_access_key",
"aws_session_token",
"aws_region_name",
"aws_session_name",
"aws_profile_name",
"aws_role_name",
"aws_web_identity_token",
"aws_sts_endpoint",
"aws_external_id",
"aws_bedrock_runtime_endpoint",
"aws_bedrock_project_id",
"aws_batch_role_arn",
"aws_workspace_id",
}
)
def _db_model_is_team_scoped(model: object) -> bool:
model_info = getattr(model, "model_info", None)
if isinstance(model_info, BaseModel):
return getattr(model_info, "team_id", None) is not None
if isinstance(model_info, str):
try:
model_info = json.loads(model_info)
except (TypeError, ValueError):
model_info = None
if isinstance(model_info, dict) and model_info.get("team_id") is not None:
return True
if getattr(model_info, "team_id", None) is not None:
return True
model_name = getattr(model, "model_name", None)
return isinstance(model_name, str) and model_name.startswith("model_name_")
def _strip_operation_id_method_suffix(operation_id: str) -> str:
@ -5009,17 +4970,12 @@ class ProxyConfig:
deleted_deployments += 1
return deleted_deployments
def _resolve_db_litellm_param(self, key: str, value: object, resolve_env_refs: bool = True) -> object:
def _resolve_db_litellm_param(self, key: str, value: object) -> object:
if not isinstance(value, str):
return value
decrypted_value = decrypt_value_helper(value=value, key=key, return_original_value=True)
if (
resolve_env_refs
and key in _DB_LITELLM_PARAM_ENV_REF_KEYS
and isinstance(decrypted_value, str)
and decrypted_value.startswith("os.environ/")
):
if isinstance(decrypted_value, str) and decrypted_value.startswith("os.environ/"):
return get_secret(decrypted_value)
return decrypted_value
@ -5040,13 +4996,10 @@ class ProxyConfig:
## ADD MODEL LOGIC
for m in db_models:
_litellm_params = m.litellm_params
resolve_env_refs = not _db_model_is_team_scoped(m)
if isinstance(_litellm_params, dict):
# decrypt values
for k, v in _litellm_params.items():
_litellm_params[k] = self._resolve_db_litellm_param(
key=k, value=v, resolve_env_refs=resolve_env_refs
)
_litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v)
_litellm_params = LiteLLM_Params(**_litellm_params)
else:
@ -5072,15 +5025,12 @@ class ProxyConfig:
_model_list: list = []
for m in new_models:
_litellm_params = m.litellm_params
resolve_env_refs = not _db_model_is_team_scoped(m)
if isinstance(_litellm_params, BaseModel):
_litellm_params = _litellm_params.model_dump()
if isinstance(_litellm_params, dict):
# decrypt values
for k, v in _litellm_params.items():
_litellm_params[k] = self._resolve_db_litellm_param(
key=k, value=v, resolve_env_refs=resolve_env_refs
)
_litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v)
_litellm_params = LiteLLM_Params(**_litellm_params)
else:
verbose_proxy_logger.error(

View file

@ -1,3 +1,4 @@
import re
from abc import ABC, abstractmethod
from typing import Any, Dict, Optional, Union
@ -5,6 +6,20 @@ import httpx
from litellm import verbose_logger
_UNSAFE_SECRET_NAME_PATTERN = re.compile(r"(^|/)\.\.(/|$)|[\x00-\x1f\x7f-\x9f…

]")
def raise_if_unsafe_secret_name(secret_name: str) -> None:
"""
Validate a secret name before it is used by a secret manager integration.
Rejects ".." only as a path segment (bounded by "/" or the start/end of the
string, e.g. "../x", "x/..", or exactly ".."), not as a plain substring, so
names like "release-1.0..2" are not rejected.
"""
if _UNSAFE_SECRET_NAME_PATTERN.search(secret_name):
raise ValueError(f"Invalid secret_name {secret_name!r}")
class BaseSecretManager(ABC):
"""

View file

@ -4,6 +4,7 @@ from typing import Any, Dict, Optional, Union
from urllib.parse import quote
import httpx
import yaml
import litellm
from litellm._logging import verbose_logger
@ -15,7 +16,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._types import KeyManagementSystem
from .base_secret_manager import BaseSecretManager
from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name
from .main import str_to_bool
@ -125,8 +126,11 @@ class CyberArkSecretManager(BaseSecretManager):
"""
# In production, we'd check if the variable exists first
# For now, we'll attempt to create it and ignore if it already exists
raise_if_unsafe_secret_name(secret_name)
policy_url = f"{self.conjur_addr}/policies/{self.conjur_account}/policy/root"
policy_yaml = f"- !variable {secret_name}\n"
# Use a real YAML serializer to build the scalar safely.
quoted_name = yaml.safe_dump(secret_name, default_style='"').strip()
policy_yaml = f"- !variable {quoted_name}\n"
try:
client = _get_httpx_client(params={"ssl_verify": self.ssl_verify})

View file

@ -14,7 +14,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._types import KeyManagementSystem
from .base_secret_manager import BaseSecretManager
from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name
class HashicorpSecretManager(BaseSecretManager):
@ -220,6 +220,7 @@ class HashicorpSecretManager(BaseSecretManager):
- With custom mount: http://127.0.0.1:8200/v1/kv/data/mykey
- With path prefix: http://127.0.0.1:8200/v1/secret/data/myapp/mykey
"""
raise_if_unsafe_secret_name(secret_name)
resolved_namespace = self._sanitize_path_component(namespace if namespace is not None else self.vault_namespace)
resolved_mount = self._sanitize_path_component(mount_name if mount_name is not None else self.vault_mount_name)
if resolved_mount is None:

View file

@ -2619,8 +2619,9 @@ _CACHE_PRICING_FIELDS = (
def _resolve_builtin_model_cost_entry(key: str, provider: str) -> Optional[Dict[str, Any]]:
"""Best-effort lookup of a built-in ``model_cost`` entry for a custom key
whose shape ``get_model_info`` cannot resolve (double provider prefixes
like ``bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region aliases).
whose shape ``get_model_info`` cannot resolve (repeated provider prefixes
like ``bedrock/bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region
aliases).
Returns a copy of the matching entry so the caller can inherit its defaults
(most importantly cache pricing) without mutating the shared built-in.
@ -5052,9 +5053,9 @@ def _get_model_info_from_generalization(
candidates = [
potential_model_names["combined_model_name"],
model,
potential_model_names["split_model"],
potential_model_names["combined_stripped_model_name"],
potential_model_names["stripped_model_name"],
potential_model_names["split_model"],
]
for candidate in candidates:
generalized_info = match_fallback_generalization(candidate)
@ -5094,6 +5095,11 @@ def _get_potential_model_names(
stripped_model_name,
)
if custom_llm_provider in ("bedrock", "bedrock_converse"):
from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix
split_model = strip_bedrock_routing_prefix(split_model)
return PotentialModelNamesAndCustomLLMProvider(
split_model=split_model,
combined_model_name=combined_model_name,
@ -5261,9 +5267,9 @@ def _get_model_info_helper(
Check if: (in order of specificity)
1. 'custom_llm_provider/model' in litellm.model_cost. Checks "groq/llama3-8b-8192" if model="llama3-8b-8192" and custom_llm_provider="groq"
2. 'model' in litellm.model_cost. Checks "gemini-1.5-pro-002" in litellm.model_cost if model="gemini-1.5-pro-002" and custom_llm_provider=None
3. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given.
4. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given.
5. 'split_model' in litellm.model_cost. Checks "llama3-8b-8192" in litellm.model_cost if model="groq/llama3-8b-8192"
3. 'split_model' in litellm.model_cost. Checks "au.anthropic.claude-opus-4-8" in litellm.model_cost if model="bedrock/au.anthropic.claude-opus-4-8"
4. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given.
5. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given.
"""
_model_info: Optional[Dict[str, Any]] = None
@ -5289,6 +5295,16 @@ def _get_model_info_helper(
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(split_model)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(combined_stripped_model_name)
if _matched_key is not None:
@ -5309,16 +5325,6 @@ def _get_model_info_helper(
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(split_model)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
generalization = _get_model_info_from_generalization(

546
tests/_ws_vcr.py Normal file
View file

@ -0,0 +1,546 @@
"""Record and replay realtime WebSocket traffic in the shared VCR Redis store.
The HTTP VCR layer (``tests/_vcr_redis_persister.py`` /
``tests/_vcr_conftest_common.py``) only intercepts httpx/aiohttp, so the
realtime suite always reached the live provider. This module intercepts at the
``websockets.connect`` boundary instead and caches whole WebSocket sessions
under a distinct ``litellm:vcr:wscassette:`` key, reusing the same Redis client,
24h TTL, save-on-pass, and best-effort degradation semantics.
Record mode logs every frame in order with its direction, a text/binary flag,
and, for each server frame, the number of client frames seen before it. That
count is the causal gate for replay: a recorded server frame is only released
once the client has sent at least that many frames, so the deterministic replay
reproduces the same interleaving without a live connection. Client frames are
matched against the recording with volatile fields (ids, timestamps) normalized
away; a structurally different client frame is contract drift and raises loudly
rather than hanging, and every replay wait is bounded by a timeout.
"""
from __future__ import annotations
import asyncio
import base64
import json
import logging
import os
import re
import warnings
from typing import AsyncIterator, Callable, Literal, Optional, Protocol, Union
from pydantic import BaseModel, ConfigDict, ValidationError
from websockets.exceptions import ConnectionClosedOK
from tests._vcr_redis_persister import (
CASSETTE_TTL_SECONDS,
VCRCassetteCacheWarning,
_build_default_client,
_record_cache_failure,
)
WS_REDIS_KEY_PREFIX = "litellm:vcr:wscassette:"
WS_MAX_SESSIONS_PER_CASSETTE = 20
WS_MAX_FRAMES_PER_SESSION = 2000
WS_REPLAY_TIMEOUT_ENV = "LITELLM_WS_VCR_REPLAY_TIMEOUT"
WS_DEFAULT_REPLAY_TIMEOUT_SECONDS = 15.0
WS_CASSETTE_SCHEMA_VERSION = 1
_log = logging.getLogger(__name__)
Message = Union[str, bytes]
Direction = Literal["client_to_server", "server_to_client"]
Opcode = Literal["text", "binary"]
class WsConnectionLike(Protocol):
async def recv(self, decode: Optional[bool] = None) -> Message: ...
async def send(self, message: Message, *args: object, **kwargs: object) -> None: ...
async def close(self, *args: object, **kwargs: object) -> None: ...
def __aiter__(self) -> AsyncIterator[Message]: ...
class WsConnectContextLike(Protocol):
async def __aenter__(self) -> WsConnectionLike: ...
async def __aexit__(self, *exc_info: object) -> Optional[bool]: ...
class RedisLike(Protocol):
def get(self, key: str) -> Optional[bytes]: ...
def set(self, key: str, value: bytes, ex: int) -> object: ...
class WsFrame(BaseModel):
model_config = ConfigDict(frozen=True)
direction: Direction
opcode: Opcode
text: Optional[str] = None
binary_b64: Optional[str] = None
client_frames_before: Optional[int] = None
class WsSession(BaseModel):
model_config = ConfigDict(frozen=True)
frames: tuple[WsFrame, ...]
class WsCassette(BaseModel):
model_config = ConfigDict(frozen=True)
schema_version: int = WS_CASSETTE_SCHEMA_VERSION
sessions: tuple[WsSession, ...]
class WsVcrReplayError(Exception): ...
class WsVcrContractDrift(WsVcrReplayError): ...
class WsVcrReplayTimeout(WsVcrReplayError): ...
_UUID_RE = re.compile(r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}")
_OPENAI_ID_RE = re.compile(
r"\b(?:evt|event|item|msg|resp|response|sess|session|call|fc|rs|conv|ce|audio)_[A-Za-z0-9]{6,}"
)
_ISO_TS_RE = re.compile(r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?")
_EPOCH_RE = re.compile(r"(?<![\d.])1[0-9]{9,12}(?![\d.])")
_VOLATILE_KEYS = frozenset(
{
"event_id",
"item_id",
"previous_item_id",
"response_id",
"id",
"session_id",
"call_id",
"conversation_id",
}
)
_BEARER_RE = re.compile(r"Bearer\s+[A-Za-z0-9._\-]+")
_OPENAI_KEY_RE = re.compile(r"sk-[A-Za-z0-9_\-]{8,}")
_XAI_KEY_RE = re.compile(r"xai-[A-Za-z0-9_\-]{8,}")
def scrub_secrets(text: str) -> str:
scrubbed = _BEARER_RE.sub("Bearer <redacted>", text)
scrubbed = _OPENAI_KEY_RE.sub("<redacted-key>", scrubbed)
scrubbed = _XAI_KEY_RE.sub("<redacted-key>", scrubbed)
return scrubbed
def _normalize_json_for_match(obj: object) -> object:
if isinstance(obj, dict):
return {
str(key): ("<vcr-id>" if key in _VOLATILE_KEYS else _normalize_json_for_match(value))
for key, value in sorted(obj.items(), key=lambda kv: str(kv[0]))
}
if isinstance(obj, list):
return [_normalize_json_for_match(item) for item in obj]
if isinstance(obj, str):
return _normalize_scalar_string(obj)
return obj
def _normalize_scalar_string(text: str) -> str:
normalized = _UUID_RE.sub("<vcr-uuid>", text)
normalized = _OPENAI_ID_RE.sub("<vcr-id>", normalized)
normalized = _ISO_TS_RE.sub("<vcr-iso-ts>", normalized)
normalized = _EPOCH_RE.sub("<vcr-epoch>", normalized)
return normalized
def normalize_text_for_match(text: str) -> str:
try:
parsed = json.loads(text)
except (ValueError, TypeError):
return _normalize_scalar_string(text)
return json.dumps(_normalize_json_for_match(parsed), sort_keys=True, separators=(",", ":"))
def text_frames_match(recorded: str, incoming: str) -> bool:
return normalize_text_for_match(recorded) == normalize_text_for_match(incoming)
def _frame_payload(message: Message) -> tuple[Opcode, Optional[str], Optional[str]]:
if isinstance(message, str):
return "text", scrub_secrets(message), None
try:
decoded = message.decode("utf-8")
except UnicodeDecodeError:
return "binary", None, base64.b64encode(message).decode("ascii")
return "text", scrub_secrets(decoded), None
def ws_redis_key_for(nodeid: str) -> str:
rel = nodeid.replace("::", "/").replace("\\", "/").lstrip("./")
return f"{WS_REDIS_KEY_PREFIX}{rel}"
def replay_timeout_seconds() -> float:
raw = os.environ.get(WS_REPLAY_TIMEOUT_ENV)
if not raw:
return WS_DEFAULT_REPLAY_TIMEOUT_SECONDS
try:
return float(raw)
except ValueError:
return WS_DEFAULT_REPLAY_TIMEOUT_SECONDS
class WsSessionRecorder:
def __init__(self) -> None:
self._frames: list[WsFrame] = []
self._client_count = 0
def record_client_frame(self, message: Message) -> None:
opcode, text, binary_b64 = _frame_payload(message)
self._frames.append(WsFrame(direction="client_to_server", opcode=opcode, text=text, binary_b64=binary_b64))
self._client_count += 1
def record_server_frame(self, message: Message) -> None:
opcode, text, binary_b64 = _frame_payload(message)
self._frames.append(
WsFrame(
direction="server_to_client",
opcode=opcode,
text=text,
binary_b64=binary_b64,
client_frames_before=self._client_count,
)
)
def to_session(self) -> WsSession:
return WsSession(frames=tuple(self._frames))
class RecordingConnection:
def __init__(self, real: WsConnectionLike, recorder: WsSessionRecorder) -> None:
self._real = real
self._recorder = recorder
async def recv(self, decode: Optional[bool] = None) -> Message:
result = await self._real.recv(decode=decode)
self._recorder.record_server_frame(result)
return result
async def send(self, message: Message, *args: object, **kwargs: object) -> None:
self._recorder.record_client_frame(message)
await self._real.send(message, *args, **kwargs)
async def close(self, *args: object, **kwargs: object) -> None:
await self._real.close(*args, **kwargs)
def __aiter__(self) -> AsyncIterator[Message]:
return self._iterate()
async def _iterate(self) -> AsyncIterator[Message]:
async for message in self._real:
self._recorder.record_server_frame(message)
yield message
class ReplayConnection:
def __init__(
self,
session: WsSession,
timeout: float,
on_error: Callable[[WsVcrReplayError], None],
) -> None:
self._server_frames = tuple(f for f in session.frames if f.direction == "server_to_client")
self._client_frames = tuple(f for f in session.frames if f.direction == "client_to_server")
self._timeout = timeout
self._on_error = on_error
self._server_cursor = 0
self._client_cursor = 0
self._client_sent = 0
self._closed = False
self._progress = asyncio.Event()
async def recv(self, decode: Optional[bool] = None) -> Message:
want_bytes = decode is False
while True:
if self._closed or self._server_cursor >= len(self._server_frames):
raise ConnectionClosedOK(None, None)
frame = self._server_frames[self._server_cursor]
needed = frame.client_frames_before or 0
if self._client_sent >= needed:
self._server_cursor += 1
return _materialize_frame(frame, want_bytes)
await self._await_client_progress(needed)
async def _await_client_progress(self, needed: int) -> None:
waiter = self._progress
try:
await asyncio.wait_for(waiter.wait(), timeout=self._timeout)
except asyncio.TimeoutError:
error = WsVcrReplayTimeout(
f"WS-VCR replay stalled: server frame #{self._server_cursor} needs "
f"{needed} client frame(s) but only {self._client_sent} were sent within "
f"{self._timeout}s. The client stopped driving the recorded session."
)
self._on_error(error)
raise error
async def send(self, message: Message, *args: object, **kwargs: object) -> None:
if self._client_cursor >= len(self._client_frames):
error = WsVcrContractDrift(
"WS-VCR contract drift: client sent frame "
f"#{self._client_cursor + 1} but the recording has only "
f"{len(self._client_frames)} client frame(s). Extra frame: {_preview(message)}"
)
self._on_error(error)
raise error
recorded = self._client_frames[self._client_cursor]
if not _client_frame_matches(recorded, message):
error = WsVcrContractDrift(
"WS-VCR contract drift on client frame "
f"#{self._client_cursor + 1}:\n recorded: {_preview_frame(recorded)}\n"
f" got: {_preview(message)}"
)
self._on_error(error)
raise error
self._client_cursor += 1
self._client_sent += 1
self._signal_progress()
async def close(self, *args: object, **kwargs: object) -> None:
self._closed = True
self._signal_progress()
def _signal_progress(self) -> None:
previous = self._progress
self._progress = asyncio.Event()
previous.set()
def __aiter__(self) -> AsyncIterator[Message]:
return self._iterate()
async def _iterate(self) -> AsyncIterator[Message]:
while True:
try:
yield await self.recv()
except ConnectionClosedOK:
return
def _materialize_frame(frame: WsFrame, want_bytes: bool) -> Message:
if frame.opcode == "text":
text = frame.text or ""
return text.encode("utf-8") if want_bytes else text
return base64.b64decode(frame.binary_b64 or "")
def _client_frame_matches(recorded: WsFrame, message: Message) -> bool:
opcode, text, binary_b64 = _frame_payload(message)
if recorded.opcode != opcode:
return False
if opcode == "text":
return text_frames_match(recorded.text or "", text or "")
return recorded.binary_b64 == binary_b64
def _preview(message: Message) -> str:
text = message if isinstance(message, str) else message.decode("utf-8", errors="replace")
return scrub_secrets(text)[:200]
def _preview_frame(frame: WsFrame) -> str:
if frame.opcode == "text":
return (frame.text or "")[:200]
return f"<binary {len(frame.binary_b64 or '')} b64 chars>"
class _RecordingConnect:
def __init__(
self,
real_context: WsConnectContextLike,
recorder: WsSessionRecorder,
on_done: Callable[[WsSessionRecorder], None],
) -> None:
self._real_context = real_context
self._recorder = recorder
self._on_done = on_done
async def __aenter__(self) -> RecordingConnection:
real = await self._real_context.__aenter__()
return RecordingConnection(real, self._recorder)
async def __aexit__(self, *exc_info: object) -> Optional[bool]:
try:
return await self._real_context.__aexit__(*exc_info)
finally:
self._on_done(self._recorder)
class _ReplayConnect:
def __init__(
self,
session: WsSession,
timeout: float,
on_error: Callable[[WsVcrReplayError], None],
) -> None:
self._session = session
self._timeout = timeout
self._on_error = on_error
async def __aenter__(self) -> ReplayConnection:
return ReplayConnection(self._session, self._timeout, self._on_error)
async def __aexit__(self, *exc_info: object) -> bool:
return False
class WsVcrController:
def __init__(
self,
original_connect: Callable[..., WsConnectContextLike],
cassette: Optional[WsCassette],
timeout: float,
) -> None:
self._original_connect = original_connect
self._cassette = cassette
self._timeout = timeout
self._replay_cursor = 0
self._recorded_sessions: list[WsSession] = []
self._errors: list[WsVcrReplayError] = []
self._replayed = False
self._recorded = False
def connect(self, *args: object, **kwargs: object) -> object:
if self._cassette is not None and self._replay_cursor < len(self._cassette.sessions):
session = self._cassette.sessions[self._replay_cursor]
self._replay_cursor += 1
self._replayed = True
return _ReplayConnect(session, self._timeout, self._errors.append)
self._recorded = True
recorder = WsSessionRecorder()
return _RecordingConnect(self._original_connect(*args, **kwargs), recorder, self._finish_recorder)
def _finish_recorder(self, recorder: WsSessionRecorder) -> None:
self._recorded_sessions.append(recorder.to_session())
@property
def errors(self) -> tuple[WsVcrReplayError, ...]:
return tuple(self._errors)
@property
def replayed(self) -> bool:
return self._replayed
@property
def recorded(self) -> bool:
return self._recorded
def built_cassette(self) -> Optional[WsCassette]:
if not self._recorded_sessions:
return None
return WsCassette(sessions=tuple(self._recorded_sessions))
def verdict(self) -> str:
if self._replayed and not self._recorded:
return f"[WS-VCR HIT] sessions={self._replay_cursor} frames={self._played_frame_count()}"
if self._recorded:
cassette = self.built_cassette()
frames = _cassette_frame_count(cassette) if cassette is not None else 0
return f"[WS-VCR MISS] recorded sessions={len(self._recorded_sessions)} frames={frames}"
return "[WS-VCR NOOP] (no websocket traffic)"
def _played_frame_count(self) -> int:
if self._cassette is None:
return 0
return sum(len(s.frames) for s in self._cassette.sessions[: self._replay_cursor])
def _cassette_frame_count(cassette: WsCassette) -> int:
return sum(len(s.frames) for s in cassette.sessions)
def load_ws_cassette(client: RedisLike, key: str) -> Optional[WsCassette]:
from redis.exceptions import RedisError
try:
data = client.get(key)
except RedisError as exc:
_record_cache_failure("load", exc)
message = f"WS-VCR redis load failed for {key}; treating as cache miss: {type(exc).__name__}: {exc}"
_log.warning(message)
warnings.warn(message, VCRCassetteCacheWarning, stacklevel=2)
return None
if data is None:
return None
try:
raw = data.decode("utf-8") if isinstance(data, (bytes, bytearray)) else data
return WsCassette.model_validate_json(raw)
except (ValidationError, ValueError, TypeError) as exc:
_record_cache_failure("load", exc)
message = (
f"WS-VCR redis load failed for {key}; cached payload is corrupt, "
f"treating as cache miss: {type(exc).__name__}: {exc}"
)
_log.warning(message)
warnings.warn(message, VCRCassetteCacheWarning, stacklevel=2)
return None
def save_ws_cassette(
client: RedisLike,
key: str,
cassette: WsCassette,
passed: bool,
ttl_seconds: int = CASSETTE_TTL_SECONDS,
) -> bool:
from redis.exceptions import RedisError
if not passed:
_log.info("WS-VCR redis save skipped for %s; test did not pass - leaving any prior cassette intact", key)
return False
if len(cassette.sessions) > WS_MAX_SESSIONS_PER_CASSETTE:
_log.warning(
"WS-VCR redis save refused for %s; %d sessions (> WS_MAX_SESSIONS_PER_CASSETTE=%d)",
key,
len(cassette.sessions),
WS_MAX_SESSIONS_PER_CASSETTE,
)
return False
if any(len(session.frames) > WS_MAX_FRAMES_PER_SESSION for session in cassette.sessions):
_log.warning(
"WS-VCR redis save refused for %s; a session exceeds WS_MAX_FRAMES_PER_SESSION=%d",
key,
WS_MAX_FRAMES_PER_SESSION,
)
return False
payload = cassette.model_dump_json().encode("utf-8")
try:
client.set(key, payload, ex=ttl_seconds)
except RedisError as exc:
_record_cache_failure("save", exc)
message = f"WS-VCR redis save failed for {key}; cassette not persisted: {type(exc).__name__}: {exc}"
_log.warning(message)
warnings.warn(message, VCRCassetteCacheWarning, stacklevel=2)
return False
return True
def build_ws_cassette_client(
builder: Callable[[], RedisLike] = _build_default_client,
) -> Optional[RedisLike]:
try:
return builder()
except Exception as exc:
_record_cache_failure("load", exc)
message = (
f"WS-VCR redis client unavailable; realtime tests fall back to live "
f"websocket traffic: {type(exc).__name__}: {exc}"
)
_log.warning(message)
warnings.warn(message, VCRCassetteCacheWarning, stacklevel=2)
return None

View file

@ -63,7 +63,7 @@ The harness is fully typed and new code must not add `Any` or widen the basedpyr
The set of tests we want is a registry checked into this repo, one row per behavior; that file is the definition of done and the denominator. Each e2e test declares what it covers with `@pytest.mark.covers("...")`, and a small collector diffs the registry against the tests and ships coverage to the existing Grafana. No Allure, no new dependencies
Coverage is organized as module > feature > test. Dashboard modules are Core LLMs, Non-Core LLMs, MCPs, Management/UI, Reliability & Performance, Logging & Guardrails, and Other. A feature is either an endpoint (`/chat/completions`) or a behavior (fallbacks, rate limits; config-driven, with no route of its own). A cell reads like `llm.chat_completions.bedrock_converse.tool_use.stream.works`
Coverage is organized as module > feature > test. Dashboard modules are `Core LLMs`, `Non-Core LLMs`, `MCPs`, `Management/UI`, `Reliability & Performance`, `Logging & Guardrails`, and `Other`. The Loki stdout formatter maps those display modules to log-safe labels (`core_llms`, `non_core_llms`, `mcp`, `management_ui`, `reliability_performance`, `logging_guardrails`, and `other`) without changing JSON or Prometheus labels. A feature is either an endpoint (`/chat/completions`) or a behavior (fallbacks, rate limits; config-driven, with no route of its own). A cell reads like `llm.chat_completions.bedrock_converse.tool_use.stream.works`
The metric is coverage: the share of registry rows that have a passing covering test, reported to Grafana per module so a gap surfaces as an uncovered row rather than a silent absence
@ -71,7 +71,7 @@ Tests do not declare a dashboard module directly. They only declare the registry
### Naming grammar per module
LLMs - endpoint features (subject = the route), seeded from the Claude Code compat matrix. `chat_completions`, `messages`, and `responses` are Core LLMs. Other LLM endpoints, including `batches` and `realtime`, roll up as Non-Core LLMs.
LLMs - endpoint features (subject = the route), seeded from the Claude Code compat matrix. `chat_completions`, `messages`, and `responses` roll up to `Core LLMs`. Other LLM endpoints, including `batches` and `realtime`, roll up to `Non-Core LLMs`.
```
llm.<endpoint>.<route>.<capability>.<streaming>.<assertion>

View file

@ -9,18 +9,18 @@ note; the naming grammar lives in `tests/e2e/CLAUDE.md`.
A **cell** is one customer-noticeable behavior a single e2e test can assert pass/fail
on, for example `llm.chat_completions.bedrock_converse.tool_use.stream.works`. Cells are
grouped `module > feature > test`, with LLM cells split into Core LLMs and Non-Core
LLMs for dashboarding. Each cell carries a tier (P0/P1/P2), a source, and a
grouped `module > feature > test`, with LLM cells split into `Core LLMs` and
`Non-Core LLMs` for dashboarding. Each cell carries a tier (P0/P1/P2), a source, and a
`fail_before_fix` flag.
The rows live in per-prefix YAML files (`llm_*.yaml`, `mgmt.yaml`, `mcp.yaml`,
`reliability.yaml`, `logging.yaml`, `guardrail.yaml`, `other.yaml`) and validate against
the discriminated union in `schema.py`, so an LLM row cannot carry a guardrail field and
vice versa. `llm` rows with `subject_endpoint` of `chat_completions`, `messages`, or
`responses` roll up to "Core LLMs"; all other LLM endpoints roll up to "Non-Core
LLMs". LLM endpoint, route, and capability values are typed in `schema.py`, so new
taxonomy values require an explicit schema change. `logging` and `guardrail` are two
id-prefixes that roll up into the single "Logging & Guardrails" dashboard module.
`responses` roll up to `Core LLMs`; all other LLM endpoints roll up to `Non-Core LLMs`.
LLM endpoint, route, and capability values are typed in `schema.py`, so new taxonomy
values require an explicit schema change. `logging` and `guardrail` are two id-prefixes
that roll up into the single `Logging & Guardrails` dashboard module.
A test declares what it covers with a marker:
@ -40,8 +40,17 @@ proxy. Whether a covered cell currently passes or fails is a separate, live conc
cd tests/e2e && PYTHONPATH=. python -m coverage_registry.collector
```
Use `--format prometheus` or `--format json` for CI jobs that publish coverage to
Grafana.
Use `--format loki` after the e2e pytest run in the same Kubernetes job/pod to print
structured stdout lines for Loki:
```
cd tests/e2e && PYTHONPATH=. python -m coverage_registry.collector --format loki --strict
```
This emits exactly one `COVERAGE_TOTAL` line and one `COVERAGE_MODULE` line per module
in `MODULE_ORDER`, in that order. Loki uses log-safe `module=` labels from
`LOKI_MODULE_LABELS` (`core_llms`, `management_ui`, etc.) so existing JSON and
Prometheus consumers keep their human-readable module names unchanged.
The headline is overall coverage. The collector also lists markers that point at ids
not in the registry, so a typo or an unenumerated behavior surfaces instead of being

View file

@ -21,7 +21,7 @@ from pathlib import Path
import pytest
from .registry import load_registry
from .schema import MODULE_ORDER, Cell, Tier, dashboard_module
from .schema import MODULE_ORDER, Cell, Tier, dashboard_module, loki_module_label
E2E_DIR = Path(__file__).resolve().parent.parent
@ -239,13 +239,31 @@ def render_prometheus(report: CoverageReport) -> str:
return "\n".join(lines)
def render_loki(report: CoverageReport) -> str:
lines = [
(
f"COVERAGE_TOTAL percent={report.coverage_percent:.1f} "
f"covered={report.covered} total={report.total}"
)
]
lines.extend(
(
f"COVERAGE_MODULE module={loki_module_label(module.module)} "
f"percent={module.coverage_percent:.1f} "
f"covered={module.covered} total={module.total}"
)
for module in report.modules
)
return "\n".join(lines)
def main() -> int:
parser = ArgumentParser()
parser.add_argument(
"--format",
choices=("text", "json", "prometheus"),
choices=("text", "json", "prometheus", "loki"),
default="text",
help="Output format. Use prometheus or json for Grafana ingestion jobs.",
help="Output format. Use loki for structured stdout lines in the e2e job.",
)
parser.add_argument(
"--strict",
@ -265,9 +283,8 @@ def main() -> int:
"text": render,
"json": render_json,
"prometheus": render_prometheus,
}[
args.format
](report)
"loki": render_loki,
}[args.format](report)
print(output) # noqa: T201 # CLI entrypoint output
if args.strict and report.orphan_markers:
return 1

View file

@ -157,6 +157,16 @@ MODULE_ORDER: tuple[str, ...] = (
"Other",
)
LOKI_MODULE_LABELS: dict[str, str] = {
"Core LLMs": "core_llms",
"Non-Core LLMs": "non_core_llms",
"MCPs": "mcp",
"Management/UI": "management_ui",
"Reliability & Performance": "reliability_performance",
"Logging & Guardrails": "logging_guardrails",
"Other": "other",
}
def dashboard_module(cell: Cell) -> str:
"""Return the Grafana/reporting module for a registry cell."""
@ -165,3 +175,8 @@ def dashboard_module(cell: Cell) -> str:
return "Core LLMs"
return "Non-Core LLMs"
return PREFIX_ROLLUP[cell.module]
def loki_module_label(module: str) -> str:
"""Return the log-safe Loki label for a dashboard module."""
return LOKI_MODULE_LABELS[module]

View file

@ -15,6 +15,7 @@ from coverage_registry.collector import (
compute_coverage,
render,
render_json,
render_loki,
render_prometheus,
)
from coverage_registry.registry import load_registry
@ -24,6 +25,7 @@ from coverage_registry.schema import (
LlmEndpoint,
LoggingCell,
Tier,
loki_module_label,
)
@ -149,6 +151,30 @@ def test_prometheus_render_exposes_module_coverage_timeseries() -> None:
assert "litellm_e2e_coverage_orphan_markers 0" in metrics
def test_loki_render_exposes_exact_stdout_lines_for_loki() -> None:
report = compute_coverage(
(_llm("llm.chat", Tier.P0), _llm("llm.batches", Tier.P0, "batches")),
frozenset({"llm.chat"}),
)
lines = render_loki(report).splitlines()
assert len(lines) == 1 + len(report.modules)
assert lines[0] == "COVERAGE_TOTAL percent=50.0 covered=1 total=2"
assert (
lines[1] == "COVERAGE_MODULE module=core_llms percent=100.0 covered=1 total=1"
)
assert (
lines[2] == "COVERAGE_MODULE module=non_core_llms percent=0.0 covered=0 total=1"
)
assert [line.split("module=", 1)[1].split(" ", 1)[0] for line in lines[1:]] == [
loki_module_label(module.module) for module in report.modules
]
assert all(
" " not in line.split("module=", 1)[1].split(" ", 1)[0] for line in lines[1:]
)
def test_real_registry_loads_and_ids_are_unique() -> None:
cells = load_registry()
ids = [c.id for c in cells]

View file

@ -5,6 +5,7 @@ Integration test for CyberArk Conjur Secret Manager.
import os
import sys
import pytest
import yaml
from dotenv import load_dotenv
load_dotenv()
@ -42,6 +43,82 @@ def create_mock_response(status_code: int, text: str = ""):
return mock_response
@pytest.mark.asyncio
async def test_cyberark_write_secret_rejects_yaml_injection():
"""
Regression test: async_write_secret must reject a secret_name that is not
safe to embed in the Conjur policy body, before any HTTP call is made.
"""
with patch("litellm.proxy.proxy_server.premium_user", True):
malicious_secret_name = "foo\n- !grant\n role: !!admin\n member: attacker"
mock_sync_client = MagicMock()
mock_async_client = AsyncMock()
with (
patch(
"litellm.secret_managers.cyberark_secret_manager._get_httpx_client",
return_value=mock_sync_client,
),
patch(
"litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client",
return_value=mock_async_client,
),
):
cyberark_manager = CyberArkSecretManager()
response = await cyberark_manager.async_write_secret(
secret_name=malicious_secret_name,
secret_value="sk-1234",
)
assert response["status"] == "error"
assert "Invalid secret_name" in response["message"]
# The malicious policy YAML must never reach the wire.
mock_sync_client.client.post.assert_not_called()
mock_async_client.post.assert_not_called()
@pytest.mark.parametrize(
"secret_name",
[
"foo: bar",
"foo # bar",
"plain-alias",
"team/user@example.com",
],
)
def test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters(secret_name):
"""
Regression test: _ensure_variable_exists must escape secret_name (not just
denylist-check it) so the policy body always parses back to exactly one
'!variable' scalar node holding the untouched secret_name.
"""
with patch("litellm.proxy.proxy_server.premium_user", True):
captured = {}
def _capture_post(url, headers=None, content=None):
captured["content"] = content
return create_mock_response(status_code=201, text="")
mock_sync_client = MagicMock()
mock_sync_client.client.post.side_effect = _capture_post
with patch(
"litellm.secret_managers.cyberark_secret_manager._get_httpx_client",
return_value=mock_sync_client,
):
cyberark_manager = CyberArkSecretManager()
cyberark_manager._ensure_variable_exists(secret_name)
policy_yaml = captured["content"]
parsed = yaml.compose(policy_yaml)
assert len(parsed.value) == 1
node = parsed.value[0]
assert node.tag == "!variable"
assert node.value == secret_name
@pytest.mark.asyncio
async def test_cyberark_write_and_read_secret():
"""

View file

@ -409,6 +409,33 @@ def test_hashicorp_custom_mount_and_prefix(hashicorp_secret_manager):
hashicorp_secret_manager.vault_namespace = original_namespace
@pytest.mark.parametrize(
"malicious_secret_name",
[
"../../../other-app/creds",
"litellm/../../secret",
"foo\nbar",
"foo
bar",
"foo
bar",
"foo\x85bar",
],
)
def test_hashicorp_get_url_rejects_path_traversal(monkeypatch, malicious_secret_name):
"""
Regression test: get_url must reject an invalid secret_name instead of
building a URL from it.
Uses monkeypatch + a directly-constructed manager (not the shared
hashicorp_secret_manager fixture) so this runs in CI without real Vault
credentials configured; get_url performs no I/O.
"""
monkeypatch.setenv("HCP_VAULT_TOKEN", "test-token-for-get-url-only")
manager = HashicorpSecretManager()
with pytest.raises(ValueError):
manager.get_url(malicious_secret_name)
mock_old_vault_response = {
"request_id": "80fafb6a-e96a-4c5b-29fa-ff505ac72201",
"lease_id": "",

View file

@ -746,7 +746,8 @@ class BaseResponsesAPITest(ABC):
E2E test for Shell tool on OpenAI Responses API.
Passes tools=[{"type": "shell", "environment": {"type": "container_auto"}}];
validates that the request is accepted and returns a valid response.
Only runs for OpenAI/Azure (Responses API with shell support).
Only runs for OpenAI; offline coverage for the Azure route lives in
tests/test_litellm/responses/test_responses_api_request_body.py.
"""
base_completion_call_args = self.get_base_completion_call_args()
model = (
@ -754,8 +755,10 @@ class BaseResponsesAPITest(ABC):
or base_completion_call_args.get("model")
or ""
)
if "openai/" not in str(model) and "azure/" not in str(model):
pytest.skip("Shell tool e2e is only run for OpenAI/Azure Responses API")
if "openai/" not in str(model):
pytest.skip(
"Shell tool e2e is OpenAI-only; no Azure deployment supports the shell tool yet, re-enable once one exists"
)
tools = [{"type": "shell", "environment": {"type": "container_auto"}}]
input_msg = "List files in /mnt/data and show python --version."
try:
@ -765,7 +768,10 @@ class BaseResponsesAPITest(ABC):
max_output_tokens=256,
tools=tools,
tool_choice="auto",
timeout=90,
)
except litellm.Timeout:
pytest.skip("Provider did not answer the shell tool request within 90s")
except litellm.InternalServerError:
pytest.skip("Skipping test due to litellm.InternalServerError")
except litellm.BadRequestError as e:

View file

@ -2,7 +2,6 @@ import os
import sys
import pytest
import asyncio
from typing import Optional
from unittest.mock import patch, AsyncMock
sys.path.insert(0, os.path.abspath("../.."))
@ -30,10 +29,6 @@ class TestAzureResponsesAPITest(BaseResponsesAPITest):
"api_version": "2025-03-01-preview",
}
def get_advanced_model_for_shell_tool(self) -> Optional[str]:
"""If specified, overrides the model used by test_responses_api_shell_tool_streaming_sees_shell_output (e.g. openai/gpt-5.2 for shell support)."""
return "azure/gpt-5-mini"
@pytest.mark.asyncio
async def test_azure_responses_api_preview_api_version():

View file

@ -41,11 +41,13 @@ def fake_openai_endpoint():
# Per-item respx detection (``apply_vcr_auto_marker_to_items``) handles
# the vast majority of respx-vs-vcrpy conflicts automatically. The only
# entry below is the persister's own unit-test file, which exercises
# ``save_cassette`` / ``load_cassette`` against fakeredis and must not
# itself run under a live cassette context.
_VCR_AUTO_MARKER_SKIP_FILES = frozenset({"test_vcr_redis_persister.py"})
# the vast majority of respx-vs-vcrpy conflicts automatically. The entries
# below are the persister's and the WebSocket VCR's own unit-test files, which
# exercise ``save_cassette`` / ``load_cassette`` against fakeredis and must not
# themselves run under a live cassette context.
_VCR_AUTO_MARKER_SKIP_FILES = frozenset(
{"test_vcr_redis_persister.py", "test_ws_vcr.py"}
)
_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = ()

View file

@ -0,0 +1,89 @@
"""WebSocket VCR wiring for the realtime suite.
This directory inherits the HTTP VCR machinery from
``tests/llm_translation/conftest.py`` (which only intercepts httpx/aiohttp and
is therefore a no-op for realtime WebSocket traffic). The autouse fixture below
adds the WebSocket layer: it patches ``websockets.connect`` for the duration of
each test so realtime frames are recorded to, or replayed from, the same
cassette Redis under a ``litellm:vcr:wscassette:`` prefix.
"""
from __future__ import annotations
import os
import sys
from typing import Optional
import pytest
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "..")))
from tests._vcr_conftest_common import ( # noqa: E402
vcr_disabled,
vcr_outcome_logging_enabled,
)
from tests._ws_vcr import ( # noqa: E402
WsVcrController,
build_ws_cassette_client,
load_ws_cassette,
replay_timeout_seconds,
save_ws_cassette,
ws_redis_key_for,
)
_ws_cassette_client: Optional[object] = None
def _get_ws_cassette_client() -> Optional[object]:
global _ws_cassette_client
if _ws_cassette_client is None:
_ws_cassette_client = build_ws_cassette_client()
return _ws_cassette_client
def _emit_verdict(request: pytest.FixtureRequest, verdict: str) -> None:
if os.environ.get("PYTEST_XDIST_WORKER"):
return
reporter = request.config.pluginmanager.getplugin("terminalreporter")
if reporter is None:
return
reporter.write_line(f"{verdict} :: {request.node.nodeid}")
@pytest.fixture(autouse=True)
def _ws_vcr(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch):
if vcr_disabled():
yield
return
import websockets
client = _get_ws_cassette_client()
if client is None:
yield
return
key = ws_redis_key_for(request.node.nodeid)
cassette = load_ws_cassette(client, key)
controller = WsVcrController(
original_connect=websockets.connect,
cassette=cassette,
timeout=replay_timeout_seconds(),
)
monkeypatch.setattr(websockets, "connect", controller.connect)
yield
rep_call = getattr(request.node, "rep_call", None)
passed = bool(rep_call and rep_call.passed)
if controller.recorded:
built = controller.built_cassette()
if built is not None:
save_ws_cassette(client, key, built, passed=passed)
if vcr_outcome_logging_enabled():
_emit_verdict(request, controller.verdict())
if controller.errors and passed:
raise controller.errors[0]

View file

@ -0,0 +1,275 @@
from __future__ import annotations
import asyncio
import os
import sys
import warnings
import fakeredis
import pytest
from websockets.exceptions import ConnectionClosedOK
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
from tests._vcr_redis_persister import ( # noqa: E402
VCRCassetteCacheWarning,
cassette_cache_health,
)
from tests._ws_vcr import ( # noqa: E402
CASSETTE_TTL_SECONDS,
RedisLike,
ReplayConnection,
WsCassette,
WsFrame,
WsSession,
WsSessionRecorder,
WsVcrContractDrift,
WsVcrReplayError,
WsVcrReplayTimeout,
build_ws_cassette_client,
load_ws_cassette,
save_ws_cassette,
scrub_secrets,
text_frames_match,
ws_redis_key_for,
)
def _server(text: str, client_frames_before: int) -> WsFrame:
return WsFrame(
direction="server_to_client",
opcode="text",
text=text,
client_frames_before=client_frames_before,
)
def _client(text: str) -> WsFrame:
return WsFrame(direction="client_to_server", opcode="text", text=text)
def _collect_errors():
errors: list[WsVcrReplayError] = []
return errors, errors.append
def test_cassette_json_roundtrip_preserves_frames_and_gate():
cassette = WsCassette(
sessions=(
WsSession(
frames=(
_server('{"type":"session.created"}', 0),
_client('{"type":"response.create"}'),
_server('{"type":"response.done"}', 1),
WsFrame(
direction="server_to_client", opcode="binary", binary_b64="dGVzdA==", client_frames_before=1
),
)
),
)
)
restored = WsCassette.model_validate_json(cassette.model_dump_json())
assert restored == cassette
assert restored.sessions[0].frames[2].client_frames_before == 1
assert restored.sessions[0].frames[3].opcode == "binary"
assert restored.sessions[0].frames[3].binary_b64 == "dGVzdA=="
def test_recorder_tracks_client_frame_count_as_causal_gate():
recorder = WsSessionRecorder()
recorder.record_server_frame('{"type":"session.created"}')
recorder.record_client_frame('{"type":"conversation.item.create"}')
recorder.record_client_frame('{"type":"response.create"}')
recorder.record_server_frame('{"type":"response.done"}')
session = recorder.to_session()
server_frames = [f for f in session.frames if f.direction == "server_to_client"]
assert server_frames[0].client_frames_before == 0
assert server_frames[1].client_frames_before == 2
async def test_replay_recv_returns_bytes_when_decode_false_and_str_otherwise():
session = WsSession(frames=(_server("hello", 0), _server("world", 0)))
_, on_error = _collect_errors()
conn = ReplayConnection(session, timeout=1.0, on_error=on_error)
as_bytes = await conn.recv(decode=False)
as_str = await conn.recv()
assert as_bytes == b"hello"
assert as_str == "world"
async def test_replay_serves_server_frame_only_after_causal_client_count_met():
session = WsSession(
frames=(
_server('{"type":"session.created"}', 0),
_client('{"type":"response.create"}'),
_server('{"type":"response.done"}', 1),
)
)
_, on_error = _collect_errors()
conn = ReplayConnection(session, timeout=2.0, on_error=on_error)
first = await conn.recv(decode=False)
assert first == b'{"type":"session.created"}'
gated = asyncio.ensure_future(conn.recv(decode=False))
await asyncio.sleep(0.1)
assert not gated.done(), "gated server frame was released before the recorded client frame was sent"
await conn.send('{"type":"response.create"}')
released = await asyncio.wait_for(gated, timeout=1.0)
assert released == b'{"type":"response.done"}'
async def test_replay_exhausted_server_frames_raise_connection_closed():
session = WsSession(frames=(_server("only", 0),))
_, on_error = _collect_errors()
conn = ReplayConnection(session, timeout=1.0, on_error=on_error)
await conn.recv(decode=False)
with pytest.raises(ConnectionClosedOK):
await conn.recv(decode=False)
async def test_replay_timeout_raises_instead_of_hanging():
session = WsSession(
frames=(
_server('{"type":"session.created"}', 0),
_server('{"type":"response.done"}', 5),
)
)
errors, on_error = _collect_errors()
conn = ReplayConnection(session, timeout=0.15, on_error=on_error)
await conn.recv(decode=False)
with pytest.raises(WsVcrReplayTimeout):
await asyncio.wait_for(conn.recv(decode=False), timeout=2.0)
assert errors and isinstance(errors[0], WsVcrReplayTimeout)
async def test_replay_accepts_client_frame_with_volatile_id_drift():
recorded_client = _client('{"type":"conversation.item.create","item":{"id":"item_ABC12345","role":"user"}}')
session = WsSession(frames=(_server("s", 0), recorded_client, _server("done", 1)))
errors, on_error = _collect_errors()
conn = ReplayConnection(session, timeout=1.0, on_error=on_error)
await conn.recv(decode=False)
await conn.send('{"type":"conversation.item.create","item":{"role":"user","id":"item_ZZ99887766"}}')
assert errors == []
assert await conn.recv(decode=False) == b"done"
async def test_replay_rejects_structurally_different_client_frame():
session = WsSession(frames=(_server("s", 0), _client('{"type":"response.create"}'), _server("done", 1)))
errors, on_error = _collect_errors()
conn = ReplayConnection(session, timeout=1.0, on_error=on_error)
await conn.recv(decode=False)
with pytest.raises(WsVcrContractDrift):
await conn.send('{"type":"session.update","session":{"voice":"alloy"}}')
assert errors and isinstance(errors[0], WsVcrContractDrift)
async def test_replay_rejects_extra_client_frame_beyond_recording():
session = WsSession(frames=(_server("s", 0), _client('{"type":"response.create"}')))
errors, on_error = _collect_errors()
conn = ReplayConnection(session, timeout=1.0, on_error=on_error)
await conn.send('{"type":"response.create"}')
with pytest.raises(WsVcrContractDrift):
await conn.send('{"type":"response.create"}')
assert errors
def test_text_frames_match_normalizes_ids_and_timestamps_but_not_structure():
assert text_frames_match(
'{"type":"x","event_id":"evt_111","ts":"2026-05-25T03:40:37.262045Z"}',
'{"type":"x","event_id":"evt_999","ts":"2026-06-01T10:00:00Z"}',
)
assert not text_frames_match('{"type":"x","text":"hi"}', '{"type":"x","text":"bye"}')
assert not text_frames_match('{"type":"x"}', '{"type":"x","extra":1}')
def test_scrub_secrets_removes_auth_material():
scrubbed = scrub_secrets("Authorization: Bearer sk-abcdef123456 and key xai-zzz99988877 raw sk-plainkey123")
assert "sk-abcdef123456" not in scrubbed
assert "xai-zzz99988877" not in scrubbed
assert "sk-plainkey123" not in scrubbed
assert "Bearer <redacted>" in scrubbed
def test_recorder_scrubs_secrets_in_stored_frames():
recorder = WsSessionRecorder()
recorder.record_client_frame('{"authorization":"Bearer sk-supersecretvalue"}')
stored = recorder.to_session().frames[0].text
assert stored is not None
assert "sk-supersecretvalue" not in stored
def _sample_cassette() -> WsCassette:
return WsCassette(sessions=(WsSession(frames=(_server('{"type":"session.created"}', 0),)),))
def test_save_sets_24h_ttl_and_load_roundtrips():
fake = fakeredis.FakeStrictRedis()
key = ws_redis_key_for("tests/llm_translation/realtime/test_x.py::test_y")
assert save_ws_cassette(fake, key, _sample_cassette(), passed=True) is True
ttl = fake.ttl(key)
assert CASSETTE_TTL_SECONDS - 5 <= ttl <= CASSETTE_TTL_SECONDS
loaded = load_ws_cassette(fake, key)
assert loaded == _sample_cassette()
def test_save_skipped_when_test_failed_leaves_no_key():
fake = fakeredis.FakeStrictRedis()
key = ws_redis_key_for("tests/llm_translation/realtime/test_x.py::test_fail")
assert save_ws_cassette(fake, key, _sample_cassette(), passed=False) is False
assert fake.get(key) is None
def test_save_skipped_when_test_failed_preserves_prior_cassette():
fake = fakeredis.FakeStrictRedis()
key = ws_redis_key_for("tests/llm_translation/realtime/test_x.py::test_keep")
save_ws_cassette(fake, key, _sample_cassette(), passed=True)
newer = WsCassette(sessions=(WsSession(frames=(_server('{"type":"other"}', 0),)),))
assert save_ws_cassette(fake, key, newer, passed=False) is False
assert load_ws_cassette(fake, key) == _sample_cassette()
def test_load_missing_key_returns_none():
fake = fakeredis.FakeStrictRedis()
assert load_ws_cassette(fake, ws_redis_key_for("never/recorded")) is None
def test_ws_redis_key_uses_distinct_prefix():
key = ws_redis_key_for("tests/llm_translation/realtime/test_x.py::TestY::test_z")
assert key.startswith("litellm:vcr:wscassette:")
assert "::" not in key
def test_build_ws_cassette_client_warns_and_counts_failure_instead_of_silently_disabling():
def _broken_builder() -> RedisLike:
raise ValueError("invalid CASSETTE_REDIS_URL")
failures_before = cassette_cache_health()["load_failures"]
with pytest.warns(VCRCassetteCacheWarning, match="fall back to live websocket traffic"):
assert build_ws_cassette_client(builder=_broken_builder) is None
assert cassette_cache_health()["load_failures"] == failures_before + 1
def test_build_ws_cassette_client_returns_built_client_without_warning():
fake = fakeredis.FakeStrictRedis()
with warnings.catch_warnings():
warnings.simplefilter("error", VCRCassetteCacheWarning)
assert build_ws_cassette_client(builder=lambda: fake) is fake

View file

@ -355,7 +355,11 @@ async def test_async_vertexai_response_basic():
user_message = "Hello, how are you?"
messages = [{"content": user_message, "role": "user"}]
response = await acompletion(
model="gemini-2.5-flash", messages=messages, temperature=0.7, timeout=5
model="gemini-3.5-flash",
messages=messages,
temperature=0.7,
timeout=5,
vertex_location="global",
)
print(f"response: {response}")
except litellm.NotFoundError as e:
@ -388,7 +392,7 @@ async def test_async_vertexai_streaming_response():
)
test_models = random.sample(list(test_models), 1)
test_models += list(litellm.vertex_language_models) # always test gemini-pro
test_models = ["gemini-2.5-flash"]
test_models = ["gemini-3.5-flash"]
for model in test_models:
if model in VERTEX_MODELS_TO_NOT_TEST or (
"gecko" in model
@ -412,6 +416,7 @@ async def test_async_vertexai_streaming_response():
temperature=0.7,
timeout=5,
stream=True,
vertex_location="global",
)
print(f"response: {response}")
complete_response: str = ""
@ -3840,10 +3845,11 @@ def test_vertex_schema_test():
}
response = litellm.completion(
model="vertex_ai/gemini-2.5-flash",
model="vertex_ai/gemini-3.5-flash",
messages=[{"role": "user", "content": "call the tool"}],
tools=[tool],
tool_choice="required",
vertex_location="global",
)
print(response)
@ -3895,10 +3901,11 @@ def test_gemini_nullable_object_tool_schema_httpx():
]
response = litellm.completion(
model="vertex_ai/gemini-2.5-flash",
model="vertex_ai/gemini-3.5-flash",
messages=[{"role": "user", "content": "call the tool"}],
tools=tools,
tool_choice="required",
vertex_location="global",
)
print(response)

View file

@ -22,10 +22,8 @@ from litellm.proxy.proxy_server import initialize_pass_through_endpoints
# Mock the async_client used in the pass_through_request function
async def mock_request(*args, **kwargs):
mock_response = httpx.Response(200, json={"message": "Mocked response"})
mock_response.request = Mock(spec=httpx.Request)
return mock_response
async def mock_request(self, request, **kwargs):
return httpx.Response(200, json={"message": "Mocked response"}, request=request)
def remove_rerank_route(app):
@ -49,8 +47,8 @@ def client():
@pytest.mark.asyncio
async def test_pass_through_endpoint_no_headers(client, monkeypatch):
# Mock the httpx.AsyncClient.request method
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
# Mock the httpx.AsyncClient.send method
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
# Define a pass-through endpoint
@ -79,8 +77,8 @@ async def test_pass_through_endpoint_no_headers(client, monkeypatch):
@pytest.mark.asyncio
async def test_pass_through_endpoint(client, monkeypatch):
# Mock the httpx.AsyncClient.request method
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
# Mock the httpx.AsyncClient.send method
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
# Define a pass-through endpoint
@ -181,7 +179,7 @@ async def test_pass_through_endpoint_rpm_limit(
expected_status_codes,
num_users,
):
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache
@ -285,7 +283,7 @@ async def test_pass_through_endpoint_rpm_limit(
async def test_pass_through_endpoint_sequential_rpm_limit(
client, monkeypatch, auth, rpm_limit, requests_to_make, expected_status_codes
):
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache
@ -504,10 +502,10 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
captured_requests = []
async def mock_bing_request(*args, **kwargs):
async def mock_bing_request(self, request, **kwargs):
captured_requests.append((args, kwargs))
mock_response = httpx.Response(
captured_requests.append(request)
return httpx.Response(
200,
json={
"_type": "SearchResponse",
@ -518,11 +516,10 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
"value": [],
},
},
request=request,
)
mock_response.request = Mock(spec=httpx.Request)
return mock_response
monkeypatch.setattr("httpx.AsyncClient.request", mock_bing_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_bing_request)
# Define a pass-through endpoint
pass_through_endpoints = [
@ -555,8 +552,8 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
client.get("/bing/search?q=bob+barker")
client.get("/bing/search-no-merge-params?q=bob+barker")
first_transformed_url = captured_requests[0][1]["url"]
second_transformed_url = captured_requests[1][1]["url"]
first_transformed_url = captured_requests[0].url
second_transformed_url = captured_requests[1].url
# Parse URLs to compare query params order-independently
# Parse first URL
@ -573,7 +570,7 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
"setLang": ["en-US"],
"mkt": ["en-US"],
}
expected_second_params = {"setLang": ["en-US"], "mkt": ["en-US"]}
expected_second_params = {"q": ["bob barker"]}
# Assert the response - compare base URL and params separately
assert (

View file

@ -0,0 +1,14 @@
{
"model": "gpt-5-mini",
"input": "List files in /mnt/data and run python --version.",
"tools": [
{
"type": "shell",
"environment": {
"type": "container_auto"
}
}
],
"tool_choice": "auto",
"max_output_tokens": 256
}

View file

@ -579,13 +579,10 @@ def _exception_event(span):
def test_error_message_recorded_as_full_exception_event_untruncated():
"""Regression for the Elasticsearch keyword/ignore_above:1024 truncation.
A long error message must survive intact on the standard ``exception``
event under ``exception.message`` — not get dropped onto a bare string
attribute that backends dynamic-map to a 1024-char ``keyword``. The SDK
must not truncate it either, so a 5000-char message stays 5000 chars.
"""
"""The ``exception`` event carries the full untruncated message under
``exception.message`` so backends that dynamic-map unknown string span
attrs to ``keyword`` (e.g. Elasticsearch with a 1024-char ``ignore_above``)
still see it in full via the semconv-recognized event field."""
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent
long_message = "boom: " + "x" * 5000
@ -596,13 +593,108 @@ def test_error_message_recorded_as_full_exception_event_untruncated():
assert len(event.attributes[ExceptionEvent.MESSAGE]) == len(long_message) > 1024
assert event.attributes[ExceptionEvent.TYPE] == "litellm.APIError"
# error.type stays a low-cardinality attribute; the message does NOT become a
# bare string attribute (which is what got truncated).
# error.type stays a low-cardinality attribute; the exception EVENT field
# ``exception.message`` never becomes a bare string attribute.
assert span.attributes[Error.TYPE] == "litellm.APIError"
assert ExceptionEvent.MESSAGE not in span.attributes
assert span.status.description == long_message
def test_error_details_stamped_as_span_attributes_for_labels_ingest():
"""OTel-defined keys and litellm-specific detail keys both ride span
attributes so backends that flatten attrs into label indexes (Elastic APM
``labels.*``, Datadog span tags) render them. The exception event with the
full untruncated message stays alongside — both places, matching v1's
shape."""
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent, LiteLLMError
from litellm.integrations.otel.emitter import SpanEmitter
cfg = OpenTelemetryV2Config(exporter="in_memory")
provider, exporter = providers.in_memory_provider(cfg)
engine = SpanEmitter(providers.get_tracer(provider, "t"), cfg)
data = LLMCallSpanData(
operation=GenAIOperation.CHAT,
provider="openai",
request_model="gpt-4o",
response_model=None,
response_id=None,
request_params=LLMRequestParams(),
usage=LLMUsage(),
finish_reasons=(),
error=SpanError(
error_type="litellm.BadRequestError",
message="400: violated moderation policy",
code="400",
stack_trace="File proxy_server.py line 8570 ...",
llm_provider="openai",
),
response_cost=None,
server=None,
identity=RequestIdentity(call_id=None),
)
engine.emit(SpanRole.LLM_CALL, data)
(span,) = exporter.get_finished_spans()
# OTel-defined keys (from the ``error.*`` semconv registry).
assert span.attributes[Error.TYPE] == "litellm.BadRequestError"
assert span.attributes[Error.MESSAGE] == "400: violated moderation policy"
# LiteLLM-specific detail keys — vendor-namespaced under ``error.*``
# for v1-parity, not defined by OTel semconv.
assert span.attributes[LiteLLMError.CODE] == "400"
assert span.attributes[LiteLLMError.STACK_TRACE] == "File proxy_server.py line 8570 ..."
assert span.attributes[LiteLLMError.LLM_PROVIDER] == "openai"
# The exception event carries the same message on the span too.
event = _exception_event(span)
assert event.attributes[ExceptionEvent.MESSAGE] == "400: violated moderation policy"
def test_error_details_omitted_when_span_error_carries_only_message():
"""A guardrail-shape error (message only, no code/traceback/provider) must
not pollute the span with empty-string detail attributes. Only the keys
that carry real data land."""
from litellm.integrations.otel.model.semconv import Error, LiteLLMError
span = _emit_error_span("guardrail rejected", error_type="ContentFilter")
assert span.attributes[Error.TYPE] == "ContentFilter"
assert span.attributes[Error.MESSAGE] == "guardrail rejected"
# LiteLLM-specific detail keys aren't stamped when the SpanError doesn't
# carry them.
assert LiteLLMError.CODE not in span.attributes
assert LiteLLMError.STACK_TRACE not in span.attributes
assert LiteLLMError.LLM_PROVIDER not in span.attributes
def test_v2_error_attribute_keys_match_v1_error_attributes_byte_for_byte():
"""v1 (``opentelemetry.py``) and v2 (``otel/`` package) stamp identical
span-attribute keys so consumers reading ``labels.error_message`` don't
care which integration produced the span. Renaming either side is a
breaking change for downstream dashboards; this test locks the vocabulary."""
from litellm.integrations._types.open_inference import ErrorAttributes
from litellm.integrations.otel.model.semconv import Error, LiteLLMError
assert Error.TYPE == ErrorAttributes.ERROR_TYPE
assert Error.MESSAGE == ErrorAttributes.ERROR_MESSAGE
assert LiteLLMError.CODE == ErrorAttributes.ERROR_CODE
assert LiteLLMError.STACK_TRACE == ErrorAttributes.ERROR_STACK_TRACE
assert LiteLLMError.LLM_PROVIDER == ErrorAttributes.ERROR_LLM_PROVIDER
def test_error_message_falls_back_to_error_type_when_message_absent():
"""A ``SpanError(error_type=..., message=None)`` still renders on the span:
the resolved message is the error_type, and it lands on ``error.message``,
the exception event, and the span-status description in lockstep so a
single-source-of-truth view isn't inconsistent."""
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent
span = _emit_error_span(message=None, error_type="RateLimitError")
assert span.attributes[Error.MESSAGE] == "RateLimitError"
assert _exception_event(span).attributes[ExceptionEvent.MESSAGE] == "RateLimitError"
assert span.status.description == "RateLimitError"
def test_success_span_records_no_exception_event():
from litellm.integrations.otel.emitter import SpanEmitter
from litellm.integrations.otel.model.semconv import ExceptionEvent

View file

@ -144,11 +144,13 @@ def _all_constants(cls):
def test_attribute_keys_are_unique_across_namespaces():
from litellm.integrations.otel import MCP, Client, JsonRpc, Network
from litellm.integrations.otel import MCP, Client, JsonRpc, LiteLLMError, Network
# prefixes are allowed to be substrings; exact keys must not collide.
# ``LiteLLMError`` shares the ``error.*`` prefix with ``Error`` by design
# (v1-parity); the assert below is the guarantee they never overlap.
exact = set()
for cls in (GenAI, Error, Server, HTTP, DB, MCP, JsonRpc, Network, Client):
for cls in (GenAI, Error, LiteLLMError, Server, HTTP, DB, MCP, JsonRpc, Network, Client):
for key in _all_constants(cls):
assert key not in exact, f"duplicate attribute key {key}"
exact.add(key)
@ -342,6 +344,47 @@ def test_llm_call_adapter_failure_path():
assert data.error.message == "429 slow down"
def test_llm_call_adapter_carries_error_detail_fields():
"""``_parse_error`` threads the full detail set from ``error_information``
(``error_code``, ``traceback``, ``llm_provider``) onto ``SpanError`` so the
emitter can stamp them as span attributes."""
payload = _sample_payload(
status="failure",
error_information={
"error_class": "BadRequestError",
"error_message": "400 violated moderation policy",
"error_code": "400",
"traceback": "File proxy_server.py line 8570 ...",
"llm_provider": "openai",
},
)
data = LLMCallSpanData.from_standard_logging_payload(payload)
assert data.error is not None
assert data.error.error_type == "BadRequestError"
assert data.error.message == "400 violated moderation policy"
assert data.error.code == "400"
assert data.error.stack_trace == "File proxy_server.py line 8570 ..."
assert data.error.llm_provider == "openai"
def test_llm_call_adapter_error_details_default_to_none_when_absent():
"""Guardrail-shape payloads carry only ``error_class`` + ``error_message``.
The detail fields must stay ``None`` so the emitter's ``if error.code:``
guards skip stamping empty attributes."""
payload = _sample_payload(
status="failure",
error_information={
"error_class": "ContentFilter",
"error_message": "guardrail rejected",
},
)
data = LLMCallSpanData.from_standard_logging_payload(payload)
assert data.error is not None
assert data.error.code is None
assert data.error.stack_trace is None
assert data.error.llm_provider is None
def test_adapter_is_resilient_to_minimal_payload():
data = LLMCallSpanData.from_standard_logging_payload({})
assert data.request_model == ""

View file

@ -0,0 +1,93 @@
"""
Unit tests for PrometheusLogger._assemble_key_object DB access.
The post-request budget metrics run for every LLM API request. Auth has
already cached the key object for any real key in the same request, so the
metrics path must read the cache only. Falling through to the DB turns every
request whose token has no DB row (e.g. master-key requests, whose token is
an alias hash that never matches a stored key) into per-request
LiteLLM_VerificationToken and LiteLLM_DeprecatedVerificationToken queries.
"""
import datetime
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from prometheus_client import REGISTRY
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.integrations.prometheus import PrometheusLogger
from litellm.proxy._types import UserAPIKeyAuth
@pytest.fixture(autouse=True)
def cleanup_prometheus_registry():
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
try:
REGISTRY.unregister(collector)
except Exception:
pass
yield
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
try:
REGISTRY.unregister(collector)
except Exception:
pass
@pytest.fixture
def prometheus_logger():
return PrometheusLogger()
@pytest.mark.asyncio
async def test_assemble_key_object_does_not_query_db_on_cache_miss(prometheus_logger):
mock_prisma = MagicMock()
mock_prisma.get_data = AsyncMock()
cache = DualCache(in_memory_cache=InMemoryCache())
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.user_api_key_cache", cache),
):
result = await prometheus_logger._assemble_key_object(
user_api_key="hashed-token-not-in-cache",
user_api_key_alias="",
key_max_budget=None,
key_spend=1.0,
response_cost=0.5,
)
mock_prisma.get_data.assert_not_called()
assert result.spend == 1.5
assert result.budget_reset_at is None
@pytest.mark.asyncio
async def test_assemble_key_object_reads_budget_reset_at_from_cache(prometheus_logger):
hashed_token = "hashed-token-in-cache"
reset_at = datetime.datetime(2026, 8, 1, tzinfo=datetime.timezone.utc)
cached_key = UserAPIKeyAuth(token=hashed_token, budget_reset_at=reset_at)
mock_prisma = MagicMock()
mock_prisma.get_data = AsyncMock()
cache = DualCache(in_memory_cache=InMemoryCache())
await cache.async_set_cache(key=hashed_token, value=cached_key)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.user_api_key_cache", cache),
):
result = await prometheus_logger._assemble_key_object(
user_api_key=hashed_token,
user_api_key_alias="alias",
key_max_budget=10.0,
key_spend=1.0,
response_cost=0.5,
)
mock_prisma.get_data.assert_not_called()
assert result.budget_reset_at == reset_at

View file

@ -3085,3 +3085,84 @@ def test_bedrock_converse_messages_pt_document_rejects_url_source():
_bedrock_converse_messages_pt(
messages, "anthropic.claude-sonnet-4-6", "bedrock"
)
def _collect_cache_points(blocks):
return [
block["cachePoint"]
for message in blocks
for block in message["content"]
if "cachePoint" in block
]
@pytest.mark.parametrize(
"messages",
[
[
{
"role": "user",
"content": [
{
"type": "text",
"text": "conversation history",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
],
},
],
[
{"role": "user", "content": "hello"},
{
"role": "assistant",
"content": [
{
"type": "text",
"text": "assistant reply",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
],
},
],
],
)
def test_bedrock_converse_message_level_cache_point_preserves_ttl(messages):
"""
Regression for https://github.com/BerriAI/litellm/issues/32154: message-level
cache_control ttl was silently dropped because the message-level
_get_cache_point_block call sites never passed model=, so multi-turn prefixes
fell back to the 5m default while the system prompt kept 1h, churning the
cache every turn on models like Opus 4.8.
"""
result = _bedrock_converse_messages_pt(
messages=messages,
model="eu.anthropic.claude-opus-4-8",
llm_provider="bedrock",
)
cache_points = _collect_cache_points(result)
assert cache_points == [{"type": "default", "ttl": "1h"}]
@pytest.mark.asyncio
async def test_bedrock_converse_message_level_cache_point_preserves_ttl_async():
messages = [
{
"role": "user",
"content": [
{
"type": "text",
"text": "conversation history",
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
],
},
]
result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
messages=messages,
model="eu.anthropic.claude-opus-4-8",
llm_provider="bedrock",
)
assert _collect_cache_points(result) == [{"type": "default", "ttl": "1h"}]

View file

@ -14,6 +14,7 @@ from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers
from litellm.main import ahealth_check
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS
def test_update_model_params_with_health_check_tracking_information():
@ -140,3 +141,139 @@ async def test_ahealth_check_failure_masks_raw_request_headers():
assert headers["Content-Type"] == "application/json"
print(f"Masked Authorization header: {headers.get('Authorization', 'NOT FOUND')}")
@pytest.mark.asyncio
async def test_batch_health_check_bridges_metadata_into_logging_obj():
"""_batch_health_check must call update_from_kwargs on the pre-injected
logging object so callbacks receive identity/tracking fields in
model_call_details["litellm_params"]["metadata"]."""
mock_logging_obj = MagicMock()
mock_logging_obj.update_from_kwargs = MagicMock()
litellm_metadata = {
"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME],
"user_api_key_alias": "health-check-key",
}
filtered_model_params = {
"model": "openai/gpt-4",
"api_base": "https://api.openai.com",
"litellm_logging_obj": mock_logging_obj,
"litellm_metadata": litellm_metadata,
}
with patch("litellm.alist_batches", new_callable=AsyncMock, return_value={}):
await HealthCheckHelpers._batch_health_check(
custom_llm_provider="openai",
model_params={"model": "openai/gpt-4"},
filtered_model_params=filtered_model_params,
)
mock_logging_obj.update_from_kwargs.assert_called_once()
call_kwargs = mock_logging_obj.update_from_kwargs.call_args[1]
assert call_kwargs["model"] == "openai/gpt-4"
assert call_kwargs["kwargs"] is filtered_model_params
assert call_kwargs["litellm_params"] == {"api_base": "https://api.openai.com"}
@pytest.mark.asyncio
async def test_batch_health_check_omits_api_base_when_absent():
"""api_base must not appear in litellm_params when the provider resolves
it implicitly (bedrock, vertex, gemini)."""
mock_logging_obj = MagicMock()
mock_logging_obj.update_from_kwargs = MagicMock()
litellm_metadata = {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]}
filtered_model_params = {
"model": "bedrock/anthropic.claude-v2",
"litellm_logging_obj": mock_logging_obj,
"litellm_metadata": litellm_metadata,
}
with patch("litellm.acompletion", new_callable=AsyncMock, return_value={}):
await HealthCheckHelpers._batch_health_check(
custom_llm_provider="bedrock",
model_params={"model": "bedrock/anthropic.claude-v2"},
filtered_model_params=filtered_model_params,
)
call_kwargs = mock_logging_obj.update_from_kwargs.call_args[1]
assert call_kwargs["litellm_params"] is None
@pytest.mark.asyncio
async def test_batch_health_check_skips_bridge_when_no_logging_obj():
"""When litellm_logging_obj is absent, dispatch still proceeds."""
litellm_metadata = {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]}
filtered_model_params = {
"model": "openai/gpt-4",
"litellm_metadata": litellm_metadata,
}
with patch(
"litellm.alist_batches", new_callable=AsyncMock, return_value={}
) as mock_alist:
await HealthCheckHelpers._batch_health_check(
custom_llm_provider="openai",
model_params={"model": "openai/gpt-4"},
filtered_model_params=filtered_model_params,
)
mock_alist.assert_called_once()
@pytest.mark.asyncio
async def test_batch_health_check_uses_alist_batches_for_supported_providers():
"""Providers in LIST_BATCHES_SUPPORTED_PROVIDERS dispatch to alist_batches."""
mock_logging_obj = MagicMock()
mock_logging_obj.update_from_kwargs = MagicMock()
litellm_metadata = {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]}
for provider in LIST_BATCHES_SUPPORTED_PROVIDERS:
filtered_model_params = {
"model": f"{provider}/some-model",
"litellm_logging_obj": mock_logging_obj,
"litellm_metadata": litellm_metadata,
}
with patch(
"litellm.alist_batches", new_callable=AsyncMock, return_value={}
) as mock_alist:
await HealthCheckHelpers._batch_health_check(
custom_llm_provider=provider,
model_params={"model": f"{provider}/some-model"},
filtered_model_params=filtered_model_params,
)
mock_alist.assert_called_once()
@pytest.mark.asyncio
async def test_batch_health_check_falls_back_to_acompletion_for_unsupported():
"""Providers not in LIST_BATCHES_SUPPORTED_PROVIDERS fall back to acompletion."""
mock_logging_obj = MagicMock()
mock_logging_obj.update_from_kwargs = MagicMock()
litellm_metadata = {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]}
filtered_model_params = {
"model": "bedrock/anthropic.claude-v2",
"litellm_logging_obj": mock_logging_obj,
"litellm_metadata": litellm_metadata,
}
model_params = {"model": "bedrock/anthropic.claude-v2", "messages": []}
with (
patch("litellm.alist_batches", new_callable=AsyncMock) as mock_alist,
patch("litellm.acompletion", new_callable=AsyncMock, return_value={}) as mock_acompletion,
):
await HealthCheckHelpers._batch_health_check(
custom_llm_provider="bedrock",
model_params=model_params,
filtered_model_params=filtered_model_params,
)
mock_alist.assert_not_called()
mock_acompletion.assert_called_once_with(**model_params)

View file

@ -387,8 +387,9 @@ async def test_pass_through_request_stream_param_no_override(
# Create mocks for the async client
mock_async_client = AsyncMock()
# Mock request to return the non-streaming response
mock_async_client.request.return_value = mock_response
# Mock build_request/send to return the non-streaming response
mock_async_client.build_request = Mock(return_value=Mock())
mock_async_client.send.return_value = mock_response
# Mock get_async_httpx_client to return our mock client
mock_client_obj = Mock()
@ -420,20 +421,19 @@ async def test_pass_through_request_stream_param_no_override(
stream=False, # Should be used since no stream in request body
)
# Verify that build_request was NOT called (no streaming path)
mock_async_client.build_request.assert_not_called()
# Verify that send was NOT called (no streaming path)
mock_async_client.send.assert_not_called()
# Verify that the non-streaming request method WAS called
mock_async_client.request.assert_called_once_with(
method="POST",
url=httpx.URL("https://api.anthropic.com/v1/messages"),
# Non-SSE requests are sent with stream semantics so large bodies can
# be relayed without buffering; the JSON response below is still
# buffered into a plain Response.
mock_async_client.request.assert_not_called()
mock_async_client.build_request.assert_called_once_with(
"POST",
httpx.URL("https://api.anthropic.com/v1/messages"),
headers={"Authorization": "Bearer test-key"},
params={},
json=request_body,
)
mock_async_client.send.assert_called_once()
assert mock_async_client.send.call_args.kwargs.get("stream") is True
# Verify response is a regular Response (not StreamingResponse)
from fastapi.responses import Response, StreamingResponse

View file

@ -3,8 +3,8 @@ import asyncio
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
InvalidatableOAuthTokenStore,
OAuthToken,
OAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
@ -16,11 +16,15 @@ class _RecordingStore:
def __init__(self, access_token: str) -> None:
self._access_token = access_token
self.calls: list[tuple[str, str]] = []
self.invalidations: list[tuple[str, str]] = []
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls.append((user_id, server_id))
return OAuthToken(access_token=self._access_token)
async def invalidate(self, user_id: str, server_id: str) -> None:
self.invalidations.append((user_id, server_id))
class _BlockingStore:
def __init__(self, access_token: str) -> None:
@ -28,6 +32,7 @@ class _BlockingStore:
self.started = asyncio.Event()
self.release = asyncio.Event()
self.calls: list[tuple[str, str]] = []
self.invalidations: list[tuple[str, str]] = []
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls.append((user_id, server_id))
@ -35,6 +40,9 @@ class _BlockingStore:
await self.release.wait()
return OAuthToken(access_token=self._access_token)
async def invalidate(self, user_id: str, server_id: str) -> None:
self.invalidations.append((user_id, server_id))
class _RedisAvailability:
def __init__(self) -> None:
@ -59,7 +67,7 @@ async def test_lazy_store_rebuilds_when_redis_becomes_available() -> None:
redis_available = _RedisAvailability()
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
if redis_available.available:
@ -94,7 +102,7 @@ async def test_lazy_store_allows_concurrent_local_fetches_without_redis() -> Non
redis_available = _RedisAvailability()
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False
@ -127,7 +135,7 @@ async def test_lazy_store_waits_for_in_flight_local_fetch_before_redis_rebuild()
redis_store = _RecordingStore("redis")
redis_available = _RedisAvailability()
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
if redis_available.available:
return redis_store, True
return local_store, False
@ -158,3 +166,83 @@ async def test_lazy_store_waits_for_in_flight_local_fetch_before_redis_rebuild()
assert second is not None and second.access_token == "redis"
assert local_store.calls == [("u", "s")]
assert redis_store.calls == [("u", "s")]
@pytest.mark.asyncio
async def test_lazy_store_invalidate_builds_chain_and_delegates() -> None:
local_store = _RecordingStore("local")
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=_RedisAvailability(),
)
await store.invalidate("u", "s")
assert build_calls == 1
assert local_store.invalidations == [("u", "s")]
@pytest.mark.asyncio
async def test_lazy_store_invalidate_reaches_the_store_fetch_reads() -> None:
local_store = _RecordingStore("local")
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=_RedisAvailability(),
)
await store.fetch("u", "s")
await store.invalidate("u", "s")
assert build_calls == 1
assert local_store.calls == [("u", "s")]
assert local_store.invalidations == [("u", "s")]
@pytest.mark.asyncio
async def test_lazy_store_invalidate_works_after_redis_chain_is_built() -> None:
redis_store = _RecordingStore("redis")
redis_available = _RedisAvailability()
redis_available.available = True
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return redis_store, True
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=redis_available,
)
await store.fetch("u", "s")
await store.invalidate("u", "s")
assert build_calls == 1
assert redis_store.invalidations == [("u", "s")]

View file

@ -4055,3 +4055,163 @@ async def test_oauth_authorization_server_404_for_unknown_server_name():
mcp_server_name="does_not_exist",
)
assert exc_info.value.status_code == 404
@pytest.mark.asyncio
async def test_store_per_user_token_server_side_invalidates_v2_token_cache():
"""A token stored by the OAuth callback (code exchange or refresh) drops the v2 per-user
token cache entry, so egress stops serving the replaced token immediately instead of
until its TTL."""
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_store_per_user_token_server_side,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="srv-cb-1",
name="cb_server",
url="https://upstream.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
)
invalidate_mock = AsyncMock(return_value=None)
cache_set_mock = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.utils.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch(
"litellm.proxy._experimental.mcp_server.db.store_user_oauth_credential",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.mcp_per_user_token_cache.set",
new=cache_set_mock,
),
patch.object(
manager_module.global_mcp_server_manager,
"invalidate_user_oauth_token_cache",
new=invalidate_mock,
),
):
await _store_per_user_token_server_side(
server=server,
user_id="user-cb-1",
token_response={"access_token": "fresh-tok", "expires_in": 3600},
)
invalidate_mock.assert_awaited_once_with("user-cb-1", "srv-cb-1")
cache_set_mock.assert_awaited_once()
@pytest.mark.asyncio
async def test_store_per_user_token_server_side_skips_invalidate_when_db_write_fails():
"""A failed DB write neither warms the v1 cache nor drops the v2 cache entry; the
previously stored token is still the truth."""
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_store_per_user_token_server_side,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="srv-cb-2",
name="cb_server_2",
url="https://upstream.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
)
invalidate_mock = AsyncMock(return_value=None)
cache_set_mock = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.utils.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch(
"litellm.proxy._experimental.mcp_server.db.store_user_oauth_credential",
new=AsyncMock(side_effect=RuntimeError("db down")),
),
patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.mcp_per_user_token_cache.set",
new=cache_set_mock,
),
patch.object(
manager_module.global_mcp_server_manager,
"invalidate_user_oauth_token_cache",
new=invalidate_mock,
),
):
await _store_per_user_token_server_side(
server=server,
user_id="user-cb-2",
token_response={"access_token": "fresh-tok", "expires_in": 3600},
)
invalidate_mock.assert_not_awaited()
cache_set_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_token_exchange_pairs_client_secret_with_server_client_id():
"""Re-auth regression: the register short-circuit hands the browser a placeholder
``client_secret: "dummy"``, which the browser echoes back to /token. The server-side
persisted client_id wins the resolution, so the secret must come from the same (server)
source; pairing the persisted public PKCE client (no stored secret) with the caller's
placeholder makes the IdP reject the exchange with 401 on every re-auth."""
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
exchange_token_with_server,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="srv-1",
name="srv-1",
server_name="srv-1",
alias="srv-1",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id="persisted-client",
client_secret=None,
authorization_url="https://provider.example/oauth/authorize",
token_url="https://provider.example/oauth/token",
)
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
mock_response = MagicMock()
mock_response.raise_for_status = MagicMock()
mock_response.json.return_value = {"access_token": "at", "token_type": "Bearer"}
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
):
await exchange_token_with_server(
request=mock_request,
mcp_server=server,
grant_type="authorization_code",
code="auth-code",
redirect_uri="https://litellm.example.com/ui/mcp/oauth/callback",
client_id="srv-1",
client_secret="dummy",
code_verifier="verifier",
)
sent = mock_async_client.post.call_args.kwargs["data"]
assert sent["client_id"] == "persisted-client"
assert "client_secret" not in sent

View file

@ -164,6 +164,25 @@ async def test_openapi_backed_server_also_respects_the_cap():
assert tracker.peak_by_server["srv-openapi"] == 2
@pytest.mark.asyncio
async def test_edited_limit_takes_effect_without_restart():
"""Editing max_concurrent_requests must rebuild the cached semaphore so the
new cap applies to subsequent calls immediately, not only after a restart."""
manager = MCPServerManager()
server = _make_server("srv-edited", max_concurrent_requests=3)
before_edit = _ConcurrencyTracker()
with _patch_client_with_tracker(manager, before_edit):
await _fire(manager, server, n=6)
assert before_edit.peak_by_server["srv-edited"] == 3
server.max_concurrent_requests = 1
after_edit = _ConcurrencyTracker()
with _patch_client_with_tracker(manager, after_edit):
await _fire(manager, server, n=6)
assert after_edit.peak_by_server["srv-edited"] == 1
def test_semaphore_is_reused_per_server_and_distinct_across_servers():
manager = MCPServerManager()
server_a = _make_server("srv-a", max_concurrent_requests=3)

View file

@ -8,8 +8,10 @@ from fastapi import HTTPException
from mcp import ReadResourceResult, Resource
from mcp.types import (
BlobResourceContents,
CallToolResult,
Prompt,
ResourceTemplate,
TextContent,
TextResourceContents,
)
@ -6598,6 +6600,308 @@ async def test_get_active_submitted_mcp_server_ids_for_user_empty_user_id_skips_
prisma_client.db.litellm_mcpservertable.find_many.assert_not_awaited()
# --------------------------------------------------------------------------- #
# MCP tool-call isError failure logging
# --------------------------------------------------------------------------- #
def _call_tool_result(is_error: bool, text: str) -> CallToolResult:
return CallToolResult(content=[TextContent(type="text", text=text)], isError=is_error)
def _mock_mcp_logging_obj() -> MagicMock:
logging_obj = MagicMock()
logging_obj.model_call_details = {}
logging_obj.async_post_mcp_tool_call_hook = AsyncMock()
logging_obj.async_success_handler = AsyncMock()
logging_obj.async_failure_handler = AsyncMock()
return logging_obj
def test_extract_mcp_tool_result_error_message():
from litellm.proxy._experimental.mcp_server.utils import (
extract_mcp_tool_result_error_message,
)
assert extract_mcp_tool_result_error_message(_call_tool_result(True, "boom")) == "boom"
assert extract_mcp_tool_result_error_message(_call_tool_result(False, "ok")) is None
assert (
extract_mcp_tool_result_error_message(CallToolResult(content=[], isError=True))
== "MCP tool call returned isError=true"
)
assert (
extract_mcp_tool_result_error_message({"isError": True, "content": [{"type": "text", "text": "denied"}]})
== "denied"
)
assert extract_mcp_tool_result_error_message({"isError": False, "content": []}) is None
assert extract_mcp_tool_result_error_message({}) is None
@pytest.mark.asyncio
async def test_fire_mcp_tool_call_logging_iserror_logs_failure():
"""Regression test: a CallToolResult with isError=True must go
down the failure logging path (async_failure_handler + post_call_failure_hook),
never async_success_handler."""
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPToolResultError
logging_obj = _mock_mcp_logging_obj()
proxy_logging_mock = MagicMock()
proxy_logging_mock.post_call_failure_hook = AsyncMock()
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock):
await _fire_mcp_tool_call_logging(
logging_obj=logging_obj,
result=_call_tool_result(True, "upstream exploded"),
start_time=datetime.now(),
end_time=datetime.now(),
user_api_key_auth=user_auth,
request_data={"litellm_call_id": "cid"},
)
logging_obj.async_success_handler.assert_not_awaited()
logging_obj.failure_handler.assert_called_once()
logging_obj.async_failure_handler.assert_awaited_once()
tool_error = logging_obj.async_failure_handler.await_args.args[0]
assert isinstance(tool_error, MCPToolResultError)
assert str(tool_error) == "upstream exploded"
logging_obj.has_run_logging.assert_any_call(event_type="sync_success")
logging_obj.has_run_logging.assert_any_call(event_type="async_success")
proxy_logging_mock.post_call_failure_hook.assert_awaited_once()
hook_kwargs = proxy_logging_mock.post_call_failure_hook.await_args.kwargs
assert hook_kwargs["route"] == "/mcp/call_tool"
assert hook_kwargs["original_exception"] is tool_error
assert hook_kwargs["user_api_key_dict"] is user_auth
logging_obj.async_post_mcp_tool_call_hook.assert_awaited_once()
@pytest.mark.asyncio
async def test_fire_mcp_tool_call_logging_success_path_unchanged():
"""isError=False must keep today's behavior: success handler fires, no
failure logging, no post_call_failure_hook."""
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
logging_obj = _mock_mcp_logging_obj()
proxy_logging_mock = MagicMock()
proxy_logging_mock.post_call_failure_hook = AsyncMock()
result = _call_tool_result(False, "all good")
with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock):
await _fire_mcp_tool_call_logging(
logging_obj=logging_obj,
result=result,
start_time=datetime.now(),
end_time=datetime.now(),
user_api_key_auth=UserAPIKeyAuth(api_key="test-key", user_id="test-user"),
request_data={},
)
logging_obj.async_success_handler.assert_awaited_once()
assert logging_obj.async_success_handler.await_args.kwargs["result"] is result
logging_obj.async_failure_handler.assert_not_awaited()
logging_obj.failure_handler.assert_not_called()
proxy_logging_mock.post_call_failure_hook.assert_not_awaited()
@pytest.mark.asyncio
async def test_fire_mcp_tool_call_logging_iserror_without_auth_skips_failure_hook():
"""Without a UserAPIKeyAuth the failure handlers still fire but the proxy
post_call_failure_hook (which requires one) is skipped."""
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
logging_obj = _mock_mcp_logging_obj()
proxy_logging_mock = MagicMock()
proxy_logging_mock.post_call_failure_hook = AsyncMock()
with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock):
await _fire_mcp_tool_call_logging(
logging_obj=logging_obj,
result={"isError": True, "content": [{"type": "text", "text": "denied"}]},
start_time=datetime.now(),
end_time=datetime.now(),
)
logging_obj.async_success_handler.assert_not_awaited()
logging_obj.async_failure_handler.assert_awaited_once()
assert str(logging_obj.async_failure_handler.await_args.args[0]) == "denied"
proxy_logging_mock.post_call_failure_hook.assert_not_awaited()
@pytest.mark.asyncio
async def test_fire_mcp_tool_call_logging_strips_credentials_from_failure_hook():
"""Credential-bearing request_data fields (raw request headers, upstream MCP
auth headers, OAuth tokens) must never reach post_call_failure_hook
callbacks; non-credential fields must survive untouched."""
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
logging_obj = _mock_mcp_logging_obj()
proxy_logging_mock = MagicMock()
proxy_logging_mock.post_call_failure_hook = AsyncMock()
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
request_data = {
"name": "explode",
"litellm_call_id": "cid",
"raw_headers": {"authorization": "Bearer sk-caller-secret"},
"mcp_auth_header": "upstream-secret",
"mcp_server_auth_headers": {"srv": {"authorization": "Bearer srv-secret"}},
"oauth2_headers": {"authorization": "Bearer oauth-secret"},
"user_api_key_auth": user_auth,
}
with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock):
await _fire_mcp_tool_call_logging(
logging_obj=logging_obj,
result=_call_tool_result(True, "boom"),
start_time=datetime.now(),
end_time=datetime.now(),
user_api_key_auth=user_auth,
request_data=request_data,
)
proxy_logging_mock.post_call_failure_hook.assert_awaited_once()
hook_request_data = proxy_logging_mock.post_call_failure_hook.await_args.kwargs["request_data"]
assert hook_request_data == {"name": "explode", "litellm_call_id": "cid"}
assert "secret" not in str(hook_request_data)
def _real_mcp_logging_obj(call_id: str):
from litellm.litellm_core_utils.litellm_logging import Logging
start_time = datetime.now()
logging_obj = Logging(
model="MCP: weather/get_forecast",
messages=[{"role": "user", "content": "tool call"}],
stream=False,
call_type="call_mcp_tool",
start_time=start_time,
litellm_call_id=call_id,
function_id="test-fn",
)
logging_obj.update_environment_variables(
model="MCP: weather/get_forecast",
user="",
optional_params={},
litellm_params={"api_base": ""},
)
logging_obj.model_call_details["mcp_tool_call_metadata"] = {
"name": "get_forecast",
"arguments": {"city": "Paris"},
"mcp_server_name": "weather",
}
return logging_obj, start_time
@pytest.mark.asyncio
async def test_fire_mcp_tool_call_logging_iserror_builds_failure_payload(monkeypatch):
"""The standard logging payload for an isError=True result must carry
status='failure' with the tool's error text, so OTel (whose _parse_error
keys off status) marks the MCP span ERROR."""
import litellm
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "_async_success_callback", [])
logging_obj, start_time = _real_mcp_logging_obj("test-mcp-iserror-payload")
await _fire_mcp_tool_call_logging(
logging_obj=logging_obj,
result=_call_tool_result(True, "upstream exploded"),
start_time=start_time,
end_time=datetime.now(),
)
payload = logging_obj.model_call_details["standard_logging_object"]
assert payload["status"] == "failure"
assert payload["error_str"] == "upstream exploded"
assert payload["error_information"]["error_class"] == "MCPToolResultError"
assert payload["metadata"]["mcp_tool_call_metadata"]["name"] == "get_forecast"
@pytest.mark.asyncio
async def test_fire_mcp_tool_call_logging_success_builds_success_payload(monkeypatch):
"""isError=False still produces a status='success' payload."""
import litellm
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "_async_success_callback", [])
logging_obj, start_time = _real_mcp_logging_obj("test-mcp-success-payload")
await _fire_mcp_tool_call_logging(
logging_obj=logging_obj,
result=_call_tool_result(False, "all good"),
start_time=start_time,
end_time=datetime.now(),
)
payload = logging_obj.model_call_details["standard_logging_object"]
assert payload["status"] == "success"
@pytest.mark.asyncio
async def test_fire_mcp_tool_call_logging_iserror_emits_otel_error_span(monkeypatch):
"""End-to-end regression for the OTel symptom: an isError=True tool
result must reach OTel as an MCP span with StatusCode.ERROR and the tool's
error message, while isError=False stays non-error."""
pytest.importorskip("opentelemetry")
from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
InMemorySpanExporter,
)
from opentelemetry.trace.status import StatusCode
import litellm
from litellm.integrations.otel import OpenTelemetryV2Config
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.plumbing import providers
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
cfg = OpenTelemetryV2Config(exporter="in_memory", legacy_compat=False)
exporter = InMemorySpanExporter()
tracer_provider = providers.build_tracer_provider(cfg, exporter=exporter)
otel_logger = OpenTelemetryV2(config=cfg, tracer_provider=tracer_provider)
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [otel_logger])
monkeypatch.setattr(litellm, "success_callback", [])
monkeypatch.setattr(litellm, "_async_success_callback", [otel_logger])
logging_obj, start_time = _real_mcp_logging_obj("test-mcp-iserror-otel")
await _fire_mcp_tool_call_logging(
logging_obj=logging_obj,
result=_call_tool_result(True, "upstream exploded"),
start_time=start_time,
end_time=datetime.now(),
)
(span,) = exporter.get_finished_spans()
assert span.name == "tools/call get_forecast"
assert span.status.status_code is StatusCode.ERROR
assert span.attributes["error.type"] == "MCPToolResultError"
assert "upstream exploded" in (span.status.description or "")
@pytest.mark.asyncio
async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow():
"""

View file

@ -2940,6 +2940,39 @@ class TestMCPServerManager:
assert await manager.has_user_oauth_token(server, user_auth) is False
assert calls == [] # short-circuited on the None spec, never hit the resolver
@pytest.mark.asyncio
async def test_invalidate_user_oauth_token_cache_delegates_to_store(self):
"""The write side's cache drop reaches the same per-user store the resolver reads."""
class _Store:
def __init__(self) -> None:
self.invalidations: list[tuple[str, str]] = []
async def fetch(self, user_id: str, server_id: str):
return None
async def invalidate(self, user_id: str, server_id: str) -> None:
self.invalidations.append((user_id, server_id))
store = _Store()
manager = MCPServerManager(per_user_oauth_token_store=store)
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
assert store.invalidations == [("alice", "srv-1")]
@pytest.mark.asyncio
async def test_invalidate_user_oauth_token_cache_swallows_store_errors(self):
"""A cache-drop failure must not fail the credential write that triggered it."""
class _Store:
async def fetch(self, user_id: str, server_id: str):
return None
async def invalidate(self, user_id: str, server_id: str) -> None:
raise RuntimeError("redis down")
manager = MCPServerManager(per_user_oauth_token_store=_Store())
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
@pytest.mark.asyncio
async def test_resolve_oauth2_headers_no_user_id(self):
"""Skip lookup entirely when user_api_key_auth has no user_id."""

View file

@ -433,7 +433,7 @@ class TestCallToolRestApiVirtualTools:
return_value=fake_result,
) as mock_execute,
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._fire_mcp_success_logging",
"litellm.proxy._experimental.mcp_server.rest_endpoints._fire_mcp_tool_call_logging",
new_callable=AsyncMock,
side_effect=RuntimeError("logging failed"),
) as mock_fire_logging,
@ -789,6 +789,47 @@ class TestCaptureHostProgressCallback:
host.request_context.session = MagicMock()
assert callable(_capture_host_progress_callback(host))
def test_returns_callable_when_token_is_integer(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progressToken = 12345
host.request_context.session = MagicMock()
assert callable(_capture_host_progress_callback(host))
def test_returns_callable_when_token_is_zero(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progressToken = 0
host.request_context.session = MagicMock()
assert callable(_capture_host_progress_callback(host))
@pytest.mark.asyncio
async def test_forwarded_progress_token_preserves_integer_value(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progressToken = 12345
session = AsyncMock()
host.request_context.session = session
callback = _capture_host_progress_callback(host)
assert callback is not None
await callback(0.5, 1.0)
session.send_progress_notification.assert_awaited_once_with(
progress_token=12345,
progress=0.5,
total=1.0,
)
class TestHandleListToolsVirtual:
"""Covers the protocol list_tools early-return when the flag is enabled."""

View file

@ -1559,7 +1559,7 @@ class TestCallToolRestAPI:
fire_logging = AsyncMock(side_effect=RuntimeError("logging failed"))
monkeypatch.setattr(
rest_endpoints,
"_fire_mcp_success_logging",
"_fire_mcp_tool_call_logging",
fire_logging,
raising=False,
)
@ -1590,13 +1590,13 @@ class TestCallToolRestAPI:
fire_logging = AsyncMock(side_effect=asyncio.CancelledError())
monkeypatch.setattr(
rest_endpoints,
"_fire_mcp_success_logging",
"_fire_mcp_tool_call_logging",
fire_logging,
raising=False,
)
with pytest.raises(asyncio.CancelledError):
await rest_endpoints._safe_fire_mcp_success_logging(
await rest_endpoints._safe_fire_mcp_tool_call_logging(
object(), {"result": "ok"}, datetime.now(), datetime.now()
)

View file

@ -516,3 +516,104 @@ async def test_full_hot_path_network_count():
assert (
summary["total_network_requests"] == 4
), f"Expected 4 total network requests on warm path, got {summary['total_network_requests']}"
# ============================================================================
# TEST: negative caching for entities that do not exist in the DB
# ============================================================================
@pytest.mark.asyncio
async def test_get_user_object_missing_user_negative_cache():
"""
A user_id with no DB row (e.g. the master key's default admin user_id)
must not trigger a DB query on every request. The first lookup hits the
DB; repeat lookups inside the db_cache_expiry window are throttled.
"""
user_id = "user-missing-negative-cache"
cache = DualCache(in_memory_cache=InMemoryCache())
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_usertable = MagicMock()
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
for _ in range(3):
with pytest.raises(ValueError):
await get_user_object(
user_id=user_id,
prisma_client=mock_prisma,
user_api_key_cache=cache,
parent_otel_span=None,
proxy_logging_obj=None,
user_id_upsert=False,
)
assert mock_prisma.db.litellm_usertable.find_unique.call_count == 1
@pytest.mark.asyncio
async def test_get_user_object_missing_user_rechecks_after_expiry():
"""
The negative cache must expire: a user created after a miss becomes
visible once the db_cache_expiry window has passed.
"""
from litellm.proxy.auth.auth_checks import db_cache_expiry, last_db_access_time
user_id = "user-missing-expiry-recheck"
cache = DualCache(in_memory_cache=InMemoryCache())
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_usertable = MagicMock()
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
with pytest.raises(ValueError):
await get_user_object(
user_id=user_id,
prisma_client=mock_prisma,
user_api_key_cache=cache,
parent_otel_span=None,
proxy_logging_obj=None,
user_id_upsert=False,
)
assert mock_prisma.db.litellm_usertable.find_unique.call_count == 1
last_db_access_time[f"user_id:{user_id}"] = (
None,
time.time() - (db_cache_expiry + 1),
)
with pytest.raises(ValueError):
await get_user_object(
user_id=user_id,
prisma_client=mock_prisma,
user_api_key_cache=cache,
parent_otel_span=None,
proxy_logging_obj=None,
user_id_upsert=False,
)
assert mock_prisma.db.litellm_usertable.find_unique.call_count == 2
def test_should_check_db_negative_entry_throttles_then_expires():
"""
A recorded miss (value=None) suppresses DB checks inside the expiry
window and allows them again after it. Exercises the timestamp element
of the stored (value, time) tuple directly.
"""
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.proxy.auth.auth_checks import (
_should_check_db,
_update_last_db_access_time,
)
tracker: LimitedSizeOrderedDict = LimitedSizeOrderedDict(max_size=10)
_update_last_db_access_time(key="k", value=None, last_db_access_time=tracker)
assert _should_check_db(key="k", last_db_access_time=tracker, db_cache_expiry=5) is False
tracker["k"] = (None, time.time() - 6)
assert _should_check_db(key="k", last_db_access_time=tracker, db_cache_expiry=5) is True

View file

@ -19,6 +19,7 @@ from litellm.proxy.auth.auth_utils import (
get_key_mcp_rpm_limit,
get_key_model_rpm_limit,
get_key_model_tpm_limit,
get_key_tag_rpm_limit,
get_model_from_request,
get_project_model_rpm_limit,
get_project_model_tpm_limit,
@ -2393,3 +2394,17 @@ class TestIsRequestBodySafeBlocksModelList:
)
is True
)
class TestGetKeyTagRateLimits:
"""Tests for get_key_tag_rpm_limit."""
def test_reads_tag_rpm_limit_from_metadata(self):
key = UserAPIKeyAuth(
api_key="sk-123", metadata={"tag_rpm_limit": {"cell-1": 5}}
)
assert get_key_tag_rpm_limit(key) == {"cell-1": 5}
def test_returns_none_when_unset(self):
key = UserAPIKeyAuth(api_key="sk-123")
assert get_key_tag_rpm_limit(key) is None

View file

@ -3573,3 +3573,136 @@ async def test_pre_call_hook_skips_reservation_when_disabled(monkeypatch):
)
assert TPM_RESERVED_TOKENS_KEY not in (data.get("metadata") or {})
@pytest.mark.asyncio
async def test_per_tag_rate_limit_independent_counters_v3(monkeypatch):
"""
A single key with per-tag RPM limits tracks each tag independently: a tag
at its limit returns 429 while a different (unlimited) tag keeps flowing,
governed only by the generous key-level limit.
"""
monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60")
_api_key = hash_token("sk-per-tag-rpm")
user_api_key_dict = UserAPIKeyAuth(
api_key=_api_key,
rpm_limit=100,
metadata={"tag_rpm_limit": {"cell-1": 2}},
)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
async def call(tag: str) -> None:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"model": "gpt-3.5-turbo", "metadata": {"tags": [tag]}},
call_type="",
)
await call("cell-1")
await call("cell-1")
with pytest.raises(HTTPException) as exc_info:
await call("cell-1")
assert exc_info.value.status_code == 429
assert "tag_per_key" in str(exc_info.value.detail)
# cell-2 has no configured tag limit, so cell-1's exhausted counter must
# not block it; only the generous key-level limit applies.
for _ in range(5):
await call("cell-2")
@pytest.mark.asyncio
async def test_per_tag_descriptor_creation_v3():
"""
_create_rate_limit_descriptors emits a tag_per_key descriptor carrying the
configured RPM limit only for request tags present in the configured map.
"""
_api_key = hash_token("sk-per-tag-desc")
user_api_key_dict = UserAPIKeyAuth(
api_key=_api_key,
metadata={"tag_rpm_limit": {"cell-1": 5}},
)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = handler._create_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1", "cell-2"]}},
rpm_limit_type=None,
tpm_limit_type=None,
model_has_failures=False,
)
tag_descriptors = [d for d in descriptors if d["key"] == "tag_per_key"]
assert len(tag_descriptors) == 1, "only the configured tag yields a descriptor"
descriptor = tag_descriptors[0]
assert descriptor["value"] == f"{_api_key}:cell-1"
assert descriptor["rate_limit"]["requests_per_unit"] == 5
@pytest.mark.asyncio
async def test_per_tag_descriptor_absent_without_config_v3():
"""No tag_per_key descriptor is created when the key has no tag limits."""
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token("sk-no-tag"),
rpm_limit=10,
)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = handler._create_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1"]}},
rpm_limit_type=None,
tpm_limit_type=None,
model_has_failures=False,
)
assert not [d for d in descriptors if d["key"] == "tag_per_key"]
@pytest.mark.asyncio
async def test_per_tag_untagged_request_governed_by_key_limit_v3(monkeypatch):
"""
Per-tag limits are opt-in sub-limits under the key-level ceiling, not a
standalone enforcement boundary: a request that carries no tag (or a tag
without a configured limit) is not rejected by any tag counter, but it is
still bounded by the key-level rpm_limit. This pins the documented
untagged-fallback behavior so a future "fail closed on missing tag" change
would fail here instead of silently breaking it.
"""
monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60")
_api_key = hash_token("sk-untagged-fallback")
user_api_key_dict = UserAPIKeyAuth(
api_key=_api_key,
rpm_limit=3,
metadata={"tag_rpm_limit": {"cell-1": 2}},
)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
async def call(metadata: dict) -> None:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"model": "gpt-3.5-turbo", "metadata": metadata},
call_type="",
)
# Untagged and unconfigured-tag requests share the key-level budget of 3
# and never hit a tag_per_key counter.
await call({})
await call({"tags": ["cell-99"]})
await call({})
with pytest.raises(HTTPException) as exc_info:
await call({"tags": ["cell-99"]})
assert exc_info.value.status_code == 429
assert "tag_per_key" not in str(exc_info.value.detail)

View file

@ -15,8 +15,11 @@ from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import HTTPException
import inspect
from litellm.proxy._types import (
GenerateKeyRequest,
NewUserRequest,
LiteLLM_BudgetTable,
LiteLLM_OrganizationTable,
LiteLLM_TeamTableCachedObj,
@ -9020,7 +9023,7 @@ class TestValidateKeyAliasFormat:
litellm.enable_key_alias_format_validation = False
def test_validation_skipped_when_flag_disabled(self):
"""When enable_key_alias_format_validation is False (default), no validation occurs."""
"""When enable_key_alias_format_validation is False (default), no charset/length validation occurs."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_key_alias_format,
)
@ -9031,6 +9034,33 @@ class TestValidateKeyAliasFormat:
_validate_key_alias_format("!invalid!")
_validate_key_alias_format("a" * 256)
@pytest.mark.parametrize(
"unsafe_alias",
[
"../../../other-app/creds",
"litellm/../../secret",
"foo\n- !grant\n role: !!admin\n member: attacker",
"foo\rbar",
"foo\x00bar",
],
)
def test_validate_key_alias_format_rejects_traversal_and_control_chars_even_when_flag_disabled(
self, unsafe_alias
):
"""
Regression test: this check must reject an invalid key_alias unconditionally,
even when enable_key_alias_format_validation (the separate, opt-in charset
rule) is disabled.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_key_alias_format,
)
with pytest.raises(ProxyException) as exc:
_validate_key_alias_format(unsafe_alias)
assert str(exc.value.code) == "400"
assert "Invalid key_alias" in str(exc.value.message)
def test_validate_key_alias_format_valid(self):
from litellm.proxy.management_endpoints.key_management_endpoints import (
_validate_key_alias_format,
@ -14480,3 +14510,23 @@ async def test_regenerate_key_non_admin_permissions_rejected_before_enterprise_g
assert int(exc.value.code) == 403
assert "permissions" in str(exc.value.message)
assert "Enterprise" not in str(exc.value.message)
def test_generate_key_helper_fn_accepts_per_tag_rate_limits():
"""
Regression: new_user / SSO sign-in forward NewUserRequest fields to
generate_key_helper_fn via `**data_json`. The per-tag limit field must be
an accepted kwarg, otherwise user creation 500s with
"generate_key_helper_fn() got an unexpected keyword argument 'tag_rpm_limit'".
"""
params = inspect.signature(generate_key_helper_fn).parameters
assert "tag_rpm_limit" in params
# The field exists on the request model that new_user forwards via **data_json.
assert "tag_rpm_limit" in NewUserRequest.model_fields
# Binding the per-tag kwarg must not raise an unexpected-keyword TypeError.
inspect.signature(generate_key_helper_fn).bind_partial(
request_type="user",
tag_rpm_limit={"cell-1": 5},
)

View file

@ -3566,6 +3566,148 @@ async def test_delete_mcp_oauth_user_credential_only_deletes_oauth():
assert result.has_credential is False
@pytest.mark.asyncio
async def test_store_mcp_oauth_user_credential_invalidates_cached_token():
"""Re-authorizing via the Tools-tab persist drops the v2 per-user token cache entry, so
egress stops serving the replaced token immediately instead of until its TTL."""
from litellm.proxy._types import MCPOAuthUserCredentialRequest
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
store_mcp_oauth_user_credential,
)
server_id = "srv-inv-1"
user_id = "user-inv-1"
invalidate_mock = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
return_value=True,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.store_user_oauth_credential",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential",
new=AsyncMock(return_value={"type": "oauth2", "access_token": "new-tok"}),
),
patch.object(
manager_module.global_mcp_server_manager,
"invalidate_user_oauth_token_cache",
new=invalidate_mock,
),
):
await store_mcp_oauth_user_credential(
server_id=server_id,
payload=MCPOAuthUserCredentialRequest(access_token="new-tok", expires_in=3600),
user_api_key_dict=_make_user_auth(user_id),
)
invalidate_mock.assert_awaited_once_with(user_id, server_id)
@pytest.mark.asyncio
async def test_delete_mcp_oauth_user_credential_invalidates_cached_token():
"""Revoking a stored OAuth credential drops the v2 per-user token cache entry, so the
revoked token stops flowing upstream immediately instead of until its TTL."""
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_oauth_user_credential,
)
server_id = "srv-inv-2"
user_id = "user-inv-2"
invalidate_mock = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential",
new=AsyncMock(return_value={"type": "oauth2", "access_token": "revoked-tok"}),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential",
new=AsyncMock(return_value=None),
),
patch.object(
manager_module.global_mcp_server_manager,
"invalidate_user_oauth_token_cache",
new=invalidate_mock,
),
):
result = await delete_mcp_oauth_user_credential(
server_id=server_id,
user_api_key_dict=_make_user_auth(user_id),
)
invalidate_mock.assert_awaited_once_with(user_id, server_id)
assert result.has_credential is False
@pytest.mark.asyncio
async def test_delete_mcp_oauth_user_credential_invalidates_when_record_already_gone():
"""A concurrent delete can remove the row between the read and the delete; the cache may
still hold the revoked token, so the invalidate must fire even on RecordNotFoundError."""
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_oauth_user_credential,
)
server_id = "srv-inv-3"
user_id = "user-inv-3"
invalidate_mock = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential",
new=AsyncMock(return_value={"type": "oauth2", "access_token": "revoked-tok"}),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential",
new=AsyncMock(side_effect=mgmt_endpoints.RecordNotFoundError({}, message="already gone")),
),
patch.object(
manager_module.global_mcp_server_manager,
"invalidate_user_oauth_token_cache",
new=invalidate_mock,
),
):
result = await delete_mcp_oauth_user_credential(
server_id=server_id,
user_api_key_dict=_make_user_auth(user_id),
)
invalidate_mock.assert_awaited_once_with(user_id, server_id)
assert result.has_credential is False
@pytest.mark.asyncio
async def test_list_mcp_user_credentials_batch_server_fetch():
"""list_mcp_user_credentials uses a single batch DB call, not N+1 queries."""

View file

@ -389,6 +389,78 @@ async def test_get_user_groups_error_handling():
assert len(result) == 0
@pytest.mark.asyncio
async def test_get_user_groups_uses_default_graph_endpoint(monkeypatch):
monkeypatch.delenv("MICROSOFT_GRAPH_ENDPOINT", raising=False)
requested_urls: list[str] = []
async def mock_get(url, *args, **kwargs):
requested_urls.append(url)
mock = MagicMock()
mock.json.return_value = {"value": []}
return mock
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_client:
mock_client.return_value = MagicMock()
mock_client.return_value.get = mock_get
await MicrosoftSSOHandler.get_user_groups_from_graph_api(access_token="mock_token")
assert requested_urls == ["https://graph.microsoft.com/v1.0/me/memberOf"]
@pytest.mark.asyncio
async def test_get_user_groups_uses_configured_graph_endpoint(monkeypatch):
monkeypatch.setenv("MICROSOFT_GRAPH_ENDPOINT", "https://graph.microsoft.us/v1.0")
requested_urls: list[str] = []
async def mock_get(url, *args, **kwargs):
requested_urls.append(url)
mock = MagicMock()
mock.json.return_value = {"value": []}
return mock
with patch(
"litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client"
) as mock_client:
mock_client.return_value = MagicMock()
mock_client.return_value.get = mock_get
await MicrosoftSSOHandler.get_user_groups_from_graph_api(access_token="mock_token")
assert requested_urls == ["https://graph.microsoft.us/v1.0/me/memberOf"]
@pytest.mark.asyncio
async def test_get_group_ids_from_service_principal_uses_configured_graph_endpoint(monkeypatch):
monkeypatch.setenv("MICROSOFT_GRAPH_ENDPOINT", "https://graph.microsoft.us/v1.0")
requested_urls: list[str] = []
async def mock_get(url, *args, **kwargs):
requested_urls.append(url)
mock = MagicMock()
mock.json.return_value = {"value": []}
return mock
async_client = MagicMock()
async_client.get = mock_get
await MicrosoftSSOHandler.get_group_ids_from_service_principal(
service_principal_id="sp-123",
async_client=async_client,
access_token="mock_token",
)
assert requested_urls == [
"https://graph.microsoft.us/v1.0/servicePrincipals/sp-123/appRoleAssignedTo"
]
def test_get_group_ids_from_graph_api_response():
# Arrange
mock_response = MicrosoftGraphAPIUserGroupResponse(
@ -7131,3 +7203,54 @@ async def test_legacy_login_page_hides_credentials_hint_via_general_settings():
assert response.status_code == 200
assert "Default Credentials" not in body
assert "MASTER_KEY" not in body
@pytest.mark.asyncio
async def test_cli_poll_key_tolerates_missing_user_row():
"""The CLI poll must still mint the JWT when the user lookup raises,
e.g. the user row was created moments ago and a negative-cache window
from the pre-creation SSO existence check is still active on this pod."""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
session_key = "cli-session-missing-user"
session_data = {
"user_id": "just-created-user",
"user_role": "internal_user",
"teams": [],
"models": ["gpt-4"],
}
mock_cache = MagicMock()
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
mock_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.missing.user"
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.prisma_client"),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value=mock_jwt_token,
),
patch(
"litellm.proxy.auth.auth_checks.get_user_object",
new=AsyncMock(side_effect=ValueError("User doesn't exist in db. 'user_id'=just-created-user")),
),
):
result = await cli_poll_key(
key_id=session_key,
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
assert result["status"] == "ready"
assert result["key"] == mock_jwt_token
assert result["user_id"] == "just-created-user"

View file

@ -1918,7 +1918,8 @@ class TestForwardHeaders:
):
# Setup mock httpx client
mock_client = MagicMock()
mock_client.request = AsyncMock(return_value=mock_httpx_response)
mock_client.build_request = MagicMock(return_value=MagicMock())
mock_client.send = AsyncMock(return_value=mock_httpx_response)
mock_client_obj = MagicMock()
mock_client_obj.client = mock_client
mock_get_client.return_value = mock_client_obj
@ -1942,10 +1943,10 @@ class TestForwardHeaders:
)
# Verify the httpx client was called
assert mock_client.request.called
assert mock_client.send.called
# Get the headers that were sent to the target
call_args = mock_client.request.call_args
call_args = mock_client.build_request.call_args
sent_headers = call_args[1]["headers"]
# Verify user headers were forwarded (except content-length and host)
@ -2019,7 +2020,8 @@ class TestForwardHeaders:
):
# Setup mock httpx client
mock_client = MagicMock()
mock_client.request = AsyncMock(return_value=mock_httpx_response)
mock_client.build_request = MagicMock(return_value=MagicMock())
mock_client.send = AsyncMock(return_value=mock_httpx_response)
mock_client_obj = MagicMock()
mock_client_obj.client = mock_client
mock_get_client.return_value = mock_client_obj
@ -2043,10 +2045,10 @@ class TestForwardHeaders:
)
# Verify the httpx client was called
assert mock_client.request.called
assert mock_client.send.called
# Get the headers that were sent to the target
call_args = mock_client.request.call_args
call_args = mock_client.build_request.call_args
sent_headers = call_args[1]["headers"]
# Verify only custom headers were sent

View file

@ -1,5 +1,6 @@
import asyncio
import json
import logging
import os
import sys
from contextlib import ExitStack
@ -1337,7 +1338,8 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream():
upstream_response.raise_for_status = MagicMock()
async_client = MagicMock()
async_client.request = AsyncMock(return_value=upstream_response)
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
async def _empty_chunks(*args, **kwargs):
@ -1361,7 +1363,7 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream():
stream=False,
)
async_client.request.assert_awaited_once()
async_client.send.assert_awaited_once()
mock_chunk_processor.assert_called_once()
logging_obj = mock_chunk_processor.call_args.kwargs[
@ -3046,7 +3048,8 @@ async def test_pass_through_request_non_streaming_uses_content_for_state_raw_bod
)
mock_async_client = AsyncMock()
mock_async_client.request = AsyncMock(return_value=upstream)
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=upstream)
mock_client_obj = MagicMock()
mock_client_obj.client = mock_async_client
@ -3082,10 +3085,12 @@ async def test_pass_through_request_non_streaming_uses_content_for_state_raw_bod
stream=False,
)
mock_async_client.request.assert_called_once()
req_kw = mock_async_client.request.call_args[1]
assert req_kw.get("content") == raw_signed
assert "json" not in req_kw
mock_async_client.build_request.assert_called_once()
build_kw = mock_async_client.build_request.call_args[1]
assert build_kw.get("content") == raw_signed
assert "json" not in build_kw
mock_async_client.send.assert_awaited_once()
assert mock_async_client.send.call_args.kwargs.get("stream") is True
@pytest.mark.asyncio
@ -3826,7 +3831,8 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan
mock_success_handler.return_value = None
async_client = MagicMock()
async_client.request = AsyncMock(return_value=upstream_response)
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
@ -3913,7 +3919,8 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa
mock_success_handler.return_value = None
async_client = MagicMock()
async_client.request = AsyncMock(return_value=upstream_response)
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
@ -4043,7 +4050,8 @@ async def test_pass_through_request_non_streaming_success_unchanged():
mock_success_handler.return_value = None
async_client = MagicMock()
async_client.request = AsyncMock(return_value=upstream_response)
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
@ -4103,3 +4111,393 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio
assert int(exc_info.value.code) == 500
assert "auth backend unavailable" in exc_info.value.message
class _RecordingUpstreamByteStream(httpx.AsyncByteStream):
def __init__(self, chunks):
self._chunks = chunks
self.chunks_served = 0
self.closed = False
async def __aiter__(self):
for chunk in self._chunks:
self.chunks_served += 1
yield chunk
async def aclose(self):
self.closed = True
class _FakeUpstreamTransport(httpx.AsyncBaseTransport):
def __init__(self, status_code, headers, stream):
self._status_code = status_code
self._headers = headers
self._stream = stream
async def handle_async_request(self, request):
return httpx.Response(
status_code=self._status_code,
headers=self._headers,
stream=self._stream,
request=request,
)
def _inject_fake_passthrough_client(transport, timeout):
"""Dependency-inject a fake upstream via the client cache that
get_async_httpx_client resolves passthrough clients from (no monkeypatching
of the HTTP layer). The cache entry is located by calling the production
get_async_httpx_client and identity-scanning the cache for the handler it
returned, so the internal cache-key format is never duplicated here. Must
run inside the test's event loop because cache keys are loop-scoped.
Returns (client, cleanup)."""
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
real_handler = get_async_httpx_client(
httpxSpecialProvider.PassThroughEndpoint,
params={"timeout": resolve_pass_through_request_timeout(timeout)},
)
cache = litellm.in_memory_llm_clients_cache
cache_key = next(
(key for key, cached in cache.cache_dict.items() if cached is real_handler),
None,
)
assert cache_key is not None, (
"PassThroughEndpoint client not found in in_memory_llm_clients_cache; "
"get_async_httpx_client may not be caching this provider."
)
fake_client = httpx.AsyncClient(transport=transport)
cache.cache_dict[cache_key] = SimpleNamespace(client=fake_client)
def _cleanup():
cache.cache_dict.pop(cache_key, None)
return fake_client, _cleanup
def _enter_relay_logging_mocks(stack, parsed_body):
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
mock_proxy_logging = stack.enter_context(
patch("litellm.proxy.proxy_server.proxy_logging_obj")
)
mock_proxy_logging.pre_call_hook = AsyncMock(return_value=parsed_body)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler = stack.enter_context(
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
)
)
mock_success_handler.return_value = None
stack.enter_context(
patch.object(
GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue", new=MagicMock()
)
)
return mock_proxy_logging, mock_success_handler
def _relay_client_request(method="GET"):
mock_request = MagicMock(spec=Request)
mock_request.method = method
mock_request.url = "http://localhost:4000/passthrough-relay/results"
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
return mock_request
@pytest.mark.asyncio
async def test_pass_through_request_relays_non_json_body_without_buffering():
"""
Regression (LIT-4009): non-SSE passthrough responses used to be fully
buffered in proxy memory (content = await response.aread()) before a single
byte reached the client, ballooning proxy RSS to a multiple of the body size
for large non-JSON downloads (e.g. Anthropic batch results .jsonl files) and
producing near-total TTFB dead air that let intermediaries kill the silent
connection mid-download.
A non-JSON 2xx body must be relayed as a StreamingResponse whose chunks are
pulled from the upstream one at a time, with zero chunks consumed before the
handler returns, upstream status/headers plus x-litellm-* headers preserved,
and the success-handler logging fired with response_body=None once the
stream completes. Pre-fix, the handler returned a plain Response after
reading the entire body, so these assertions fail on the old code.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_chunks = (
b'{"custom_id": "a", "result": {}}\n',
b'{"custom_id": "b", "result": {}}\n',
b'{"custom_id": "c", "result": {}}\n',
)
upstream_stream = _RecordingUpstreamByteStream(upstream_chunks)
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=200,
headers={
"content-type": "application/x-jsonl",
"x-upstream-marker": "batch-results",
"content-length": str(sum(len(c) for c in upstream_chunks)),
},
stream=upstream_stream,
),
timeout=311.0,
)
try:
with ExitStack() as stack:
_, mock_success_handler = _enter_relay_logging_mocks(stack, {})
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/messages/batches/b1/results",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=311.0,
)
assert isinstance(response, StreamingResponse)
assert upstream_stream.chunks_served == 0
mock_success_handler.assert_not_called()
iterator = response.body_iterator
first_chunk = await iterator.__anext__()
assert first_chunk == upstream_chunks[0]
assert upstream_stream.chunks_served == 1
remaining = [chunk async for chunk in iterator]
assert b"".join([first_chunk, *remaining]) == b"".join(upstream_chunks)
assert upstream_stream.closed is True
assert response.status_code == 200
assert response.headers["x-upstream-marker"] == "batch-results"
assert "x-litellm-call-id" in response.headers
assert "content-length" not in response.headers
mock_success_handler.assert_called_once()
success_kwargs = mock_success_handler.call_args.kwargs
assert success_kwargs["response_body"] is None
assert (
success_kwargs["url_route"]
== "http://upstream.test/v1/messages/batches/b1/results"
)
finally:
cleanup()
await fake_client.aclose()
@pytest.mark.asyncio
async def test_pass_through_request_json_response_stays_buffered_for_logging():
"""
JSON responses (content-type application/json) must keep the buffered
behavior: spend logging and guardrails inspect the parsed body, so the
handler reads the full upstream body and passes the parsed dict to the
success handler.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_chunks = (b'{"id": "file-123"', b', "status": "processed"}')
upstream_stream = _RecordingUpstreamByteStream(upstream_chunks)
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=200,
headers={"content-type": "application/json"},
stream=upstream_stream,
),
timeout=312.0,
)
try:
with ExitStack() as stack:
_, mock_success_handler = _enter_relay_logging_mocks(stack, {})
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/files/file-123",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=312.0,
)
assert not isinstance(response, StreamingResponse)
assert response.status_code == 200
assert response.body == b"".join(upstream_chunks)
assert upstream_stream.chunks_served == len(upstream_chunks)
mock_success_handler.assert_called_once()
success_kwargs = mock_success_handler.call_args.kwargs
assert success_kwargs["response_body"] == {
"id": "file-123",
"status": "processed",
}
finally:
cleanup()
await fake_client.aclose()
@pytest.mark.asyncio
async def test_pass_through_request_upstream_error_body_stays_buffered():
"""
Upstream errors are never relayed as a stream, whatever their content-type:
the body must stay available for the failure hook and reach the client
buffered with the upstream status code, exactly as before the fix.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_stream = _RecordingUpstreamByteStream((b"upstream ", b"exploded"))
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=502,
headers={"content-type": "application/x-jsonl"},
stream=upstream_stream,
),
timeout=313.0,
)
try:
with ExitStack() as stack:
mock_proxy_logging, mock_success_handler = _enter_relay_logging_mocks(
stack, {}
)
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/messages/batches/b1/results",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=313.0,
)
assert not isinstance(response, StreamingResponse)
assert response.status_code == 502
assert response.body == b"upstream exploded"
mock_proxy_logging.post_call_failure_hook.assert_called_once()
mock_success_handler.assert_not_called()
finally:
cleanup()
await fake_client.aclose()
_PARTIAL_RELAY_WARNING_MARKER = "ended before upstream body was fully relayed"
@pytest.mark.asyncio
async def test_pass_through_relay_client_disconnect_logs_partial_relay_warning(caplog):
"""
Regression: when the client disconnects mid-relay (GeneratorExit), the
proxy log must record that the upstream body was only partially delivered,
including the route and the byte count that reached the client, while the
success handler still fires so the partial delivery produces a spend-log
row. Pre-fix, the finally block fired the success handler silently and a
partial delivery was indistinguishable from a complete one.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_chunks = (b'{"custom_id": "a"}\n', b'{"custom_id": "b"}\n')
upstream_stream = _RecordingUpstreamByteStream(upstream_chunks)
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=200,
headers={"content-type": "application/x-jsonl"},
stream=upstream_stream,
),
timeout=314.0,
)
try:
with ExitStack() as stack:
_, mock_success_handler = _enter_relay_logging_mocks(stack, {})
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/messages/batches/b1/results",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=314.0,
)
assert isinstance(response, StreamingResponse)
iterator = response.body_iterator
first_chunk = await iterator.__anext__()
assert first_chunk == upstream_chunks[0]
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await iterator.aclose()
partial_relay_warnings = [
record.getMessage()
for record in caplog.records
if record.levelno == logging.WARNING
and _PARTIAL_RELAY_WARNING_MARKER in record.getMessage()
]
assert len(partial_relay_warnings) == 1
assert (
"http://upstream.test/v1/messages/batches/b1/results"
in partial_relay_warnings[0]
)
assert (
f"{len(first_chunk)} bytes were sent to the client"
in partial_relay_warnings[0]
)
assert upstream_stream.closed is True
mock_success_handler.assert_called_once()
assert mock_success_handler.call_args.kwargs["response_body"] is None
finally:
cleanup()
await fake_client.aclose()
@pytest.mark.asyncio
async def test_pass_through_relay_full_consumption_logs_no_partial_relay_warning(caplog):
"""
A fully consumed relay must not be reported as a partial delivery: the
success handler fires and no partial-relay warning is logged.
"""
from fastapi.responses import StreamingResponse
from litellm.proxy._types import UserAPIKeyAuth
upstream_chunks = (b'{"custom_id": "a"}\n', b'{"custom_id": "b"}\n')
upstream_stream = _RecordingUpstreamByteStream(upstream_chunks)
fake_client, cleanup = _inject_fake_passthrough_client(
_FakeUpstreamTransport(
status_code=200,
headers={"content-type": "application/x-jsonl"},
stream=upstream_stream,
),
timeout=315.0,
)
try:
with ExitStack() as stack:
_, mock_success_handler = _enter_relay_logging_mocks(stack, {})
response = await pass_through_request(
request=_relay_client_request(),
target="http://upstream.test/v1/messages/batches/b1/results",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"),
timeout=315.0,
)
assert isinstance(response, StreamingResponse)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
relayed = [chunk async for chunk in response.body_iterator]
assert b"".join(relayed) == b"".join(upstream_chunks)
assert not any(
_PARTIAL_RELAY_WARNING_MARKER in record.getMessage()
for record in caplog.records
)
mock_success_handler.assert_called_once()
finally:
cleanup()
await fake_client.aclose()

View file

@ -1102,6 +1102,11 @@ def test_ProxyConfig__add_deployment_invalid_litellm_params_skips(monkeypatch):
def test_ProxyConfig__add_deployment_resolves_env_refs_after_db_decrypt(monkeypatch):
"""Every ``os.environ/`` value on an admin-scoped DB row resolves at
load time, regardless of the field name. Replaces the earlier
behavior where only fields in ``_DB_LITELLM_PARAM_ENV_REF_KEYS``
resolved: the whitelist has been removed so the resolver applies to
every string field."""
monkeypatch.setenv("LITELLM_DB_MODEL_API_KEY", "resolved-secret")
monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret")
monkeypatch.setattr(
@ -1129,19 +1134,21 @@ def test_ProxyConfig__add_deployment_resolves_env_refs_after_db_decrypt(monkeypa
assert added == 1
assert deployment.litellm_params.api_key == "resolved-secret"
assert deployment.litellm_params.api_base == "os.environ/LITELLM_MASTER_KEY"
assert deployment.litellm_params.api_base == "master-secret"
def test_ProxyConfig__add_deployment_keeps_team_env_refs_literal(monkeypatch):
def fail_on_call(secret_name, *args, **kwargs):
raise AssertionError("team DB models should not resolve env refs")
def test_ProxyConfig__add_deployment_resolves_team_env_refs(monkeypatch):
"""Team-scoped DB rows now resolve ``os.environ/`` refs the same way
admin rows do. The prior team-scoped short-circuit and the
field-by-field whitelist have both been removed; the write-side team
auth check in ``ModelManagementAuthChecks.can_user_make_model_call``
remains the single trust boundary. A literal (non-``os.environ/``)
value still passes through unchanged."""
monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret")
monkeypatch.setattr(
"litellm.proxy.proxy_server.decrypt_value_helper",
lambda value, key, return_original_value: value,
)
monkeypatch.setattr("litellm.proxy.proxy_server.get_secret", fail_on_call)
fake_router = MagicMock()
fake_router.upsert_deployment = MagicMock(return_value=True)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router)
@ -1153,7 +1160,7 @@ def test_ProxyConfig__add_deployment_keeps_team_env_refs_literal(monkeypatch):
litellm_params={
"model": "openai/gpt-4o-mini",
"api_key": "os.environ/LITELLM_MASTER_KEY",
"api_base": "https://attacker.example",
"api_base": "https://team.example",
},
blocked=False,
)
@ -1162,8 +1169,8 @@ def test_ProxyConfig__add_deployment_keeps_team_env_refs_literal(monkeypatch):
deployment = fake_router.upsert_deployment.call_args.kwargs["deployment"]
assert added == 1
assert deployment.litellm_params.api_key == "os.environ/LITELLM_MASTER_KEY"
assert deployment.litellm_params.api_base == "https://attacker.example"
assert deployment.litellm_params.api_key == "master-secret"
assert deployment.litellm_params.api_base == "https://team.example"
def test_ProxyConfig__resolve_db_litellm_param_skips_non_string_values(monkeypatch):
@ -1242,31 +1249,26 @@ def test_ProxyConfig__add_deployment_resolves_env_refs_for_aws_bedrock_auth_para
assert getattr(deployment.litellm_params, key) == expected, key
def test_ProxyConfig__add_deployment_keeps_team_aws_env_refs_literal(monkeypatch):
"""Team-scoped DB models must NOT resolve env refs even for AWS auth
params: this is the LIT-3831 defense-in-depth path where a team admin
could otherwise craft a DB entry that reads the process environment."""
def fail_on_call(secret_name, *args, **kwargs):
raise AssertionError("team DB models should not resolve env refs")
monkeypatch.setenv("BEDROCK_ASSUME_ROLE_ARN", "arn:aws:iam::123:role/should-not-leak")
def test_ProxyConfig__add_deployment_resolves_env_refs_on_arbitrary_field(monkeypatch):
"""A made-up field name that was never on the removed whitelist still
resolves ``os.environ/`` refs. Pins the "no whitelist" invariant:
the resolver applies to every string field, not a curated list."""
monkeypatch.setenv("SOME_CUSTOM_ENV", "resolved-custom-value")
monkeypatch.setattr(
"litellm.proxy.proxy_server.decrypt_value_helper",
lambda value, key, return_original_value: value,
)
monkeypatch.setattr("litellm.proxy.proxy_server.get_secret", fail_on_call)
fake_router = MagicMock()
fake_router.upsert_deployment = MagicMock(return_value=True)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router)
pc = ProxyConfig()
db_model = SimpleNamespace(
model_id="model-1",
model_name="model_name_team-1_bedrock",
model_info={"id": "model-1", "team_id": "team-1"},
model_name="custom-field-model",
model_info={"id": "model-1"},
litellm_params={
"model": "bedrock/anthropic.claude-v2",
"aws_role_name": "os.environ/BEDROCK_ASSUME_ROLE_ARN",
"model": "openai/gpt-4o-mini",
"some_future_field": "os.environ/SOME_CUSTOM_ENV",
},
blocked=False,
)
@ -1275,7 +1277,7 @@ def test_ProxyConfig__add_deployment_keeps_team_aws_env_refs_literal(monkeypatch
deployment = fake_router.upsert_deployment.call_args.kwargs["deployment"]
assert added == 1
assert deployment.litellm_params.aws_role_name == "os.environ/BEDROCK_ASSUME_ROLE_ARN"
assert deployment.litellm_params.some_future_field == "resolved-custom-value"
# ---------------------------------------------------------------------------
@ -1313,6 +1315,9 @@ def test_ProxyConfig_decrypt_model_list_from_db_returns_decrypted(monkeypatch):
def test_ProxyConfig_decrypt_model_list_from_db_resolves_env_refs_after_db_decrypt(
monkeypatch,
):
"""Path B (feeding /v2/model/info fallback and /model/info fallback)
resolves every ``os.environ/`` field on admin-scoped rows, mirroring
path A. Both paths now share the same universal-resolution shape."""
monkeypatch.setenv("LITELLM_DB_MODEL_API_KEY", "resolved-secret")
monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret")
monkeypatch.setattr(
@ -1341,21 +1346,21 @@ def test_ProxyConfig_decrypt_model_list_from_db_resolves_env_refs_after_db_decry
out = pc.decrypt_model_list_from_db(new_models=[m])
assert out[0]["litellm_params"]["api_key"] == "resolved-secret"
assert out[0]["litellm_params"]["api_base"] == "os.environ/LITELLM_MASTER_KEY"
assert out[0]["litellm_params"]["api_base"] == "master-secret"
def test_ProxyConfig_decrypt_model_list_from_db_keeps_team_env_refs_literal_after_db_decrypt(
def test_ProxyConfig_decrypt_model_list_from_db_resolves_team_env_refs_after_db_decrypt(
monkeypatch,
):
def fail_on_call(secret_name, *args, **kwargs):
raise AssertionError("team DB models should not resolve env refs")
"""Team-scoped rows on path B resolve ``os.environ/`` refs just like
admin rows do. Pairs with
``test_ProxyConfig__add_deployment_resolves_team_env_refs`` on path
A — both paths now agree on the trust model."""
monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret")
monkeypatch.setattr(
"litellm.proxy.proxy_server.decrypt_value_helper",
lambda value, key, return_original_value: "os.environ/LITELLM_MASTER_KEY" if key == "api_key" else value,
)
monkeypatch.setattr("litellm.proxy.proxy_server.get_secret", fail_on_call)
pc = ProxyConfig()
m = SimpleNamespace(
model_id="model-1",
@ -1363,7 +1368,7 @@ def test_ProxyConfig_decrypt_model_list_from_db_keeps_team_env_refs_literal_afte
model_info={"id": "model-1", "team_id": "team-1"},
litellm_params={
"api_key": "encrypted-env-ref",
"api_base": "https://attacker.example",
"api_base": "https://team.example",
"model": "openai/gpt-4o-mini",
},
blocked=False,
@ -1371,8 +1376,8 @@ def test_ProxyConfig_decrypt_model_list_from_db_keeps_team_env_refs_literal_afte
out = pc.decrypt_model_list_from_db(new_models=[m])
assert out[0]["litellm_params"]["api_key"] == "os.environ/LITELLM_MASTER_KEY"
assert out[0]["litellm_params"]["api_base"] == "https://attacker.example"
assert out[0]["litellm_params"]["api_key"] == "master-secret"
assert out[0]["litellm_params"]["api_base"] == "https://team.example"
def test_ProxyConfig_decrypt_model_list_from_db_invalid_params_skips():

View file

@ -1,6 +1,7 @@
"""
Test that litellm.responses() / litellm.aresponses() send the expected request body
over the wire. Expected JSON bodies are stored in expected_responses_api_request/.
over the wire and surface provider errors correctly. Expected JSON bodies are stored
in expected_responses_api_request/.
"""
import json
@ -18,24 +19,20 @@ def _expected_dir() -> Path:
return Path(__file__).resolve().parent.parent / "expected_responses_api_request"
@pytest.mark.asyncio
async def test_aresponses_context_management_and_shell_request_body_matches_expected():
"""
Call litellm.aresponses() with context_management and shell tool;
assert the httpx POST request body matches the expected JSON.
"""
expected_path = _expected_dir() / "context_management_and_shell.json"
def _load_expected_body(filename: str) -> dict:
expected_path = _expected_dir() / filename
assert expected_path.exists(), f"Expected file not found: {expected_path}"
with open(expected_path) as f:
expected_body = json.load(f)
return json.load(f)
# Minimal Responses API response so parsing succeeds
mock_response = {
"id": "resp_ctx_shell_test",
def _minimal_responses_api_payload(response_id: str, model: str) -> dict:
return {
"id": response_id,
"object": "response",
"created_at": 1734366691,
"status": "completed",
"model": "gpt-4o",
"model": model,
"output": [
{
"type": "message",
@ -69,21 +66,41 @@ async def test_aresponses_context_management_and_shell_request_body_matches_expe
"user": None,
}
class MockResponse:
def __init__(self, json_data, status_code=200):
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = httpx.Headers({})
def json(self):
return self._json_data
class MockResponse:
def __init__(self, json_data, status_code=200):
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = httpx.Headers({})
def json(self):
return self._json_data
def _assert_request_body_matches(request_body: dict, expected_body: dict) -> None:
for key, expected_value in expected_body.items():
assert key in request_body, f"Missing key in request body: {key}"
assert (
request_body[key] == expected_value
), f"Mismatch for key {key}: got {request_body[key]!r}, expected {expected_value!r}"
@pytest.mark.asyncio
async def test_aresponses_context_management_and_shell_request_body_matches_expected():
"""
Call litellm.aresponses() with context_management and shell tool;
assert the httpx POST request body matches the expected JSON.
"""
expected_body = _load_expected_body("context_management_and_shell.json")
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = MockResponse(mock_response, 200)
mock_post.return_value = MockResponse(
_minimal_responses_api_payload("resp_ctx_shell_test", "gpt-4o"), 200
)
await litellm.aresponses(
model="openai/gpt-4o",
@ -95,10 +112,87 @@ async def test_aresponses_context_management_and_shell_request_body_matches_expe
)
mock_post.assert_called_once()
request_body = mock_post.call_args.kwargs["json"]
_assert_request_body_matches(mock_post.call_args.kwargs["json"], expected_body)
for key, expected_value in expected_body.items():
assert key in request_body, f"Missing key in request body: {key}"
assert (
request_body[key] == expected_value
), f"Mismatch for key {key}: got {request_body[key]!r}, expected {expected_value!r}"
@pytest.mark.asyncio
async def test_aresponses_azure_shell_tool_request_body_matches_expected():
"""
Call litellm.aresponses() on the Azure route with the shell tool;
assert the httpx POST request body carries the shell tool verbatim.
"""
expected_body = _load_expected_body("azure_shell_tool.json")
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = MockResponse(
_minimal_responses_api_payload("resp_azure_shell_test", "gpt-5-mini"), 200
)
await litellm.aresponses(
model="azure/gpt-5-mini",
api_base="https://fake-resource.openai.azure.com",
api_key="fake-api-key",
api_version="2025-03-01-preview",
input=expected_body["input"],
tools=expected_body["tools"],
tool_choice=expected_body["tool_choice"],
max_output_tokens=expected_body["max_output_tokens"],
)
mock_post.assert_called_once()
_assert_request_body_matches(mock_post.call_args.kwargs["json"], expected_body)
@pytest.mark.asyncio
async def test_aresponses_azure_shell_tool_400_maps_to_bad_request_error():
"""
Azure rejects the shell tool for unsupported deployments with a 400;
litellm must surface that as litellm.BadRequestError carrying the provider message.
"""
error_body = {
"error": {
"message": "Tool of type 'shell' is not supported with this model.",
"type": "invalid_request_error",
"param": "tools",
"code": None,
}
}
def _raise_azure_400(*args, **kwargs):
response = httpx.Response(
status_code=400,
json=error_body,
request=httpx.Request(
"POST",
kwargs.get(
"url",
"https://fake-resource.openai.azure.com/openai/responses",
),
),
)
response.raise_for_status()
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.side_effect = _raise_azure_400
with pytest.raises(litellm.BadRequestError) as excinfo:
await litellm.aresponses(
model="azure/gpt-5-mini",
api_base="https://fake-resource.openai.azure.com",
api_key="fake-api-key",
api_version="2025-03-01-preview",
input="List files in /mnt/data and run python --version.",
tools=[{"type": "shell", "environment": {"type": "container_auto"}}],
tool_choice="auto",
max_output_tokens=256,
)
assert excinfo.value.status_code == 400
assert "shell" in str(excinfo.value).lower()
assert "not supported" in str(excinfo.value).lower()

View file

@ -0,0 +1,59 @@
"""
Test raise_if_unsafe_secret_name, the shared guard applied before secret_name
reaches a secret manager backend.
"""
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path
from litellm.secret_managers.base_secret_manager import raise_if_unsafe_secret_name
@pytest.mark.parametrize(
"secret_name",
[
"..",
"../../../other-app/creds",
"litellm/../../secret",
"foo/../bar",
"foo/..",
"../foo",
"foo\nbar",
"foo\rbar",
"foo\x00bar",
"foo\x7fbar",
"foo\x85bar",
"foo
bar",
"foo
bar",
],
)
def test_raise_if_unsafe_secret_name_rejects_traversal_and_line_breaks(secret_name):
with pytest.raises(ValueError):
raise_if_unsafe_secret_name(secret_name)
@pytest.mark.parametrize(
"secret_name",
[
"plain-alias",
"my-key-123",
"prod/my-service-key",
"team/user@example.com",
"foo: bar",
"foo # bar",
"foo?evil=1",
"foo#bar",
"a" * 500,
"release-1.0..2",
"my..key",
"..foo",
"foo..",
"v2.0..1-beta",
],
)
def test_raise_if_unsafe_secret_name_allows_legitimate_aliases(secret_name):
raise_if_unsafe_secret_name(secret_name)

View file

@ -320,8 +320,9 @@ def test_register_model_strips_none_litellm_provider_from_get_model_info(monkeyp
def test_register_model_inherits_builtin_cache_pricing_for_unmapped_key():
"""Registering a custom override under a key shape that
``get_model_info`` cannot resolve (e.g. a double provider prefix like
``bedrock/bedrock/us.anthropic.claude-sonnet-4-6``) must still inherit
``get_model_info`` cannot resolve (e.g. a triple provider prefix like
``bedrock/bedrock/bedrock/us.anthropic.claude-sonnet-4-6``; a double
prefix now resolves like a routing prefix) must still inherit
the built-in cache pricing for the underlying model.
Before the fix ``register_model`` fell back to an empty ``existing_model``
@ -341,7 +342,7 @@ def test_register_model_inherits_builtin_cache_pricing_for_unmapped_key():
litellm.model_cost = litellm.get_model_cost_map(url="")
builtin_key = "us.anthropic.claude-sonnet-4-6"
registered_key = f"bedrock/bedrock/{builtin_key}"
registered_key = f"bedrock/bedrock/bedrock/{builtin_key}"
builtin = litellm.model_cost[builtin_key]
assert builtin["cache_creation_input_token_cost"] > 0

View file

@ -1063,6 +1063,43 @@ def test_get_model_info_gemini():
assert info.get("rpm") is not None, f"{model} does not have rpm"
def test_get_model_info_bedrock_regional_inference_profile_pricing(local_model_cost_map):
"""Regression LIT-4056: with the bedrock/ routing prefix (plain, converse/, or
invoke/), the exact regional cost-map entry must win over the region-stripped
base entry, matching the unprefixed control form."""
regional = litellm.model_cost["au.anthropic.claude-opus-4-8"]
base = litellm.model_cost["anthropic.claude-opus-4-8"]
assert regional["input_cost_per_token"] > base["input_cost_per_token"]
for model in (
"bedrock/au.anthropic.claude-opus-4-8",
"bedrock/converse/au.anthropic.claude-opus-4-8",
"bedrock/invoke/au.anthropic.claude-opus-4-8",
):
info = litellm.get_model_info(model=model)
assert info["key"] == "au.anthropic.claude-opus-4-8", model
assert info["input_cost_per_token"] == regional["input_cost_per_token"], model
assert info["output_cost_per_token"] == regional["output_cost_per_token"], model
control = litellm.get_model_info(model="au.anthropic.claude-opus-4-8", custom_llm_provider="bedrock")
assert control["key"] == "au.anthropic.claude-opus-4-8"
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 info["key"] == "anthropic.claude-opus-4-8"
def test_get_model_info_bedrock_double_provider_prefix_resolves(local_model_cost_map):
"""A doubled bedrock/ prefix routes at runtime via strip_bedrock_routing_prefix,
so model info must resolve it to the same entry the request actually bills as."""
info = litellm.get_model_info(model="bedrock/bedrock/us.anthropic.claude-sonnet-4-6")
assert info["key"] == "us.anthropic.claude-sonnet-4-6"
def test_openai_models_in_model_info():
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")

View file

@ -6,7 +6,7 @@
"limit": 27522
},
"LIT003": {
"limit": 422
"limit": 292
},
"LIT004": {
"limit": 44

View file

@ -2,5 +2,7 @@
"@typescript-eslint/no-explicit-any": { "max": 2040, "target": 1500 },
"no-console": { "max": 484, "target": 0 },
"complexity": { "max": 140, "target": 80 },
"max-depth": { "max": 70, "target": 30 }
"max-depth": { "max": 70, "target": 30 },
"local/no-large-inline-object-arg": { "max": 560, "target": 300 },
"local/no-long-condition-chain": { "max": 265, "target": 120 }
}

View file

@ -1,6 +1,8 @@
{
"@typescript-eslint/no-explicit-any": 1990,
"@typescript-eslint/no-explicit-any": 1988,
"complexity": 127,
"local/no-large-inline-object-arg": 514,
"local/no-long-condition-chain": 233,
"max-depth": 59,
"no-console": 15
}

File diff suppressed because it is too large Load diff

View file

@ -3,6 +3,7 @@ import tseslint from "typescript-eslint";
import nextCoreWebVitals from "eslint-config-next/core-web-vitals";
import prettier from "eslint-config-prettier/flat";
import unusedImports from "eslint-plugin-unused-imports";
import local from "./scripts/eslint-rules/index.mjs";
const eslintConfig = [
{
@ -13,9 +14,11 @@ const eslintConfig = [
...nextCoreWebVitals,
prettier,
{
plugins: { "unused-imports": unusedImports },
plugins: { "unused-imports": unusedImports, local },
rules: {
"unused-imports/no-unused-imports": "error",
"local/no-large-inline-object-arg": "warn",
"local/no-long-condition-chain": "warn",
"@typescript-eslint/no-explicit-any": "warn",
"no-console": ["warn", { allow: ["warn", "error"] }],
"@typescript-eslint/no-unused-vars": "off",
@ -28,6 +31,7 @@ const eslintConfig = [
"no-useless-escape": "off",
"no-self-assign": "error",
"no-var": "error",
"no-nested-ternary": "error",
"react/no-danger": "error",
complexity: ["warn", 20],
"max-depth": ["warn", 4],

View file

@ -0,0 +1,11 @@
import noLargeInlineObjectArg from "./no-large-inline-object-arg.mjs";
import noLongConditionChain from "./no-long-condition-chain.mjs";
const plugin = {
rules: {
"no-large-inline-object-arg": noLargeInlineObjectArg,
"no-long-condition-chain": noLongConditionChain,
},
};
export default plugin;

View file

@ -0,0 +1,41 @@
const DEFAULT_MIN_PROPERTIES = 4;
const isArgumentOf = (node) => {
const parent = node.parent;
if (parent == null) return false;
if (parent.type !== "CallExpression" && parent.type !== "NewExpression") return false;
return parent.arguments.includes(node);
};
const rule = {
meta: {
type: "suggestion",
docs: {
description:
"Disallow passing a large object literal inline as a call argument; assign it to a named variable first.",
},
schema: [
{
type: "object",
properties: { minProperties: { type: "integer", minimum: 1 } },
additionalProperties: false,
},
],
messages: {
tooLarge:
"Object literal with {{count}} properties passed inline as an argument; assign it to a named variable first.",
},
},
create(context) {
const minProperties = context.options[0]?.minProperties ?? DEFAULT_MIN_PROPERTIES;
return {
ObjectExpression(node) {
if (!isArgumentOf(node)) return;
if (node.properties.length < minProperties) return;
context.report({ node, messageId: "tooLarge", data: { count: node.properties.length } });
},
};
},
};
export default rule;

View file

@ -0,0 +1,41 @@
const DEFAULT_MIN_CONDITIONS = 4;
const isBooleanLogical = (node) =>
node?.type === "LogicalExpression" && (node.operator === "&&" || node.operator === "||");
const countConditions = (node) =>
isBooleanLogical(node) ? countConditions(node.left) + countConditions(node.right) : 1;
const rule = {
meta: {
type: "suggestion",
docs: {
description:
"Disallow logical expressions that combine many conditions; extract the condition into a named boolean.",
},
schema: [
{
type: "object",
properties: { minConditions: { type: "integer", minimum: 2 } },
additionalProperties: false,
},
],
messages: {
tooMany: "Boolean expression combines {{count}} conditions; extract it into a named variable.",
},
},
create(context) {
const minConditions = context.options[0]?.minConditions ?? DEFAULT_MIN_CONDITIONS;
return {
LogicalExpression(node) {
if (!isBooleanLogical(node)) return;
if (isBooleanLogical(node.parent)) return;
const count = countConditions(node);
if (count < minConditions) return;
context.report({ node, messageId: "tooMany", data: { count } });
},
};
},
};
export default rule;

View file

@ -164,6 +164,7 @@ const TopKeyView: React.FC<TopKeyViewProps> = ({ topKeys, teams, showTags = fals
const spendColumn = {
header: "Spend (USD)",
accessorKey: "spend",
meta: { numeric: true },
cell: (info: any) => {
const value = info.getValue();
return value > 0 && value < 0.01 ? "<$0.01" : `$${formatNumberWithCommas(value, 2)}`;
@ -247,13 +248,7 @@ const TopKeyView: React.FC<TopKeyViewProps> = ({ topKeys, teams, showTags = fals
</div>
) : (
<div className="border rounded-lg overflow-hidden max-h-[600px] overflow-y-auto">
<DataTable
columns={columns}
data={topKeys}
renderSubComponent={() => <></>}
getRowCanExpand={() => false}
isLoading={false}
/>
<DataTable columns={columns} data={topKeys} isLoading={false} />
</div>
)}

View file

@ -30,6 +30,7 @@ export default function TopModelView({ topModels, topModelsLimit, setTopModelsLi
{
header: "Spend (USD)",
accessorKey: "spend",
meta: { numeric: true },
cell: (info: any) => {
const value = info.getValue();
return `$${formatNumberWithCommas(value, 2)}`;
@ -38,16 +39,19 @@ export default function TopModelView({ topModels, topModelsLimit, setTopModelsLi
{
header: "Successful",
accessorKey: "successful_requests",
meta: { numeric: true },
cell: (info: any) => <span className="text-green-600">{info.getValue()?.toLocaleString() || 0}</span>,
},
{
header: "Failed",
accessorKey: "failed_requests",
meta: { numeric: true },
cell: (info: any) => <span className="text-red-600">{info.getValue()?.toLocaleString() || 0}</span>,
},
{
header: "Tokens",
accessorKey: "tokens",
meta: { numeric: true },
cell: (info: any) => info.getValue()?.toLocaleString() || 0,
},
];
@ -99,13 +103,7 @@ export default function TopModelView({ topModels, topModelsLimit, setTopModelsLi
</div>
) : (
<div className="border rounded-lg overflow-hidden max-h-[600px] overflow-y-auto">
<DataTable
columns={columns}
data={processedTopModels}
renderSubComponent={() => <></>}
getRowCanExpand={() => false}
isLoading={false}
/>
<DataTable columns={columns} data={processedTopModels} isLoading={false} />
</div>
)}
</>

View file

@ -0,0 +1,103 @@
import { Button, Input, InputNumber } from "antd";
import React from "react";
export interface TagRateLimitEntry {
// Stable identity for React list keys so deleting a middle row doesn't shift
// the controlled inputs of the rows below it.
id: string;
tag: string;
rpm_limit: number | null;
}
let nextRowId = 0;
const newRowId = (): string => `tag-row-${nextRowId++}`;
export interface TagRateLimits {
tag_rpm_limit: Record<string, number>;
}
// Build the rpm limit map from editor rows. A tag only enters the map when its
// name is non-empty and the RPM cell holds a number.
export const tagRowsToLimits = (rows: TagRateLimitEntry[]): TagRateLimits => {
const tag_rpm_limit: Record<string, number> = {};
rows.forEach(({ tag, rpm_limit }) => {
const name = tag.trim();
if (!name) return;
if (typeof rpm_limit === "number") tag_rpm_limit[name] = rpm_limit;
});
return { tag_rpm_limit };
};
// Coerce an untyped metadata value into a {tag: number} map, dropping anything
// that isn't a numeric entry. Key metadata is loosely typed, so validate here.
const toNumberMap = (raw: unknown): Record<string, number> => {
if (!raw || typeof raw !== "object") return {};
const out: Record<string, number> = {};
Object.entries(raw as Record<string, unknown>).forEach(([tag, limit]) => {
if (typeof limit === "number") out[tag] = limit;
});
return out;
};
// Reconstruct editor rows from the stored rpm map.
export const tagLimitsToRows = (tagRpmLimit?: unknown): TagRateLimitEntry[] => {
const rpm = toNumberMap(tagRpmLimit);
return Object.keys(rpm).map((tag) => ({
id: newRowId(),
tag,
rpm_limit: rpm[tag],
}));
};
interface TagRateLimitEditorProps {
value: TagRateLimitEntry[];
onChange: (v: TagRateLimitEntry[]) => void;
}
export function TagRateLimitEditor({ value, onChange }: TagRateLimitEditorProps) {
const addRow = () => {
onChange([...value, { id: newRowId(), tag: "", rpm_limit: null }]);
};
const removeRow = (idx: number) => {
onChange(value.filter((_, i) => i !== idx));
};
const updateRow = (idx: number, field: keyof TagRateLimitEntry, fieldValue: string | number | null) => {
onChange(value.map((row, i) => (i === idx ? { ...row, [field]: fieldValue } : row)));
};
return (
<div>
{value.map((row, idx) => (
<div key={row.id} style={{ display: "flex", gap: 8, alignItems: "center", marginBottom: 12 }}>
<Input
value={row.tag}
onChange={(e) => updateRow(idx, "tag", e.target.value)}
placeholder="Tag (e.g. cell-1)"
style={{ width: 180 }}
/>
<InputNumber
min={0}
value={row.rpm_limit ?? undefined}
onChange={(v) => updateRow(idx, "rpm_limit", v ?? null)}
placeholder="RPM"
style={{ width: 120 }}
/>
<Button type="text" danger size="small" onClick={() => removeRow(idx)} style={{ padding: "0 4px" }}>
✕
</Button>
</div>
))}
<Button
size="small"
onClick={(e) => {
e.preventDefault();
addRow();
}}
>
+ Add Tag Limit
</Button>
</div>
);
}

View file

@ -509,8 +509,6 @@ export function MCPToolsetsTab({ accessToken, userRole }: MCPToolsetsTabProps) {
<DataTable
data={toolsets}
columns={columns}
renderSubComponent={() => <div />}
getRowCanExpand={() => false}
isLoading={isLoading}
noDataMessage="No toolsets yet. Click 'New Toolset' to create one."
loadingMessage="Loading toolsets..."

View file

@ -388,16 +388,58 @@ describe("CreateMCPServer", () => {
expect(screen.queryByText("Subject Token Type (optional)")).not.toBeInTheDocument();
});
it("routes OAuth Token Exchange (OBO) config to the backend payload", async () => {
it("sends max_concurrent_requests in the create payload when set", async () => {
await selectHttpTransport();
const user = userEvent.setup({ delay: null });
const nameInput = getServerNameInput();
await user.type(nameInput, "TE_Server");
await user.type(nameInput, "Limited_Server");
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
await user.type(urlInput, "https://upstream.example.com/mcp");
await user.type(urlInput, "https://example.com/mcp");
await selectAntOption("Authentication", "None");
const limitInput = screen.getByPlaceholderText("e.g. 10");
await user.type(limitInput, "5");
vi.mocked(networking.createMCPServer).mockResolvedValue({
server_id: "new-server-1",
server_name: "Limited_Server",
alias: "Limited_Server",
url: "https://example.com/mcp",
transport: "http",
auth_type: "none",
created_at: "2024-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2024-01-01T00:00:00Z",
updated_by: "user-1",
});
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];
expect(payload.max_concurrent_requests).toBe(5);
});
it("routes OAuth Token Exchange (OBO) config to the backend payload", async () => {
await selectHttpTransport();
// fireEvent.change over user.type: this test asserts payload shape, not
// keystroke behavior, and char-by-char typing re-renders the whole form
// per character, which pushed this test past the 30s CI timeout.
fireEvent.change(getServerNameInput(), { target: { value: "TE_Server" } });
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
fireEvent.change(urlInput, { target: { value: "https://upstream.example.com/mcp" } });
await selectAntOption("Authentication", "OAuth Token Exchange (OBO)");
@ -405,12 +447,15 @@ describe("CreateMCPServer", () => {
expect(screen.getByPlaceholderText("https://idp.example.com/oauth2/token")).toBeInTheDocument();
});
await user.type(
screen.getByPlaceholderText("https://idp.example.com/oauth2/token"),
"https://idp.example.com/oauth2/token",
);
await user.type(screen.getByPlaceholderText("Enter OAuth client ID"), "te-client-id");
await user.type(screen.getByPlaceholderText("Enter OAuth client secret"), "te-client-secret");
fireEvent.change(screen.getByPlaceholderText("https://idp.example.com/oauth2/token"), {
target: { value: "https://idp.example.com/oauth2/token" },
});
fireEvent.change(screen.getByPlaceholderText("Enter OAuth client ID"), {
target: { value: "te-client-id" },
});
fireEvent.change(screen.getByPlaceholderText("Enter OAuth client secret"), {
target: { value: "te-client-secret" },
});
vi.mocked(networking.createMCPServer).mockResolvedValue({
server_id: "new-server-te",
@ -447,10 +492,13 @@ describe("CreateMCPServer", () => {
it("makes scope required when the Entra OBO profile is selected", async () => {
await selectHttpTransport();
const user = userEvent.setup({ delay: null });
await user.type(getServerNameInput(), "Entra_Server");
await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://upstream.example.com/mcp");
// fireEvent.change over user.type for the same reason as the payload
// test above: char-by-char typing re-renders the whole form per
// character and pushes this test toward the 30s CI timeout.
fireEvent.change(getServerNameInput(), { target: { value: "Entra_Server" } });
fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), {
target: { value: "https://upstream.example.com/mcp" },
});
await selectAntOption("Authentication", "OAuth Token Exchange (OBO)");
await waitFor(() => {
@ -459,12 +507,15 @@ describe("CreateMCPServer", () => {
await selectAntOption("Profile", "Microsoft Entra OBO");
await user.type(
screen.getByPlaceholderText("https://idp.example.com/oauth2/token"),
"https://login.microsoftonline.com/tenant/oauth2/v2.0/token",
);
await user.type(screen.getByPlaceholderText("Enter OAuth client ID"), "entra-client");
await user.type(screen.getByPlaceholderText("Enter OAuth client secret"), "entra-secret");
fireEvent.change(screen.getByPlaceholderText("https://idp.example.com/oauth2/token"), {
target: { value: "https://login.microsoftonline.com/tenant/oauth2/v2.0/token" },
});
fireEvent.change(screen.getByPlaceholderText("Enter OAuth client ID"), {
target: { value: "entra-client" },
});
fireEvent.change(screen.getByPlaceholderText("Enter OAuth client secret"), {
target: { value: "entra-secret" },
});
// Selecting Entra OBO makes the scope required; submitting without one is blocked by validation
// (rfc8693 would not require it), which confirms the profile selection took effect.
@ -1067,6 +1118,22 @@ describe("CreateMCPServer", () => {
const reopenedUrlInput = screen.getByPlaceholderText("https://your-mcp-server.com") as HTMLInputElement;
expect(reopenedUrlInput.value).toBe("");
});
it("does not reset an in-flight OAuth resume when mounted with the modal closed (post-redirect restore)", () => {
// After the "Authorize & Fetch Token" redirect the page reloads and this
// component mounts with isModalVisible=false while useMcpOAuthFlow is still
// exchanging the authorization code. Calling reset() during that mount bumps
// the hook's reset version and the fetched token is silently discarded, so
// the user sees no Connection Status / Tool Configuration and must authorize
// again after saving.
const { rerender } = render(<CreateMCPServer {...defaultProps} isModalVisible={false} />);
expect(oauthHook.reset).not.toHaveBeenCalled();
// A real open -> closed transition must still reset (the #30000 leak fix).
rerender(<CreateMCPServer {...defaultProps} isModalVisible={true} />);
rerender(<CreateMCPServer {...defaultProps} isModalVisible={false} />);
expect(oauthHook.reset).toHaveBeenCalled();
});
});
describe("when stdio transport is selected", () => {

View file

@ -1,5 +1,5 @@
import React, { useState } from "react";
import { Modal, Tooltip, Form, Select, Input, Switch, Collapse } from "antd";
import { Modal, Tooltip, Form, Select, Input, InputNumber, Switch, Collapse } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { Button, TextInput } from "@tremor/react";
import { createMCPServer, registerMCPServer, storeMCPOAuthUserCredential } from "../networking";
@ -626,9 +626,15 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
// Clear form, tools, and OAuth state when the modal closes so a previous server's
// authorization, credentials, or tool list never bleed into the next "Add New MCP
// Server" session, including when a parent dismisses the modal without routing
// through handleCancel or handleCreate.
// through handleCancel or handleCreate. Only a real open -> closed transition may
// trigger this: on the post-OAuth-redirect remount the modal starts closed while
// resumeOAuthFlow's token exchange is in flight, and resetting then discards the
// fetched token.
const wasModalVisibleRef = React.useRef(isModalVisible);
React.useEffect(() => {
if (!isModalVisible) {
const wasVisible = wasModalVisibleRef.current;
wasModalVisibleRef.current = isModalVisible;
if (!isModalVisible && wasVisible) {
form.resetFields();
setFormValues({});
setOauthAccessToken(null);
@ -923,6 +929,26 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
</>
)}
<Form.Item
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Max Concurrent Requests (optional)
<Tooltip title="Maximum number of tool calls LiteLLM will run against this server at the same time. Additional calls wait for a free slot. Leave blank for no limit.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
</span>
}
name="max_concurrent_requests"
>
<InputNumber
min={1}
precision={0}
placeholder="e.g. 10"
style={{ width: "100%" }}
className="rounded-lg"
/>
</Form.Item>
{/* Authentication - show for HTTP, SSE, and OpenAPI */}
{transportType !== "stdio" && transportType !== "" && (
<Collapse

View file

@ -1305,3 +1305,83 @@ describe("MCPServerEdit OAuth flow prefill display", () => {
});
});
});
describe("MCPServerEdit (max concurrent requests)", () => {
beforeEach(() => {
vi.clearAllMocks();
});
const limitedServer = {
...interactiveOAuthServer,
auth_type: "none",
max_concurrent_requests: 5,
};
it("prefills the existing limit and sends an updated value in the payload", async () => {
vi.mocked(networking.updateMCPServer).mockResolvedValue({
...limitedServer,
max_concurrent_requests: 2,
});
render(
<MCPServerEdit
mcpServer={limitedServer}
accessToken="access-token"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
const limitInput = screen.getByPlaceholderText("e.g. 10") as HTMLInputElement;
expect(limitInput.value).toBe("5");
fireEvent.change(limitInput, { target: { value: "2" } });
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];
expect(payload.max_concurrent_requests).toBe(2);
});
it("sends null when the limit is cleared so the backend unsets it", async () => {
vi.mocked(networking.updateMCPServer).mockResolvedValue({
...limitedServer,
max_concurrent_requests: null,
});
render(
<MCPServerEdit
mcpServer={limitedServer}
accessToken="access-token"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
const limitInput = screen.getByPlaceholderText("e.g. 10") as HTMLInputElement;
expect(limitInput.value).toBe("5");
fireEvent.change(limitInput, { target: { value: "" } });
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];
expect(payload.max_concurrent_requests).toBeNull();
});
});

View file

@ -852,6 +852,26 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
</Form.Item>
)}
<Form.Item
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Max Concurrent Requests (optional)
<Tooltip title="Maximum number of tool calls LiteLLM will run against this server at the same time. Additional calls wait for a free slot. Leave blank for no limit.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
</span>
}
name="max_concurrent_requests"
>
<InputNumber
min={1}
precision={0}
placeholder="e.g. 10"
style={{ width: "100%" }}
className="rounded-lg"
/>
</Form.Item>
{/* Authentication - for HTTP, SSE, and OpenAPI */}
{!isStdioTransport && (
<Form.Item label="Authentication" name="auth_type" rules={[{ required: true }]}>

View file

@ -261,6 +261,7 @@ export interface MCPServer {
available_on_public_internet?: boolean;
delegate_auth_to_upstream?: boolean;
oauth_passthrough?: boolean;
max_concurrent_requests?: number | null;
/** Stdio-only fields (present when transport === 'stdio') */
command?: string | null;

View file

@ -30,6 +30,7 @@ import ProjectDropdown from "../common_components/ProjectDropdown";
import { CreateUserButton } from "../CreateUserButton";
import { BudgetFallbacksEditor } from "../key_team_helpers/BudgetFallbacksEditor";
import { BudgetWindowEntry, BudgetWindowsEditor } from "../key_team_helpers/BudgetWindowsEditor";
import { TagRateLimitEditor, TagRateLimitEntry, tagRowsToLimits } from "../key_team_helpers/TagRateLimitEditor";
import {
excludeProxyWideSentinel,
getModelDisplayName,
@ -202,6 +203,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
const [rotationInterval, setRotationInterval] = useState<string>("30d");
const [routerSettings, setRouterSettings] = useState<RouterSettingsAccordionValue | null>(null);
const [budgetLimits, setBudgetLimits] = useState<BudgetWindowEntry[]>([]);
const [tagRateLimits, setTagRateLimits] = useState<TagRateLimitEntry[]>([]);
const [budgetFallbacks, setBudgetFallbacks] = useState<Record<string, string[]>>({});
const [budgetFallbacksKey, setBudgetFallbacksKey] = useState<number>(0);
const [routerSettingsKey, setRouterSettingsKey] = useState<number>(0);
@ -223,6 +225,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
setSelectedOrganizationId(null);
setSelectedProjectId(null);
setBudgetLimits([]);
setTagRateLimits([]);
setBudgetFallbacks({});
setBudgetFallbacksKey((k) => k + 1);
};
@ -244,6 +247,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
setSelectedOrganizationId(null);
setSelectedProjectId(null);
setBudgetLimits([]);
setTagRateLimits([]);
setBudgetFallbacks({});
setBudgetFallbacksKey((k) => k + 1);
};
@ -543,6 +547,12 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
formValues.budget_limits = validWindows;
}
// Add per-tag rate limits (only when at least one row is configured)
const { tag_rpm_limit } = tagRowsToLimits(tagRateLimits);
if (Object.keys(tag_rpm_limit).length > 0) {
formValues.tag_rpm_limit = tag_rpm_limit;
}
if (Object.keys(budgetFallbacks).length > 0) {
formValues.budget_fallbacks = budgetFallbacks;
}
@ -567,6 +577,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
NotificationsManager.success("Virtual Key Created");
form.resetFields();
setBudgetLimits([]);
setTagRateLimits([]);
setBudgetFallbacks({});
setBudgetFallbacksKey((k) => k + 1);
localStorage.removeItem("userData" + userID);
@ -1177,6 +1188,19 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
form={form}
showDetailedDescriptions={true}
/>
<Form.Item
className="mt-4"
label={
<span>
Per-Tag Rate Limits{" "}
<Tooltip title="Scope rate limits to a request tag so each tag (e.g. a cell or group) gets its own RPM counter. Requests without a matching tag fall back to the key-level limit.">
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
</Tooltip>
</span>
}
>
<TagRateLimitEditor value={tagRateLimits} onChange={setTagRateLimits} />
</Form.Item>
<Form.Item
className="mt-4"
label={

View file

@ -263,8 +263,6 @@ const PassThroughSettings: React.FC<GeneralSettingsPageProps> = ({
<DataTable
data={generalSettings}
columns={columns}
renderSubComponent={() => <div></div>}
getRowCanExpand={() => false}
isLoading={false}
noDataMessage="No pass-through endpoints configured"
/>

View file

@ -18,6 +18,12 @@ import OrganizationDropdown from "../common_components/OrganizationDropdown";
import { extractLoggingSettings, formatMetadataForDisplay, stripTagsFromMetadata } from "../key_info_utils";
import { BudgetFallbacksEditor } from "../key_team_helpers/BudgetFallbacksEditor";
import { BudgetWindowEntry, BudgetWindowsEditor } from "../key_team_helpers/BudgetWindowsEditor";
import {
TagRateLimitEditor,
TagRateLimitEntry,
tagLimitsToRows,
tagRowsToLimits,
} from "../key_team_helpers/TagRateLimitEditor";
import { excludeProxyWideSentinel, hasAllModelsSentinel } from "../key_team_helpers/fetch_available_models_team_key";
import { KeyResponse } from "../key_team_helpers/key_list";
import MCPServerSelector from "../mcp_server_management/MCPServerSelector";
@ -110,6 +116,9 @@ export function KeyEditView({
const [budgetLimits, setBudgetLimits] = useState<BudgetWindowEntry[]>(
Array.isArray(keyData.budget_limits) ? keyData.budget_limits : [],
);
const [tagRateLimits, setTagRateLimits] = useState<TagRateLimitEntry[]>(
tagLimitsToRows(keyData.metadata?.tag_rpm_limit),
);
const [budgetFallbacks, setBudgetFallbacks] = useState<Record<string, string[]>>(
keyData.budget_fallbacks && typeof keyData.budget_fallbacks === "object" ? keyData.budget_fallbacks : {},
);
@ -311,6 +320,11 @@ export function KeyEditView({
values.budget_limits = [];
}
// Always send the current per-tag limit map so removing every row
// clears the stored limits ({} overwrites the metadata field).
const { tag_rpm_limit } = tagRowsToLimits(tagRateLimits);
values.tag_rpm_limit = tag_rpm_limit;
const hadExistingFallbacks = keyData.budget_fallbacks != null && Object.keys(keyData.budget_fallbacks).length > 0;
if (Object.keys(budgetFallbacks).length > 0) {
values.budget_fallbacks = budgetFallbacks;
@ -553,6 +567,19 @@ export function KeyEditView({
<Input.TextArea rows={4} placeholder='{"gpt-4": 100, "claude-v1": 200}' />
</Form.Item>
<Form.Item
label={
<span>
Per-Tag Rate Limits{" "}
<Tooltip title="Scope rate limits to a request tag so each tag (e.g. a cell or group) gets its own RPM counter. Requests without a matching tag fall back to the key-level limit.">
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
</Tooltip>
</span>
}
>
<TagRateLimitEditor value={tagRateLimits} onChange={setTagRateLimits} />
</Form.Item>
<Form.Item label="Guardrails" name="guardrails">
{accessToken && (
<GuardrailSelector

View file

@ -869,6 +869,13 @@ export default function KeyInfoView({
? JSON.stringify(currentKeyData.metadata.model_rpm_limit)
: "Unlimited"}
</Text>
<Text>
Tag RPM Limits:{" "}
{currentKeyData.metadata?.tag_rpm_limit &&
Object.keys(currentKeyData.metadata.tag_rpm_limit).length > 0
? JSON.stringify(currentKeyData.metadata.tag_rpm_limit)
: "Unlimited"}
</Text>
</div>
<div>

View file

@ -0,0 +1,120 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import { LogDetailsDrawer } from "./LogDetailsDrawer";
import { sessionSpendLogsCall } from "../../networking";
import { LogEntry } from "../columns";
vi.mock("../../networking", () => ({
sessionSpendLogsCall: vi.fn(),
}));
vi.mock("@/app/(dashboard)/hooks/logDetails/useLogDetails", () => ({
useLogDetails: () => ({ data: null, isLoading: false }),
}));
vi.mock("./LogDetailContent", () => ({
LogDetailContent: () => null,
GuardrailJumpLink: () => null,
}));
vi.mock("./DrawerHeader", () => ({
DrawerHeader: () => null,
}));
const makeLog = (overrides: Partial<LogEntry>): LogEntry => ({
request_id: "req",
api_key: "",
team_id: "",
model: "",
model_id: "",
call_type: "acompletion",
spend: 0,
total_tokens: 0,
prompt_tokens: 0,
completion_tokens: 0,
startTime: "2026-07-08T10:00:00.000Z",
endTime: "2026-07-08T10:00:01.000Z",
cache_hit: "false",
messages: [],
response: {},
...overrides,
});
const sessionLogs = [
makeLog({
request_id: "llm-early",
model: "llm-early",
startTime: "2026-07-08T10:00:00.000Z",
endTime: "2026-07-08T10:00:02.000Z",
}),
makeLog({
request_id: "mcp-early",
model: "tool-early",
call_type: "call_mcp_tool",
startTime: "2026-07-08T10:00:01.000Z",
endTime: "2026-07-08T10:00:06.000Z",
}),
makeLog({
request_id: "llm-late",
model: "llm-late",
startTime: "2026-07-08T10:00:02.000Z",
endTime: "2026-07-08T10:00:05.000Z",
}),
makeLog({
request_id: "mcp-late",
model: "tool-late",
call_type: "call_mcp_tool",
startTime: "2026-07-08T10:00:03.000Z",
endTime: "2026-07-08T10:00:03.500Z",
}),
];
const renderSessionDrawer = () => {
vi.mocked(sessionSpendLogsCall).mockResolvedValue({ data: sessionLogs, total: 4, total_pages: 1 });
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
const drawer = (open: boolean) => (
<QueryClientProvider client={queryClient}>
<LogDetailsDrawer open={open} onClose={() => {}} logEntry={null} sessionId="session-1" accessToken="token" />
</QueryClientProvider>
);
const { rerender } = render(drawer(true));
return { rerender, drawer };
};
const sidebarEventNames = () =>
screen.queryAllByText(/^(llm-early|llm-late|tool-early|tool-late)$/).map((el) => el.textContent);
describe("LogDetailsDrawer session sidebar sorting", () => {
it("defaults to duration order, longest call first across LLM and MCP calls", async () => {
renderSessionDrawer();
await waitFor(() => expect(sidebarEventNames()).toHaveLength(4));
expect(sidebarEventNames()).toEqual(["tool-early", "llm-late", "llm-early", "tool-late"]);
});
it("switches to chronological order across LLM and MCP calls when Start time is selected", async () => {
renderSessionDrawer();
await waitFor(() => expect(sidebarEventNames()).toHaveLength(4));
fireEvent.click(screen.getByText("Start time"));
await waitFor(() => expect(sidebarEventNames()).toEqual(["llm-early", "tool-early", "llm-late", "tool-late"]));
fireEvent.click(screen.getByText("Duration"));
await waitFor(() => expect(sidebarEventNames()).toEqual(["tool-early", "llm-late", "llm-early", "tool-late"]));
});
it("resets the sort mode back to duration when the drawer is closed and reopened", async () => {
const { rerender, drawer } = renderSessionDrawer();
await waitFor(() => expect(sidebarEventNames()).toHaveLength(4));
fireEvent.click(screen.getByText("Start time"));
await waitFor(() => expect(sidebarEventNames()).toEqual(["llm-early", "tool-early", "llm-late", "tool-late"]));
rerender(drawer(false));
rerender(drawer(true));
await waitFor(() => expect(sidebarEventNames()).toEqual(["tool-early", "llm-late", "llm-early", "tool-late"]));
});
});

View file

@ -1,5 +1,5 @@
import { useEffect, useMemo, useState } from "react";
import { Button, Drawer } from "antd";
import { Button, Drawer, Segmented } from "antd";
import { CheckOutlined, CopyOutlined, LeftOutlined, RightOutlined } from "@ant-design/icons";
import { Bot, Sparkles, Wrench } from "lucide-react";
import { LogEntry } from "../columns";
@ -11,7 +11,7 @@ import { LogDetailContent, GuardrailJumpLink } from "./LogDetailContent";
import { sessionSpendLogsCall } from "../../networking";
import { useQuery } from "@tanstack/react-query";
import { getSpendString } from "@/utils/dataUtils";
import { normalizeGuardrailEntries } from "./utils";
import { normalizeGuardrailEntries, sortSessionLogs, SessionLogSortMode } from "./utils";
import { DRAWER_WIDTH } from "./constants";
import { useLogDetails } from "@/app/(dashboard)/hooks/logDetails/useLogDetails";
@ -117,6 +117,7 @@ export function LogDetailsDrawer({
}: LogDetailsDrawerProps) {
const isSessionMode = Boolean(sessionId);
const [selectedSessionRequestId, setSelectedSessionRequestId] = useState<string | null>(null);
const [sessionSortMode, setSessionSortMode] = useState<SessionLogSortMode>("duration");
const [isSidebarCollapsed, setIsSidebarCollapsed] = useState(false);
const [copiedLeftPanelId, setCopiedLeftPanelId] = useState(false);
@ -152,35 +153,29 @@ export function LogDetailsDrawer({
// backend omits total, so the truncation note reflects what was fetched.
const total: number = firstPage.total ?? rows.length;
const logs = rows
.map((row) => ({
...row,
request_duration_ms: row.request_duration_ms ?? Date.parse(row.endTime) - Date.parse(row.startTime),
}))
.sort((a, b) => {
const aIsMcp = MCP_CALL_TYPES.includes(a.call_type) ? 1 : 0;
const bIsMcp = MCP_CALL_TYPES.includes(b.call_type) ? 1 : 0;
if (aIsMcp !== bIsMcp) return aIsMcp - bIsMcp;
// Newest first, matching the all-sessions logs overview. MCP calls
// stay grouped last (above), newest-first within that group too.
return new Date(b.startTime).getTime() - new Date(a.startTime).getTime();
});
const logs = rows.map((row) => ({
...row,
request_duration_ms: row.request_duration_ms ?? Date.parse(row.endTime) - Date.parse(row.startTime),
}));
return { logs, total };
},
enabled: Boolean(open && isSessionMode && sessionId && accessToken),
});
const sessionLogs: LogEntry[] = sessionData?.logs ?? [];
const sessionLogs: LogEntry[] = useMemo(
() => sortSessionLogs(sessionData?.logs ?? [], sessionSortMode),
[sessionData, sessionSortMode],
);
// total reported by the backend; when the page cap truncates the fetch this
// exceeds sessionLogs.length, which drives the "showing most recent" note.
const sessionTotalCount = sessionData?.total ?? sessionLogs.length;
const sessionTruncated = sessionTotalCount > sessionLogs.length;
// Default selection for a freshly opened session: the most recent log (latest
// startTime). The list is sorted newest-first, but MCP calls are grouped last,
// so the latest log by time is not necessarily sessionLogs[0]; compute it
// explicitly. A clicked/remembered log still wins over this default.
// startTime). The list is ordered by the selected sort mode, so the latest
// log by time is not necessarily sessionLogs[0]; compute it explicitly.
// A clicked/remembered log still wins over this default.
const mostRecentLog = useMemo<LogEntry | null>(
() =>
sessionLogs.reduce<LogEntry | null>(
@ -222,6 +217,7 @@ export function LogDetailsDrawer({
setIsSidebarCollapsed(false);
} else {
if (isSessionMode) setSelectedSessionRequestId(null);
setSessionSortMode("duration");
setCopiedLeftPanelId(false);
}
}, [open, isSessionMode]);
@ -391,6 +387,19 @@ export function LogDetailsDrawer({
Showing most recent {logsForList.length} of {sessionTotalCount}
</div>
)}
{isSessionMode && (
<Segmented
block
size="small"
className="mt-1.5 [&_.ant-segmented-item-label]:text-[11px]"
options={[
{ label: "Duration", value: "duration" },
{ label: "Start time", value: "start_time" },
]}
value={sessionSortMode}
onChange={(value) => setSessionSortMode(value as SessionLogSortMode)}
/>
)}
</div>
<div className="flex-1 overflow-y-auto">

View file

@ -0,0 +1,44 @@
import { describe, expect, it } from "vitest";
import { sortSessionLogs } from "./utils";
const log = (id: string, startTime: string, endTime: string, request_duration_ms?: number) => ({
request_id: id,
startTime,
endTime,
request_duration_ms,
});
const ids = (rows: { request_id: string }[]) => rows.map((row) => row.request_id);
describe("sortSessionLogs", () => {
const rows = [
log("mid-duration", "2026-07-08T10:00:01.000Z", "2026-07-08T10:00:01.500Z", 2000),
log("longest", "2026-07-08T10:00:02.000Z", "2026-07-08T10:00:02.500Z", 5000),
log("shortest", "2026-07-08T10:00:03.000Z", "2026-07-08T10:00:03.500Z", 300),
log("earliest-no-duration-field", "2026-07-08T10:00:00.000Z", "2026-07-08T10:00:04.000Z"),
];
it("duration mode sorts longest call first, deriving duration from timestamps when the field is missing", () => {
expect(ids(sortSessionLogs(rows, "duration"))).toEqual([
"longest",
"earliest-no-duration-field",
"mid-duration",
"shortest",
]);
});
it("start_time mode sorts calls in the order they started", () => {
expect(ids(sortSessionLogs(rows, "start_time"))).toEqual([
"earliest-no-duration-field",
"mid-duration",
"longest",
"shortest",
]);
});
it("does not mutate the input array", () => {
const input = [...rows];
sortSessionLogs(input, "duration");
expect(ids(input)).toEqual(ids(rows));
});
});

View file

@ -3,6 +3,20 @@
* These functions handle data formatting, validation, and guardrail calculations.
*/
export type SessionLogSortMode = "duration" | "start_time";
type SortableSessionLog = { startTime: string; endTime: string; request_duration_ms?: number };
const durationMs = (row: SortableSessionLog): number =>
row.request_duration_ms ?? Date.parse(row.endTime) - Date.parse(row.startTime);
export function sortSessionLogs<T extends SortableSessionLog>(rows: T[], mode: SessionLogSortMode): T[] {
if (mode === "start_time") {
return [...rows].sort((a, b) => new Date(a.startTime).getTime() - new Date(b.startTime).getTime());
}
return [...rows].sort((a, b) => durationMs(b) - durationMs(a));
}
/**
* Formats data for display. If input is a string, attempts to parse as JSON.
* @param input - Data to format (string or object)

View file

@ -231,13 +231,14 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef<LogEntry>[]
: "Cost",
accessorKey: "spend",
size: 110,
meta: { numeric: true },
cell: (info: any) => {
const row = info.row.original;
const mcpCount = row.mcp_tool_call_count || 0;
const mcpSpend = row.mcp_tool_call_spend || 0;
return (
<div className="flex flex-col">
<div className="flex flex-col items-end">
<Tooltip title={`$${String(info.getValue() || 0)}`}>
<span>{getSpendString(info.getValue() || 0)}</span>
</Tooltip>
@ -263,13 +264,14 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef<LogEntry>[]
)
: "Duration (s)",
accessorKey: "request_duration_ms",
meta: { numeric: true },
cell: (info: any) => {
const ms = info.getValue();
if (ms == null) return <span>-</span>;
const seconds = (ms / 1000).toFixed(2);
return (
<Tooltip title={`${ms}ms`}>
<span className="max-w-[15ch] truncate block">{seconds}</span>
<span className="max-w-[15ch] truncate inline-block">{seconds}</span>
</Tooltip>
);
},
@ -287,6 +289,7 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef<LogEntry>[]
)
: "TTFT (s)",
accessorKey: "completionStartTime",
meta: { numeric: true },
cell: (info: any) => {
const row = info.row.original;
const completionStartTime = info.getValue();
@ -298,7 +301,7 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef<LogEntry>[]
const ttftSeconds = (ttftMs / 1000).toFixed(2);
return (
<Tooltip title={`${ttftMs}ms`}>
<span className="max-w-[15ch] truncate block">{ttftSeconds}</span>
<span className="max-w-[15ch] truncate inline-block">{ttftSeconds}</span>
</Tooltip>
);
},
@ -395,6 +398,7 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef<LogEntry>[]
: "Tokens",
accessorKey: "total_tokens",
size: 140,
meta: { numeric: true },
cell: (info: any) => {
const row = info.row.original;
return (

View file

@ -287,6 +287,7 @@ export default function SpendLogsTable({ accessToken, token, userRole, userID, p
<DataTable
columns={columns}
data={deferredData}
getRowId={(row) => row.request_id}
onRowClick={handleRowClick}
isLoading={isLogsLoading}
/>

Some files were not shown because too many files have changed in this diff Show more