Merge remote-tracking branch 'origin/main' into litellm_spend_log_cleanup_cancel_outcome
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Modules / fmt, validate, test (aws) (push) Has been cancelled
Terraform Modules / fmt, validate, test (gcp) (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

This commit is contained in:
jesus 2026-09-18 21:03:54 +00:00
commit 90582ec122
95 changed files with 11471 additions and 694 deletions

View file

@ -41,4 +41,5 @@ jobs:
"$RUNNER_TEMP/osv-scanner" scan source \
--config osv-scanner.toml \
-L uv.lock \
-L ui/litellm-dashboard/package-lock.json
-L ui/litellm-dashboard/package-lock.json \
-L vscode-extension/package-lock.json

View file

@ -0,0 +1,65 @@
name: VS Code Extension
permissions:
contents: read
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "vscode-extension/**"
- ".github/workflows/test-vscode-extension.yml"
push:
branches:
- main
paths:
- "vscode-extension/**"
- ".github/workflows/test-vscode-extension.yml"
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
vscode-extension:
runs-on: ubuntu-latest
timeout-minutes: 10
defaults:
run:
working-directory: vscode-extension
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 1
persist-credentials: false
- name: Set up Node.js
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version: "24"
cache: npm
cache-dependency-path: vscode-extension/package-lock.json
- name: Install dependencies
run: npm ci
- name: Typecheck
run: npm run typecheck
- name: Unit tests
run: npm test
- name: Package extension
run: npm run package
- name: Upload VSIX
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: litellm-vscode
path: vscode-extension/*.vsix
if-no-files-found: error

View file

@ -183,6 +183,9 @@ MCP_TOOL_LISTING_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_TOOL_LISTING_TIME
MCP_METADATA_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_METADATA_TIMEOUT", "10.0"))
MCP_HEALTH_CHECK_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_HEALTH_CHECK_TIMEOUT", "10.0"))
MCP_TOOL_LISTING_MAX_PAGES: Final = 1000
MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH: Final = 8
MCP_BYOK_CREDENTIAL_CACHE_TTL_SECONDS: Final = 60
MCP_BYOK_CREDENTIAL_CACHE_MAX_SIZE: Final = 4096
# Allowlist of commands permitted for MCP stdio transport.
# Prevents arbitrary command execution via /mcp-rest/test/* endpoints or server creation.
@ -1657,6 +1660,11 @@ LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_INTERVAL_SECONDS: Final = int(
LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE: Final = int(
os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE", 1000)
)
LOGIN_THROTTLE_CACHE_KEY_PREFIX: Final = "login_fail"
LOGIN_THROTTLE_UNKNOWN_SOURCE: Final = "unknown"
LOGIN_THROTTLE_MAX_TRACKED_COUNTERS: Final = 20_000
LOGIN_THROTTLE_MAX_TRACKED_BLOCKS: Final = 10_000
LOGIN_THROTTLE_NOT_BLOCKED: Final = (0, 0)
LITELLM_PROXY_ADMIN_NAME: Final = "default_user_id"
LITELLM_PROXY_BUDGET_NAME: Final = "litellm-proxy-budget"
GLOBAL_PROXY_SPEND_CACHE_KEY: Final = f"{LITELLM_PROXY_ADMIN_NAME}:spend"
@ -2049,6 +2057,7 @@ MCP_SPEND_LOG_MODEL_PREFIX: Final[str] = "MCP: "
PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__"
PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job"
PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900
USAGE_TOP_API_KEYS_LIMIT: Final[int] = int(os.getenv("USAGE_TOP_API_KEYS_LIMIT", "100"))
# Furthest back the catch-up pass looks for unpriced PTU days when a deployment
# declares no ptu_effective_from, bounding the scan for an open-ended window.
PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90

View file

@ -101,7 +101,8 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig):
extra_body: dict[str, object] | None = None,
) -> tuple[str, dict]:
encoded_vector_store_id: Final = encode_url_path_segment(vector_store_id, field_name="vector_store_id")
url: Final = f"{api_base}/{encoded_vector_store_id}/search"
base_url, query_separator, query_string = api_base.partition("?")
url: Final = f"{base_url}/{encoded_vector_store_id}/search{query_separator}{query_string}"
typed_request_body: Final = VectorStoreSearchRequest(
query=query,
filters=vector_store_search_optional_params.get("filters", None),

View file

@ -16690,6 +16690,46 @@
"supports_tool_choice": true,
"supports_vision": true
},
"dashscope/qwen3.8-flash": {
"cache_creation_input_token_cost": 2e-07,
"cache_read_input_token_cost": 1.6e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
"dashscope/qwen3.8-omni-flash": {
"cache_read_input_token_cost": 1.6e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
"dashscope/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "dashscope",
@ -18594,6 +18634,46 @@
"supports_tool_choice": true,
"supports_vision": true
},
"qwen_ai_platform/qwen3.8-flash": {
"cache_creation_input_token_cost": 2e-07,
"cache_read_input_token_cost": 1.6e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "qwen_ai_platform",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
"qwen_ai_platform/qwen3.8-omni-flash": {
"cache_read_input_token_cost": 1.6e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "qwen_ai_platform",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
"qwen_ai_platform/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "qwen_ai_platform",
@ -22090,8 +22170,8 @@
"embed-english-light-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "cohere",
"max_input_tokens": 1024,
"max_tokens": 1024,
"max_input_tokens": 512,
"max_tokens": 512,
"mode": "embedding",
"output_cost_per_token": 0.0
},
@ -22108,8 +22188,8 @@
"input_cost_per_image": 0.0001,
"input_cost_per_token": 1e-07,
"litellm_provider": "cohere",
"max_input_tokens": 1024,
"max_tokens": 1024,
"max_input_tokens": 512,
"max_tokens": 512,
"metadata": {
"notes": "'supports_image_input' is a deprecated field. Use 'supports_embedding_image_input' instead."
},
@ -22130,8 +22210,8 @@
"embed-multilingual-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "cohere",
"max_input_tokens": 1024,
"max_tokens": 1024,
"max_input_tokens": 512,
"max_tokens": 512,
"mode": "embedding",
"output_cost_per_token": 0.0,
"supports_embedding_image_input": true
@ -22139,8 +22219,8 @@
"embed-multilingual-light-v3.0": {
"input_cost_per_token": 0.0001,
"litellm_provider": "cohere",
"max_input_tokens": 1024,
"max_tokens": 1024,
"max_input_tokens": 512,
"max_tokens": 512,
"mode": "embedding",
"output_cost_per_token": 0.0,
"supports_embedding_image_input": true
@ -57910,14 +57990,14 @@
"supports_tool_choice": true
},
"bedrock_mantle/openai.gpt-5.6-sol": {
"input_cost_per_token": 5.5e-06,
"input_cost_per_token_above_272k_tokens": 1.1e-05,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
"cache_read_input_token_cost": 5.5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"input_cost_per_token": 4.4e-06,
"input_cost_per_token_above_272k_tokens": 8.8e-06,
"cache_creation_input_token_cost": 5.5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1.1e-05,
"cache_read_input_token_cost": 4.4e-07,
"cache_read_input_token_cost_above_272k_tokens": 8.8e-07,
"output_cost_per_token": 2.2e-05,
"output_cost_per_token_above_272k_tokens": 3.3e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
@ -65816,9 +65896,9 @@
"supports_web_search": false
},
"openrouter/z-ai/glm-5.3": {
"input_cost_per_token": 1.4e-06,
"output_cost_per_token": 4.4e-06,
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 9.1e-07,
"output_cost_per_token": 2.86e-06,
"cache_read_input_token_cost": 1.69e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
"max_output_tokens": 943717,
@ -70539,14 +70619,14 @@
"supports_web_search": false
},
"openrouter/~deepseek/deepseek-flash-latest": {
"cache_read_input_token_cost": 1.5e-08,
"input_cost_per_token": 1.5e-07,
"cache_read_input_token_cost": 4.2e-09,
"input_cost_per_token": 1.4e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 6e-07,
"output_cost_per_token": 4.2e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -70831,14 +70911,14 @@
"supports_web_search": false
},
"openrouter/~z-ai/glm-latest": {
"cache_read_input_token_cost": 1.5e-07,
"cache_read_input_token_cost": 1.46625e-07,
"input_cost_per_token": 9e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
"max_output_tokens": 235929,
"max_tokens": 235929,
"mode": "chat",
"output_cost_per_token": 3e-06,
"output_cost_per_token": 2.805e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -73982,14 +74062,14 @@
"supports_web_search": false
},
"openrouter/tencent/hy3": {
"cache_read_input_token_cost": 3.3e-08,
"input_cost_per_token": 1.32e-07,
"cache_read_input_token_cost": 2.0625e-08,
"input_cost_per_token": 8.25e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5.28e-07,
"output_cost_per_token": 3.3e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,

View file

@ -1,8 +1,15 @@
import sys
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from pydantic import TypeAdapter
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS: Final = 600.0
_SECONDS: Final = TypeAdapter(float)
_NO_PARAMS: Final[Mapping[str, object]] = MappingProxyType({})
def resolve_pass_through_request_timeout(
endpoint_timeout: float | None = None,
@ -31,26 +38,41 @@ def resolve_pass_through_request_timeout(
def resolve_llm_passthrough_timeout(
kwargs: dict | None = None,
litellm_params: dict | None = None,
router_timeout: float | None = None,
kwargs: Mapping[str, object] | None = None,
litellm_params: Mapping[str, object] | None = None,
router_timeout: float | str | None = None,
router_stream_timeout: float | str | None = None,
) -> float:
"""
Resolve upstream httpx timeout for SDK native passthrough (e.g. Bedrock /converse).
Resolve upstream httpx timeout for SDK native passthrough (e.g. Bedrock /converse,
Anthropic /v1/messages).
Precedence: kwargs timeout/request_timeout -> litellm_params timeout/request_timeout
-> router_timeout -> general_settings.pass_through_request_timeout -> 600s default.
Non-streaming precedence: kwargs timeout/request_timeout -> litellm_params
timeout/request_timeout -> router_timeout -> general_settings.pass_through_request_timeout
-> 600s default.
Streaming (``kwargs["stream"]`` truthy) resolves ``stream_timeout`` at every level before
any generic timeout, matching ``Router._get_stream_timeout`` on the completion route:
kwargs stream_timeout -> litellm_params stream_timeout -> router_stream_timeout, then the
non-streaming chain above.
Only the first set value is validated as seconds, so a value in a lower-precedence
field never fails the call.
"""
kwargs = kwargs or {}
litellm_params = litellm_params or {}
for source in (kwargs, litellm_params):
for key in ("timeout", "request_timeout"):
val = source.get(key)
if val is not None:
return float(val)
if router_timeout is not None:
return float(router_timeout)
return resolve_pass_through_request_timeout()
request: Final = kwargs if kwargs is not None else _NO_PARAMS
deployment: Final = litellm_params if litellm_params is not None else _NO_PARAMS
stream_candidates: Final = (
(request.get("stream_timeout"), deployment.get("stream_timeout"), router_stream_timeout)
if request.get("stream")
else ()
)
candidates: Final = (
*stream_candidates,
request.get("timeout"),
request.get("request_timeout"),
deployment.get("timeout"),
deployment.get("request_timeout"),
router_timeout,
)
winner: Final = next((val for val in candidates if val is not None), None)
return resolve_pass_through_request_timeout() if winner is None else _SECONDS.validate_python(winner)

View file

@ -0,0 +1,38 @@
"""Per-worker cache of stored BYOK credentials, keyed so peer workers can evict it over the auth cache pub/sub."""
from dataclasses import dataclass
from typing import Final
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import MCP_BYOK_CREDENTIAL_CACHE_MAX_SIZE, MCP_BYOK_CREDENTIAL_CACHE_TTL_SECONDS
_CACHE_KEY_PREFIX: Final = "mcp_byok_credential"
@dataclass(frozen=True, slots=True)
class CachedByokCredential:
credential: str | None
byok_credential_cache: Final = InMemoryCache(
max_size_in_memory=MCP_BYOK_CREDENTIAL_CACHE_MAX_SIZE,
default_ttl=MCP_BYOK_CREDENTIAL_CACHE_TTL_SECONDS,
)
def byok_credential_cache_key(user_id: str, server_id: str) -> str:
return f"{_CACHE_KEY_PREFIX}:{user_id}:{server_id}"
def get_cached_byok_credential(user_id: str, server_id: str) -> CachedByokCredential | None:
cached: Final = byok_credential_cache.get_cache( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # InMemoryCache is untyped
byok_credential_cache_key(user_id, server_id)
)
return cached if isinstance(cached, CachedByokCredential) else None
def cache_byok_credential(user_id: str, server_id: str, credential: str | None) -> None:
byok_credential_cache.set_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped
byok_credential_cache_key(user_id, server_id),
CachedByokCredential(credential=credential),
)

View file

@ -865,7 +865,7 @@ async def byok_token(
_invalidate_byok_cred_cache,
)
_invalidate_byok_cred_cache(user_id, server_id)
await _invalidate_byok_cred_cache(user_id, server_id)
except Exception as exc:
verbose_proxy_logger.error(
"byok_token: failed to store user credential for user=%s server=%s: %s",

View file

@ -24,6 +24,7 @@ from litellm.proxy._types import (
MCPApprovalStatus,
MCPEnvVar,
MCPEnvVarScope,
MCPServerUserCredentialListItem,
MCPSubmissionsSummary,
NewMCPServerRequest,
SpecialMCPServerName,
@ -1504,6 +1505,37 @@ async def get_user_oauth_credential(
return _parse_oauth_payload(decoded)
def _server_user_credential_item(
row: "prisma_db_models.LiteLLM_MCPUserCredentials",
) -> MCPServerUserCredentialListItem:
oauth_payload: Final = _decode_oauth_payload(row.credential_b64)
if oauth_payload is None:
return MCPServerUserCredentialListItem(
user_id=row.user_id,
credential_type="byok",
updated_at=row.updated_at.isoformat(),
)
return MCPServerUserCredentialListItem(
user_id=row.user_id,
credential_type="oauth2",
expires_at=oauth_payload.get("expires_at"),
connected_at=oauth_payload.get("connected_at"),
updated_at=row.updated_at.isoformat(),
)
async def list_server_user_credentials(
prisma_client: PrismaClient,
server_id: str,
) -> tuple[MCPServerUserCredentialListItem, ...]:
"""Every user's stored credential for one server, typed but without the secret, for admins."""
rows: Final = await _db_find_user_credential_rows(
prisma_client,
{"server_id": server_id}, # mutable-ok: prisma where-inputs must be plain dicts
)
return tuple(_server_user_credential_item(row) for row in rows)
async def list_user_oauth_credentials(
prisma_client: PrismaClient,
user_id: str,

View file

@ -295,12 +295,15 @@ class MCPPerUserTokenCache:
)
async def delete(self, user_id: str, server_id: str) -> None:
"""Invalidate the cached token (removes from both in-memory and Redis layers)."""
"""Invalidate the cached token in Redis, here, and in every peer worker's in-memory layer."""
try:
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( # noqa: PLC0415 # proxy import cycle
evict_and_broadcast,
)
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415
key: Final = self._cache_key(user_id, server_id)
await user_api_key_cache.async_delete_cache(key)
await evict_and_broadcast((key,), user_api_key_cache)
except Exception as exc:
verbose_logger.debug(
"MCPPerUserTokenCache.delete failed for user=%s server=%s: %s",

View file

@ -28,7 +28,10 @@ from starlette.types import Message, Receive, Scope, Send
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.constants import (
MAXIMUM_TRACEBACK_LINES_TO_LOG,
MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -38,6 +41,12 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
byok_credential_cache,
byok_credential_cache_key,
cache_byok_credential,
get_cached_byok_credential,
)
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
get_request_base_url,
)
@ -82,6 +91,9 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
publish_auth_cache_invalidation,
)
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
get_chain_id_from_headers,
@ -91,6 +103,7 @@ from litellm.types.mcp import (
MCPGatewaySession,
MCPGatewaySessionGroupCount,
MCPGatewaySessionsResponse,
MCPGatewaySessionsTerminateResponse,
MCPSpecVersion,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
@ -102,13 +115,6 @@ if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload
# Short-lived in-memory cache for BYOK credentials.
# Keyed by (user_id, server_id); value is (credential_or_None, monotonic_timestamp).
# Storing the credential value (not just a bool) means _get_byok_credential and
# _check_byok_credential share a single DB round-trip per TTL window.
_byok_cred_cache: Final[dict[tuple[str, str], tuple[str | None, float]]] = {}
_BYOK_CRED_CACHE_TTL: Final = 60 # seconds
_BYOK_CRED_CACHE_MAX_SIZE: Final = 4096 # cap to prevent unbounded growth
_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60
# Upper bound on concurrent stateful sessions a single caller may hold. Each
# `initialize` creates a session that survives until the idle timeout, so
@ -127,20 +133,11 @@ _MCP_TRANSPORT_SPAN_SCOPE_KEY: Final = "litellm_otel_transport_span"
_MCP_DESTINATIONS_SCOPE_KEY: Final = "litellm_otel_request_destinations"
def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None:
"""Remove a (user_id, server_id) entry from the BYOK credential cache.
Call this after storing or deleting a credential so subsequent calls
see the fresh value rather than a stale cached result.
"""
_byok_cred_cache.pop((user_id, server_id), None)
def _write_byok_cred_cache(user_id: str, server_id: str, credential: str | None) -> None:
"""Write a credential value to the cache, evicting all entries if at capacity."""
if len(_byok_cred_cache) >= _BYOK_CRED_CACHE_MAX_SIZE:
_byok_cred_cache.clear()
_byok_cred_cache[(user_id, server_id)] = (credential, time.monotonic())
async def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None:
"""Drop a stored-or-deleted BYOK credential from this worker's cache and from every peer worker's."""
cache_key: Final = byok_credential_cache_key(user_id, server_id)
byok_credential_cache.delete_cache(cache_key)
await publish_auth_cache_invalidation(cache_key=cache_key)
# Check if MCP is available
@ -618,6 +615,7 @@ if MCP_AVAILABLE:
_stateful_session_locks: Final[dict[str, asyncio.Lock]] = {}
_stateful_session_active_request_counts: Final[dict[str, int]] = {}
_stateful_session_client_info: Final[dict[str, Implementation]] = {} # mutable-ok: cleared on session teardown
_admin_terminated_session_ids: Final[dict[str, float]] = {} # mutable-ok: admin-closed id -> last replay
class _TerminableTransport(Protocol):
async def terminate(self) -> None: ...
@ -689,6 +687,7 @@ if MCP_AVAILABLE:
for session_id in list(_stateful_session_auth_context_last_seen):
if session_id not in _stateful_session_auth_contexts:
_remove_stateful_session_tracking(session_id)
_forget_expired_admin_terminated_session_ids(now)
async def _enforce_stateful_session_cap_for_owner(owner: str) -> bool:
"""
@ -2811,35 +2810,28 @@ if MCP_AVAILABLE:
mcp_server: MCPServer,
user_api_key_auth: UserAPIKeyAuth | None,
) -> str | None:
"""Retrieve the stored BYOK credential for a user+server pair.
Uses the shared _byok_cred_cache to avoid a DB round-trip on every
tool call within the TTL window.
"""
"""Retrieve the stored BYOK credential for a user+server pair, served from the worker cache within its TTL."""
if not mcp_server.is_byok:
return None
user_id: Final = (user_api_key_auth.user_id if user_api_key_auth else None) or ""
if not user_id:
return None
cache_key: Final = (user_id, mcp_server.server_id)
cached: Final = _byok_cred_cache.get(cache_key)
cached: Final = get_cached_byok_credential(user_id, mcp_server.server_id)
if cached is not None:
credential, ts = cached
if time.monotonic() - ts < _BYOK_CRED_CACHE_TTL:
return credential
return cached.credential
from litellm.proxy._experimental.mcp_server.db import get_user_credential
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return None
credential = await get_user_credential(
credential: Final = await get_user_credential(
prisma_client=prisma_client,
user_id=user_id,
server_id=mcp_server.server_id,
)
_write_byok_cred_cache(user_id, mcp_server.server_id, credential)
cache_byok_credential(user_id, mcp_server.server_id, credential)
return credential
async def _check_byok_credential(
@ -2868,27 +2860,23 @@ if MCP_AVAILABLE:
headers={"WWW-Authenticate": get_byok_www_authenticate()},
)
# Check shared credential cache before hitting the DB.
cache_key: Final = (user_id, mcp_server.server_id)
cached: Final = _byok_cred_cache.get(cache_key)
cached: Final = get_cached_byok_credential(user_id, mcp_server.server_id)
if cached is not None:
cached_cred, ts = cached
if time.monotonic() - ts < _BYOK_CRED_CACHE_TTL:
if cached_cred is None:
raise HTTPException(
status_code=401,
detail={
"error": "byok_auth_required",
"server_id": mcp_server.server_id,
"server_name": mcp_server.server_name or mcp_server.name,
"message": (
"No stored credential found for this BYOK server. "
"Complete the OAuth authorization flow to provide your API key."
),
},
headers={"WWW-Authenticate": get_byok_www_authenticate()},
)
return
if cached.credential is None:
raise HTTPException(
status_code=401,
detail={
"error": "byok_auth_required",
"server_id": mcp_server.server_id,
"server_name": mcp_server.server_name or mcp_server.name,
"message": (
"No stored credential found for this BYOK server. "
"Complete the OAuth authorization flow to provide your API key."
),
},
headers={"WWW-Authenticate": get_byok_www_authenticate()},
)
return
from litellm.proxy._experimental.mcp_server.db import get_user_credential
from litellm.proxy.proxy_server import prisma_client
@ -2912,7 +2900,7 @@ if MCP_AVAILABLE:
user_id=user_id,
server_id=mcp_server.server_id,
)
_write_byok_cred_cache(user_id, mcp_server.server_id, credential)
cache_byok_credential(user_id, mcp_server.server_id, credential)
if credential is None:
raise HTTPException(
status_code=401,
@ -3850,7 +3838,7 @@ if MCP_AVAILABLE:
client_info: Final = _stateful_session_client_info.get(session_id)
key_auth: Final = auth_user.user_api_key_auth
return MCPGatewaySession(
session_id_prefix=session_id[:8],
session_id_prefix=session_id[:MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH],
client_name=client_info.name if client_info is not None else None,
client_version=client_info.version if client_info is not None else None,
user_id=key_auth.user_id if key_auth is not None else None,
@ -3885,6 +3873,72 @@ if MCP_AVAILABLE:
sessions=sessions,
)
def _session_matches_admin_selector(
session_id: str,
auth_user: MCPAuthenticatedUser,
session_id_prefix: str | None,
user_id: str | None,
) -> bool:
if session_id_prefix is not None and not session_id.startswith(session_id_prefix):
return False
if user_id is None:
return True
key_auth: Final = auth_user.user_api_key_auth
return key_auth is not None and key_auth.user_id == user_id
def _forget_expired_admin_terminated_session_ids(now: float) -> None:
for session_id in [
session_id
for session_id, last_replayed in _admin_terminated_session_ids.items()
if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
]:
del _admin_terminated_session_ids[session_id]
def _is_admin_terminated_session_id(session_id: str, now: float) -> bool:
last_replayed: Final = _admin_terminated_session_ids.get(session_id)
if last_replayed is None:
return False
if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS:
del _admin_terminated_session_ids[session_id]
return False
_admin_terminated_session_ids[session_id] = now
return True
async def terminate_mcp_gateway_sessions(
*,
session_id_prefix: str | None = None,
user_id: str | None = None,
) -> MCPGatewaySessionsTerminateResponse:
"""Force-close every live stateful session on this worker matching the selector.
The transport is terminated (open streams close), all per-session
tracking is dropped, and the id is remembered so a client that keeps
sending it receives 404 and has to ``initialize`` again, which re-runs
admission. Only sessions held by this worker process are affected.
"""
now: Final = time.monotonic()
_forget_expired_admin_terminated_session_ids(now)
server_instances: Final = _stateful_server_instances()
targets: Final = tuple(
(session_id, auth_user)
for session_id, auth_user in tuple(_stateful_session_auth_contexts.items())
if session_id in server_instances
and _session_matches_admin_selector(session_id, auth_user, session_id_prefix, user_id)
)
terminated: Final = tuple(_gateway_session_for(session_id, auth_user, now) for session_id, auth_user in targets)
for session_id, _ in targets:
_admin_terminated_session_ids[session_id] = now
transport = server_instances.pop(session_id, None)
_remove_stateful_session_tracking(session_id)
if transport is not None:
await transport.terminate()
verbose_logger.warning("MCP session '%s' terminated by an administrator.", session_id)
return MCPGatewaySessionsTerminateResponse(
worker_pid=os.getpid(),
terminated_sessions=len(terminated),
sessions=terminated,
)
async def _read_request_body_for_routing(
receive: Receive,
) -> tuple[list[Message], bytes]:
@ -4009,6 +4063,17 @@ if MCP_AVAILABLE:
await success_response(scope, receive, send)
return True
if _is_admin_terminated_session_id(_session_id, time.monotonic()):
terminated_response: Final = JSONResponse(
status_code=404,
content={ # mutable-ok: JSONResponse content must be a plain dict
"error": "Not Found",
"details": "mcp-session-id was terminated by an administrator. Send initialize to start a new session.",
},
)
await terminated_response(scope, receive, send)
return True
# Non-DELETE: strip stale session ID to allow new session creation
verbose_logger.warning(
"MCP session ID '%s' not found in this worker's memory. "

View file

@ -3050,6 +3050,18 @@
},
"DailySpendMetadata": {
"properties": {
"api_key_limit": {
"anyOf": [
{
"type": "integer"
},
{
"type": "null"
}
],
"description": "When set, api_keys and every api_key_breakdown list at most this many keys, ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.",
"title": "Api Key Limit"
},
"has_more": {
"default": false,
"title": "Has More",
@ -3060,6 +3072,18 @@
"title": "Page",
"type": "integer"
},
"total_api_keys": {
"anyOf": [
{
"type": "integer"
},
{
"type": "null"
}
],
"description": "Distinct API keys matching the filters. When this exceeds api_key_limit, the per-key lists are truncated to the highest-spend keys.",
"title": "Total Api Keys"
},
"total_api_requests": {
"default": 0,
"title": "Total Api Requests",
@ -27965,6 +27989,32 @@
"title": "MCPGatewaySessionsResponse",
"type": "object"
},
"MCPGatewaySessionsTerminateResponse": {
"description": "Stateful sessions an administrator force-closed on this proxy worker.",
"properties": {
"sessions": {
"items": {
"$ref": "#/components/schemas/MCPGatewaySession"
},
"title": "Sessions",
"type": "array"
},
"terminated_sessions": {
"title": "Terminated Sessions",
"type": "integer"
},
"worker_pid": {
"title": "Worker Pid",
"type": "integer"
}
},
"required": [
"worker_pid",
"terminated_sessions"
],
"title": "MCPGatewaySessionsTerminateResponse",
"type": "object"
},
"MCPOAuthUserCredentialRequest": {
"description": "Stores a user's OAuth2 token for an OpenAPI MCP server.",
"properties": {
@ -28061,6 +28111,56 @@
"title": "MCPOAuthUserCredentialStatus",
"type": "object"
},
"MCPServerUserCredentialListItem": {
"description": "One user's stored credential for an MCP server, as an admin sees it. Never carries the secret.",
"properties": {
"connected_at": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Connected At"
},
"credential_type": {
"enum": [
"oauth2",
"byok"
],
"title": "Credential Type",
"type": "string"
},
"expires_at": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Expires At"
},
"updated_at": {
"title": "Updated At",
"type": "string"
},
"user_id": {
"title": "User Id",
"type": "string"
}
},
"required": [
"user_id",
"credential_type",
"updated_at"
],
"title": "MCPServerUserCredentialListItem",
"type": "object"
},
"MCPSubmissionsSummary": {
"properties": {
"active": {
@ -30237,7 +30337,7 @@
},
"/v1/mcp/server/{server_id}/oauth-user-credential": {
"delete": {
"description": "Revoke the calling user's stored OAuth2 token for an MCP server",
"description": "Revoke the calling user's stored OAuth2 token for an MCP server. A proxy admin may pass user_id to revoke another user's stored token.",
"operationId": "delete_mcp_oauth_user_credential_v1_mcp_server__server_id__oauth_user_credential_delete",
"parameters": [
{
@ -30248,6 +30348,23 @@
"title": "Server Id",
"type": "string"
}
},
{
"in": "query",
"name": "user_id",
"required": false,
"schema": {
"anyOf": [
{
"minLength": 1,
"type": "string"
},
{
"type": "null"
}
],
"title": "User Id"
}
}
],
"responses": {
@ -30447,7 +30564,7 @@
},
"/v1/mcp/server/{server_id}/user-credential": {
"delete": {
"description": "Delete the calling user's stored API key for a BYOK MCP server",
"description": "Delete the calling user's stored API key for a BYOK MCP server. A proxy admin may pass user_id to revoke another user's stored key.",
"operationId": "delete_mcp_user_credential_v1_mcp_server__server_id__user_credential_delete",
"parameters": [
{
@ -30458,6 +30575,23 @@
"title": "Server Id",
"type": "string"
}
},
{
"in": "query",
"name": "user_id",
"required": false,
"schema": {
"anyOf": [
{
"minLength": 1,
"type": "string"
},
{
"type": "null"
}
],
"title": "User Id"
}
}
],
"responses": {
@ -30549,6 +30683,58 @@
]
}
},
"/v1/mcp/server/{server_id}/user-credentials": {
"get": {
"description": "List every user's stored BYOK or OAuth2 credential for an MCP server (admin only, no secrets)",
"operationId": "list_mcp_server_user_credentials_v1_mcp_server__server_id__user_credentials_get",
"parameters": [
{
"in": "path",
"name": "server_id",
"required": true,
"schema": {
"title": "Server Id",
"type": "string"
}
}
],
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"items": {
"$ref": "#/components/schemas/MCPServerUserCredentialListItem"
},
"title": "Response List Mcp Server User Credentials V1 Mcp Server Server Id User Credentials Get",
"type": "array"
}
}
},
"description": "Successful Response"
},
"422": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
},
"description": "Validation Error"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "List Mcp Server User Credentials",
"tags": [
"mcp_management"
]
}
},
"/v1/mcp/server/{server_id}/user-env-vars": {
"delete": {
"description": "Clear the calling user's per-user MCP env var values for this server.",
@ -30700,6 +30886,77 @@
}
},
"/v1/mcp/sessions": {
"delete": {
"description": "Force-close live stateful MCP gateway sessions on this proxy worker, selected by session id prefix and/or by the LiteLLM user that opened them (proxy admin only).",
"operationId": "delete_mcp_gateway_sessions_v1_mcp_sessions_delete",
"parameters": [
{
"in": "query",
"name": "session_id_prefix",
"required": false,
"schema": {
"anyOf": [
{
"minLength": 8,
"type": "string"
},
{
"type": "null"
}
],
"title": "Session Id Prefix"
}
},
{
"in": "query",
"name": "user_id",
"required": false,
"schema": {
"anyOf": [
{
"minLength": 1,
"type": "string"
},
{
"type": "null"
}
],
"title": "User Id"
}
}
],
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/MCPGatewaySessionsTerminateResponse"
}
}
},
"description": "Successful Response"
},
"422": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
},
"description": "Validation Error"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Delete Mcp Gateway Sessions",
"tags": [
"mcp_management"
]
},
"get": {
"description": "Live stateful MCP gateway sessions on this proxy worker, grouped by AI client and by user.",
"operationId": "get_mcp_gateway_sessions_v1_mcp_sessions_get",
@ -38542,6 +38799,7 @@
"type": "object"
},
"SCIMMultiValuedAttribute": {
"additionalProperties": true,
"properties": {
"display": {
"anyOf": [
@ -38577,13 +38835,17 @@
"title": "Type"
},
"value": {
"title": "Value",
"type": "string"
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Value"
}
},
"required": [
"value"
],
"title": "SCIMMultiValuedAttribute",
"type": "object"
},

View file

@ -1726,6 +1726,16 @@ class MCPUserCredentialListItem(LiteLLMPydanticObjectBase):
connected_at: str | None = None # ISO-8601
class MCPServerUserCredentialListItem(LiteLLMPydanticObjectBase):
"""One user's stored credential for an MCP server, as an admin sees it. Never carries the secret."""
user_id: str
credential_type: Literal["oauth2", "byok"]
expires_at: str | None = None
connected_at: str | None = None
updated_at: str
class MCPUserEnvVarsRequest(LiteLLMPydanticObjectBase):
"""Payload for storing the calling user's per-user env var values."""
@ -2768,6 +2778,25 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
description="sends alerts if requests hang for 5min+",
)
ui_access_mode: Literal["admin_only", "all"] | None = Field("all", description="Control access to the Proxy UI")
max_failed_login_attempts_per_source: int | None = Field(
None,
ge=1,
description="Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded down but at least 1, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10",
)
max_failed_login_attempts_per_source_overrides: dict[str, int] | None = Field(
None,
description="Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins (between equivalent keys such as '1.2.3.4' and '1.2.3.4/32', an exemption wins, then the higher limit), and the per-username allowance for that address follows as half the override. A value of 0 exempts the address from both limits. Set under `general_settings` in config.yaml",
)
failed_login_window_seconds: int | None = Field(
None,
ge=1,
description="Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Set under `general_settings` in config.yaml. Defaults to 60",
)
failed_login_block_seconds: int | None = Field(
None,
ge=1,
description="How long a blocked source address, or source address and username, stays blocked. Every attempt from a blocked key, right or wrong, is refused with 429 before the password is checked; the block is not extended by refused attempts. Set under `general_settings` in config.yaml. Defaults to 300",
)
allowed_routes: list | None = Field(None, description="Proxy API Endpoints you want users to be able to access")
reject_clientside_metadata_tags: bool | None = Field(
None,
@ -2881,7 +2910,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
)
trusted_proxy_ranges: list[str] | None = Field(
None,
description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler.",
description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, or containing an entry that is not an address or CIDR range, the per-source sign-in limit is off.",
)
store_model_in_db: bool | None = Field(
None,

View file

@ -0,0 +1,445 @@
"""Failed-login accounting for the Admin UI sign-in path.
Wrong passwords are counted over a short window per source address and per source-and-username
pair; too many in one window blocks that key for a fixed time. While a key is blocked every attempt
from it, right or wrong, is refused with 429 before the password is checked. A blocked pair stops
counting against its source, so one script stuck on one account does not block the whole office.
Recovery is the master key over the API, which never passes through here, or waiting out the block.
"""
from __future__ import annotations
import asyncio
import hashlib
import ipaddress
import math
import time
from collections.abc import Mapping
from dataclasses import dataclass
from functools import cache
from typing import Final, Literal, NamedTuple, Protocol, TypeAlias
from fastapi import Request, status
from pydantic import TypeAdapter, ValidationError
from redis.exceptions import RedisError
from litellm._logging import verbose_proxy_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache, RedisCircuitBreakerOpenError
from litellm.constants import (
EMPTY_MAPPING,
LOGIN_THROTTLE_CACHE_KEY_PREFIX,
LOGIN_THROTTLE_MAX_TRACKED_BLOCKS,
LOGIN_THROTTLE_MAX_TRACKED_COUNTERS,
LOGIN_THROTTLE_NOT_BLOCKED,
LOGIN_THROTTLE_UNKNOWN_SOURCE,
)
from litellm.proxy._types import ProxyErrorTypes, ProxyException
from litellm.proxy.auth.network import TrustedProxyConfig, resolve_client_ip
from litellm.secret_managers.main import get_secret_bool
DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE: Final = 10
DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 60
DEFAULT_FAILED_LOGIN_BLOCK_SECONDS: Final = 300
IPV6_SOURCE_PREFIX_LENGTH: Final = 64
EXEMPT: Final = 0
SOURCE_LIMIT_KEY: Final = "max_failed_login_attempts_per_source"
SOURCE_LIMIT_OVERRIDES_KEY: Final = "max_failed_login_attempts_per_source_overrides"
WINDOW_KEY: Final = "failed_login_window_seconds"
BLOCK_KEY: Final = "failed_login_block_seconds"
TRUSTED_PROXY_RANGES_KEY: Final = "trusted_proxy_ranges"
_REDIS_FAILURES: Final = (RedisError, RedisCircuitBreakerOpenError, OSError, asyncio.TimeoutError)
_LOCAL_BLOCK_EXPIRY: Final = TypeAdapter[float | None](float | None)
_SOURCE_LIMIT_OVERRIDES: Final = TypeAdapter[Mapping[str, object]](Mapping[str, object])
_RANGE_ENTRIES: Final = TypeAdapter[tuple[object, ...]](tuple[object, ...])
Scope: TypeAlias = Literal["user", "source"]
_BlockTtls: TypeAlias = tuple[int, int]
_LUA_BLOCK_TTLS: Final = TypeAdapter[_BlockTtls](_BlockTtls)
_Network: TypeAlias = ipaddress.IPv4Network | ipaddress.IPv6Network
class LocalStore(Protocol):
"""The per-worker store behind the counters and blocks; ``InMemoryCache`` satisfies it."""
def get_cache(self, key: str) -> object: ...
def set_cache(self, key: str, value: float, *, ttl: int) -> None: ...
def increment_cache(self, key: str, value: float, *, ttl: int) -> float: ...
def delete_cache(self, key: str) -> None: ...
# KEYS: pair counter, pair block, source counter, source block (one cluster slot via the source hash tag)
# ARGV: pair limit, source limit (0 = source scope off), window seconds, block seconds
# Both scripts return {pair block TTL, source block TTL}; 0 or below means not blocked
_BLOCK_TTLS_LUA: Final = "return {redis.call('TTL', KEYS[2]), redis.call('TTL', KEYS[4])}"
_RECORD_FAILURE_LUA: Final = (
"local function bump(count_key, block_key, limit) "
"local blocked = redis.call('TTL', block_key) "
"if blocked > 0 then return blocked end "
"local count = redis.call('INCR', count_key) "
"if redis.call('TTL', count_key) < 0 then redis.call('EXPIRE', count_key, ARGV[3]) end "
"if count > limit then redis.call('SET', block_key, '1', 'EX', ARGV[4]) return tonumber(ARGV[4]) end "
"return 0 end "
"local user_block = bump(KEYS[1], KEYS[2], tonumber(ARGV[1])) "
"local source_block = 0 "
"if tonumber(ARGV[2]) > 0 and user_block == 0 then "
"source_block = bump(KEYS[3], KEYS[4], tonumber(ARGV[2])) end "
"return {user_block, source_block}"
)
_COUNTERS: Final = InMemoryCache(
max_size_in_memory=LOGIN_THROTTLE_MAX_TRACKED_COUNTERS, default_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS
)
_BLOCKS: Final = InMemoryCache(
max_size_in_memory=LOGIN_THROTTLE_MAX_TRACKED_BLOCKS, default_ttl=DEFAULT_FAILED_LOGIN_BLOCK_SECONDS
)
@cache
def _rate_limit_disabled() -> bool:
return get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT", default_value=False) is True
@cache
def warn_login_counters_are_per_worker(num_workers: str) -> None:
verbose_proxy_logger.warning(
"Running %s workers but Redis is not configured. Failed Admin UI sign-in attempts are counted "
"per worker, so the effective limits are %s times the configured values. Configure Redis "
"to share one count across workers.",
num_workers,
num_workers,
)
@cache
def warn_source_login_limit_is_off() -> None:
verbose_proxy_logger.warning(
"%s is not set or not a valid list of ranges, so failed Admin UI sign-in attempts are limited per "
"source address and username only. Set it to the address ranges of the proxies in front of LiteLLM, "
"or to an empty list when clients connect directly, to also limit each source address across usernames.",
TRUSTED_PROXY_RANGES_KEY,
)
def declared_proxy_ranges(settings: Mapping[str, object]) -> tuple[str, ...] | None:
"""What the operator says fronts LiteLLM: the proxy ranges, an empty tuple for none, None when unsaid.
Only a declared topology makes the source address trustworthy enough to limit across usernames.
An unset key, a value that is not a list of ranges, or a list with an entry that is not an address
or range leaves it unknown and the source scope off.
"""
entries: Final = _configured_range_entries(settings.get(TRUSTED_PROXY_RANGES_KEY))
if entries is None or any(_parse_network(entry, TRUSTED_PROXY_RANGES_KEY) is None for entry in entries):
return None
return entries
def _configured_range_entries(raw_ranges: object) -> tuple[str, ...] | None:
"""Every configured entry, blanks included, so a stray empty string fails validation like any other typo."""
if raw_ranges is None:
return None
if isinstance(raw_ranges, str):
return tuple(part.strip() for part in raw_ranges.split(","))
try:
return tuple(str(entry).strip() for entry in _RANGE_ENTRIES.validate_python(raw_ranges))
except ValidationError:
verbose_proxy_logger.warning(
"Invalid %s value: expected a list of address ranges, got %s",
TRUSTED_PROXY_RANGES_KEY,
type(raw_ranges).__name__,
)
return None
def _positive_int(raw: object, key: str, default: int) -> int:
if raw is None:
return default
try:
value: Final = int(str(raw))
except (TypeError, ValueError):
verbose_proxy_logger.warning("Invalid %s value %r; using %s", key, raw, default)
return default
if value < 1:
verbose_proxy_logger.warning("Invalid %s value %s (must be >= 1); using %s", key, value, default)
return default
return value
def _int_setting(settings: Mapping[str, object], key: str, default: int) -> int:
return _positive_int(settings.get(key), key, default)
def _override_limit(raw: object, default: int) -> int:
"""A per-address override: a limit of 1 or more, or ``EXEMPT`` (0) to leave that address unlimited."""
if str(raw).strip() == str(EXEMPT):
return EXEMPT
return _positive_int(raw, SOURCE_LIMIT_OVERRIDES_KEY, default)
def _parse_address(client_ip: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address | None:
"""The address as it is limited and counted: an IPv4-mapped IPv6 address is its IPv4 address."""
try:
address: Final = ipaddress.ip_address(client_ip)
except ValueError:
return None
if isinstance(address, ipaddress.IPv6Address) and address.ipv4_mapped is not None:
return address.ipv4_mapped
return address
def _parse_network(raw_range: str, setting_name: str = SOURCE_LIMIT_OVERRIDES_KEY) -> _Network | None:
try:
return ipaddress.ip_network(raw_range.strip(), strict=False)
except ValueError:
verbose_proxy_logger.warning("Invalid address or range %r in %s; skipping", raw_range, setting_name)
return None
def _precedence(network: _Network, limit: int) -> tuple[int, bool, int]:
"""Sort key for competing overrides: the longest prefix wins, then an exemption, then the higher limit."""
return (network.prefixlen, limit == EXEMPT, limit)
def _source_limit(settings: Mapping[str, object], client_ip: str) -> int:
"""Failure allowance for this address: the most specific configured range containing it, else the default.
``EXEMPT`` (0) means the operator opted this address out of both limits. Between equivalent keys such as
``1.2.3.4`` and ``1.2.3.4/32`` an exemption wins, then the higher limit.
"""
default: Final = _int_setting(settings, SOURCE_LIMIT_KEY, DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE)
raw_overrides: Final = settings.get(SOURCE_LIMIT_OVERRIDES_KEY)
if raw_overrides is None:
return default
try:
overrides: Final = _SOURCE_LIMIT_OVERRIDES.validate_python(raw_overrides)
except ValidationError:
verbose_proxy_logger.warning(
"Invalid %s value; expected a mapping of address or range to limit", SOURCE_LIMIT_OVERRIDES_KEY
)
return default
address: Final = _parse_address(client_ip)
if address is None:
return default
matches: Final = sorted(
_precedence(network, _override_limit(raw_limit, default))
for raw_range, raw_limit in overrides.items()
if (network := _parse_network(raw_range)) is not None and address in network
)
return matches[-1][-1] if matches else default
def user_limit_for(source_limit: int) -> int:
"""Failures allowed for one username from one address: half the address allowance, rounded down, at least 1."""
return max(source_limit // 2, 1)
def source_group(client_ip: str) -> str:
"""The bucket an address is counted in: IPv4 as is, IPv6 by its /64, so one prefix holder cannot rotate."""
address: Final = _parse_address(client_ip)
if address is None:
return client_ip
if isinstance(address, ipaddress.IPv6Address):
return str(ipaddress.ip_network((address, IPV6_SOURCE_PREFIX_LENGTH), strict=False))
return str(address)
class _Keys(NamedTuple):
pair_counter: str
pair_block: str
source_counter: str
source_block: str
@dataclass(frozen=True, slots=True)
class Block:
scope: Scope
retry_after: int
@dataclass(frozen=True, slots=True)
class LoginThrottle:
"""Failed-login limits for one request's source address.
``source_limit`` is None when the source scope is off: ``trusted_proxy_ranges`` is unset, so the peer
address may be a shared ingress. An empty list means clients connect directly and the peer is the source.
``user_limit`` is derived from the address allowance either way, see ``user_limit_for``. An address whose
override is ``EXEMPT`` gets a disabled throttle: nothing is counted or blocked for it.
"""
client_ip: str
source_limit: int | None
user_limit: int
window_seconds: int
block_seconds: int
counters: LocalStore
blocks: LocalStore
redis_cache: RedisCache | None = None
enabled: bool = True
@classmethod
def from_request(
cls,
request: Request,
general_settings: Mapping[str, object] | None,
redis_cache: RedisCache | None,
) -> LoginThrottle:
settings: Final[Mapping[str, object]] = general_settings if general_settings is not None else EMPTY_MAPPING
proxies: Final = declared_proxy_ranges(settings)
resolved, _ = resolve_client_ip(
request, TrustedProxyConfig(use_forwarded_for=bool(proxies), trusted_proxy_cidrs=proxies or ())
)
source_limit: Final = _source_limit(settings, resolved or LOGIN_THROTTLE_UNKNOWN_SOURCE)
exempt: Final = source_limit == EXEMPT
return cls(
client_ip=resolved or LOGIN_THROTTLE_UNKNOWN_SOURCE,
source_limit=source_limit if proxies is not None and resolved is not None and not exempt else None,
user_limit=user_limit_for(source_limit),
window_seconds=_int_setting(settings, WINDOW_KEY, DEFAULT_FAILED_LOGIN_WINDOW_SECONDS),
block_seconds=_int_setting(settings, BLOCK_KEY, DEFAULT_FAILED_LOGIN_BLOCK_SECONDS),
counters=_COUNTERS,
blocks=_BLOCKS,
redis_cache=redis_cache,
enabled=not exempt and not _rate_limit_disabled(),
)
def _keys(self, username: str) -> _Keys:
group: Final = source_group(self.client_ip)
user: Final = hashlib.sha256(username.casefold().encode("utf-8")).hexdigest()
return _Keys(
pair_counter=f"{LOGIN_THROTTLE_CACHE_KEY_PREFIX}:{{{group}}}:user:{user}",
pair_block=f"{LOGIN_THROTTLE_CACHE_KEY_PREFIX}:{{{group}}}:block:user:{user}",
source_counter=f"{LOGIN_THROTTLE_CACHE_KEY_PREFIX}:{{{group}}}:source",
source_block=f"{LOGIN_THROTTLE_CACHE_KEY_PREFIX}:{{{group}}}:block:source",
)
async def attempt(self, username: str) -> LoginAttempt:
"""Refuses a blocked key before any credential is looked at; otherwise hands back the attempt to settle."""
if not self.enabled:
return LoginAttempt(throttle=self, username=username)
block: Final = await self._active_block(self._keys(username))
if block is None:
return LoginAttempt(throttle=self, username=username)
verbose_proxy_logger.warning(
"Admin UI sign-in refused: the %s is blocked for %s more seconds; username=%r source=%s",
block.scope,
block.retry_after,
username,
self.client_ip,
)
raise self.refused(block.retry_after)
async def _active_block(self, keys: _Keys) -> Block | None:
local: Final = self._local_block_ttls(keys)
shared: Final = await self._shared_block_ttls(keys)
user_ttl: Final = max(local[0], shared[0])
source_ttl: Final = max(local[1], shared[1])
if self.source_limit is not None and source_ttl > 0:
return Block(scope="source", retry_after=source_ttl)
if user_ttl > 0:
return Block(scope="user", retry_after=user_ttl)
return None
async def _shared_block_ttls(self, keys: _Keys) -> _BlockTtls:
if self.redis_cache is None:
return LOGIN_THROTTLE_NOT_BLOCKED
try:
return _LUA_BLOCK_TTLS.validate_python(
await self.redis_cache.async_register_script(_BLOCK_TTLS_LUA)(keys, ())
)
except _REDIS_FAILURES as err:
self._warn_redis(err)
return LOGIN_THROTTLE_NOT_BLOCKED
def _local_block_ttls(self, keys: _Keys) -> _BlockTtls:
return self._local_block_ttl(keys.pair_block), self._local_block_ttl(keys.source_block)
def _local_block_ttl(self, block_key: str) -> int:
expires_at: Final = _LOCAL_BLOCK_EXPIRY.validate_python(self.blocks.get_cache(block_key))
if expires_at is None:
return 0
return max(math.ceil(expires_at - time.time()), 0)
async def record_failure(self, username: str) -> _BlockTtls:
keys: Final = self._keys(username)
source_limit: Final = self.source_limit or 0
if self.redis_cache is not None:
try:
return _LUA_BLOCK_TTLS.validate_python(
await self.redis_cache.async_register_script(_RECORD_FAILURE_LUA)(
keys, (self.user_limit, source_limit, self.window_seconds, self.block_seconds)
)
)
except _REDIS_FAILURES as err:
self._warn_redis(err)
user_block: Final = self._local_bump(keys.pair_counter, keys.pair_block, self.user_limit)
if source_limit == 0 or user_block > 0:
return user_block, 0
return user_block, self._local_bump(keys.source_counter, keys.source_block, source_limit)
def _local_bump(self, count_key: str, block_key: str, limit: int) -> int:
blocked: Final = self._local_block_ttl(block_key)
if blocked > 0:
return blocked
count: Final = int(self.counters.increment_cache(count_key, 1, ttl=self.window_seconds))
if count <= limit:
return 0
self.blocks.set_cache(block_key, time.time() + self.block_seconds, ttl=self.block_seconds)
return self.block_seconds
async def clear_pair(self, username: str) -> None:
pair_counter: Final = self._keys(username).pair_counter
if self.redis_cache is not None:
try:
await self.redis_cache.async_delete_cache(pair_counter)
except _REDIS_FAILURES as err:
self._warn_redis(err)
self.counters.delete_cache(pair_counter)
def _warn_redis(self, err: Exception) -> None:
verbose_proxy_logger.warning(
"Redis failed while counting Admin UI sign-in attempts; using this worker's own counters "
"until it recovers: %s",
err,
)
@staticmethod
def refused(retry_after: int) -> ProxyException:
return ProxyException(
message="Too many failed sign-in attempts. Try again later.",
type=ProxyErrorTypes.auth_error,
param="username",
code=status.HTTP_429_TOO_MANY_REQUESTS,
headers={"Retry-After": str(retry_after)}, # mutable-ok: ProxyException writes into its headers dict
)
@dataclass(frozen=True, slots=True)
class LoginAttempt:
throttle: LoginThrottle
username: str
async def succeeded(self) -> None:
if not self.throttle.enabled:
return
await self.throttle.clear_pair(self.username)
async def failed(self) -> None:
if not self.throttle.enabled:
return
user_block, source_block = await self.throttle.record_failure(self.username)
if user_block == 0 and source_block == 0:
return
verbose_proxy_logger.warning(
"Admin UI sign-in blocked for %s seconds after too many failures; scope=%s username=%r source=%s",
user_block or source_block,
"user" if user_block else "source",
self.username,
self.throttle.client_ip,
)

View file

@ -27,6 +27,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured
from litellm.proxy.auth.login_throttle import LoginAttempt, LoginThrottle
from litellm.proxy.management_endpoints.internal_user_endpoints import user_update
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_helper_fn,
@ -44,6 +45,11 @@ from litellm.repositories.user_repository import UserRepository
from litellm.secret_managers.main import get_secret_bool
from litellm.types.proxy.ui_sso import ReturnedUITokenObject
INVALID_UI_CREDENTIALS_MESSAGE: Final = (
"Invalid credentials used to access UI. Check 'UI_USERNAME' and 'UI_PASSWORD', or the password set for your user"
)
INVALID_USER_PASSWORD_MESSAGE: Final = "Invalid credentials used to access UI. Check the password set for your user"
async def _rehash_password_if_needed(user_id: str, password: str, stored: str) -> None:
"""Rehash legacy password (SHA256) to scrypt on successful login."""
@ -92,6 +98,21 @@ def _matches_env_credentials(username: str, password: str, master_key: str | Non
)
def _admin_credentials_match(
username: str, password: str, master_key: str, general_settings: Mapping[str, object]
) -> bool:
return general_settings.get("disable_env_credential_login") is not True and _matches_env_credentials(
username, password, master_key
)
def _invalid_credentials_message(general_settings: Mapping[str, object]) -> str:
"""One rejection message for unknown usernames and wrong passwords alike, so neither can be enumerated."""
if is_env_credential_login_enabled(general_settings):
return INVALID_UI_CREDENTIALS_MESSAGE
return INVALID_USER_PASSWORD_MESSAGE
def is_env_credential_login_enabled(general_settings: Mapping[str, object]) -> bool:
"""Whether a login with UI_USERNAME/UI_PASSWORD (or the master-key fallback) can succeed.
@ -137,6 +158,7 @@ async def authenticate_user(
password: str,
master_key: str | None,
prisma_client: PrismaClient | None,
throttle: LoginThrottle,
general_settings: Mapping[str, object] = MappingProxyType({}),
) -> LoginResult:
"""
@ -151,6 +173,7 @@ async def authenticate_user(
password: Password from the login form
master_key: Master key for the proxy (required)
prisma_client: Prisma database client (optional)
throttle: Failed sign-in accounting for this request's source address
general_settings: Proxy general_settings, checked for
`disable_password_login_when_sso_enabled` and
`disable_env_credential_login`
@ -163,9 +186,11 @@ async def authenticate_user(
or if username/password login is disabled while SSO is configured
Recovery: an admin locked out of the UI by
`disable_password_login_when_sso_enabled` can still administer the proxy over
the API with the master key (Authorization: Bearer <master_key>), which never
goes through this function. To restore UI username/password login, unset the
`disable_password_login_when_sso_enabled`, or by the failed sign-in block in
`throttle`, can still administer the proxy over the API with the master key
(Authorization: Bearer <master_key>), which never goes through this function.
No credential, the env admin credentials and the master key included, is
exempt from the block. To restore UI username/password login, unset the
setting in config.yaml (or the DB-persisted general_settings) and restart the
proxy; this is a deliberate, auditable config change rather than a hidden
bypass.
@ -194,6 +219,19 @@ async def authenticate_user(
code=500,
)
attempt: Final = await throttle.attempt(username)
return await _sign_in(username, password, master_key, prisma_client, attempt, general_settings)
async def _sign_in(
username: str,
password: str,
master_key: str,
prisma_client: PrismaClient | None,
attempt: LoginAttempt,
general_settings: Mapping[str, object],
) -> LoginResult:
admin_credentials_match: Final = _admin_credentials_match(username, password, master_key, general_settings)
# Check if we can find the `username` in the db. On the UI, users can enter username=their email
_user_row: LiteLLM_UserTable | None = None
user_role: (
@ -219,20 +257,13 @@ async def authenticate_user(
- Login with UI_USERNAME and UI_PASSWORD
- Login with Invite Link `user_email` and `password` combination
"""
if general_settings.get("disable_env_credential_login") is not True and _matches_env_credentials(
username, password, master_key
):
if admin_credentials_match:
# Non SSO -> If user is using UI_USERNAME and UI_PASSWORD they are Proxy admin
user_role = LitellmUserRoles.PROXY_ADMIN
user_id = LITELLM_PROXY_ADMIN_NAME
# we want the key created to have PROXY_ADMIN_PERMISSIONS
key_user_id = LITELLM_PROXY_ADMIN_NAME
if (
os.getenv("PROXY_ADMIN_ID", None) is not None and os.environ["PROXY_ADMIN_ID"] == user_id
) or user_id == LITELLM_PROXY_ADMIN_NAME:
# checks if user is admin
key_user_id = os.getenv("PROXY_ADMIN_ID", LITELLM_PROXY_ADMIN_NAME)
key_user_id: Final = os.getenv("PROXY_ADMIN_ID", LITELLM_PROXY_ADMIN_NAME)
# Admin is Authe'd in - generate key for the UI to access Proxy
@ -294,6 +325,8 @@ async def authenticate_user(
key = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(user_info)
await attempt.succeeded()
return LoginResult(
user_id=user_id,
key=key,
@ -349,6 +382,8 @@ async def authenticate_user(
key = response["token"]
await attempt.succeeded()
return LoginResult(
user_id=user_id,
key=key,
@ -357,20 +392,17 @@ async def authenticate_user(
login_method="username_password",
)
else:
await attempt.failed()
raise ProxyException(
message=f"Invalid credentials used to access UI.\nNot valid credentials for {username}",
message=_invalid_credentials_message(general_settings),
type=ProxyErrorTypes.auth_error,
param="invalid_credentials",
code=401,
)
else:
env_credentials_hint: Final = (
"\nCheck 'UI_USERNAME', 'UI_PASSWORD' in .env file"
if is_env_credential_login_enabled(general_settings)
else ""
)
await attempt.failed()
raise ProxyException(
message=f"Invalid credentials used to access UI.{env_credentials_hint}",
message=_invalid_credentials_message(general_settings),
type=ProxyErrorTypes.auth_error,
param="invalid_credentials",
code=401,

View file

@ -1,6 +1,7 @@
from __future__ import annotations
import ipaddress
from collections.abc import Sequence
from typing import Any, Final
from fastapi import Request
@ -19,7 +20,7 @@ class NetworkContext(BaseModel):
class TrustedProxyConfig(BaseModel):
use_forwarded_for: bool = False
trusted_proxy_cidrs: list[str] = Field(default_factory=list)
trusted_proxy_cidrs: Sequence[str] = Field(default_factory=tuple)
def normalize_cidr_ranges(configured_ranges: Any, *, setting_name: str = "trusted_proxy_cidrs") -> list[str]:
@ -49,6 +50,12 @@ def parse_trusted_proxy_ranges(
return networks
def _unmapped(addr: ipaddress.IPv4Address | ipaddress.IPv6Address) -> ipaddress.IPv4Address | ipaddress.IPv6Address:
if isinstance(addr, ipaddress.IPv6Address) and addr.ipv4_mapped is not None:
return addr.ipv4_mapped
return addr
def ip_in_networks(client_ip: str | None, networks: list[TrustedProxyNetwork]) -> bool:
if not client_ip or not networks:
return False
@ -56,7 +63,8 @@ def ip_in_networks(client_ip: str | None, networks: list[TrustedProxyNetwork]) -
addr: Final = ipaddress.ip_address(client_ip.strip())
except ValueError:
return False
return any(addr in network for network in networks)
candidates: Final = (addr, _unmapped(addr))
return any(candidate in network for candidate in candidates for network in networks)
def _is_valid_ip(value: str) -> bool:

View file

@ -87,6 +87,12 @@ class SettingsStore(MutableMapping[str, JsonValue]):
)
self._deleted_runtime_keys = self._deleted_runtime_keys | frozenset((key,))
def clear(self) -> None:
self._deleted_runtime_keys = frozenset(key for key in self._keys() if not self.owned_by_config(key))
self._runtime_values = MappingProxyType(
{key: value for key, value in self._runtime_values.items() if self.owned_by_config(key)}
)
def __iter__(self) -> Iterator[str]:
return iter(
key

View file

@ -9,7 +9,7 @@ from fastapi import HTTPException, status
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.constants import PTU_SENTINEL_API_KEY
from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.spend_tracking.key_metadata_recovery import (
attach_user_emails,
@ -146,15 +146,9 @@ class _AggregatedSpendData(TypedDict):
totals: SpendMetrics
class _GroupingSetsRow(SimpleNamespace):
class _RollupMetricsRow(SimpleNamespace):
date: str
api_key: str | None
model: str | None
model_group: str | None
custom_llm_provider: str | None
mcp_namespaced_tool_name: str | None
endpoint: str | None
group_level: int
spend: float | None
prompt_tokens: int | None
completion_tokens: int | None
@ -172,12 +166,46 @@ class _GroupingSetsRow(SimpleNamespace):
timed_requests: int | None
class _EntityRollupRow(_GroupingSetsRow):
class _GroupingSetsRow(_RollupMetricsRow):
model: str | None
model_group: str | None
custom_llm_provider: str | None
mcp_namespaced_tool_name: str | None
endpoint: str | None
group_level: int
distinct_api_keys: int | None
class _EntityRollupRow(_RollupMetricsRow):
entity_id: str | None
api_key_rolled: int
def _reported_flat_cost(record: DailySpendRecord | _GroupingSetsRow) -> float:
class _AggregatedQueryKwargs(TypedDict):
table_name: ReadOnly[str]
entity_id_field: ReadOnly[str]
entity_id: ReadOnly[str | list[str] | None]
start_date: ReadOnly[str]
end_date: ReadOnly[str]
model: ReadOnly[str | None]
api_key: ReadOnly[str | list[str] | None]
exclude_entity_ids: ReadOnly[list[str] | None]
timezone_offset_minutes: ReadOnly[int | None]
include_current_utc_day: ReadOnly[bool]
_SqlQuery = tuple[str, list[str]]
async def _query_raw_optional(
prisma_client: PrismaClient, query: _SqlQuery | None
) -> list[dict[str, object]] | None: # mutable-ok: prisma query_raw return shape
if query is None:
return None
return await prisma_client.db.query_raw(query[0], *query[1])
def _reported_flat_cost(record: DailySpendRecord | _RollupMetricsRow) -> float:
"""Flat cost a daily row reports, which is zero unless PTU cost attribution is enabled.
Both read paths funnel through here: the paginated path reads the ``ptu_flat_cost``
@ -699,71 +727,8 @@ def _ptu_flat_cost_select(table_name: str) -> str:
return "0::float AS ptu_flat_cost"
def _build_aggregated_sql_query(
*,
table_name: str,
entity_id_field: str,
entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
start_date: str,
end_date: str,
model: str | None,
api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path
timezone_offset_minutes: int | None = None,
include_current_utc_day: bool = False,
) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params
"""Build a parameterized SQL GROUP BY query for aggregated daily activity.
Groups by (date, api_key, model, model_group, custom_llm_provider,
mcp_namespaced_tool_name, endpoint) with SUMs on all metric columns.
The entity_id column is intentionally omitted from GROUP BY to collapse
rows across entities — this is where the biggest row reduction comes from.
Returns:
Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw().
"""
pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name)
if pg_table is None:
raise ValueError(f"Unknown table name: {table_name}")
adjusted_start, adjusted_end = _adjust_dates_for_timezone(
start_date, end_date, timezone_offset_minutes, include_current_utc_day
)
where_clause, sql_params = _build_aggregated_where_clause(
entity_id_field=entity_id_field,
entity_id=entity_id,
adjusted_start=adjusted_start,
adjusted_end=adjusted_end,
model=model,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
)
# Postgres computes every rollup level the response needs — per-date
# totals, per-(date, model), per-(date, model, api_key), per-provider,
# etc. — in a single pass via GROUPING SETS. The GROUPING() bitmask
# encodes which level a row belongs to so Python can dispatch rows
# straight into their buckets without re-summing. The leaf grouping
# is omitted on purpose: nothing in the response shape needs it once
# all the rollups are present.
#
# TODO: drop the successful_requests/failed_requests aggregates (and the
# total_successful_requests metadata they feed) once the admin UI reads SGR
# only from LiteLLM_DailyGatewayRequests. The remaining spend, token and
# api_requests rollups are still served from here.
sql_query: Final = f"""
SELECT
date,
api_key,
model,
COALESCE(NULLIF(model_group, ''), model) AS model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
GROUPING(date, api_key, model, COALESCE(NULLIF(model_group, ''), model),
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,
def _rollup_metric_select(table_name: str) -> str:
return f"""
SUM(spend)::float AS spend,
{_ptu_flat_cost_select(table_name)},
SUM(prompt_tokens)::bigint AS prompt_tokens,
@ -779,27 +744,113 @@ def _build_aggregated_sql_query(
SUM(successful_requests)::bigint AS successful_requests,
SUM(failed_requests)::bigint AS failed_requests,
SUM(total_response_time_ms)::bigint AS total_response_time_ms,
SUM(timed_requests)::bigint AS timed_requests
SUM(timed_requests)::bigint AS timed_requests"""
_MODEL_GROUP_EXPR: Final = "COALESCE(NULLIF(model_group, ''), model)"
def _build_aggregated_sql_query(
*,
table_name: str,
entity_id_field: str,
entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
start_date: str,
end_date: str,
model: str | None,
api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path
timezone_offset_minutes: int | None = None,
include_current_utc_day: bool = False,
) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params
"""Build the GROUPING SETS query for aggregated daily activity.
Returns:
Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw().
"""
pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name)
if pg_table is None:
raise ValueError(f"Unknown table name: {table_name}")
adjusted_start, adjusted_end = _adjust_dates_for_timezone(
start_date, end_date, timezone_offset_minutes, include_current_utc_day
)
where_clause, where_params = _build_aggregated_where_clause(
entity_id_field=entity_id_field,
entity_id=entity_id,
adjusted_start=adjusted_start,
adjusted_end=adjusted_end,
model=model,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
)
sentinel_param: Final = f"${len(where_params) + 1}"
metric_select: Final = _rollup_metric_select(table_name)
# TODO: drop the successful_requests/failed_requests aggregates (and the
# total_successful_requests metadata they feed) once the admin UI reads SGR
# only from LiteLLM_DailyGatewayRequests. The remaining spend, token and
# api_requests rollups are still served from here.
sql_query: Final = f"""
(SELECT
date,
NULL::text AS api_key,
model,
{_MODEL_GROUP_EXPR} AS model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
(GROUPING(date) << 6) | {_API_KEY_ROLLED_UP_BIT}
| GROUPING(model, {_MODEL_GROUP_EXPR},
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,
NULL::bigint AS distinct_api_keys,{metric_select}
FROM "{pg_table}"
WHERE {where_clause}
GROUP BY GROUPING SETS (
(date),
(date, api_key),
(date, model),
(date, model, api_key),
(date, COALESCE(NULLIF(model_group, ''), model)),
(date, COALESCE(NULLIF(model_group, ''), model), api_key),
(date, {_MODEL_GROUP_EXPR}),
(date, custom_llm_provider),
(date, custom_llm_provider, api_key),
(date, mcp_namespaced_tool_name),
(date, mcp_namespaced_tool_name, api_key),
(date, endpoint),
(date, endpoint, api_key),
()
))
UNION ALL
(WITH top_api_keys AS (
SELECT api_key, COUNT(*) OVER () AS distinct_api_keys
FROM "{pg_table}"
WHERE {where_clause} AND api_key <> {sentinel_param}
GROUP BY api_key
ORDER BY SUM(spend) DESC, api_key
LIMIT {USAGE_TOP_API_KEYS_LIMIT}
)
SELECT
date,
api_key,
model,
{_MODEL_GROUP_EXPR} AS model_group,
custom_llm_provider,
mcp_namespaced_tool_name,
endpoint,
GROUPING(date, api_key, model, {_MODEL_GROUP_EXPR},
custom_llm_provider, mcp_namespaced_tool_name,
endpoint) AS group_level,
MAX(top_api_keys.distinct_api_keys) AS distinct_api_keys,{metric_select}
FROM "{pg_table}" JOIN top_api_keys USING (api_key)
WHERE {where_clause}
GROUP BY GROUPING SETS (
(date, api_key),
(date, model, api_key),
(date, {_MODEL_GROUP_EXPR}, api_key),
(date, custom_llm_provider, api_key),
(date, mcp_namespaced_tool_name, api_key),
(date, endpoint, api_key)
))
"""
return sql_query, sql_params
return sql_query, [*where_params, PTU_SENTINEL_API_KEY]
def _build_entity_rollup_sql_query(
@ -844,23 +895,7 @@ def _build_entity_rollup_sql_query(
"{entity_id_field}" AS entity_id,
date,
api_key,
GROUPING(api_key) AS api_key_rolled,
SUM(spend)::float AS spend,
{_ptu_flat_cost_select(table_name)},
SUM(prompt_tokens)::bigint AS prompt_tokens,
SUM(completion_tokens)::bigint AS completion_tokens,
SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens,
SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens,
SUM(compression_saved_tokens)::bigint AS compression_saved_tokens,
SUM(compression_savings_spend)::float AS compression_savings_spend,
SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend,
SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend,
SUM(autorouter_savings_spend)::float AS autorouter_savings_spend,
SUM(api_requests)::bigint AS api_requests,
SUM(successful_requests)::bigint AS successful_requests,
SUM(failed_requests)::bigint AS failed_requests,
SUM(total_response_time_ms)::bigint AS total_response_time_ms,
SUM(timed_requests)::bigint AS timed_requests
GROUPING(api_key) AS api_key_rolled,{_rollup_metric_select(table_name)}
FROM "{pg_table}"
WHERE {where_clause}
GROUP BY GROUPING SETS (
@ -962,6 +997,7 @@ async def _aggregate_spend_records(
# current grouping set's key), 0 when the column is part of the key.
_GROUP_GRAND_TOTAL: Final = 127 # 0b1111111 — all rolled up
_GROUP_DATE: Final = 63 # 0b0111111 — only date kept
_API_KEY_ROLLED_UP_BIT: Final = 32 # 0b0100000
_GROUP_DATE_API_KEY: Final = 31 # 0b0011111
_GROUP_DATE_MODEL: Final = 47 # 0b0101111
_GROUP_DATE_MODEL_API_KEY: Final = 15 # 0b0001111
@ -975,7 +1011,7 @@ _GROUP_DATE_ENDPOINT: Final = 62 # 0b0111110
_GROUP_DATE_ENDPOINT_API_KEY: Final = 30 # 0b0011110
def _record_to_spend_metrics(record: _GroupingSetsRow) -> SpendMetrics:
def _record_to_spend_metrics(record: _RollupMetricsRow) -> SpendMetrics:
"""Build a SpendMetrics directly from one already-aggregated rollup row.
SUM() over zero rows is SQL NULL, so rollup rows (notably the grand-total
@ -1329,10 +1365,6 @@ async def get_daily_activity_aggregated(
) -> SpendAnalyticsPaginatedResponse:
"""Aggregated variant that returns the full result set (no pagination).
Uses SQL GROUP BY to aggregate rows in the database rather than fetching
all individual rows into Python. This collapses rows across entities
(users/teams/orgs), reducing ~150k rows to ~2-3k grouped rows.
include_entity_breakdown runs a small companion rollup query and folds
`breakdown.entities` onto the response, as entity-scoped views like Team Usage need.
@ -1351,7 +1383,7 @@ async def get_daily_activity_aggregated(
)
try:
sql_query, sql_params = _build_aggregated_sql_query(
query_kwargs: Final = _AggregatedQueryKwargs(
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
@ -1363,36 +1395,16 @@ async def get_daily_activity_aggregated(
timezone_offset_minutes=timezone_offset_minutes,
include_current_utc_day=include_current_utc_day,
)
sql_query, sql_params = _build_aggregated_sql_query(**query_kwargs)
entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None
entity_query: Final = (
_build_entity_rollup_sql_query(
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
start_date=start_date,
end_date=end_date,
model=model,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
timezone_offset_minutes=timezone_offset_minutes,
include_current_utc_day=include_current_utc_day,
)
if include_entity_breakdown
else None
raw_rows, raw_entity_rows = await asyncio.gather(
prisma_client.db.query_raw(sql_query, *sql_params),
_query_raw_optional(prisma_client, entity_query),
)
# Execute the GROUPING SETS query (one row per rollup level), alongside
# the per-entity companion rollup when the caller wants entities.
raw_rows, raw_entity_rows = (
await asyncio.gather(
prisma_client.db.query_raw(sql_query, *sql_params),
prisma_client.db.query_raw(entity_query[0], *entity_query[1]),
)
if entity_query is not None
else (await prisma_client.db.query_raw(sql_query, *sql_params), None)
)
records: Final = [_GroupingSetsRow(**row) for row in (raw_rows or [])]
records: Final = [_GroupingSetsRow(**row) for row in (raw_rows or ())]
total_api_keys: Final = next((r.distinct_api_keys for r in records if r.distinct_api_keys is not None), 0)
# The grouping-sets dispatcher places each row directly in its bucket
# using the row's GROUPING() bitmask. No Python-side summing needed.
@ -1446,6 +1458,8 @@ async def get_daily_activity_aggregated(
page=1,
total_pages=1,
has_more=False,
api_key_limit=USAGE_TOP_API_KEYS_LIMIT,
total_api_keys=total_api_keys,
),
)

View file

@ -47,7 +47,7 @@ except ImportError:
import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.constants import LITELLM_PROXY_ADMIN_NAME, MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH
from litellm.proxy._experimental.mcp_server.utils import (
LITELLM_MCP_SERVER_DESCRIPTION,
LITELLM_MCP_SERVER_NAME,
@ -145,6 +145,7 @@ if MCP_AVAILABLE:
get_user_env_vars,
get_user_env_vars_bulk,
get_user_oauth_credential,
list_server_user_credentials,
list_user_oauth_credentials,
mcp_oauth_token_identity,
merge_user_env_vars,
@ -180,6 +181,7 @@ if MCP_AVAILABLE:
MCPApprovalStatus,
MCPOAuthUserCredentialRequest,
MCPOAuthUserCredentialStatus,
MCPServerUserCredentialListItem,
MCPSubmissionsSummary,
MCPTransport,
MCPUserCredentialListItem,
@ -221,6 +223,7 @@ if MCP_AVAILABLE:
MCPAuth,
MCPCredentials,
MCPGatewaySessionsResponse,
MCPGatewaySessionsTerminateResponse,
normalize_upstream_header_name,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -662,6 +665,31 @@ if MCP_AVAILABLE:
"""
return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
def _resolve_credential_target_user_id(user_api_key_dict: UserAPIKeyAuth, requested_user_id: str | None) -> str:
"""The user whose stored MCP credential a request acts on.
Defaults to the caller. Naming another user is a revocation and needs
``PROXY_ADMIN``; a read-only admin or a regular user gets 403.
"""
caller_user_id: Final = user_api_key_dict.user_id or ""
if requested_user_id is not None and requested_user_id != caller_user_id:
if not _user_is_full_admin(user_api_key_dict):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict
"error": "Proxy admin access required to revoke another user's MCP credential.",
},
)
return requested_user_id
if not caller_user_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "User ID not found in token"
}, # mutable-ok: FastAPI HTTPException detail requires a plain dict
)
return caller_user_id
def _is_restricted_virtual_key_request(user_api_key_dict: UserAPIKeyAuth) -> bool:
"""Best-effort detection for route-restricted virtual keys.
@ -1373,6 +1401,41 @@ if MCP_AVAILABLE:
return get_mcp_gateway_sessions_report()
@router.delete(
"/sessions",
description=(
"Force-close live stateful MCP gateway sessions on this proxy worker, selected by session id prefix "
"and/or by the LiteLLM user that opened them (proxy admin only)."
),
dependencies=(Depends(user_api_key_auth),),
response_model=MCPGatewaySessionsTerminateResponse,
)
@management_endpoint_wrapper
async def delete_mcp_gateway_sessions(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
session_id_prefix: Annotated[str | None, Query(min_length=MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH)] = None,
user_id: Annotated[str | None, Query(min_length=1)] = None,
) -> MCPGatewaySessionsTerminateResponse:
if not _user_is_full_admin(user_api_key_dict):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict
"error": "Proxy admin access required to terminate MCP gateway sessions.",
},
)
if session_id_prefix is None and user_id is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict
"error": "Provide session_id_prefix and/or user_id to select the sessions to terminate.",
},
)
from litellm.proxy._experimental.mcp_server.server import (
terminate_mcp_gateway_sessions,
)
return await terminate_mcp_gateway_sessions(session_id_prefix=session_id_prefix, user_id=user_id)
@router.get(
"/server/submissions",
description="Returns all MCP servers submitted by non-admin users (admin review queue). Mirrors GET /guardrails/submissions.",
@ -2254,14 +2317,17 @@ if MCP_AVAILABLE:
_invalidate_byok_cred_cache,
)
_invalidate_byok_cred_cache(user_id, server_id)
await _invalidate_byok_cred_cache(user_id, server_id)
return MCPUserCredentialResponse(server_id=server_id, has_credential=True)
# save=False: credential not persisted
return MCPUserCredentialResponse(server_id=server_id, has_credential=False)
@router.delete(
"/server/{server_id}/user-credential",
description="Delete the calling user's stored API key for a BYOK MCP server",
description=(
"Delete the calling user's stored API key for a BYOK MCP server. "
"A proxy admin may pass user_id to revoke another user's stored key."
),
dependencies=[Depends(user_api_key_auth)],
response_model=MCPUserCredentialResponse,
)
@ -2269,24 +2335,20 @@ if MCP_AVAILABLE:
async def delete_mcp_user_credential(
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
user_id: Annotated[str | None, Query(min_length=1)] = None,
):
"""Remove the calling user's BYOK credential."""
"""Remove the target user's BYOK credential (the caller unless an admin names another user)."""
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
user_id: Final = user_api_key_dict.user_id or ""
if not user_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "User ID not found in token"},
)
target_user_id: Final = _resolve_credential_target_user_id(user_api_key_dict, user_id)
try:
await delete_user_credential(prisma_client, user_id, server_id)
await delete_user_credential(prisma_client, target_user_id, server_id)
except RecordNotFoundError:
pass # Already deleted or didn't exist
from litellm.proxy._experimental.mcp_server.server import (
_invalidate_byok_cred_cache,
)
_invalidate_byok_cred_cache(user_id, server_id)
await _invalidate_byok_cred_cache(target_user_id, server_id)
return MCPUserCredentialResponse(server_id=server_id, has_credential=False)
# ── OAuth2 user-credential endpoints ──────────────────────────────────────
@ -2362,7 +2424,10 @@ if MCP_AVAILABLE:
@router.delete(
"/server/{server_id}/oauth-user-credential",
description="Revoke the calling user's stored OAuth2 token for an MCP server",
description=(
"Revoke the calling user's stored OAuth2 token for an MCP server. "
"A proxy admin may pass user_id to revoke another user's stored token."
),
dependencies=[Depends(user_api_key_auth)],
response_model=MCPOAuthUserCredentialStatus,
)
@ -2370,29 +2435,25 @@ if MCP_AVAILABLE:
async def delete_mcp_oauth_user_credential(
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
user_id: Annotated[str | None, Query(min_length=1)] = None,
):
"""Revoke/delete the user's OAuth2 credential."""
"""Revoke the target user's OAuth2 credential (the caller unless an admin names another user)."""
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
user_id: Final = user_api_key_dict.user_id or ""
if not user_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "User ID not found in token"},
)
target_user_id: Final = _resolve_credential_target_user_id(user_api_key_dict, user_id)
# Only delete if the stored credential is actually an OAuth2 token.
# This prevents accidentally deleting a BYOK credential if one exists
# for the same (user_id, server_id) pair.
cred_to_delete: Final = await get_user_oauth_credential(prisma_client, user_id, server_id)
cred_to_delete: Final = await get_user_oauth_credential(prisma_client, target_user_id, server_id)
if cred_to_delete is not None:
try:
await delete_user_credential(prisma_client, user_id, server_id)
await delete_user_credential(prisma_client, target_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)
await global_mcp_server_manager.invalidate_user_oauth_token_cache(target_user_id, server_id)
return MCPOAuthUserCredentialStatus(
server_id=server_id,
has_credential=False,
@ -2481,6 +2542,30 @@ if MCP_AVAILABLE:
)
return items
@router.get(
"/server/{server_id}/user-credentials",
description="List every user's stored BYOK or OAuth2 credential for an MCP server (admin only, no secrets)",
dependencies=(Depends(user_api_key_auth),),
response_model=list[MCPServerUserCredentialListItem],
)
@management_endpoint_wrapper
async def list_mcp_server_user_credentials(
server_id: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> tuple[MCPServerUserCredentialListItem, ...]:
if user_api_key_dict.user_role not in (
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict
"error": "Admin access required to view MCP server user credentials.",
},
)
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
return await list_server_user_credentials(prisma_client, server_id)
# ── Per-user MCP env var endpoints ────────────────────────────────────────
async def _authorize_and_fetch_mcp_server(

View file

@ -2107,7 +2107,7 @@ def _handle_multi_valued_attribute_update(path: str, op_type: str, value: object
except ValidationError:
raise HTTPException(
status_code=400,
detail={"error": f"Invalid value for {base}: expected a list of objects with a 'value' sub-attribute"},
detail={"error": f"Invalid value for {base}: expected a list of objects or strings"},
)
dumped: Final = [attr.model_dump(exclude_none=True) for attr in attrs]

View file

@ -1411,6 +1411,8 @@ def run_server(
# DO NOT DELETE - enables global variables to work across files
from litellm.proxy.proxy_server import app
os.environ["NUM_WORKERS"] = str(num_workers)
# Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups
prometheus_multiproc_dir: Final = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir(
num_workers=num_workers,

View file

@ -309,6 +309,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._experimental.mcp_server.byok_credential_cache import byok_credential_cache
from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot
from litellm.proxy._types import *
from litellm.proxy.analytics_endpoints.analytics_endpoints import (
@ -330,6 +331,12 @@ from litellm.proxy.auth.fallback_budget import router_fallback_budget_check
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_REMEDY, LicenseCheck
from litellm.proxy.auth.login_throttle import (
LoginThrottle,
declared_proxy_ranges,
warn_login_counters_are_per_worker,
warn_source_login_limit_is_off,
)
from litellm.proxy.auth.model_checks import (
expand_wildcard_deployments_for_model_info,
get_all_fallbacks,
@ -832,6 +839,7 @@ from fastapi.openapi.docs import get_swagger_ui_html
from fastapi.openapi.utils import get_openapi
from fastapi.responses import (
FileResponse,
HTMLResponse,
JSONResponse,
ORJSONResponse,
RedirectResponse,
@ -6067,6 +6075,12 @@ class ProxyConfig:
general_settings = config.get("general_settings", {})
if general_settings is None:
general_settings = {}
if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None:
warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1"))
if declared_proxy_ranges(general_settings) is None:
warn_source_login_limit_is_off()
_bg_hc_model_groups: Final = parse_background_health_check_model_groups(general_settings)
_enable_hc_routing = False
_hc_staleness = None
@ -7554,7 +7568,7 @@ class ProxyConfig:
subscriber: Final = AuthCacheInvalidationSubscriber(
redis_cache=redis_cache,
user_api_key_cache=user_api_key_cache,
additional_in_memory_caches=(spend_counter_cache.in_memory_cache,),
additional_in_memory_caches=(spend_counter_cache.in_memory_cache, byok_credential_cache),
)
self.auth_cache_invalidation_subscriber = subscriber
subscriber.start()
@ -15918,8 +15932,6 @@ async def fallback_login(request: Request):
else:
redirect_url += "/sso/callback"
from fastapi.responses import HTMLResponse
hide_default_credentials_hint: Final = should_hide_default_credentials_hint(general_settings)
return HTMLResponse(
content=build_ui_login_form(
@ -15941,13 +15953,27 @@ async def login(request: Request):
password: Final = str(form.get("password"))
# Authenticate user and get login result
login_result: Final = await authenticate_user(
username=username,
password=password,
master_key=master_key,
prisma_client=prisma_client,
general_settings=general_settings,
)
try:
login_result: Final = await authenticate_user(
username=username,
password=password,
master_key=master_key,
prisma_client=prisma_client,
throttle=LoginThrottle.from_request(request, general_settings, redis_usage_cache),
general_settings=general_settings,
)
except ProxyException as exc:
if int(exc.code) != status.HTTP_429_TOO_MANY_REQUESTS:
raise
retry_after: Final = exc.headers.get("Retry-After", "30")
return HTMLResponse(
content=(
"<html><body><h1>Too many sign-in attempts</h1>"
f"<p>Try again in about {retry_after} seconds</p></body></html>"
),
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
headers=exc.headers,
)
# Create UI token object
returned_ui_token_object: Final = create_ui_token_object(
@ -16026,6 +16052,7 @@ async def login_v2(request: Request):
password=password,
master_key=master_key,
prisma_client=prisma_client,
throttle=LoginThrottle.from_request(request, general_settings, redis_usage_cache),
general_settings=general_settings,
)
@ -16097,6 +16124,7 @@ async def login_v3(request: Request):
password=password,
master_key=master_key,
prisma_client=prisma_client,
throttle=LoginThrottle.from_request(request, general_settings, redis_usage_cache),
general_settings=general_settings,
)

View file

@ -3940,12 +3940,24 @@ class Router:
)
_router_timeout: Final = (
float(self._explicit_timeout) if isinstance(self._explicit_timeout, (int, float)) else None
self.request_timeout
if self.request_timeout is not None
else float(self._explicit_timeout)
if isinstance(self._explicit_timeout, (int, float))
else None
)
_router_stream_timeout: Final = (
self.stream_timeout
if self.stream_timeout is not None
else self.request_timeout
if self.request_timeout is not None
else self.default_litellm_params.get("stream_timeout")
)
kwargs["timeout"] = resolve_llm_passthrough_timeout(
kwargs=kwargs,
litellm_params=deployment["litellm_params"],
router_timeout=_router_timeout,
router_stream_timeout=_router_stream_timeout,
)
else:
kwargs["timeout"] = self._get_timeout(kwargs=kwargs, data=deployment["litellm_params"])

View file

@ -464,3 +464,11 @@ class MCPGatewaySessionsResponse(BaseModel):
by_client: list[MCPGatewaySessionGroupCount] = Field(default_factory=list)
by_user: list[MCPGatewaySessionGroupCount] = Field(default_factory=list)
sessions: list[MCPGatewaySession] = Field(default_factory=list)
class MCPGatewaySessionsTerminateResponse(BaseModel):
"""Stateful sessions an administrator force-closed on this proxy worker."""
worker_pid: int
terminated_sessions: int
sessions: list[MCPGatewaySession] = Field(default_factory=list)

View file

@ -100,6 +100,16 @@ class DailySpendMetadata(BaseModel):
page: int = Field(default=1)
total_pages: int = Field(default=1)
has_more: bool = Field(default=False)
api_key_limit: int | None = Field(
default=None,
description="When set, api_keys and every api_key_breakdown list at most this many keys, "
"ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.",
)
total_api_keys: int | None = Field(
default=None,
description="Distinct API keys matching the filters. When this exceeds api_key_limit, the per-key "
"lists are truncated to the highest-spend keys.",
)
class SpendAnalyticsPaginatedResponse(BaseModel):

View file

@ -61,7 +61,9 @@ class SCIMUserGroup(BaseModel):
class SCIMMultiValuedAttribute(BaseModel):
value: str
model_config = ConfigDict(extra="allow")
value: str | None = None
display: str | None = None
type: str | None = None
primary: bool | None = None

View file

@ -16690,6 +16690,46 @@
"supports_tool_choice": true,
"supports_vision": true
},
"dashscope/qwen3.8-flash": {
"cache_creation_input_token_cost": 2e-07,
"cache_read_input_token_cost": 1.6e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
"dashscope/qwen3.8-omni-flash": {
"cache_read_input_token_cost": 1.6e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
"dashscope/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "dashscope",
@ -18594,6 +18634,46 @@
"supports_tool_choice": true,
"supports_vision": true
},
"qwen_ai_platform/qwen3.8-flash": {
"cache_creation_input_token_cost": 2e-07,
"cache_read_input_token_cost": 1.6e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "qwen_ai_platform",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
"qwen_ai_platform/qwen3.8-omni-flash": {
"cache_read_input_token_cost": 1.6e-08,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "qwen_ai_platform",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing",
"supports_audio_input": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
"qwen_ai_platform/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "qwen_ai_platform",
@ -22090,8 +22170,8 @@
"embed-english-light-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "cohere",
"max_input_tokens": 1024,
"max_tokens": 1024,
"max_input_tokens": 512,
"max_tokens": 512,
"mode": "embedding",
"output_cost_per_token": 0.0
},
@ -22108,8 +22188,8 @@
"input_cost_per_image": 0.0001,
"input_cost_per_token": 1e-07,
"litellm_provider": "cohere",
"max_input_tokens": 1024,
"max_tokens": 1024,
"max_input_tokens": 512,
"max_tokens": 512,
"metadata": {
"notes": "'supports_image_input' is a deprecated field. Use 'supports_embedding_image_input' instead."
},
@ -22130,8 +22210,8 @@
"embed-multilingual-v3.0": {
"input_cost_per_token": 1e-07,
"litellm_provider": "cohere",
"max_input_tokens": 1024,
"max_tokens": 1024,
"max_input_tokens": 512,
"max_tokens": 512,
"mode": "embedding",
"output_cost_per_token": 0.0,
"supports_embedding_image_input": true
@ -22139,8 +22219,8 @@
"embed-multilingual-light-v3.0": {
"input_cost_per_token": 0.0001,
"litellm_provider": "cohere",
"max_input_tokens": 1024,
"max_tokens": 1024,
"max_input_tokens": 512,
"max_tokens": 512,
"mode": "embedding",
"output_cost_per_token": 0.0,
"supports_embedding_image_input": true
@ -57910,14 +57990,14 @@
"supports_tool_choice": true
},
"bedrock_mantle/openai.gpt-5.6-sol": {
"input_cost_per_token": 5.5e-06,
"input_cost_per_token_above_272k_tokens": 1.1e-05,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
"cache_read_input_token_cost": 5.5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"input_cost_per_token": 4.4e-06,
"input_cost_per_token_above_272k_tokens": 8.8e-06,
"cache_creation_input_token_cost": 5.5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1.1e-05,
"cache_read_input_token_cost": 4.4e-07,
"cache_read_input_token_cost_above_272k_tokens": 8.8e-07,
"output_cost_per_token": 2.2e-05,
"output_cost_per_token_above_272k_tokens": 3.3e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
@ -65816,9 +65896,9 @@
"supports_web_search": false
},
"openrouter/z-ai/glm-5.3": {
"input_cost_per_token": 1.4e-06,
"output_cost_per_token": 4.4e-06,
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 9.1e-07,
"output_cost_per_token": 2.86e-06,
"cache_read_input_token_cost": 1.69e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
"max_output_tokens": 943717,
@ -70539,14 +70619,14 @@
"supports_web_search": false
},
"openrouter/~deepseek/deepseek-flash-latest": {
"cache_read_input_token_cost": 1.5e-08,
"input_cost_per_token": 1.5e-07,
"cache_read_input_token_cost": 4.2e-09,
"input_cost_per_token": 1.4e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 6e-07,
"output_cost_per_token": 4.2e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -70831,14 +70911,14 @@
"supports_web_search": false
},
"openrouter/~z-ai/glm-latest": {
"cache_read_input_token_cost": 1.5e-07,
"cache_read_input_token_cost": 1.46625e-07,
"input_cost_per_token": 9e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1310720,
"max_output_tokens": 235929,
"max_tokens": 235929,
"mode": "chat",
"output_cost_per_token": 3e-06,
"output_cost_per_token": 2.805e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -73982,14 +74062,14 @@
"supports_web_search": false
},
"openrouter/tencent/hy3": {
"cache_read_input_token_cost": 3.3e-08,
"input_cost_per_token": 1.32e-07,
"cache_read_input_token_cost": 2.0625e-08,
"input_cost_per_token": 8.25e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5.28e-07,
"output_cost_per_token": 3.3e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,

View file

@ -51,9 +51,7 @@ try:
if general_settings_section:
# Extract the table rows, which contain the documented keys
table_content = general_settings_section.group(1)
doc_key_pattern = re.compile(
r"\|\s*([^\|]+?)\s*\|"
) # Capture the key from each row of the table
doc_key_pattern = re.compile(r"^\|\s*([^\|]+?)\s*\|", re.MULTILINE)
documented_keys.update(doc_key_pattern.findall(table_content))
except Exception as e:
raise Exception(

View file

@ -331,9 +331,7 @@ class TestMCPPerUserTokenCache:
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_dual_cache):
await cache.delete("alice", "slack-test")
mock_dual_cache.async_delete_cache.assert_called_once_with(
"mcp:per_user_token:alice:slack-test"
)
mock_dual_cache.async_delete_cache.assert_called_once_with(key="mcp:per_user_token:alice:slack-test")
mock_dual_cache.async_set_cache.assert_not_called()
@pytest.mark.asyncio

View file

@ -0,0 +1,20 @@
from litellm.llms.azure.vector_stores.transformation import AzureOpenAIVectorStoreConfig
def test_transform_search_vector_store_request_preserves_azure_query_string():
config = AzureOpenAIVectorStoreConfig()
api_base = config.get_complete_url(
api_base="https://x.openai.azure.com",
litellm_params={"api_version": "2024-10-21"},
)
url, _ = config.transform_search_vector_store_request(
vector_store_id="vs_1",
query="hello",
vector_store_search_optional_params={},
api_base=api_base,
litellm_logging_obj=None,
litellm_params={"api_version": "2024-10-21"},
)
assert url == "https://x.openai.azure.com/openai/vector_stores/vs_1/search?api-version=2024-10-21"

View file

@ -1839,6 +1839,27 @@ class TestBedrockMantleResponsesSigV4:
class TestBedrockMantleResponsesPricing:
@pytest.mark.parametrize(
"model",
["openai.gpt-5.6-sol", "openai.gpt-5.6-terra", "openai.gpt-5.6-luna"],
)
def test_mantle_matches_in_region_converse_pricing(self, local_cost_map, model):
"""bedrock-mantle serves these models In-Region only, and the AWS model
cards price In-Region and Geo CRIS identically -- so every cost field on
the mantle key must equal the `us.` converse key. A price change applied
to one namespace but not the other shows up here.
"""
mantle = litellm.model_cost[f"bedrock_mantle/{model}"]
converse = litellm.model_cost[f"us.{model}"]
cost_fields = [k for k in converse if "cost" in k and k != "search_context_cost_per_query"]
assert cost_fields, "expected cost fields on the converse entry"
for field in cost_fields:
assert mantle.get(field) == pytest.approx(converse[field]), (
f"{model}: {field} is {mantle.get(field)} on bedrock_mantle "
f"but {converse[field]} on us. (bedrock_converse)"
)
def test_models_registered(self, local_cost_map):
assert "bedrock_mantle/openai.gpt-5.5" in litellm.bedrock_mantle_models
assert "bedrock_mantle/openai.gpt-5.4" in litellm.bedrock_mantle_models

View file

@ -7,11 +7,8 @@ from litellm.types.vector_stores import (
class TestOpenAIVectorStoreAPIConfig:
@pytest.mark.parametrize("metadata", [{}, None])
def test_transform_create_vector_store_request_with_metadata_empty_or_none(
self, metadata
):
def test_transform_create_vector_store_request_with_metadata_empty_or_none(self, metadata):
"""
Test transform_create_vector_store_request when metadata is None or empty dict.
"""
@ -24,9 +21,7 @@ class TestOpenAIVectorStoreAPIConfig:
"metadata": metadata,
}
url, request_body = config.transform_create_vector_store_request(
vector_store_create_params, api_base
)
url, request_body = config.transform_create_vector_store_request(vector_store_create_params, api_base)
assert url == api_base
assert request_body["name"] == "test-vector-store"
@ -50,9 +45,7 @@ class TestOpenAIVectorStoreAPIConfig:
"metadata": large_metadata,
}
url, request_body = config.transform_create_vector_store_request(
vector_store_create_params, api_base
)
url, request_body = config.transform_create_vector_store_request(vector_store_create_params, api_base)
assert url == api_base
assert request_body["name"] == "test-vector-store"
@ -77,8 +70,19 @@ class TestOpenAIVectorStoreAPIConfig:
litellm_params={},
)
assert (
url
== "https://api.openai.com/v1/vector_stores/..%2F..%2Ffiles%3Fx%3D1%23frag/search"
)
assert url == "https://api.openai.com/v1/vector_stores/..%2F..%2Ffiles%3Fx%3D1%23frag/search"
assert request_body["query"] == "hello"
def test_transform_search_vector_store_request_preserves_query_string(self):
config = OpenAIVectorStoreConfig()
url, _ = config.transform_search_vector_store_request(
vector_store_id="vs_1",
query="hello",
vector_store_search_optional_params={},
api_base="https://x.openai.azure.com/openai/vector_stores?api-version=2024-10-21",
litellm_logging_obj=None,
litellm_params={},
)
assert url == "https://x.openai.azure.com/openai/vector_stores/vs_1/search?api-version=2024-10-21"

View file

@ -0,0 +1,57 @@
import json
import pytest
from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
CachedByokCredential,
byok_credential_cache,
byok_credential_cache_key,
cache_byok_credential,
get_cached_byok_credential,
)
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
class _FakeRedisCache:
namespace = None
def init_async_client(self) -> object:
return object()
@pytest.fixture(autouse=True)
def _empty_cache():
byok_credential_cache.flush_cache()
yield
byok_credential_cache.flush_cache()
def test_a_cached_negative_lookup_is_distinguishable_from_a_miss():
assert get_cached_byok_credential("u-1", "srv-1") is None
cache_byok_credential("u-1", "srv-1", None)
assert get_cached_byok_credential("u-1", "srv-1") == CachedByokCredential(credential=None)
cache_byok_credential("u-1", "srv-1", "sk-stored")
assert get_cached_byok_credential("u-1", "srv-1") == CachedByokCredential(credential="sk-stored")
assert get_cached_byok_credential("u-1", "srv-2") is None
def test_peer_worker_invalidation_message_evicts_the_cached_credential():
"""The key a mutating worker broadcasts must be the key every other worker caches under."""
cache_byok_credential("mallory", "srv-byok", "sk-revoked")
cache_byok_credential("alice", "srv-byok", "sk-kept")
subscriber = AuthCacheInvalidationSubscriber(
redis_cache=_FakeRedisCache(), # pyright: ignore[reportArgumentType] # subscriber is never started; only its message handler runs
user_api_key_cache=UserApiKeyCache(),
additional_in_memory_caches=(byok_credential_cache,),
)
subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler
{
"type": "message",
"data": json.dumps({"cache_key": byok_credential_cache_key("mallory", "srv-byok")}).encode(),
}
)
assert get_cached_byok_credential("mallory", "srv-byok") is None
assert get_cached_byok_credential("alice", "srv-byok") == CachedByokCredential(credential="sk-kept")

View file

@ -592,7 +592,7 @@ async def test_check_byok_credential_missing_credential(monkeypatch):
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
monkeypatch.setattr(server_module, "_byok_cred_cache", {})
server_module.byok_credential_cache.flush_cache()
mock_prisma = MagicMock()
with (
@ -628,7 +628,7 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com/proxy")
monkeypatch.setattr(mcp_module, "_byok_cred_cache", {})
mcp_module.byok_credential_cache.flush_cache()
server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True)
prisma = MagicMock()
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None)
@ -677,6 +677,40 @@ async def test_check_byok_credential_has_credential():
await _check_byok_credential(server, user_auth)
@pytest.mark.asyncio
async def test_invalidate_byok_cred_cache_evicts_locally_and_broadcasts_the_same_key():
"""A revoked credential must stop being served here and on every peer worker within the TTL."""
from litellm.proxy._experimental.mcp_server import server as server_module
from litellm.proxy._experimental.mcp_server.byok_credential_cache import byok_credential_cache_key
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(server_id="byok-revoke", name="byok-server", transport=MCPTransport.http, is_byok=True)
user_auth = UserAPIKeyAuth(user_id="mallory", api_key="sk-test")
server_module.byok_credential_cache.flush_cache()
db_lookup = AsyncMock(side_effect=["sk-before-revoke", None])
publish = AsyncMock()
with (
patch( # test-quality-ok: the DB row lookup is the only seam below the credential resolver; no Prisma fake exists
"litellm.proxy._experimental.mcp_server.db.get_user_credential", new=db_lookup
),
patch( # test-quality-ok: the resolver reads the module-level prisma_client singleton; the suite's only seam
"litellm.proxy.proxy_server.prisma_client", MagicMock()
),
patch.object( # test-quality-ok: the redis publisher is module-level; asserting the broadcast without a redis
server_module, "publish_auth_cache_invalidation", new=publish
),
):
assert await server_module._get_byok_credential(server, user_auth) == "sk-before-revoke"
assert await server_module._get_byok_credential(server, user_auth) == "sk-before-revoke"
await server_module._invalidate_byok_cred_cache("mallory", "byok-revoke")
assert await server_module._get_byok_credential(server, user_auth) is None
assert db_lookup.await_count == 2
publish.assert_awaited_once_with(cache_key=byok_credential_cache_key("mallory", "byok-revoke"))
@pytest.mark.asyncio
async def test_check_byok_credential_db_unavailable_fails_closed():
"""BYOK server with no prisma_client → 503, not silent pass.

View file

@ -212,6 +212,53 @@ async def test_purge_user_oauth_credentials_for_server_invalidates_each_user():
assert set(invalidations) == {("alice", "srv-1"), ("bob", "srv-1")}
@pytest.mark.asyncio
async def test_list_server_user_credentials_types_each_row_without_leaking_the_secret():
"""The admin view of one server's stored credentials names the user and the kind of
credential (OAuth2 vs BYOK) and echoes OAuth expiry, but never the token or key itself."""
from litellm.proxy._experimental.mcp_server.db import list_server_user_credentials
oauth_row = _legacy_row(
json.dumps(
{
"type": "oauth2",
"access_token": "tok-alice",
"expires_at": "2026-12-31T00:00:00+00:00",
"connected_at": "2026-01-01T00:00:00+00:00",
}
)
)
oauth_row.user_id = "alice"
oauth_row.updated_at = datetime(2026, 1, 1, tzinfo=timezone.utc)
byok_row = _byok_row("carol")
byok_row.updated_at = datetime(2026, 2, 1, tzinfo=timezone.utc)
prisma = MagicMock()
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[oauth_row, byok_row])
items = await list_server_user_credentials(prisma, "srv-1")
prisma.db.litellm_mcpusercredentials.find_many.assert_awaited_once_with(where={"server_id": "srv-1"})
assert [item.model_dump() for item in items] == [
{
"user_id": "alice",
"credential_type": "oauth2",
"expires_at": "2026-12-31T00:00:00+00:00",
"connected_at": "2026-01-01T00:00:00+00:00",
"updated_at": "2026-01-01T00:00:00+00:00",
},
{
"user_id": "carol",
"credential_type": "byok",
"expires_at": None,
"connected_at": None,
"updated_at": "2026-02-01T00:00:00+00:00",
},
]
serialized = "".join(item.model_dump_json() for item in items)
assert "tok-alice" not in serialized
assert "sk-byok-carol" not in serialized
@pytest.mark.asyncio
async def test_purge_user_oauth_credentials_for_server_spares_byok_rows():
"""Regression: the purge used to delete_many on server_id alone, wiping BYOK API keys that share

View file

@ -2871,6 +2871,255 @@ def test_remove_stateful_session_tracking_drops_client_info():
assert session_id not in mcp_server._stateful_session_client_info
def _admin_terminate_fixture(mcp_server):
def auth_user(user_id: str):
return mcp_server.MCPAuthenticatedUser(
user_api_key_auth=UserAPIKeyAuth(api_key=f"key-{user_id}", user_id=user_id),
)
contexts = {
"alice-session-1": auth_user("alice"),
"alice-session-2": auth_user("alice"),
"bob-session-1": auth_user("bob"),
"anon-session-1": mcp_server.MCPAuthenticatedUser(user_api_key_auth=None),
"gone-session-1": auth_user("alice"),
}
transports = {
session_id: MagicMock(terminate=AsyncMock())
for session_id in ("alice-session-1", "alice-session-2", "bob-session-1", "anon-session-1")
}
return contexts, transports
@pytest.mark.asyncio
async def test_terminate_mcp_gateway_sessions_by_user_closes_every_live_session_of_that_user():
try:
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.server import session_manager_stateful
except ImportError:
pytest.skip("MCP server not available")
contexts, transports = _admin_terminate_fixture(mcp_server)
live_transports = dict(transports)
last_seen = {session_id: 100.0 for session_id in contexts}
locks = {session_id: asyncio.Lock() for session_id in contexts}
with (
patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam
session_manager_stateful, "_server_instances", live_transports
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_contexts, contexts, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_context_last_seen, last_seen, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_locks, locks, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_owners, {session_id: "owner" for session_id in contexts}, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_active_request_counts, {}, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_client_info, {}, clear=True
),
):
result = await mcp_server.terminate_mcp_gateway_sessions(user_id="alice")
assert set(live_transports) == {"bob-session-1", "anon-session-1"}
assert set(mcp_server._stateful_session_auth_contexts) == {"bob-session-1", "anon-session-1", "gone-session-1"}
assert set(mcp_server._stateful_session_locks) == {"bob-session-1", "anon-session-1", "gone-session-1"}
assert set(mcp_server._stateful_session_owners) == {"bob-session-1", "anon-session-1", "gone-session-1"}
assert set(mcp_server._stateful_session_auth_context_last_seen) == {
"bob-session-1",
"anon-session-1",
"gone-session-1",
}
transports["alice-session-1"].terminate.assert_awaited_once()
transports["alice-session-2"].terminate.assert_awaited_once()
transports["bob-session-1"].terminate.assert_not_awaited()
transports["anon-session-1"].terminate.assert_not_awaited()
assert result.terminated_sessions == 2
assert sorted(session.session_id_prefix for session in result.sessions) == ["alice-se", "alice-se"]
assert {session.user_id for session in result.sessions} == {"alice"}
assert "key-alice" not in result.model_dump_json()
assert "alice-session-1" not in result.model_dump_json()
@pytest.mark.asyncio
async def test_terminate_mcp_gateway_sessions_prefix_and_user_must_both_match():
try:
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.server import session_manager_stateful
except ImportError:
pytest.skip("MCP server not available")
contexts, transports = _admin_terminate_fixture(mcp_server)
live_transports = dict(transports)
with (
patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam
session_manager_stateful, "_server_instances", live_transports
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_contexts, contexts, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_client_info, {}, clear=True
),
):
mismatch = await mcp_server.terminate_mcp_gateway_sessions(session_id_prefix="alice-session-1", user_id="bob")
assert mismatch.terminated_sessions == 0
assert set(live_transports) == set(transports)
stale = await mcp_server.terminate_mcp_gateway_sessions(session_id_prefix="gone-session-1")
assert stale.terminated_sessions == 0
exact = await mcp_server.terminate_mcp_gateway_sessions(session_id_prefix="alice-session-1", user_id="alice")
assert exact.terminated_sessions == 1
assert set(live_transports) == {"alice-session-2", "bob-session-1", "anon-session-1"}
@pytest.mark.asyncio
async def test_admin_terminated_session_id_gets_404_instead_of_a_fresh_stateless_session():
"""Once an admin closes a session, a client replaying its id must not be silently upgraded to a
new stateless session by the stale-header path; it gets 404 and has to initialize again."""
try:
from starlette.types import Scope
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.server import session_manager_stateful
except ImportError:
pytest.skip("MCP server not available")
session_id = "admin-closed-session-1"
live_transports = {session_id: MagicMock(terminate=AsyncMock())}
contexts = {
session_id: mcp_server.MCPAuthenticatedUser(
user_api_key_auth=UserAPIKeyAuth(api_key="key-alice", user_id="alice"),
)
}
def scope_with_session_header() -> Scope:
return {
"type": "http",
"method": "POST",
"headers": [(b"content-type", b"application/json"), (b"mcp-session-id", session_id.encode())],
}
try:
with (
patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam
session_manager_stateful, "_server_instances", live_transports
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_contexts, contexts, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_client_info, {}, clear=True
),
):
await mcp_server.terminate_mcp_gateway_sessions(session_id_prefix=session_id)
terminated_scope = scope_with_session_header()
send = AsyncMock()
handled = await mcp_server._handle_stale_mcp_session(
terminated_scope, AsyncMock(), send, session_manager_stateful
)
assert handled is True
statuses = [m["status"] for (m,), _ in send.await_args_list if m["type"] == "http.response.start"]
assert statuses == [404]
assert [k for k, _ in terminated_scope["headers"]] == [b"content-type", b"mcp-session-id"]
unknown_scope = scope_with_session_header()
unknown_scope["headers"][1] = (b"mcp-session-id", b"never-seen-session")
assert (
await mcp_server._handle_stale_mcp_session(
unknown_scope, AsyncMock(), AsyncMock(), session_manager_stateful
)
is False
)
assert [k for k, _ in unknown_scope["headers"]] == [b"content-type"]
finally:
mcp_server._admin_terminated_session_ids.clear()
@pytest.mark.asyncio
async def test_admin_terminated_session_id_stays_refused_while_replayed_and_is_forgotten_like_an_idle_session():
"""The refusal window slides on every replay, so a client that keeps retrying is never silently
upgraded to a stateless session no matter how many other sessions an admin closes later; an id
nobody has replayed for a full idle timeout is dropped from the table by the idle sweep."""
try:
from starlette.types import Scope
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.server import session_manager_stateful
except ImportError:
pytest.skip("MCP server not available")
idle_timeout = mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
retrying_id, silent_id = "admin-closed-retrying", "admin-closed-silent"
contexts = {
session_id: mcp_server.MCPAuthenticatedUser(
user_api_key_auth=UserAPIKeyAuth(api_key="key-alice", user_id="alice"),
)
for session_id in (retrying_id, silent_id)
}
live_transports = {session_id: MagicMock(terminate=AsyncMock()) for session_id in contexts}
async def replay(session_id: str, now: float) -> tuple[bool, list[bytes]]:
scope: Scope = {
"type": "http",
"method": "POST",
"headers": [(b"content-type", b"application/json"), (b"mcp-session-id", session_id.encode())],
}
with patch.object( # test-quality-ok: the stale-session handler reads the clock directly; no injectable now
mcp_server.time, "monotonic", return_value=now
):
handled = await mcp_server._handle_stale_mcp_session(
scope, AsyncMock(), AsyncMock(), session_manager_stateful
)
return handled, [k for k, _ in scope["headers"]]
try:
with (
patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam
session_manager_stateful, "_server_instances", live_transports
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_contexts, contexts, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_client_info, {}, clear=True
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_context_last_seen, {}, clear=True
),
):
with patch.object( # test-quality-ok: termination stamps the tombstone from the clock directly; no injectable now
mcp_server.time, "monotonic", return_value=1000.0
):
closed = await mcp_server.terminate_mcp_gateway_sessions(user_id="alice")
assert closed.terminated_sessions == 2
for elapsed in (idle_timeout - 1, 2 * idle_timeout - 2, 3 * idle_timeout - 3):
assert await replay(retrying_id, 1000.0 + elapsed) == (True, [b"content-type", b"mcp-session-id"])
await mcp_server._purge_expired_stateful_session_auth_contexts(now=1000.0 + idle_timeout)
assert set(mcp_server._admin_terminated_session_ids) == {retrying_id}
assert await replay(silent_id, 1000.0 + idle_timeout) == (False, [b"content-type"])
assert await replay(retrying_id, 1000.0 + 4 * idle_timeout) == (False, [b"content-type"])
assert mcp_server._admin_terminated_session_ids == {}
finally:
mcp_server._admin_terminated_session_ids.clear()
@pytest.mark.asyncio
async def test_initialize_request_with_existing_session_tracks_new_session():
try:

View file

@ -395,6 +395,34 @@ async def test_invalidate_clears_every_identity_for_a_server():
assert mock_client.post.call_count == 3
@pytest.mark.asyncio
async def test_per_user_token_delete_evicts_locally_and_broadcasts_to_peer_workers():
"""Revoking a user's OAuth token must not leave peer workers serving it from their in-memory layer."""
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import MCPPerUserTokenCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
local_cache = UserApiKeyCache()
publish = AsyncMock()
token_cache = MCPPerUserTokenCache()
key = token_cache._cache_key("mallory", "srv-oauth") # pyright: ignore[reportPrivateUsage] # asserting the broadcast names the stored key
local_cache.in_memory_cache.set_cache(key, "encrypted-token")
with (
patch.object( # test-quality-ok: the token cache reads the module-level user_api_key_cache singleton; the suite's only seam
proxy_server, "user_api_key_cache", local_cache
),
patch( # test-quality-ok: the redis publisher is module-level; asserting the broadcast without a redis
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation",
new=publish,
),
):
await token_cache.delete("mallory", "srv-oauth")
assert local_cache.in_memory_cache.get_cache(key) is None
publish.assert_awaited_once_with(cache_key=key)
@pytest.mark.asyncio
async def test_m2m_mint_uses_admin_entered_token_url_when_issuer_yield_empties_resolved():
"""A pinned issuer empties the resolved token_url while configured_token_url keeps the

File diff suppressed because it is too large Load diff

View file

@ -57,6 +57,21 @@ def test_xff_honored_from_trusted_peer():
assert via_proxy is True
def test_ipv4_mapped_peer_and_hop_match_ipv4_trusted_ranges():
request = make_request(headers={"x-forwarded-for": "203.0.113.9, ::ffff:10.0.0.5"}, client=("::ffff:10.0.0.1", 1))
ip, via_proxy = resolve_client_ip(request, TRUSTED)
assert ip == "203.0.113.9"
assert via_proxy is True
def test_ipv4_mapped_peer_still_matches_mapped_notation_trusted_range():
config = TrustedProxyConfig(use_forwarded_for=True, trusted_proxy_cidrs=["::ffff:10.0.0.0/104"])
request = make_request(headers={"x-forwarded-for": "203.0.113.9"}, client=("::ffff:10.0.0.1", 1))
ip, via_proxy = resolve_client_ip(request, config)
assert ip == "203.0.113.9"
assert via_proxy is True
def test_spoofed_xff_from_untrusted_peer_is_ignored():
request = make_request(
headers={"x-forwarded-for": "203.0.113.9"}, client=("8.8.8.8", 1)

View file

@ -1,6 +1,7 @@
from __future__ import annotations
from typing import Final
from unittest.mock import patch
import pytest
@ -168,6 +169,54 @@ def test_settings_store_refuses_a_runtime_write_to_a_config_owned_key() -> None:
assert store.source("max_parallel_requests") == "config"
@pytest.mark.timeout(10)
def test_settings_store_clear_removes_every_key_the_config_file_does_not_own() -> None:
store: Final = SettingsStore("general_settings")
store.load_yaml({"master_key": "os.environ/MASTER_KEY"})
store.apply_db_row("general_settings", {"max_parallel_requests": 3, "alerting": ["slack"]})
store.apply_runtime_values({"master_key": "sk-resolved", "alerting": ["slack"]})
store["allow_requests_on_db_unavailable"] = True
del store["alerting"]
store.clear()
assert dict(store) == {"master_key": "sk-resolved"}
assert "alerting" not in store
with pytest.raises(KeyError):
store["max_parallel_requests"]
@pytest.mark.timeout(10)
def test_settings_store_clear_then_refill_matches_a_plain_dict() -> None:
refilled: Final[dict[str, JsonValue]] = {"alerting": ["email"], "max_parallel_requests": 11}
store: Final = SettingsStore("general_settings")
store.update({"max_parallel_requests": 3, "alerting": ["slack"]})
store.clear()
store.update(refilled)
assert dict(store) == refilled
assert tuple(store) == tuple(refilled)
assert len(store) == len(refilled)
@pytest.mark.timeout(10)
@pytest.mark.parametrize("clear", (False, True))
def test_settings_store_survives_a_patch_dict_round_trip_when_the_config_file_owns_a_key(clear: bool) -> None:
store: Final = SettingsStore("general_settings")
store.load_yaml({"master_key": "os.environ/MASTER_KEY"})
store.apply_db_row("general_settings", {"max_parallel_requests": 3})
store.apply_runtime_values({"master_key": "sk-resolved", "max_parallel_requests": 3})
before: Final = dict(store)
with patch.dict(store, {"allow_requests_on_db_unavailable": True}, clear=clear):
assert store["allow_requests_on_db_unavailable"] is True
assert store["master_key"] == "sk-resolved"
assert ("max_parallel_requests" in store) is not clear
assert dict(store) == before
def test_settings_store_reports_the_config_owned_keys_a_write_would_change() -> None:
store: Final = SettingsStore("general_settings")
store.load_yaml({"max_parallel_requests": 3, "ui_access_mode": "admin_only"})

View file

@ -422,7 +422,7 @@ def test_apply_patch_ops_invalid_entitlements_value_raises_400():
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(
op="replace", path="entitlements", value=[{"display": "no value"}]
op="replace", path="entitlements", value=[42]
)
]
)
@ -433,6 +433,22 @@ def test_apply_patch_ops_invalid_entitlements_value_raises_400():
assert exc_info.value.status_code == 400
def test_apply_patch_ops_replace_entitlements_without_value_member_is_stored_as_sent():
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(
op="replace", path="entitlements", value=[{"groups": ["S0506MKA55L"]}]
)
]
)
update_data, _ = _apply_patch_ops(
existing_user=_user_with_metadata({}), patch_ops=patch_ops
)
assert update_data["metadata"]["scim_entitlements"] == [{"groups": ["S0506MKA55L"]}]
def test_apply_patch_ops_add_without_value_raises_400_naming_value_member():
patch_ops = SCIMPatchOp(
Operations=[SCIMPatchOperation(op="add", path="entitlements")]

View file

@ -1,3 +1,4 @@
import json
import logging
import time
from collections.abc import Callable, Mapping, Sequence
@ -1303,6 +1304,75 @@ async def test_update_user_success(mocker):
assert call_args[1]["data"]["teams"] == ["new-team"]
@pytest.mark.asyncio
async def test_update_user_put_with_valueless_entitlements_deactivates_user(scim_test_client, mocker):
existing_user = mocker.MagicMock()
existing_user.teams = []
existing_user.metadata = {"scim_active": True}
updated_user = {
"user_id": "suspend-me",
"user_email": "suspend@example.com",
"user_alias": None,
"teams": [],
"metadata": "{}",
}
response_scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="suspend-me",
userName="suspend-me",
active=False,
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user)
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists",
AsyncMock(return_value=existing_user),
)
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes",
AsyncMock(),
)
set_keys_blocked_mock = mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2._set_user_keys_blocked",
AsyncMock(return_value=1),
)
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(return_value=response_scim_user),
)
async with scim_test_client as client:
response = await client.put(
"/scim/v2/Users/suspend-me",
json={
"schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"],
"userName": "suspend-me",
"emails": [{"value": "suspend@example.com", "primary": True}],
"entitlements": [{"groups": ["S0506MKA55L", "S0506MKA56M"]}],
"roles": [{"display": "Viewer"}],
"active": False,
},
)
assert response.status_code == 200, response.text
assert response.json()["active"] is False
written_metadata = json.loads(mock_prisma_client.db.litellm_usertable.update.call_args.kwargs["data"]["metadata"])
assert written_metadata["scim_active"] is False
assert written_metadata["scim_entitlements"] == [{"groups": ["S0506MKA55L", "S0506MKA56M"]}]
assert written_metadata["scim_roles"] == [{"display": "Viewer"}]
set_keys_blocked_mock.assert_awaited_once_with(user_id="suspend-me", blocked=True)
@pytest.mark.asyncio
@pytest.mark.parametrize("groups", [None, []], ids=["groups-omitted", "groups-empty"])
async def test_update_user_without_groups_preserves_memberships_and_role(mocker, monkeypatch, groups):

View file

@ -1,13 +1,19 @@
import re
from collections.abc import Sequence
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import psycopg
import pytest
from psycopg.rows import dict_row
from pytest_postgresql import factories
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
from litellm.proxy.management_endpoints.common_daily_activity import (
_adjust_dates_for_timezone,
_build_aggregated_sql_query,
@ -169,6 +175,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"endpoint": "/v1/chat/completions",
"api_key": None,
"group_level": 62,
"distinct_api_keys": None,
"spend": 15.0,
"prompt_tokens": 150,
"completion_tokens": 75,
@ -181,31 +188,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"endpoint": "/v1/embeddings",
"api_key": None,
"group_level": 62,
"spend": 3.0,
"prompt_tokens": 30,
"completion_tokens": 0,
"api_requests": 1,
"successful_requests": 1,
},
# (date, endpoint, api_key) — populates the per-key sub-bucket
{
**base,
"date": "2024-01-01",
"endpoint": "/v1/chat/completions",
"api_key": "key-1",
"group_level": 30,
"spend": 15.0,
"prompt_tokens": 150,
"completion_tokens": 75,
"api_requests": 2,
"successful_requests": 2,
},
{
**base,
"date": "2024-01-01",
"endpoint": "/v1/embeddings",
"api_key": "key-2",
"group_level": 30,
"distinct_api_keys": None,
"spend": 3.0,
"prompt_tokens": 30,
"completion_tokens": 0,
@ -219,6 +202,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"endpoint": None,
"api_key": None,
"group_level": 63,
"distinct_api_keys": None,
"spend": 18.0,
"prompt_tokens": 180,
"completion_tokens": 75,
@ -232,12 +216,40 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown():
"endpoint": None,
"api_key": None,
"group_level": 127,
"distinct_api_keys": None,
"spend": 18.0,
"prompt_tokens": 180,
"completion_tokens": 75,
"api_requests": 3,
"successful_requests": 3,
},
# (date, endpoint, api_key) — populates the per-key sub-bucket
{
**base,
"date": "2024-01-01",
"endpoint": "/v1/chat/completions",
"api_key": "key-1",
"group_level": 30,
"distinct_api_keys": 2,
"spend": 15.0,
"prompt_tokens": 150,
"completion_tokens": 75,
"api_requests": 2,
"successful_requests": 2,
},
{
**base,
"date": "2024-01-01",
"endpoint": "/v1/embeddings",
"api_key": "key-2",
"group_level": 30,
"distinct_api_keys": 2,
"spend": 3.0,
"prompt_tokens": 30,
"completion_tokens": 0,
"api_requests": 1,
"successful_requests": 1,
},
]
mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows)
@ -474,9 +486,7 @@ async def test_get_api_key_metadata_recovers_double_hashed_key_via_reverse_hash(
return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com")]
)
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{"digest": double_hashed, "key_alias": "batch-worker", "team_id": "team-1", "user_id": "alice"}
]
return_value=[{"digest": double_hashed, "key_alias": "batch-worker", "team_id": "team-1", "user_id": "alice"}]
)
result = await get_api_key_metadata(
@ -835,6 +845,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
"endpoint": "/v1/chat/completions",
"api_key": None,
"group_level": 62,
"distinct_api_keys": None,
"spend": 10.0,
"prompt_tokens": 100,
"completion_tokens": 50,
@ -847,6 +858,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys():
"endpoint": "/v1/chat/completions",
"api_key": "deleted-key-hash",
"group_level": 30,
"distinct_api_keys": 1,
"spend": 10.0,
"prompt_tokens": 100,
"completion_tokens": 50,
@ -1230,42 +1242,11 @@ class TestBuildAggregatedSqlQuery:
"user-1",
"bedrock/global.anthropic.claude-opus-4-8",
"sk-test",
PTU_SENTINEL_API_KEY,
]
assert "model = $4" in sql
assert "api_key = $5" in sql
def test_model_group_rollups_fall_back_to_model_name(self):
"""Aggregated model_groups rollups must fall back to model for group-less rows.
The (date, model_group) grouping level cannot recover the model column
after the fact (it is rolled up), so the fallback has to happen in SQL;
without it, group-less rows silently vanish from the model_groups
breakdown that the usage UI now renders by default. Group-less rows are
stored as empty strings, not NULL (spend_tracking_utils defaults
model_group to ""), so a plain COALESCE is not enough: the fallback must
be NULLIF-wrapped to catch both
"""
sql, _ = _build_aggregated_sql_query(
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
start_date="2026-07-01",
end_date="2026-07-01",
model=None,
api_key=None,
)
normalized = " ".join(sql.split())
fallback = "COALESCE(NULLIF(model_group, ''), model)"
assert f"{fallback} AS model_group" in normalized
assert (
f"GROUPING(date, api_key, model, {fallback}, "
"custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level" in normalized
)
assert f"(date, {fallback}), (date, {fallback}, api_key)," in normalized
assert "(date, model_group)" not in normalized
assert "COALESCE(model_group, model)" not in normalized
class TestAggregatedEmptyEntityFilter:
_BUILDERS: Final = (_build_aggregated_sql_query, _build_entity_rollup_sql_query)
@ -1285,7 +1266,8 @@ class TestAggregatedEmptyEntityFilter:
normalized = " ".join(sql.split())
assert "IN ()" not in normalized
assert '"team_id" IN' not in normalized
assert params == ["2026-08-01", "2026-08-19"]
sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_aggregated_sql_query else []
assert params == ["2026-08-01", "2026-08-19", *sentinel_params]
@pytest.mark.parametrize("build", _BUILDERS)
def test_empty_entity_list_matches_nothing_rather_than_everything(self, build):
@ -1316,7 +1298,8 @@ class TestAggregatedEmptyEntityFilter:
normalized = " ".join(sql.split())
assert '"team_id" IN ($3, $4)' in normalized
assert "FALSE" not in normalized
assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta"]
sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_aggregated_sql_query else []
assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta", *sentinel_params]
@pytest.mark.asyncio
@ -1341,6 +1324,7 @@ async def test_get_daily_activity_aggregated_empty_result_set():
"mcp_namespaced_tool_name": None,
"endpoint": None,
"group_level": 127,
"distinct_api_keys": None,
"spend": None,
"prompt_tokens": None,
"completion_tokens": None,
@ -1385,6 +1369,305 @@ async def test_get_daily_activity_aggregated_empty_result_set():
assert result.metadata.total_compression_saved_tokens == 0
_aggregated_postgresql_proc: Final = factories.postgresql_proc()
_aggregated_postgresql: Final = factories.postgresql("_aggregated_postgresql_proc")
_DAILY_USER_SPEND_DDL: Final = """
CREATE TABLE "LiteLLM_DailyUserSpend" (
id TEXT PRIMARY KEY,
user_id TEXT,
date TEXT NOT NULL,
api_key TEXT NOT NULL,
model TEXT,
model_group TEXT,
custom_llm_provider TEXT,
mcp_namespaced_tool_name TEXT,
endpoint TEXT,
prompt_tokens BIGINT DEFAULT 0,
completion_tokens BIGINT DEFAULT 0,
cache_read_input_tokens BIGINT DEFAULT 0,
cache_creation_input_tokens BIGINT DEFAULT 0,
compression_saved_tokens BIGINT DEFAULT 0,
compression_savings_spend DOUBLE PRECISION DEFAULT 0,
prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
autorouter_savings_spend DOUBLE PRECISION DEFAULT 0,
spend DOUBLE PRECISION DEFAULT 0,
api_requests BIGINT DEFAULT 0,
successful_requests BIGINT DEFAULT 0,
failed_requests BIGINT DEFAULT 0,
total_response_time_ms BIGINT DEFAULT 0,
timed_requests BIGINT DEFAULT 0
)
"""
def _seed_daily_user_spend(conn: psycopg.Connection, rows: Sequence[tuple[object, ...]]) -> None:
with conn.cursor() as cur:
cur.execute(_DAILY_USER_SPEND_DDL)
cur.executemany(
"""
INSERT INTO "LiteLLM_DailyUserSpend"
(id, user_id, date, api_key, model, model_group, custom_llm_provider,
endpoint, prompt_tokens, spend, api_requests, successful_requests)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
""",
rows,
)
conn.commit()
def _psycopg_query_raw(conn: psycopg.Connection, row_counts: list[int]):
"""Run the proxy's $N-parameterized SQL through psycopg, recording each result size."""
async def query_raw(sql: str, *params: str) -> list[dict[str, object]]:
converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql)
with conn.cursor(row_factory=dict_row) as cur:
cur.execute(
converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query
{f"p{i}": v for i, v in enumerate(params, start=1)},
)
rows: Final = cur.fetchall()
row_counts.append(len(rows))
return rows
return query_raw
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_bounds_api_key_rollups(
_aggregated_postgresql: psycopg.Connection,
):
"""Run the GROUPING SETS statement against real Postgres with more keys than the cap.
key-004 and key-005 tie on spend exactly at the USAGE_TOP_API_KEYS_LIMIT
cutoff; the api_key tiebreaker must keep key-004 and drop key-005. The PTU
sentinel outspends every key but must not take a slot. Excluded keys and the
sentinel still count toward the totals and the model rollup, which come from
the key-free arm.
"""
n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 5
key_rows: Final = [
(
f"row-{i:03d}",
f"user-{i:03d}",
"2026-06-01",
f"key-{i:03d}",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
6.0 if i == 4 else float(i + 1),
1,
1,
)
for i in range(n_keys)
]
sentinel_row: Final = (
"row-ptu",
None,
"2026-06-01",
PTU_SENTINEL_API_KEY,
"gpt-5",
"",
"azure",
None,
0,
1000.0,
0,
0,
)
_seed_daily_user_spend(_aggregated_postgresql, [*key_rows, sentinel_row])
key_spend: Final = sum(6.0 if i == 4 else float(i + 1) for i in range(n_keys))
row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key=None,
)
# Key-free arm: (), (date), (date, model), (date, model_group), two providers,
# one mcp NULL bucket, endpoint plus its NULL bucket = 9 rows regardless of key count.
# Per-key arm: six per-key grouping sets, each capped at the limit.
assert row_counts == [9 + 6 * USAGE_TOP_API_KEYS_LIMIT]
assert result.metadata.total_spend == pytest.approx(key_spend + 1000.0)
assert result.metadata.total_api_requests == n_keys
assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT
assert result.metadata.total_api_keys == n_keys
expected_top: Final = {f"key-{i:03d}" for i in range(6, n_keys)} | {"key-004"}
day: Final = result.results[0]
assert day.metrics.spend == pytest.approx(key_spend + 1000.0)
assert set(day.breakdown.api_keys) == expected_top
assert day.breakdown.api_keys["key-004"].metrics.spend == 6.0
assert "key-005" not in day.breakdown.api_keys
assert PTU_SENTINEL_API_KEY not in day.breakdown.api_keys
assert day.breakdown.models["gpt-5"].metrics.spend == pytest.approx(key_spend + 1000.0)
assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == expected_top
assert day.breakdown.providers["openai"].metrics.spend == pytest.approx(key_spend)
assert set(day.breakdown.providers["openai"].api_key_breakdown) == expected_top
assert day.breakdown.endpoints["/v1/chat/completions"].metrics.api_requests == n_keys
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both_arms(
_aggregated_postgresql: psycopg.Connection,
):
"""An explicit api_key filter must scope the key-free totals and the per-key
rollups to that key alone, so the two arms never disagree."""
rows: Final = [
(
f"row-{i}",
f"user-{i}",
"2026-06-01",
f"key-{i}",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
float(i + 1),
1,
1,
)
for i in range(3)
]
_seed_daily_user_spend(_aggregated_postgresql, rows)
row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key="key-1",
)
assert result.metadata.total_spend == 2.0
assert result.metadata.total_api_keys == 1
day: Final = result.results[0]
assert set(day.breakdown.api_keys) == {"key-1"}
assert day.breakdown.api_keys["key-1"].metrics.spend == 2.0
assert day.breakdown.models["gpt-5"].metrics.spend == 2.0
assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"}
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_reports_exact_limit_key_count_as_complete(
_aggregated_postgresql: psycopg.Connection,
):
"""With exactly USAGE_TOP_API_KEYS_LIMIT keys nothing is dropped, and the
response must say so: total_api_keys equals the limit rather than exceeding it."""
rows: Final = [
(
f"row-{i:03d}",
f"user-{i:03d}",
"2026-06-01",
f"key-{i:03d}",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
float(i + 1),
1,
1,
)
for i in range(USAGE_TOP_API_KEYS_LIMIT)
]
_seed_daily_user_spend(_aggregated_postgresql, rows)
row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key=None,
)
assert result.metadata.total_api_keys == USAGE_TOP_API_KEYS_LIMIT
assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT
assert len(result.results[0].breakdown.api_keys) == USAGE_TOP_API_KEYS_LIMIT
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_model_name(
_aggregated_postgresql: psycopg.Connection,
):
"""Rows stored with an empty or NULL model_group must land in the model_groups
breakdown under their model name instead of vanishing from the usage UI."""
rows: Final = [
("row-0", "user-0", "2026-06-01", "key-0", "gpt-5", "gpt-5-eu", "openai", "/v1/chat/completions", 10, 7.0, 1, 1),
("row-1", "user-1", "2026-06-01", "key-1", "gpt-5", "", "openai", "/v1/chat/completions", 10, 3.0, 1, 1),
("row-2", "user-2", "2026-06-01", "key-2", "claude-x", None, "anthropic", "/v1/messages", 10, 2.0, 1, 1),
]
_seed_daily_user_spend(_aggregated_postgresql, rows)
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, [])
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key=None,
)
breakdown: Final = result.results[0].breakdown
assert set(breakdown.model_groups) == {"gpt-5-eu", "gpt-5", "claude-x"}
assert breakdown.model_groups["gpt-5-eu"].metrics.spend == 7.0
assert breakdown.model_groups["gpt-5"].metrics.spend == 3.0
assert breakdown.model_groups["claude-x"].metrics.spend == 2.0
assert set(breakdown.model_groups["gpt-5"].api_key_breakdown) == {"key-1"}
assert set(breakdown.models) == {"gpt-5", "claude-x"}
assert breakdown.models["gpt-5"].metrics.spend == 10.0
def _no_spend_record():
"""A rollup row for a key with no spend, where SUM() returns NULL (None)."""
return SimpleNamespace(
@ -2170,7 +2453,7 @@ def test_entity_rollup_sql_query_and_api_key_list_filter():
api_key=[],
)
assert "FALSE" in empty_sql
assert empty_params == ["2024-01-01", "2024-01-31"]
assert empty_params == ["2024-01-01", "2024-01-31", PTU_SENTINEL_API_KEY]
@pytest.mark.asyncio
@ -2204,10 +2487,10 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown():
"successful_requests": 0,
}
main_rows = [
{**base, "date": None, "group_level": 127, "spend": 18.0},
{**base, "date": "2024-01-01", "group_level": 63, "spend": 18.0},
{**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "spend": 18.0},
{**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0},
{**base, "date": None, "group_level": 127, "distinct_api_keys": None, "spend": 18.0},
{**base, "date": "2024-01-01", "group_level": 63, "distinct_api_keys": None, "spend": 18.0},
{**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "distinct_api_keys": None, "spend": 18.0},
{**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "distinct_api_keys": 1, "spend": 12.0},
]
entity_base = {
key: value

View file

@ -24,6 +24,7 @@ from litellm.proxy._types import (
LiteLLM_MCPServerTable,
LitellmUserRoles,
MCPTransport,
MCPUserCredentialResponse,
NewMCPServerRequest,
UpdateMCPServerRequest,
UserAPIKeyAuth,
@ -5136,6 +5137,266 @@ async def test_delete_mcp_oauth_user_credential_invalidates_when_record_already_
assert result.has_credential is False
def _make_admin_auth(role: LitellmUserRoles = LitellmUserRoles.PROXY_ADMIN) -> "UserAPIKeyAuth":
return UserAPIKeyAuth(api_key="sk-admin", user_id="admin-user", user_role=role)
@pytest.mark.asyncio
async def test_admin_revokes_another_users_byok_credential():
"""A proxy admin naming user_id deletes and cache-invalidates that user's stored key, not their own."""
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_user_credential,
)
delete_mock = AsyncMock(return_value=None)
invalidate_mock = AsyncMock()
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch( # test-quality-ok: endpoint test stubs the credential row delete
"litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential",
new=delete_mock,
),
patch.object( # test-quality-ok: the cache invalidator is module scoped; the suite's only seam
mcp_server, "_invalidate_byok_cred_cache", new=invalidate_mock
),
):
result = await delete_mcp_user_credential(
server_id="srv-byok-admin",
user_api_key_dict=_make_admin_auth(),
user_id="mallory",
)
delete_mock.assert_awaited_once()
assert delete_mock.await_args.args[1:] == ("mallory", "srv-byok-admin")
invalidate_mock.assert_awaited_once_with("mallory", "srv-byok-admin")
assert result.has_credential is False
@pytest.mark.asyncio
@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
async def test_non_full_admin_cannot_revoke_another_users_byok_credential(role):
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_user_credential,
)
delete_mock = AsyncMock(return_value=None)
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch( # test-quality-ok: endpoint test stubs the credential row delete
"litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential",
new=delete_mock,
),
):
with pytest.raises(HTTPException) as exc_info:
await delete_mcp_user_credential(
server_id="srv-byok-forbidden",
user_api_key_dict=_make_admin_auth(role),
user_id="mallory",
)
assert exc_info.value.status_code == 403
delete_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_user_naming_themselves_still_deletes_own_byok_credential():
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_user_credential,
)
deleted_rows: list[tuple[str, str]] = [] # mutable-ok: test-local recorder for the fake delete boundary
async def _fake_delete_user_credential(_prisma_client: object, user_id: str, server_id: str) -> None:
deleted_rows.append((user_id, server_id))
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch( # test-quality-ok: endpoint test stubs the credential row delete
"litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential",
new=_fake_delete_user_credential,
),
patch.object( # test-quality-ok: the cache invalidator is module scoped; the suite's only seam
mcp_server, "_invalidate_byok_cred_cache", new=AsyncMock()
),
):
result = await delete_mcp_user_credential(
server_id="srv-byok-self",
user_api_key_dict=_make_user_auth("user-self"),
user_id="user-self",
)
assert deleted_rows == [("user-self", "srv-byok-self")]
assert result == MCPUserCredentialResponse(server_id="srv-byok-self", has_credential=False)
@pytest.mark.asyncio
async def test_admin_revokes_another_users_oauth_credential():
"""A proxy admin naming user_id reads, deletes, and cache-invalidates that user's OAuth token."""
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,
)
get_mock = AsyncMock(return_value={"type": "oauth2", "access_token": "mallory-tok"})
delete_mock = AsyncMock(return_value=None)
invalidate_mock = AsyncMock(return_value=None)
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch( # test-quality-ok: endpoint test stubs the stored OAuth token read
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential",
new=get_mock,
),
patch( # test-quality-ok: endpoint test stubs the credential row delete
"litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential",
new=delete_mock,
),
patch.object( # test-quality-ok: the OAuth cache lives on the global manager; the suite's only seam
manager_module.global_mcp_server_manager,
"invalidate_user_oauth_token_cache",
new=invalidate_mock,
),
):
result = await delete_mcp_oauth_user_credential(
server_id="srv-oauth-admin",
user_api_key_dict=_make_admin_auth(),
user_id="mallory",
)
assert get_mock.await_args.args[1:] == ("mallory", "srv-oauth-admin")
assert delete_mock.await_args.args[1:] == ("mallory", "srv-oauth-admin")
invalidate_mock.assert_awaited_once_with("mallory", "srv-oauth-admin")
assert result.has_credential is False
@pytest.mark.asyncio
@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
async def test_non_full_admin_cannot_revoke_another_users_oauth_credential(role):
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_oauth_user_credential,
)
get_mock = AsyncMock(return_value={"type": "oauth2", "access_token": "mallory-tok"})
delete_mock = AsyncMock(return_value=None)
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch( # test-quality-ok: endpoint test stubs the stored OAuth token read
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential",
new=get_mock,
),
patch( # test-quality-ok: endpoint test stubs the credential row delete
"litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential",
new=delete_mock,
),
):
with pytest.raises(HTTPException) as exc_info:
await delete_mcp_oauth_user_credential(
server_id="srv-oauth-forbidden",
user_api_key_dict=_make_admin_auth(role),
user_id="mallory",
)
assert exc_info.value.status_code == 403
get_mock.assert_not_awaited()
delete_mock.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
async def test_admin_lists_every_users_credential_for_a_server(role):
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy._types import MCPServerUserCredentialListItem
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
list_mcp_server_user_credentials,
)
items = (
MCPServerUserCredentialListItem(user_id="alice", credential_type="byok", updated_at="2026-01-01T00:00:00"),
MCPServerUserCredentialListItem(user_id="bob", credential_type="oauth2", updated_at="2026-01-02T00:00:00"),
)
list_mock = AsyncMock(return_value=items)
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch( # test-quality-ok: endpoint test stubs the credential row listing
"litellm.proxy.management_endpoints.mcp_management_endpoints.list_server_user_credentials",
new=list_mock,
),
):
result = await list_mcp_server_user_credentials(
server_id="srv-list-admin",
user_api_key_dict=_make_admin_auth(role),
)
assert list_mock.await_args.args[1:] == ("srv-list-admin",)
assert [(item.user_id, item.credential_type) for item in result] == [("alice", "byok"), ("bob", "oauth2")]
@pytest.mark.asyncio
async def test_non_admin_cannot_list_a_servers_user_credentials():
if not mgmt_endpoints.MCP_AVAILABLE:
pytest.skip("MCP module not installed")
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
list_mcp_server_user_credentials,
)
list_mock = AsyncMock(return_value=())
with (
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=_make_prisma_client(),
),
patch( # test-quality-ok: endpoint test stubs the credential row listing
"litellm.proxy.management_endpoints.mcp_management_endpoints.list_server_user_credentials",
new=list_mock,
),
):
with pytest.raises(HTTPException) as exc_info:
await list_mcp_server_user_credentials(
server_id="srv-list-forbidden",
user_api_key_dict=_make_user_auth("user-plain"),
)
assert exc_info.value.status_code == 403
list_mock.assert_not_awaited()
@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."""
@ -7321,3 +7582,146 @@ class TestGetMCPGatewaySessions:
assert [(group.label, group.count) for group in result.by_client] == [("cursor", 1)]
assert [(group.label, group.count) for group in result.by_user] == [("alice", 1)]
assert "sk-live-secret" not in result.model_dump_json()
class TestDeleteMCPGatewaySessions:
@pytest.fixture(autouse=True)
def _forget_admin_terminated_ids(self):
from litellm.proxy._experimental.mcp_server import server as mcp_server
yield
mcp_server._admin_terminated_session_ids.clear()
@pytest.mark.asyncio
@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
async def test_non_full_admin_forbidden_before_any_session_is_touched(self, role):
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_gateway_sessions,
)
session_id = "gateway-terminate-forbidden-1"
transport = MagicMock(terminate=AsyncMock())
auth_user = mcp_server.MCPAuthenticatedUser(
user_api_key_auth=UserAPIKeyAuth(api_key="sk-live", user_id="alice"),
)
with (
patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam
mcp_server.session_manager_stateful, "_server_instances", {session_id: transport}
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_contexts, {session_id: auth_user}, clear=True
),
):
with pytest.raises(HTTPException) as exc_info:
await delete_mcp_gateway_sessions(
user_api_key_dict=generate_mock_user_api_key_auth(user_role=role),
session_id_prefix=session_id,
user_id=None,
)
assert exc_info.value.status_code == 403
transport.terminate.assert_not_awaited()
assert session_id in mcp_server._stateful_session_auth_contexts
@pytest.mark.asyncio
async def test_requires_a_selector(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_gateway_sessions,
)
with pytest.raises(HTTPException) as exc_info:
await delete_mcp_gateway_sessions(
user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN),
session_id_prefix=None,
user_id=None,
)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_admin_terminates_only_the_selected_session(self):
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_gateway_sessions,
)
from litellm.types.mcp import MCPGatewaySessionsTerminateResponse
target_id = "11111111-target-session"
other_id = "22222222-other-session"
target_transport = MagicMock(terminate=AsyncMock())
other_transport = MagicMock(terminate=AsyncMock())
transports = {target_id: target_transport, other_id: other_transport}
contexts = {
target_id: mcp_server.MCPAuthenticatedUser(
user_api_key_auth=UserAPIKeyAuth(api_key="sk-live-target", user_id="alice"),
),
other_id: mcp_server.MCPAuthenticatedUser(
user_api_key_auth=UserAPIKeyAuth(api_key="sk-live-other", user_id="bob"),
),
}
with (
patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam
mcp_server.session_manager_stateful, "_server_instances", transports
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_contexts, contexts, clear=True
),
):
result = await delete_mcp_gateway_sessions(
user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN),
session_id_prefix=target_id[:8],
user_id=None,
)
assert target_id not in transports
assert other_id in transports
assert target_id not in mcp_server._stateful_session_auth_contexts
assert other_id in mcp_server._stateful_session_auth_contexts
target_transport.terminate.assert_awaited_once()
other_transport.terminate.assert_not_awaited()
assert isinstance(result, MCPGatewaySessionsTerminateResponse)
assert result.terminated_sessions == 1
assert [(s.session_id_prefix, s.user_id) for s in result.sessions] == [(target_id[:8], "alice")]
assert target_id not in result.model_dump_json()
assert "sk-live-target" not in result.model_dump_json()
@pytest.mark.asyncio
async def test_admin_terminates_every_session_of_the_selected_user(self):
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
delete_mcp_gateway_sessions,
)
def auth_user(user_id: str):
return mcp_server.MCPAuthenticatedUser(
user_api_key_auth=UserAPIKeyAuth(api_key=f"sk-live-{user_id}", user_id=user_id),
)
transports = {
"bob-session-1": MagicMock(terminate=AsyncMock()),
"bob-session-2": MagicMock(terminate=AsyncMock()),
"alice-session-1": MagicMock(terminate=AsyncMock()),
}
contexts = {
"bob-session-1": auth_user("bob"),
"bob-session-2": auth_user("bob"),
"alice-session-1": auth_user("alice"),
}
with (
patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam
mcp_server.session_manager_stateful, "_server_instances", transports
),
patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam
mcp_server._stateful_session_auth_contexts, contexts, clear=True
),
):
result = await delete_mcp_gateway_sessions(
user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN),
session_id_prefix=None,
user_id="bob",
)
assert set(transports) == {"alice-session-1"}
assert set(mcp_server._stateful_session_auth_contexts) == {"alice-session-1"}
assert result.terminated_sessions == 2
assert {s.user_id for s in result.sessions} == {"bob"}
assert "sk-live-bob" not in result.model_dump_json()

View file

@ -12,6 +12,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import Request, Response, UploadFile
from pydantic import ValidationError
from starlette.datastructures import FormData, Headers, QueryParams
from starlette.datastructures import UploadFile as StarletteUploadFile
@ -1155,6 +1156,95 @@ def test_resolve_llm_passthrough_timeout_precedence():
assert resolve_llm_passthrough_timeout() == 6.0
def test_resolve_llm_passthrough_timeout_stream_timeout_precedence():
assert (
resolve_llm_passthrough_timeout(
kwargs={"stream": True, "stream_timeout": 1800, "timeout": 45},
)
== 1800.0
)
assert (
resolve_llm_passthrough_timeout(
kwargs={"stream": True, "timeout": 45},
litellm_params={"stream_timeout": 1800, "timeout": 90},
)
== 1800.0
)
assert (
resolve_llm_passthrough_timeout(
kwargs={"stream": True, "timeout": 45},
litellm_params={"timeout": 90},
router_timeout=120,
router_stream_timeout=1800,
)
== 1800.0
)
assert (
resolve_llm_passthrough_timeout(
kwargs={"stream": True},
router_stream_timeout="1800",
)
== 1800.0
)
assert (
resolve_llm_passthrough_timeout(
kwargs={"stream": True},
litellm_params={"timeout": 90},
router_timeout=120,
)
== 90.0
)
assert (
resolve_llm_passthrough_timeout(
kwargs={"stream": False, "stream_timeout": 1800},
litellm_params={"stream_timeout": 1800, "timeout": 90},
router_stream_timeout=1800,
)
== 90.0
)
assert (
resolve_llm_passthrough_timeout(
litellm_params={"stream_timeout": 1800},
router_timeout=120,
router_stream_timeout=1800,
)
== 120.0
)
@pytest.mark.parametrize(
"stream, expected",
[(None, 90.0), (0, 90.0), ("", 90.0), (1, 1800.0), ("yes", 1800.0)],
)
def test_resolve_llm_passthrough_timeout_reads_stream_by_truthiness(stream: object, expected: float):
assert (
resolve_llm_passthrough_timeout(
kwargs={"stream": stream},
litellm_params={"stream_timeout": 1800, "timeout": 90},
)
== expected
)
@pytest.mark.parametrize(
"kwargs, litellm_params, expected",
[
({"stream": True, "stream_timeout": 1800, "timeout": httpx.Timeout(30.0)}, {}, 1800.0),
({"stream": False}, {"stream_timeout": httpx.Timeout(30.0), "timeout": 90}, 90.0),
({"timeout": 45}, {"request_timeout": httpx.Timeout(30.0)}, 45.0),
],
)
def test_resolve_llm_passthrough_timeout_validates_only_the_winning_value(
kwargs: dict[str, object], litellm_params: dict[str, object], expected: float
):
assert resolve_llm_passthrough_timeout(kwargs=kwargs, litellm_params=litellm_params) == expected
def test_resolve_llm_passthrough_timeout_rejects_a_non_numeric_winner():
with pytest.raises(ValidationError):
resolve_llm_passthrough_timeout(kwargs={"timeout": httpx.Timeout(30.0)})
@pytest.mark.asyncio
async def test_pass_through_request_uses_resolved_timeout():
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:

View file

@ -511,3 +511,27 @@ def make_key(
max_budget=max_budget,
**kwargs,
)
@pytest.fixture(autouse=True)
def reset_login_throttle(monkeypatch):
"""Clear the Admin UI failed-login counters between tests.
`client` is session scoped and the counters live in shared module stores with a 300s block
window, so without this a failed sign-in test could block unrelated tests later.
Only the throttle's own keys are removed, so other cache entries remain untouched.
"""
from litellm.constants import LOGIN_THROTTLE_CACHE_KEY_PREFIX
from litellm.proxy import proxy_server as ps
from litellm.proxy.auth.login_throttle import _BLOCKS, _COUNTERS
def _drop_throttle_keys() -> None:
for store in (_COUNTERS, _BLOCKS):
for key in tuple(store.cache_dict) + tuple(store.ttl_dict):
if key.startswith(LOGIN_THROTTLE_CACHE_KEY_PREFIX):
store.delete_cache(key)
monkeypatch.setattr(ps, "redis_usage_cache", None)
_drop_throttle_keys()
yield _drop_throttle_keys
_drop_throttle_keys()

View file

@ -12,8 +12,6 @@ from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock
import pytest
from .conftest import normalize
# ---------------------------------------------------------------------------
@ -29,7 +27,7 @@ def _install_login_mocks(monkeypatch, raise_on_auth: bool = False) -> None:
"""
from litellm.proxy import proxy_server as ps
async def _fake_auth(username, password, master_key, prisma_client, general_settings=None):
async def _fake_auth(username, password, master_key, prisma_client, throttle=None, general_settings=None):
if raise_on_auth:
raise Exception("boom-auth-failure")
fake = MagicMock()
@ -471,3 +469,222 @@ def test_login_form_ignores_open_redirect_return_to(client, monkeypatch):
location = response.headers.get("location", "")
assert "evil.example.com" not in location
assert "/ui" in location # dashboard fallback
# ---------------------------------------------------------------------------
# Failed-login accounting across the login routes (LIT-5285)
# ---------------------------------------------------------------------------
def _install_real_auth(monkeypatch, **settings):
"""Run the real authenticate_user so the throttle inside it is exercised.
prisma_client stays None, so every guess falls through to the credential rejection.
"""
from litellm.proxy import proxy_server as ps
monkeypatch.setenv("UI_USERNAME", "admin")
monkeypatch.setenv("UI_PASSWORD", "right-password")
monkeypatch.setattr(ps, "master_key", "sk-test-master")
monkeypatch.setattr(ps, "prisma_client", None)
monkeypatch.setattr(ps, "premium_user", False)
monkeypatch.setattr(ps, "general_settings", dict(settings))
def _form_login(client, username="admin", password="wrong"):
return client.post("/login", data={"username": username, "password": password}, follow_redirects=False).status_code
def _json_login(client, path, username="admin", password="wrong"):
return client.post(path, json={"username": username, "password": password}).status_code
def _db_user(monkeypatch, email: str):
"""A database user with a stored hash, faked so the route reaches the known-user branch without Postgres."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy import proxy_server as ps
user = MagicMock()
user.user_id = "u-1"
user.user_email = email
user.user_role = "internal_user"
user.password = "scrypt:stored"
repo = MagicMock()
repo.return_value.table.find_first = AsyncMock(return_value=user)
monkeypatch.setattr(ps, "prisma_client", MagicMock())
monkeypatch.setattr("litellm.proxy.auth.login_utils.UserRepository", repo)
monkeypatch.setattr("litellm.proxy.auth.login_utils._rehash_password_if_needed", AsyncMock())
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.verify_password", lambda given, stored: given == "right-db-password"
)
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.generate_key_helper_fn", AsyncMock(return_value={"token": "sk-ui"})
)
monkeypatch.setenv("DATABASE_URL", "postgresql://stub")
def test_budget_is_shared_across_every_login_endpoint(client, monkeypatch, reset_login_throttle):
"""The endpoint is not part of the key, so spending the budget on one route blocks the rest.
Partitioning the counter per endpoint would silently triple the real allowance.
"""
_install_real_auth(
monkeypatch,
max_failed_login_attempts_per_source=20,
control_plane_url="https://cp.example.com",
)
assert [_form_login(client) for _ in range(5)] == [401] * 5
assert [_json_login(client, "/v2/login") for _ in range(5)] == [401] * 5
assert _json_login(client, "/v3/login") == 401, "the eleventh failure crosses the limit and installs the block"
assert _json_login(client, "/v3/login") == 429, "the twelfth attempt must be refused on a third route"
def test_budget_is_shared_across_username_casing(client, monkeypatch, reset_login_throttle):
"""The database lookup is case-insensitive, so casing must not partition the counter."""
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=6)
assert [_json_login(client, "/v2/login", username="admin@corp.com") for _ in range(2)] == [401] * 2
assert [_json_login(client, "/v2/login", username="ADMIN@corp.com") for _ in range(2)] == [401] * 2
assert _json_login(client, "/v2/login", username="Admin@corp.com") == 429
def test_a_refused_attempt_carries_retry_after(client, monkeypatch, reset_login_throttle):
"""The 429 tells the caller how long the block has left."""
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2, failed_login_block_seconds=77)
assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401]
refused = client.post("/v2/login", json={"username": "admin", "password": "wrong"})
assert refused.status_code == 429
assert refused.headers.get("retry-after") == "77"
def test_the_form_returns_a_human_readable_lockout_page(client, monkeypatch, reset_login_throttle):
"""The no-JavaScript form must render a wait page when its POST is throttled."""
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2, failed_login_block_seconds=77)
assert [_form_login(client) for _ in range(2)] == [401, 401]
refused = client.post("/login", data={"username": "admin", "password": "wrong"})
assert refused.status_code == 429
assert refused.headers.get("content-type", "").startswith("text/html")
assert "Try again in about 77 seconds" in refused.text
assert refused.headers.get("retry-after") == "77"
def test_a_second_username_from_the_same_source_still_gets_through(client, monkeypatch, reset_login_throttle):
"""The pair block is per username, so one account's block cannot take the office down with it."""
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2)
assert [_json_login(client, "/v2/login", username="admin") for _ in range(3)] == [401, 401, 429]
assert _json_login(client, "/v2/login", username="someone-else@example.com") == 401
def test_a_spray_across_usernames_is_blocked_on_the_source_when_the_source_is_attributable(
client, monkeypatch, reset_login_throttle
):
"""A fresh username per guess keeps every pair at one, so the address is what stops it."""
_install_real_auth(monkeypatch, trusted_proxy_ranges=["10.0.0.0/8"], max_failed_login_attempts_per_source=4)
sprayed = [_json_login(client, "/v2/login", username=f"sprayed-{i}@corp.com") for i in range(5)]
assert sprayed == [401] * 5
assert _json_login(client, "/v2/login", username="sprayed-6@corp.com") == 429
def test_a_spray_across_usernames_is_not_blocked_without_trusted_proxy_ranges(
client, monkeypatch, reset_login_throttle
):
"""Without a configured proxy range the peer address is whoever fronts the proxy, shared by every
client, so a source-wide block would block them all and the source scope stays off."""
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=4)
sprayed = [_json_login(client, "/v2/login", username=f"sprayed-{i}@corp.com") for i in range(8)]
assert sprayed == [401] * 8
def test_a_spray_across_usernames_is_blocked_on_the_source_with_an_empty_trusted_proxy_ranges(
client, monkeypatch, reset_login_throttle
):
"""An explicit empty list says nothing fronts the proxy, so the peer address is the client and the
source scope is on. A forwarded header from an untrusted peer is ignored rather than trusted."""
_install_real_auth(monkeypatch, trusted_proxy_ranges=[], max_failed_login_attempts_per_source=4)
sprayed = [
client.post(
"/v2/login",
json={"username": f"sprayed-{i}@corp.com", "password": "wrong"},
headers={"x-forwarded-for": f"203.0.113.{i}"},
).status_code
for i in range(5)
]
assert sprayed == [401] * 5
assert _json_login(client, "/v2/login", username="sprayed-6@corp.com") == 429
def test_the_configured_admin_password_is_refused_while_blocked(client, monkeypatch, reset_login_throttle):
"""The env credentials get no bypass: a bypass would make them the one password worth guessing without
limit. An operator who is blocked administers the proxy with the master key over the API meanwhile."""
from unittest.mock import AsyncMock, patch
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2)
monkeypatch.setenv("DATABASE_URL", "postgresql://stub")
assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429]
with (
patch( # test-quality-ok: the admin sign-in upserts the admin row; faked so no DB is needed
"litellm.proxy.auth.login_utils.user_update", new=AsyncMock()
),
patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed
"litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"})
),
):
assert _json_login(client, "/v2/login", password="right-password") == 429
reset_login_throttle()
assert _json_login(client, "/v2/login", password="right-password") == 200
def test_the_master_key_as_a_bearer_token_still_works_while_the_ui_password_is_blocked(
client, monkeypatch, reset_login_throttle
):
"""Lockout recovery: the API path with the master key never enters the sign-in throttle."""
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2)
assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429]
assert client.get("/models", headers={"Authorization": "Bearer sk-not-the-master"}).status_code >= 400
assert client.get("/models", headers={"Authorization": "Bearer sk-test-master"}).status_code == 200
assert _json_login(client, "/v2/login", password="right-password") == 429, "the UI block is unaffected"
def test_a_database_users_correct_password_is_refused_while_blocked(client, monkeypatch, reset_login_throttle):
"""The block is hard: while it lasts, nothing from that source signs in as that user, right password or not,
and the block is not extended by the refused attempts."""
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2, failed_login_block_seconds=64)
_db_user(monkeypatch, "user@corp.com")
assert [_json_login(client, "/v2/login", username="user@corp.com") for _ in range(3)] == [401, 401, 429]
refused = client.post("/v2/login", json={"username": "user@corp.com", "password": "right-db-password"})
assert refused.status_code == 429
assert refused.headers.get("retry-after") == "64"
reset_login_throttle()
assert _json_login(client, "/v2/login", username="user@corp.com", password="right-db-password") == 200
def test_sign_in_succeeds_again_once_the_block_is_cleared(client, monkeypatch, reset_login_throttle):
"""A cleared store lets the same username straight back to a plain credential check."""
_install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2)
assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429]
reset_login_throttle()
assert _json_login(client, "/v2/login") == 401

View file

@ -26,7 +26,6 @@ from fastapi.encoders import jsonable_encoder
from fastapi.staticfiles import StaticFiles
from fastapi.testclient import TestClient
import litellm
import litellm.proxy.proxy_server as proxy_server_module
from litellm.caching.caching import RedisCache
@ -41,6 +40,7 @@ from litellm.proxy._types import (
TokenCountRequest,
UserAPIKeyAuth,
)
from litellm.proxy.auth.login_throttle import LoginThrottle
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.hooks.parallel_request_limiter_v3 import RequestRateLimiterStash
from litellm.proxy.proxy_server import app, initialize, openai_exception_handler
@ -139,13 +139,14 @@ def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch):
}
assert response.cookies.get("token") == "signed-token"
mock_authenticate_user.assert_awaited_once_with(
username="alice",
password="secret",
master_key="test-master-key",
prisma_client=mock_prisma_client,
general_settings={},
)
mock_authenticate_user.assert_awaited_once()
auth_kwargs = mock_authenticate_user.call_args.kwargs
assert auth_kwargs["username"] == "alice"
assert auth_kwargs["password"] == "secret"
assert auth_kwargs["master_key"] == "test-master-key"
assert auth_kwargs["prisma_client"] is mock_prisma_client
assert auth_kwargs["general_settings"] == {}
assert isinstance(auth_kwargs["throttle"], LoginThrottle), "the endpoint must thread a throttle through"
mock_create_ui_token_object.assert_called_once_with(
login_result=mock_login_result,
general_settings={},
@ -3410,6 +3411,60 @@ async def test_load_config_user_url_validation_handles_null_and_string_false(tmp
assert litellm.user_url_validation is False
@pytest.mark.asyncio
async def test_load_config_warns_per_worker_login_counters_without_general_settings(tmp_path, monkeypatch, caplog):
"""Regression: the failed-login throttle is on by default, so a multi-worker proxy with no
Redis must hear that its counters are per worker even when the config has no general_settings."""
import logging
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.auth.login_throttle import warn_login_counters_are_per_worker
from litellm.proxy.proxy_server import ProxyConfig
for redis_var in ("REDIS_HOST", "REDIS_URL", "REDIS_CLUSTER_NODES", "REDIS_SENTINEL_NODES"):
monkeypatch.delenv(redis_var, raising=False)
monkeypatch.setenv("NUM_WORKERS", "4")
monkeypatch.setattr(proxy_server, "redis_usage_cache", None)
warn_login_counters_are_per_worker.cache_clear()
config_file = tmp_path / "config.yaml"
config_file.write_text("model_list: []\n")
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file))
assert "Running 4 workers but Redis is not configured" in caplog.text
@pytest.mark.asyncio
async def test_load_config_warns_that_the_source_login_limit_is_off_without_trusted_proxy_ranges(
tmp_path, monkeypatch, caplog
):
"""The per-source failed-login limit is skipped when the source cannot be attributed, and the
operator must be told so at startup. Both a configured range and an explicit empty list (no
proxies, the peer is the source) silence it, since both keep the limit on."""
import logging
from litellm.proxy.auth.login_throttle import warn_source_login_limit_is_off
from litellm.proxy.proxy_server import ProxyConfig
monkeypatch.setenv("NUM_WORKERS", "1")
warn_source_login_limit_is_off.cache_clear()
config_file = tmp_path / "config.yaml"
config_file.write_text("model_list: []\n")
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file))
assert "trusted_proxy_ranges is not set" in caplog.text
for configured in ("['10.0.0.0/8']", "[]"):
caplog.clear()
warn_source_login_limit_is_off.cache_clear()
config_file.write_text(f"model_list: []\ngeneral_settings:\n trusted_proxy_ranges: {configured}\n")
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file))
assert "trusted_proxy_ranges is not set" not in caplog.text, configured
@pytest.mark.asyncio
async def test_load_environment_variables_direct_and_os_environ():
"""
@ -7536,25 +7591,35 @@ async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_servi
deleted. The proxy's own registry of live pass-through routes is what decides whether
a request is routed upstream or falls through to the auth error, so it has to lose the
entry on the reload rather than at the next process restart."""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import InitPassThroughEndpointHelpers
from litellm.proxy.proxy_server import ProxyConfig
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
_registered_pass_through_routes,
)
from litellm.proxy.proxy_server import ProxyConfig, app
path: Final = f"/v1/deleted-{uuid.uuid4().hex[:8]}"
db_endpoint: Final = {"id": "db-1", "path": path, "target": "https://example.com/post"}
prior_routes: Final = list(app.routes)
prior_registry: Final = dict(_registered_pass_through_routes)
def live_routes() -> set[str]:
return {route for route in InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() if path in route}
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {}) # test-quality-ok: the method reads this module global; no injection seam
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", None) # test-quality-ok: module global holding the YAML endpoints; this case has none
with settings, yaml_endpoints:
pc = ProxyConfig()
await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
assert live_routes(), "the stored endpoint should be serving before the row is deleted"
try:
with settings, yaml_endpoints:
pc = ProxyConfig()
await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
assert live_routes(), "the stored endpoint should be serving before the row is deleted"
await pc._update_general_settings(db_general_settings={})
await pc._update_general_settings(db_general_settings={})
assert live_routes() == set()
assert live_routes() == set()
finally:
app.routes[:] = prior_routes
_registered_pass_through_routes.clear()
_registered_pass_through_routes.update(prior_registry)
@pytest.mark.asyncio
@ -7564,15 +7629,18 @@ async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_rout
serving untouched. The stored entry never gets a route of its own."""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
_registered_pass_through_routes,
initialize_pass_through_endpoints,
)
from litellm.proxy.proxy_server import ProxyConfig
from litellm.proxy.proxy_server import ProxyConfig, app
marker: Final = uuid.uuid4().hex[:8]
config_path: Final = f"/v1/kept-{marker}"
db_path: Final = f"/v1/ignored-{marker}"
config_endpoint: Final = {"id": f"cfg-{marker}", "path": config_path, "target": "https://example.com/post"}
db_endpoint: Final = {"id": f"db-{marker}", "path": db_path, "target": "https://example.com/post"}
prior_routes: Final = list(app.routes)
prior_registry: Final = dict(_registered_pass_through_routes)
def live_paths() -> set[str]:
registered: Final = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes()
@ -7580,17 +7648,22 @@ async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_rout
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [config_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [config_endpoint]) # test-quality-ok: module global holding the YAML endpoints the reload merges in
with settings, yaml_endpoints:
await initialize_pass_through_endpoints(pass_through_endpoints=[config_endpoint])
assert live_paths() == {config_path}
try:
with settings, yaml_endpoints:
await initialize_pass_through_endpoints(pass_through_endpoints=[config_endpoint])
assert live_paths() == {config_path}
pc = ProxyConfig()
await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
assert live_paths() == {config_path}
pc = ProxyConfig()
await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
assert live_paths() == {config_path}
await pc._update_general_settings(db_general_settings={})
await pc._update_general_settings(db_general_settings={})
assert live_paths() == {config_path}
assert live_paths() == {config_path}
finally:
app.routes[:] = prior_routes
_registered_pass_through_routes.clear()
_registered_pass_through_routes.update(prior_registry)
def _fill_user_api_key_cache(cache: DualCache, count: int) -> None:
@ -13952,6 +14025,35 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the
)
@pytest.mark.asyncio
async def test_login_throttle_settings_are_not_hot_applied_from_the_database():
"""LIT-5285: a stored sign-in limit does not take effect on a live worker.
_update_general_settings copies an allowlist of keys out of the DB row on every config
poll. Adding these to it would let a stored value outrank config.yaml without a restart,
so an operator locked out by a bad value could not fix it by editing YAML and restarting.
"""
import litellm.proxy.proxy_server as ps
from litellm.proxy.proxy_server import ProxyConfig
original = dict(ps.general_settings)
try:
ps.general_settings.clear()
await ProxyConfig()._update_general_settings(
db_general_settings={
"max_failed_login_attempts_per_source": 999,
"failed_login_window_seconds": 1,
"failed_login_block_seconds": 1,
}
)
assert "max_failed_login_attempts_per_source" not in ps.general_settings
assert "failed_login_window_seconds" not in ps.general_settings
assert "failed_login_block_seconds" not in ps.general_settings
finally:
ps.general_settings.clear()
ps.general_settings.update(original)
@pytest.mark.asyncio
async def test_load_config_router_authorizes_fallback_targets_against_the_calling_key(tmp_path):
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
@ -14284,3 +14386,74 @@ async def test_token_counter_loads_a_custom_tokenizer_once_per_identifier_revisi
]
finally:
litellm.utils._select_custom_tokenizer_helper.cache_clear()
@pytest.mark.asyncio
async def test_auth_cache_invalidation_subscriber_evicts_byok_credentials_cached_by_this_worker():
"""A peer worker's BYOK revocation broadcast must reach this worker's BYOK credential cache."""
from redis.asyncio import Redis
from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
byok_credential_cache,
byok_credential_cache_key,
cache_byok_credential,
get_cached_byok_credential,
)
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
class _QueuePubSub:
def __init__(self, messages: list[object]) -> None:
self.queue: asyncio.Queue[object] = asyncio.Queue()
for message in messages:
self.queue.put_nowait(message)
async def subscribe(self, *channels: str) -> None:
return None
async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> object | None:
try:
return await asyncio.wait_for(self.queue.get(), timeout)
except asyncio.TimeoutError:
return None
async def aclose(self) -> None:
return None
class _PubSubRedisClient(Redis):
def __init__(self, pubsub: _QueuePubSub) -> None:
self._scripted_pubsub = pubsub
def pubsub(self) -> _QueuePubSub:
return self._scripted_pubsub
class _FakeRedisCache:
namespace = None
def __init__(self, client: object) -> None:
self._client = client
def init_async_client(self) -> object:
return self._client
byok_credential_cache.flush_cache()
cache_byok_credential("mallory", "srv-byok", "sk-revoked-elsewhere")
message: Final = {
"type": "message",
"data": json.dumps({"cache_key": byok_credential_cache_key("mallory", "srv-byok")}).encode(),
}
proxy_config: Final = proxy_server_module.ProxyConfig()
proxy_config.start_auth_cache_invalidation_subscriber(
redis_cache=_FakeRedisCache(_PubSubRedisClient(_QueuePubSub([message]))), # pyright: ignore[reportArgumentType] # fake pub/sub capable redis; no live redis in this unit test
user_api_key_cache=UserApiKeyCache(),
)
try:
for _ in range(200):
if get_cached_byok_credential("mallory", "srv-byok") is None:
break
await asyncio.sleep(0.01)
evicted: Final = get_cached_byok_credential("mallory", "srv-byok") is None
finally:
await proxy_config.stop_auth_cache_invalidation_subscriber()
byok_credential_cache.flush_cache()
assert evicted, "the subscriber does not evict the BYOK credential cache on a peer worker's broadcast"

View file

@ -5538,6 +5538,65 @@ def test_update_kwargs_with_deployment_uses_pass_through_request_timeout():
assert kwargs["timeout"] == 6.0
def _passthrough_timeout(router: litellm.Router, deployment: dict, stream: bool) -> float:
kwargs: Final[dict] = {"stream": stream}
router._update_kwargs_with_deployment(
deployment=deployment,
kwargs=kwargs,
function_name="_ageneric_api_call_with_fallbacks",
)
return kwargs["timeout"]
def test_update_kwargs_with_deployment_passthrough_honors_stream_timeout():
router = litellm.Router(
model_list=[
{
"model_name": "anthropic-with-stream-timeout",
"litellm_params": {
"model": "anthropic/claude-sonnet-4-5",
"api_key": "fake-key",
"timeout": 60,
"stream_timeout": 1800,
},
},
{
"model_name": "anthropic-router-default",
"litellm_params": {
"model": "anthropic/claude-sonnet-4-5",
"api_key": "fake-key",
"timeout": 60,
},
},
],
timeout=120,
stream_timeout=900,
)
per_deployment, router_default = router.model_list
assert _passthrough_timeout(router, per_deployment, stream=True) == 1800.0
assert _passthrough_timeout(router, router_default, stream=True) == 900.0
assert _passthrough_timeout(router, per_deployment, stream=False) == 60.0
assert _passthrough_timeout(router, router_default, stream=False) == 60.0
def test_update_kwargs_with_deployment_passthrough_router_stream_timeout_sources():
deployment: Final[dict] = {
"model_name": "anthropic-router-default",
"litellm_params": {"model": "anthropic/claude-sonnet-4-5", "api_key": "fake-key"},
}
string_router = litellm.Router(model_list=[deployment], timeout=120, stream_timeout="900")
default_router = litellm.Router(
model_list=[deployment],
timeout=120,
default_litellm_params={"stream_timeout": 700},
)
assert _passthrough_timeout(string_router, string_router.model_list[0], stream=True) == 900.0
assert _passthrough_timeout(default_router, default_router.model_list[0], stream=True) == 700.0
assert _passthrough_timeout(default_router, default_router.model_list[0], stream=False) == 120.0
@pytest.mark.asyncio
async def test_router_acompletion_with_unknown_model_and_default_fallback():
"""
@ -8229,6 +8288,16 @@ class TestRouterRequestTimeoutPropagation:
== 60
)
def test_passthrough_prefers_request_timeout_over_router_timeout(self, explicit_request_timeout):
router = self._make_router(timeout=330)
deployment: Final = router.model_list[0]
assert _passthrough_timeout(router, deployment, stream=False) == 300.0
assert _passthrough_timeout(router, deployment, stream=True) == 300.0
def test_passthrough_stream_timeout_still_wins_over_request_timeout(self, explicit_request_timeout):
router = self._make_router(timeout=330, stream_timeout=45)
assert _passthrough_timeout(router, router.model_list[0], stream=True) == 45.0
# ---------------------------------------------------------------------------
# Deferred-stream eager-fetch tests

View file

@ -178,4 +178,28 @@ describe("CacheLeakageCard", () => {
screen.queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."),
).not.toBeInTheDocument();
});
it("says which keys are missing from the key ranking when the proxy capped the per-key lists", () => {
const day = dayWithKeys("2026-07-12", {
"hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }),
});
renderWith([day], { apiKeyTruncation: { limit: 100, total: 3000 } });
expect(screen.getByRole("note")).toHaveTextContent(
"Only the 100 highest-spend keys of 3,000 are loaded, so a lower-spend key that leaks more is not listed here.",
);
fireEvent.click(screen.getByRole("tab", { name: "By model" }));
expect(screen.queryByRole("note")).not.toBeInTheDocument();
});
it("keeps the key ranking note off when every key was loaded", () => {
const day = dayWithKeys("2026-07-12", {
"hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }),
});
renderWith([day]);
expect(screen.queryByRole("note")).not.toBeInTheDocument();
});
});

View file

@ -81,7 +81,7 @@ const SortableHead = ({
};
const CacheLeakageCard: React.FC<CacheLeakageCardProps> = ({ activity }) => {
const { dateValue, onDateChange, results, loading, isFetchingMore } = activity;
const { dateValue, onDateChange, results, loading, isFetchingMore, apiKeyTruncation } = activity;
const [dimension, setDimension] = useState<CacheLeakageDimension>("key");
const [sort, setSort] = useState<SortState>({ column: "potentialSavings", dir: "desc" });
const leakage = useMemo(() => computeCacheLeakage(results, dimension), [results, dimension]);
@ -123,6 +123,13 @@ const CacheLeakageCard: React.FC<CacheLeakageCardProps> = ({ activity }) => {
</Tabs>
</CardHeader>
<CardContent>
{dimension === "key" && apiKeyTruncation !== undefined && (
<p className="mb-2 text-sm text-muted-foreground" role="note">
Only the {apiKeyTruncation.limit.toLocaleString()} highest-spend keys of{" "}
{apiKeyTruncation.total.toLocaleString()} are loaded, so a lower-spend key that leaks more is not listed
here. Raise USAGE_TOP_API_KEYS_LIMIT on the proxy to load more keys.
</p>
)}
{rows.length > 0 && isFetchingMore && (
<p className="mb-2 text-sm text-muted-foreground">
Data is still loading; rows and totals will update as the rest of the range arrives.

View file

@ -4,12 +4,13 @@ import { describe, expect, it, vi } from "vitest";
const mockUsePaginatedDailyActivity = vi.fn();
const mockCancel = vi.fn();
let mockMetadata: Record<string, number> = {};
vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", () => ({
usePaginatedDailyActivity: (args: unknown) => {
mockUsePaginatedDailyActivity(args);
return {
data: { results: [] },
data: { results: [], metadata: mockMetadata },
loading: false,
isFetchingMore: false,
progress: { currentPage: 4, totalPages: 9 },
@ -80,4 +81,18 @@ describe("useDailyActivityRange", () => {
expect(mockUsePaginatedDailyActivity).toHaveBeenLastCalledWith(expect.objectContaining({ enabled: false }));
});
it("reports how many keys the proxy left out of the per-key lists", () => {
mockMetadata = { api_key_limit: 100, total_api_keys: 3000 };
const { result } = renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin"));
expect(result.current.apiKeyTruncation).toEqual({ limit: 100, total: 3000 });
});
it("reports no key truncation when every key fit under the proxy limit", () => {
mockMetadata = { api_key_limit: 100, total_api_keys: 100 };
const { result } = renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin"));
expect(result.current.apiKeyTruncation).toBeUndefined();
});
});

View file

@ -1,6 +1,7 @@
import { useMemo, useState } from "react";
import { userDailyActivityAggregatedCall, userDailyActivityCall } from "@/components/networking";
import { ApiKeyTruncation, getApiKeyTruncation } from "@/components/EntityUsageExport/exportBlockedReason";
import { DailyData } from "@/components/UsagePage/types";
import { spendScopeUserId } from "@/utils/roles";
import { usePaginatedDailyActivity } from "@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity";
@ -22,6 +23,7 @@ export interface DailyActivityRange {
cancelled: boolean;
failed: boolean;
cancel: () => void;
apiKeyTruncation?: ApiKeyTruncation;
}
/**
@ -78,6 +80,7 @@ export const useScopedDailyActivityRange = (
cancelled,
failed,
cancel,
apiKeyTruncation: getApiKeyTruncation(data.metadata?.api_key_limit, data.metadata?.total_api_keys),
};
};

View file

@ -1,13 +1,15 @@
import React from "react";
import { render, screen, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { MCPGatewaySessionsTab, formatIdleSeconds } from "./MCPGatewaySessionsTab";
import { MCPGatewaySessionsTab, describeTerminateResult, formatIdleSeconds } from "./MCPGatewaySessionsTab";
import * as networking from "@/components/networking";
import type { MCPGatewaySessionsResponse } from "@/components/mcp_tools/types";
import type { MCPGatewaySessionsResponse, MCPGatewaySessionsTerminateResponse } from "@/components/mcp_tools/types";
vi.mock("@/components/networking", () => ({
fetchMCPGatewaySessions: vi.fn(),
terminateMCPGatewaySessions: vi.fn(),
}));
const REPORT: MCPGatewaySessionsResponse = {
@ -64,11 +66,11 @@ const REPORT: MCPGatewaySessionsResponse = {
],
};
const renderTab = () => {
const renderTab = ({ canTerminate = false }: { canTerminate?: boolean } = {}) => {
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } });
return render(
<QueryClientProvider client={queryClient}>
<MCPGatewaySessionsTab accessToken="token" />
<MCPGatewaySessionsTab accessToken="token" canTerminate={canTerminate} />
</QueryClientProvider>,
);
};
@ -83,6 +85,17 @@ describe("formatIdleSeconds", () => {
});
});
describe("describeTerminateResult", () => {
it("pluralizes the session count and names the worker", () => {
expect(describeTerminateResult({ worker_pid: 9, terminated_sessions: 1, sessions: [] })).toBe(
"Disconnected 1 session on worker pid 9.",
);
expect(describeTerminateResult({ worker_pid: 9, terminated_sessions: 0, sessions: [] })).toBe(
"Disconnected 0 sessions on worker pid 9.",
);
});
});
describe("MCPGatewaySessionsTab", () => {
beforeEach(() => {
vi.clearAllMocks();
@ -135,4 +148,73 @@ describe("MCPGatewaySessionsTab", () => {
expect(alert).toHaveTextContent("Could not load live connections");
expect(alert).toHaveTextContent("Admin access required");
});
it("hides every disconnect control from a read-only admin", async () => {
vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(REPORT);
renderTab({ canTerminate: false });
await screen.findByRole("region", { name: "Live sessions" });
expect(screen.queryByRole("button", { name: /^Disconnect/ })).not.toBeInTheDocument();
});
it("disconnects one session by its displayed prefix after confirmation and refetches", async () => {
const user = userEvent.setup();
const terminated: MCPGatewaySessionsTerminateResponse = {
worker_pid: 4242,
terminated_sessions: 1,
sessions: [REPORT.sessions[1]],
};
vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(REPORT);
vi.mocked(networking.terminateMCPGatewaySessions).mockResolvedValue(terminated);
renderTab({ canTerminate: true });
await user.click(await screen.findByRole("button", { name: "Disconnect session bbbb2222" }));
expect(networking.terminateMCPGatewaySessions).not.toHaveBeenCalled();
const dialog = await screen.findByRole("alertdialog");
expect(dialog).toHaveTextContent("session bbbb2222");
await user.click(within(dialog).getByRole("button", { name: "Disconnect" }));
const status = await screen.findByText("Disconnected 1 session on worker pid 4242.", { exact: false });
expect(status).toBeInTheDocument();
expect(networking.terminateMCPGatewaySessions).toHaveBeenCalledWith("token", { session_id_prefix: "bbbb2222" });
expect(networking.fetchMCPGatewaySessions).toHaveBeenCalledTimes(2);
});
it("disconnects every session of a user from the by-user table", async () => {
const user = userEvent.setup();
vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(REPORT);
vi.mocked(networking.terminateMCPGatewaySessions).mockResolvedValue({
worker_pid: 4242,
terminated_sessions: 2,
sessions: [REPORT.sessions[0], REPORT.sessions[1]],
});
renderTab({ canTerminate: true });
const byUser = await screen.findByRole("region", { name: "Sessions by user" });
expect(within(byUser).queryByRole("button", { name: /\(unknown\)/ })).not.toBeInTheDocument();
await user.click(within(byUser).getByRole("button", { name: "Disconnect all sessions for user alice" }));
const dialog = await screen.findByRole("alertdialog");
expect(dialog).toHaveTextContent("every live session opened by user alice");
await user.click(within(dialog).getByRole("button", { name: "Disconnect" }));
expect(await screen.findByText(/Disconnected 2 sessions on worker pid 4242\./)).toBeInTheDocument();
expect(networking.terminateMCPGatewaySessions).toHaveBeenCalledWith("token", { user_id: "alice" });
});
it("shows the API error when a disconnect is refused", async () => {
const user = userEvent.setup();
vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(REPORT);
vi.mocked(networking.terminateMCPGatewaySessions).mockRejectedValue(
new Error("Proxy admin access required to terminate MCP gateway sessions."),
);
renderTab({ canTerminate: true });
await user.click(await screen.findByRole("button", { name: "Disconnect session aaaa1111" }));
await user.click(within(await screen.findByRole("alertdialog")).getByRole("button", { name: "Disconnect" }));
const alert = await screen.findByRole("alert");
expect(alert).toHaveTextContent("Could not disconnect");
expect(alert).toHaveTextContent("Proxy admin access required to terminate MCP gateway sessions.");
expect(screen.getByRole("region", { name: "Live sessions" })).toBeInTheDocument();
});
});

View file

@ -1,14 +1,27 @@
"use client";
import React from "react";
import { useQuery } from "@tanstack/react-query";
import { RefreshCw } from "lucide-react";
import React, { useState } from "react";
import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query";
import { RefreshCw, Unplug } from "lucide-react";
import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert";
import {
AlertDialog,
AlertDialogContent,
AlertDialogDescription,
AlertDialogFooter,
AlertDialogHeader,
AlertDialogTitle,
} from "@/components/ui/alert-dialog";
import { Button } from "@/components/ui/button";
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { fetchMCPGatewaySessions } from "@/components/networking";
import type { MCPGatewaySessionGroupCount, MCPGatewaySessionsResponse } from "@/components/mcp_tools/types";
import { fetchMCPGatewaySessions, terminateMCPGatewaySessions } from "@/components/networking";
import type {
MCPGatewaySessionGroupCount,
MCPGatewaySessionSelector,
MCPGatewaySessionsResponse,
MCPGatewaySessionsTerminateResponse,
} from "@/components/mcp_tools/types";
import { createQueryKeys } from "@/app/(dashboard)/hooks/common/queryKeysFactory";
const mcpGatewaySessionKeys = createQueryKeys("mcpGatewaySessions");
@ -28,6 +41,16 @@ function groupLabel(label: string | null): string {
return label === "" ? '""' : label;
}
export function describeSelector(selector: MCPGatewaySessionSelector): string {
if (selector.user_id !== undefined) return `every live session opened by user ${groupLabel(selector.user_id)}`;
return `session ${selector.session_id_prefix}`;
}
export function describeTerminateResult(result: MCPGatewaySessionsTerminateResponse): string {
const noun = result.terminated_sessions === 1 ? "session" : "sessions";
return `Disconnected ${result.terminated_sessions} ${noun} on worker pid ${result.worker_pid}.`;
}
function StatCard({ label, value }: { label: string; value: number }) {
return (
<div className="bg-card border border-border rounded-lg px-4 py-3">
@ -37,14 +60,37 @@ function StatCard({ label, value }: { label: string; value: number }) {
);
}
function DisconnectUserButton({
userId,
onDisconnectUser,
}: {
userId: string | null;
onDisconnectUser: (userId: string) => void;
}) {
if (userId === null || userId === "") return null;
return (
<Button
variant="outline"
size="sm"
onClick={() => onDisconnectUser(userId)}
aria-label={`Disconnect all sessions for user ${groupLabel(userId)}`}
>
<Unplug className="size-4" />
Disconnect all
</Button>
);
}
function GroupCountTable({
title,
groups,
labelHeader,
onDisconnectUser,
}: {
title: string;
groups: MCPGatewaySessionGroupCount[];
labelHeader: string;
onDisconnectUser?: (userId: string) => void;
}) {
return (
<section aria-label={title} className="rounded-lg border border-border bg-card">
@ -54,6 +100,7 @@ function GroupCountTable({
<TableRow>
<TableHead>{labelHeader}</TableHead>
<TableHead className="text-right">Sessions</TableHead>
{onDisconnectUser ? <TableHead className="text-right">Actions</TableHead> : null}
</TableRow>
</TableHeader>
<TableBody>
@ -61,6 +108,11 @@ function GroupCountTable({
<TableRow key={group.label ?? "__unknown__"}>
<TableCell className="font-mono text-xs">{groupLabel(group.label)}</TableCell>
<TableCell className="text-right">{group.count}</TableCell>
{onDisconnectUser ? (
<TableCell className="text-right">
<DisconnectUserButton userId={group.label} onDisconnectUser={onDisconnectUser} />
</TableCell>
) : null}
</TableRow>
))}
</TableBody>
@ -73,10 +125,12 @@ function SessionsBody({
data,
error,
isLoading,
onDisconnect,
}: {
data: MCPGatewaySessionsResponse | undefined;
error: Error | null;
isLoading: boolean;
onDisconnect: ((selector: MCPGatewaySessionSelector) => void) | null;
}) {
if (isLoading) {
return (
@ -117,7 +171,12 @@ function SessionsBody({
</div>
<div className="grid grid-cols-1 gap-4 lg:grid-cols-2">
<GroupCountTable title="Sessions by AI client" labelHeader="Client" groups={data.by_client} />
<GroupCountTable title="Sessions by user" labelHeader="User" groups={data.by_user} />
<GroupCountTable
title="Sessions by user"
labelHeader="User"
groups={data.by_user}
onDisconnectUser={onDisconnect ? (userId) => onDisconnect({ user_id: userId }) : undefined}
/>
</div>
<section aria-label="Live sessions" className="rounded-lg border border-border bg-card">
<h3 className="border-b border-border px-4 py-2 text-sm font-semibold text-foreground">
@ -134,11 +193,12 @@ function SessionsBody({
<TableHead>Client IP</TableHead>
<TableHead className="text-right">Idle</TableHead>
<TableHead className="text-right">In flight</TableHead>
{onDisconnect ? <TableHead className="text-right">Actions</TableHead> : null}
</TableRow>
</TableHeader>
<TableBody>
{data.sessions.map((session) => (
<TableRow key={session.session_id_prefix}>
{data.sessions.map((session, index) => (
<TableRow key={`${session.session_id_prefix}-${index}`}>
<TableCell className="font-mono text-xs">{session.session_id_prefix}</TableCell>
<TableCell>
{session.client_name === null ? (
@ -169,6 +229,19 @@ function SessionsBody({
<TableCell className="font-mono text-xs">{session.client_ip || "-"}</TableCell>
<TableCell className="text-right text-xs">{formatIdleSeconds(session.idle_seconds)}</TableCell>
<TableCell className="text-right text-xs">{session.in_flight_requests}</TableCell>
{onDisconnect ? (
<TableCell className="text-right">
<Button
variant="outline"
size="sm"
onClick={() => onDisconnect({ session_id_prefix: session.session_id_prefix })}
aria-label={`Disconnect session ${session.session_id_prefix}`}
>
<Unplug className="size-4" />
Disconnect
</Button>
</TableCell>
) : null}
</TableRow>
))}
</TableBody>
@ -180,9 +253,12 @@ function SessionsBody({
interface MCPGatewaySessionsTabProps {
accessToken: string | null;
canTerminate: boolean;
}
export function MCPGatewaySessionsTab({ accessToken }: MCPGatewaySessionsTabProps) {
export function MCPGatewaySessionsTab({ accessToken, canTerminate }: MCPGatewaySessionsTabProps) {
const queryClient = useQueryClient();
const [pendingSelector, setPendingSelector] = useState<MCPGatewaySessionSelector | null>(null);
const queryOptions = {
queryKey: mcpGatewaySessionKeys.lists(),
queryFn: () => fetchMCPGatewaySessions(accessToken!),
@ -190,6 +266,15 @@ export function MCPGatewaySessionsTab({ accessToken }: MCPGatewaySessionsTabProp
refetchInterval: REFETCH_INTERVAL_MS,
};
const { data, error, isLoading, isFetching, refetch } = useQuery<MCPGatewaySessionsResponse, Error>(queryOptions);
const terminate = useMutation<MCPGatewaySessionsTerminateResponse, Error, MCPGatewaySessionSelector>({
mutationFn: (selector) => terminateMCPGatewaySessions(accessToken!, selector),
onSettled: () => queryClient.invalidateQueries({ queryKey: mcpGatewaySessionKeys.lists() }),
});
const confirmDisconnect = () => {
if (pendingSelector === null) return;
terminate.mutate(pendingSelector);
setPendingSelector(null);
};
return (
<div className="mt-4 space-y-4" data-testid="mcp-gateway-sessions-tab">
@ -214,7 +299,48 @@ export function MCPGatewaySessionsTab({ accessToken }: MCPGatewaySessionsTabProp
</Button>
</div>
<SessionsBody data={data} error={error} isLoading={isLoading} />
{terminate.isError ? (
<Alert variant="destructive">
<AlertTitle>Could not disconnect</AlertTitle>
<AlertDescription>{terminate.error.message}</AlertDescription>
</Alert>
) : null}
{terminate.isSuccess ? (
<Alert>
<AlertTitle>Disconnected</AlertTitle>
<AlertDescription>
{describeTerminateResult(terminate.data)} Clients holding those sessions must send a new initialize request,
which re-runs authentication. Sessions on other proxy workers are not affected.
</AlertDescription>
</Alert>
) : null}
<SessionsBody
data={data}
error={error}
isLoading={isLoading}
onDisconnect={canTerminate ? setPendingSelector : null}
/>
<AlertDialog open={pendingSelector !== null} onOpenChange={(open) => !open && setPendingSelector(null)}>
<AlertDialogContent>
<AlertDialogHeader>
<AlertDialogTitle>Disconnect MCP session</AlertDialogTitle>
<AlertDialogDescription>
{pendingSelector ? `This force-closes ${describeSelector(pendingSelector)} on this proxy worker. ` : ""}
In-flight requests fail and the client must initialize again before it can call tools.
</AlertDialogDescription>
</AlertDialogHeader>
<AlertDialogFooter>
<Button variant="outline" onClick={() => setPendingSelector(null)}>
Cancel
</Button>
<Button variant="destructive" onClick={confirmDisconnect} disabled={terminate.isPending}>
Disconnect
</Button>
</AlertDialogFooter>
</AlertDialogContent>
</AlertDialog>
</div>
);
}

View file

@ -0,0 +1,102 @@
import React from "react";
import { render, screen, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { MCPServerUserCredentialsPanel } from "./MCPServerUserCredentialsPanel";
import * as networking from "@/components/networking";
import type { MCPServerUserCredentialListItem } from "@/components/mcp_tools/types";
vi.mock("@/components/networking", () => ({
fetchMCPServerUserCredentials: vi.fn(),
revokeMCPServerUserCredential: vi.fn(),
}));
const ITEMS: MCPServerUserCredentialListItem[] = [
{
user_id: "alice",
credential_type: "oauth2",
expires_at: "2026-12-31T00:00:00+00:00",
connected_at: "2026-01-01T00:00:00+00:00",
updated_at: "2026-01-01T00:00:00+00:00",
},
{
user_id: "carol",
credential_type: "byok",
expires_at: null,
connected_at: null,
updated_at: "2026-02-01T00:00:00+00:00",
},
];
const renderPanel = ({ canRevoke = false }: { canRevoke?: boolean } = {}) => {
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } });
return render(
<QueryClientProvider client={queryClient}>
<MCPServerUserCredentialsPanel serverId="srv-1" accessToken="token" canRevoke={canRevoke} />
</QueryClientProvider>,
);
};
describe("MCPServerUserCredentialsPanel", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("lists each user's credential type without a revoke control for a read-only admin", async () => {
vi.mocked(networking.fetchMCPServerUserCredentials).mockResolvedValue(ITEMS);
renderPanel({ canRevoke: false });
const table = await screen.findByRole("region", { name: "Stored user credentials" });
expect(within(table).getByRole("row", { name: /alice/ })).toHaveTextContent("OAuth2");
expect(within(table).getByRole("row", { name: /carol/ })).toHaveTextContent("BYOK API key");
expect(screen.queryByRole("button", { name: /^Revoke credential/ })).not.toBeInTheDocument();
expect(networking.fetchMCPServerUserCredentials).toHaveBeenCalledWith("token", "srv-1");
});
it("revokes the selected user's credential through the route for its type and refetches", async () => {
const user = userEvent.setup();
vi.mocked(networking.fetchMCPServerUserCredentials).mockResolvedValueOnce(ITEMS).mockResolvedValueOnce([ITEMS[1]]);
vi.mocked(networking.revokeMCPServerUserCredential).mockResolvedValue(undefined);
renderPanel({ canRevoke: true });
await user.click(await screen.findByRole("button", { name: "Revoke credential for user alice" }));
expect(networking.revokeMCPServerUserCredential).not.toHaveBeenCalled();
const dialog = await screen.findByRole("alertdialog");
expect(dialog).toHaveTextContent("OAuth2 credential stored for user alice");
await user.click(within(dialog).getByRole("button", { name: "Revoke" }));
expect(await screen.findByText(/OAuth2 credential for user alice was deleted/)).toBeInTheDocument();
expect(networking.revokeMCPServerUserCredential).toHaveBeenCalledWith("token", "srv-1", "alice", "oauth2");
const table = await screen.findByRole("region", { name: "Stored user credentials" });
expect(within(table).queryByRole("row", { name: /alice/ })).not.toBeInTheDocument();
expect(within(table).getByRole("row", { name: /carol/ })).toBeInTheDocument();
});
it("shows the API error when a revoke is refused and keeps the list", async () => {
const user = userEvent.setup();
vi.mocked(networking.fetchMCPServerUserCredentials).mockResolvedValue(ITEMS);
vi.mocked(networking.revokeMCPServerUserCredential).mockRejectedValue(
new Error("Proxy admin access required to revoke another user's MCP credential."),
);
renderPanel({ canRevoke: true });
await user.click(await screen.findByRole("button", { name: "Revoke credential for user carol" }));
await user.click(within(await screen.findByRole("alertdialog")).getByRole("button", { name: "Revoke" }));
const alert = await screen.findByRole("alert");
expect(alert).toHaveTextContent("Could not revoke credential");
expect(alert).toHaveTextContent("Proxy admin access required to revoke another user's MCP credential.");
expect(networking.revokeMCPServerUserCredential).toHaveBeenCalledWith("token", "srv-1", "carol", "byok");
expect(screen.getByRole("region", { name: "Stored user credentials" })).toBeInTheDocument();
});
it("shows the API error when the list cannot be loaded", async () => {
vi.mocked(networking.fetchMCPServerUserCredentials).mockRejectedValue(new Error("Admin access required"));
renderPanel();
const alert = await screen.findByRole("alert");
expect(alert).toHaveTextContent("Could not load user credentials");
expect(alert).toHaveTextContent("Admin access required");
});
});

View file

@ -0,0 +1,212 @@
"use client";
import React, { useState } from "react";
import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query";
import { RefreshCw, ShieldOff } from "lucide-react";
import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert";
import {
AlertDialog,
AlertDialogContent,
AlertDialogDescription,
AlertDialogFooter,
AlertDialogHeader,
AlertDialogTitle,
} from "@/components/ui/alert-dialog";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { fetchMCPServerUserCredentials, revokeMCPServerUserCredential } from "@/components/networking";
import type { MCPServerUserCredentialListItem } from "@/components/mcp_tools/types";
import { createQueryKeys } from "@/app/(dashboard)/hooks/common/queryKeysFactory";
const mcpServerUserCredentialKeys = createQueryKeys("mcpServerUserCredentials");
export function credentialTypeLabel(credentialType: MCPServerUserCredentialListItem["credential_type"]): string {
return credentialType === "oauth2" ? "OAuth2" : "BYOK API key";
}
export function formatTimestamp(value: string | null): string {
if (value === null) return "-";
const parsed = new Date(value);
return Number.isNaN(parsed.getTime()) ? value : parsed.toLocaleString();
}
function CredentialsBody({
items,
error,
isLoading,
onRevoke,
}: {
items: MCPServerUserCredentialListItem[] | undefined;
error: Error | null;
isLoading: boolean;
onRevoke: ((item: MCPServerUserCredentialListItem) => void) | null;
}) {
if (isLoading) {
return (
<div
role="status"
className="flex items-center justify-center gap-3 rounded-lg border border-dashed border-border bg-card p-12"
>
<UiLoadingSpinner className="size-6 text-muted-foreground" />
<p className="text-sm text-muted-foreground">Loading user credentials...</p>
</div>
);
}
if (error) {
return (
<Alert variant="destructive">
<AlertTitle>Could not load user credentials</AlertTitle>
<AlertDescription>{error.message}</AlertDescription>
</Alert>
);
}
if (!items) return null;
if (items.length === 0) {
return (
<div className="rounded-lg border border-dashed border-border bg-card p-12 text-center">
<p className="text-sm text-muted-foreground">No user has a stored credential for this server.</p>
</div>
);
}
return (
<section aria-label="Stored user credentials" className="rounded-lg border border-border bg-card">
<Table>
<TableHeader>
<TableRow>
<TableHead>User</TableHead>
<TableHead>Type</TableHead>
<TableHead>Connected</TableHead>
<TableHead>Expires</TableHead>
<TableHead>Updated</TableHead>
{onRevoke ? <TableHead className="text-right">Actions</TableHead> : null}
</TableRow>
</TableHeader>
<TableBody>
{items.map((item) => (
<TableRow key={item.user_id}>
<TableCell className="font-mono text-xs">{item.user_id}</TableCell>
<TableCell>
<Badge variant="secondary">{credentialTypeLabel(item.credential_type)}</Badge>
</TableCell>
<TableCell className="text-xs">{formatTimestamp(item.connected_at)}</TableCell>
<TableCell className="text-xs">{formatTimestamp(item.expires_at)}</TableCell>
<TableCell className="text-xs">{formatTimestamp(item.updated_at)}</TableCell>
{onRevoke ? (
<TableCell className="text-right">
<Button
variant="outline"
size="sm"
onClick={() => onRevoke(item)}
aria-label={`Revoke credential for user ${item.user_id}`}
>
<ShieldOff className="size-4" />
Revoke
</Button>
</TableCell>
) : null}
</TableRow>
))}
</TableBody>
</Table>
</section>
);
}
interface MCPServerUserCredentialsPanelProps {
serverId: string;
accessToken: string | null;
canRevoke: boolean;
}
export function MCPServerUserCredentialsPanel({
serverId,
accessToken,
canRevoke,
}: MCPServerUserCredentialsPanelProps) {
const queryClient = useQueryClient();
const [pendingItem, setPendingItem] = useState<MCPServerUserCredentialListItem | null>(null);
const queryKey = mcpServerUserCredentialKeys.detail(serverId);
const { data, error, isLoading, isFetching, refetch } = useQuery<MCPServerUserCredentialListItem[], Error>({
queryKey,
queryFn: () => fetchMCPServerUserCredentials(accessToken!, serverId),
enabled: !!accessToken,
});
const revoke = useMutation<void, Error, MCPServerUserCredentialListItem>({
mutationFn: (item) => revokeMCPServerUserCredential(accessToken!, serverId, item.user_id, item.credential_type),
onSettled: () => queryClient.invalidateQueries({ queryKey }),
});
const confirmRevoke = () => {
if (pendingItem === null) return;
revoke.mutate(pendingItem);
setPendingItem(null);
};
return (
<div className="space-y-4" data-testid="mcp-server-user-credentials-panel">
<div className="flex flex-wrap items-start justify-between gap-3">
<div>
<h2 className="text-lg font-medium">User Credentials</h2>
<p className="text-sm text-muted-foreground">
Per-user OAuth2 tokens and BYOK API keys stored for this server. Revoking one deletes it from the database
and clears the cached copy, so the user must connect again before the gateway will call this server for
them.
</p>
</div>
<Button
variant="outline"
size="sm"
onClick={() => refetch()}
disabled={isFetching}
aria-label="Refresh user credentials"
>
<RefreshCw className={`size-4 ${isFetching ? "animate-spin" : ""}`} />
Refresh
</Button>
</div>
{revoke.isError ? (
<Alert variant="destructive">
<AlertTitle>Could not revoke credential</AlertTitle>
<AlertDescription>{revoke.error.message}</AlertDescription>
</Alert>
) : null}
{revoke.isSuccess ? (
<Alert>
<AlertTitle>Credential revoked</AlertTitle>
<AlertDescription>
The stored {credentialTypeLabel(revoke.variables.credential_type)} credential for user{" "}
{revoke.variables.user_id} was deleted.
</AlertDescription>
</Alert>
) : null}
<CredentialsBody items={data} error={error} isLoading={isLoading} onRevoke={canRevoke ? setPendingItem : null} />
<AlertDialog open={pendingItem !== null} onOpenChange={(open) => !open && setPendingItem(null)}>
<AlertDialogContent>
<AlertDialogHeader>
<AlertDialogTitle>Revoke stored credential</AlertDialogTitle>
<AlertDialogDescription>
{pendingItem
? `This deletes the ${credentialTypeLabel(pendingItem.credential_type)} credential stored for user ${pendingItem.user_id}. `
: ""}
Their next MCP request to this server fails until they connect again.
</AlertDialogDescription>
</AlertDialogHeader>
<AlertDialogFooter>
<Button variant="outline" onClick={() => setPendingItem(null)}>
Cancel
</Button>
<Button variant="destructive" onClick={confirmRevoke} disabled={revoke.isPending}>
Revoke
</Button>
</AlertDialogFooter>
</AlertDialogContent>
</AlertDialog>
</div>
);
}
export default MCPServerUserCredentialsPanel;

View file

@ -1,7 +1,9 @@
import { render, screen } from "@testing-library/react";
import { render, screen, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { MCPServerView } from "./mcp_server_view";
import * as networking from "@/components/networking";
import type { MCPServer } from "@/components/mcp_tools/types";
vi.mock(".", () => ({
@ -13,6 +15,12 @@ vi.mock("./mcp_server_edit", () => ({
EDIT_OAUTH_UI_STATE_KEY: "litellm-mcp-oauth-edit-state",
}));
vi.mock("@/components/networking", async (importOriginal) => ({
...(await importOriginal<typeof import("@/components/networking")>()),
fetchMCPServerUserCredentials: vi.fn(),
revokeMCPServerUserCredential: vi.fn(),
}));
const baseServer = {
server_id: "srv-1",
server_name: "demo server",
@ -25,19 +33,38 @@ const baseServer = {
const renderView = (overrides: Partial<MCPServer> = {}, props: Record<string, unknown> = {}) =>
render(
<MCPServerView
mcpServer={{ ...baseServer, ...overrides } as MCPServer}
onBack={vi.fn()}
isProxyAdmin
isEditing={false}
accessToken="tok"
userRole="Admin"
userID="u1"
availableAccessGroups={[]}
{...props}
/>,
<QueryClientProvider client={new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } })}>
<MCPServerView
mcpServer={{ ...baseServer, ...overrides } as MCPServer}
onBack={vi.fn()}
isProxyAdmin
isEditing={false}
accessToken="tok"
userRole="Admin"
userID="u1"
availableAccessGroups={[]}
{...props}
/>
</QueryClientProvider>,
);
const openUserCredentials = async (props: Record<string, unknown>) => {
vi.mocked(networking.fetchMCPServerUserCredentials).mockResolvedValue([
{
user_id: "alice",
credential_type: "byok",
expires_at: null,
connected_at: null,
updated_at: "2026-01-01T00:00:00+00:00",
},
]);
renderView({}, props);
await userEvent.click(screen.getByRole("tab", { name: "User Credentials" }));
return within(await screen.findByRole("region", { name: "Stored user credentials" })).getByRole("row", {
name: /alice/,
});
};
describe("MCPServerView", () => {
beforeEach(() => {
vi.clearAllMocks();
@ -149,4 +176,15 @@ describe("MCPServerView", () => {
expect(await screen.findByText("All tools enabled")).toBeInTheDocument();
});
it("lets a full admin revoke a stored user credential", async () => {
const row = await openUserCredentials({});
expect(within(row).getByRole("button", { name: "Revoke credential for user alice" })).toBeInTheDocument();
});
it("shows stored credentials to a view-only admin session without a revoke control", async () => {
const row = await openUserCredentials({ isViewOnly: true });
expect(row).toHaveTextContent("BYOK API key");
expect(within(row).queryByRole("button", { name: /^Revoke credential/ })).not.toBeInTheDocument();
});
});

View file

@ -9,7 +9,9 @@ import { MCPServer, handleTransport, handleAuth } from "@/components/mcp_tools/t
// TODO: Move Tools viewer from index file
import { MCPToolsViewer } from ".";
import MCPServerEdit, { EDIT_OAUTH_UI_STATE_KEY } from "./mcp_server_edit";
import { MCPServerUserCredentialsPanel } from "./MCPServerUserCredentialsPanel";
import { getSecureItem } from "@/utils/secureStorage";
import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles";
import MCPServerCostDisplay from "./mcp_server_cost_display";
import { getMaskedAndFullUrl } from "./utils";
import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils";
@ -23,6 +25,7 @@ interface MCPServerViewProps {
accessToken: string | null;
userRole: string | null;
userID: string | null;
isViewOnly?: boolean;
availableAccessGroups: string[];
initialTabIndex?: number;
}
@ -53,6 +56,7 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
accessToken,
userRole,
userID,
isViewOnly = false,
availableAccessGroups,
initialTabIndex = 0,
}) => {
@ -63,6 +67,8 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
const [showFullUrl, setShowFullUrl] = useState(false);
const [copiedStates, setCopiedStates] = useState<Record<string, boolean>>({});
const [selectedTabIndex, setSelectedTabIndex] = useState(returningFromEditOAuth ? 2 : initialTabIndex);
const canViewUserCredentials = userRole !== null && isProxyAdminTierRole(userRole);
const canRevokeUserCredentials = userRole !== null && isProxyAdminRole(userRole) && !isViewOnly;
const handleSuccess = (updated: MCPServer) => {
setEditing(false);
@ -142,6 +148,11 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
Settings
</TabsTrigger>
)}
{canViewUserCredentials && (
<TabsTrigger value="3" className="flex-none rounded-none px-4 py-2">
User Credentials
</TabsTrigger>
)}
</TabsList>
{/* Overview Panel */}
@ -387,6 +398,18 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
)}
</Card>
</TabsContent>
{canViewUserCredentials && (
<TabsContent value="3">
<Card className="p-6">
<MCPServerUserCredentialsPanel
serverId={mcpServer.server_id}
accessToken={accessToken}
canRevoke={canRevokeUserCredentials}
/>
</Card>
</TabsContent>
)}
</Tabs>
</div>
);

View file

@ -17,6 +17,8 @@ vi.mock("@/components/networking", () => ({
updateConfigFieldSetting: vi.fn().mockResolvedValue(undefined),
deleteConfigFieldSetting: vi.fn().mockResolvedValue(undefined),
listMCPUserEnvVarStatus: vi.fn().mockResolvedValue([]),
fetchMCPGatewaySessions: vi.fn(),
terminateMCPGatewaySessions: vi.fn(),
}));
const createQueryClient = () =>
@ -400,4 +402,50 @@ describe("MCPServers", () => {
// The server list refresh must NOT trigger a second health check
expect(networking.fetchMCPServerHealth).toHaveBeenCalledTimes(1);
});
const liveSessionsReport = {
worker_pid: 4242,
total_sessions: 1,
by_client: [{ label: "claude-code", count: 1 }],
by_user: [{ label: "alice", count: 1 }],
sessions: [
{
session_id_prefix: "aaaa1111",
client_name: "claude-code",
client_version: "1.0.0",
user_id: "alice",
user_email: "alice@example.com",
key_alias: "alice-key",
team_id: null,
team_alias: null,
client_ip: "10.0.0.1",
idle_seconds: 5,
in_flight_requests: 0,
},
],
};
const openLiveConnections = async (props: { isViewOnly?: boolean }) => {
vi.mocked(networking.fetchMCPServers).mockResolvedValue([]);
vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(liveSessionsReport);
render(
<QueryClientProvider client={createQueryClient()}>
<MCPServers {...defaultProps} {...props} />
</QueryClientProvider>,
);
await userEvent.click(await screen.findByRole("tab", { name: "Live Connections" }));
return within(await screen.findByRole("region", { name: "Live sessions" })).getByRole("row", { name: /aaaa1111/ });
};
it("lets a full admin disconnect a live session", async () => {
const row = await openLiveConnections({ isViewOnly: false });
expect(within(row).getByRole("button", { name: "Disconnect session aaaa1111" })).toBeInTheDocument();
});
it("shows live sessions to a view-only admin session without any disconnect control", async () => {
const row = await openLiveConnections({ isViewOnly: true });
expect(row).toHaveTextContent("alice@example.com");
expect(within(row).queryByRole("button", { name: /^Disconnect/ })).not.toBeInTheDocument();
expect(screen.queryByRole("button", { name: /^Disconnect all/ })).not.toBeInTheDocument();
});
});

View file

@ -1,4 +1,4 @@
import { isAdminRole, isProxyAdminTierRole } from "@/utils/roles";
import { isAdminRole, isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles";
import { CircleHelp, Search } from "lucide-react";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
@ -109,7 +109,7 @@ const readToolsOAuthServerId = (): string | null => {
}
};
const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID }) => {
const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID, isViewOnly = false }) => {
const { data: mcpServers, isLoading: isLoadingServers, refetch } = useMCPServers();
// Fetch health status for all servers
@ -578,6 +578,7 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
accessToken={accessToken}
userID={userID}
userRole={userRole}
isViewOnly={isViewOnly}
availableAccessGroups={uniqueMcpAccessGroups}
initialTabIndex={selectedServerId === toolsTabServerId ? 1 : 0}
/>
@ -755,7 +756,10 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
)}
{isProxyAdminTierRole(userRole) && (
<TabsContent value="connections">
<MCPGatewaySessionsTab accessToken={accessToken} />
<MCPGatewaySessionsTab
accessToken={accessToken}
canTerminate={isProxyAdminRole(userRole) && !isViewOnly}
/>
</TabsContent>
)}
</Tabs>

View file

@ -4,6 +4,6 @@ import { MCPServers } from "./_components";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
export default function McpServers() {
const { accessToken, userRole, userId } = useAuthorized();
return <MCPServers accessToken={accessToken} userRole={userRole} userID={userId} />;
const { accessToken, userRole, userId, isViewOnly } = useAuthorized();
return <MCPServers accessToken={accessToken} userRole={userRole} userID={userId} isViewOnly={isViewOnly} />;
}

View file

@ -569,6 +569,23 @@ describe("EntityUsage", () => {
expect(screen.getAllByText("Activity Metrics")[1]).toBeInTheDocument();
});
it("tells the team view how many keys the proxy left out of the per-key lists", async () => {
mockTeamDailyActivityAggregatedCall.mockResolvedValue({
...mockSpendData,
metadata: { ...mockSpendData.metadata, api_key_limit: 100, total_api_keys: 3000 },
});
render(<EntityUsage {...defaultProps} entityType="team" />);
await waitFor(() => {
expect(mockTeamDailyActivityAggregatedCall).toHaveBeenCalled();
});
act(() => {
fireEvent.click(screen.getByText("Key Activity"));
});
expect(await screen.findByRole("note")).toHaveTextContent("Only the 100 highest-spend keys of 3,000 are loaded");
});
// An inactive tab panel is marked aria-selected="false" by one tab library and hidden by the
// other, so treat either as "not on screen" and the assertion holds whichever one is rendering.
const isShowing = (element: HTMLElement): boolean => {

View file

@ -25,7 +25,7 @@ import TeamMultiSelect from "@/components/common_components/team_multi_select";
import UserDropdown from "@/components/common_components/UserDropdown";
import { ActivityMetrics, processActivityData } from "@/components/activity_metrics";
import { UsageExportHeader } from "@/components/EntityUsageExport";
import { getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason";
import { getApiKeyTruncation, getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason";
import type { EntityType } from "@/components/EntityUsageExport/types";
import {
agentDailyActivityCall,
@ -71,6 +71,8 @@ interface EntitySpendData {
total_successful_requests: number;
total_failed_requests: number;
total_tokens: number;
api_key_limit?: number | null;
total_api_keys?: number | null;
};
}
@ -160,6 +162,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({
});
const spendData = spendDataRaw as unknown as EntitySpendData;
const apiKeyTruncation = getApiKeyTruncation(spendData.metadata?.api_key_limit, spendData.metadata?.total_api_keys);
const {
data: agentSpendDataRaw,
@ -659,12 +662,18 @@ const EntityUsage: React.FC<EntityUsageProps> = ({
{
key: "keys",
label: "Key Activity",
content: <KeyActivityPanel keyMetrics={keyMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
content: (
<KeyActivityPanel
keyMetrics={keyMetrics}
hidePromptCachingMetrics={entityType === "agent"}
apiKeyTruncation={apiKeyTruncation}
/>
),
},
{ key: "endpoints", label: "Endpoint Activity", content: <EndpointUsage userSpendData={spendData} /> },
];
const spendFetchState = { coversRange, cancelled, failed };
const spendFetchState = { coversRange, cancelled, failed, apiKeyTruncation };
return (
<div style={{ width: "100%" }} className="relative">

View file

@ -30,7 +30,7 @@ import { ActivityMetrics, processActivityData } from "@/components/activity_metr
import CloudZeroExportModal from "@/components/cloudzero_export_modal";
import UserDropdown from "@/components/common_components/UserDropdown";
import EntityUsageExportModal from "@/components/EntityUsageExport";
import { getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason";
import { getApiKeyTruncation, getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason";
import KeyActivityPanel from "@/components/UsagePage/components/KeyActivityPanel";
import { Team } from "@/components/key_team_helpers/key_list";
import {
@ -256,6 +256,10 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
coversRange: activeAggregated !== null || paginatedResult.coversRange,
cancelled: paginatedResult.cancelled,
failed: paginatedResult.failed,
apiKeyTruncation: getApiKeyTruncation(
userSpendData.metadata?.api_key_limit,
userSpendData.metadata?.total_api_keys,
),
};
const exportBlockedReason = getExportBlockedReason(spendFetchState);
@ -904,7 +908,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
<ActivityMetrics modelMetrics={modelMetrics} />
</TabsContent>
<TabsContent value="keys" keepMounted>
<KeyActivityPanel keyMetrics={keyMetrics} />
<KeyActivityPanel keyMetrics={keyMetrics} apiKeyTruncation={spendFetchState.apiKeyTruncation} />
</TabsContent>
<TabsContent value="mcp" keepMounted>
<ActivityMetrics modelMetrics={mcpServerMetrics} />

View file

@ -1,11 +1,12 @@
import { describe, expect, it } from "vitest";
import { getExportBlockedReason, type UsageFetchState } from "./exportBlockedReason";
import { getApiKeyTruncation, getExportBlockedReason, type UsageFetchState } from "./exportBlockedReason";
const state = (overrides: Partial<UsageFetchState> = {}): UsageFetchState => ({
coversRange: true,
cancelled: false,
failed: false,
apiKeyTruncation: undefined,
...overrides,
});
@ -31,4 +32,27 @@ describe("getExportBlockedReason", () => {
expect(reason).toMatch(/failed to load/i);
expect(reason).not.toMatch(/stopped/i);
});
it("blocks when the aggregated endpoint dropped keys, since a per-team CSV would miss them", () => {
const reason = getExportBlockedReason(state({ apiKeyTruncation: { limit: 100, total: 3000 } }));
expect(reason).toMatch(/100 highest-spend keys of 3000/);
expect(reason).toMatch(/USAGE_TOP_API_KEYS_LIMIT/);
});
});
describe("getApiKeyTruncation", () => {
it("reports truncation once the proxy saw more keys than it returned", () => {
expect(getApiKeyTruncation(100, 101)).toEqual({ limit: 100, total: 101 });
});
it("stays quiet when exactly the cap exists, since every key is on screen", () => {
expect(getApiKeyTruncation(100, 100)).toBeUndefined();
expect(getApiKeyTruncation(100, 7)).toBeUndefined();
});
it("stays quiet when the response carries no cap, as the paginated fallback does", () => {
expect(getApiKeyTruncation(undefined, undefined)).toBeUndefined();
expect(getApiKeyTruncation(100, null)).toBeUndefined();
});
});

View file

@ -1,13 +1,31 @@
export interface ApiKeyTruncation {
limit: number;
total: number;
}
export interface UsageFetchState {
coversRange: boolean;
cancelled: boolean;
failed: boolean;
apiKeyTruncation: ApiKeyTruncation | undefined;
}
export const getExportBlockedReason = ({ coversRange, cancelled, failed }: UsageFetchState): string | undefined => {
export const getApiKeyTruncation = (apiKeyLimit: unknown, totalApiKeys: unknown): ApiKeyTruncation | undefined => {
if (typeof apiKeyLimit !== "number" || typeof totalApiKeys !== "number") return undefined;
return totalApiKeys > apiKeyLimit ? { limit: apiKeyLimit, total: totalApiKeys } : undefined;
};
export const getExportBlockedReason = ({
coversRange,
cancelled,
failed,
apiKeyTruncation,
}: UsageFetchState): string | undefined => {
if (failed) return "Some spend data failed to load, so an export would under-report. Reload the page to try again.";
if (cancelled)
return "Loading was stopped before the whole range arrived, so an export would under-report. Reload the page to load it all.";
if (!coversRange) return "Spend data is still loading, so an export would under-report. Wait for it to finish.";
if (apiKeyTruncation !== undefined)
return `Only the ${apiKeyTruncation.limit} highest-spend keys of ${apiKeyTruncation.total} were loaded, so a per-team export would under-report. Raise USAGE_TOP_API_KEYS_LIMIT on the proxy to load more keys.`;
return undefined;
};

View file

@ -68,4 +68,14 @@ describe("KeyActivityPanel", () => {
expect(screen.getByLabelText("Search keys")).toHaveValue("");
expect(screen.getByTestId("rendered-keys")).toHaveTextContent("hash-alicehash-bob");
});
it("says how many keys the proxy left out when only the top spenders were loaded", () => {
render(<KeyActivityPanel keyMetrics={keyMetrics} apiKeyTruncation={{ limit: 2, total: 3000 }} />);
expect(screen.getByRole("note")).toHaveTextContent("Only the 2 highest-spend keys of 3,000 are loaded");
});
it("shows no truncation note when every key is loaded", () => {
render(<KeyActivityPanel keyMetrics={keyMetrics} />);
expect(screen.queryByRole("note")).not.toBeInTheDocument();
});
});

View file

@ -2,6 +2,7 @@ import { Search, X } from "lucide-react";
import React, { useMemo, useState } from "react";
import { ActivityMetrics } from "@/components/activity_metrics";
import type { ApiKeyTruncation } from "@/components/EntityUsageExport/exportBlockedReason";
import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group";
import { filterKeyActivity } from "../keyActivityFilter";
@ -10,9 +11,14 @@ import type { ModelActivityData } from "../types";
interface KeyActivityPanelProps {
keyMetrics: Record<string, ModelActivityData>;
hidePromptCachingMetrics?: boolean;
apiKeyTruncation?: ApiKeyTruncation;
}
const KeyActivityPanel: React.FC<KeyActivityPanelProps> = ({ keyMetrics, hidePromptCachingMetrics = false }) => {
const KeyActivityPanel: React.FC<KeyActivityPanelProps> = ({
keyMetrics,
hidePromptCachingMetrics = false,
apiKeyTruncation,
}) => {
const [query, setQuery] = useState("");
const filtered = useMemo(() => filterKeyActivity(keyMetrics, query), [keyMetrics, query]);
const totalKeys = Object.keys(keyMetrics).length;
@ -43,6 +49,12 @@ const KeyActivityPanel: React.FC<KeyActivityPanelProps> = ({ keyMetrics, hidePro
<span className="text-sm text-muted-foreground">
Showing {shownKeys.toLocaleString()} of {totalKeys.toLocaleString()} keys
</span>
{apiKeyTruncation !== undefined && (
<span className="text-sm text-muted-foreground" role="note">
Only the {apiKeyTruncation.limit.toLocaleString()} highest-spend keys of{" "}
{apiKeyTruncation.total.toLocaleString()} are loaded
</span>
)}
</div>
{isFiltering && totalKeys > 0 && shownKeys === 0 ? (
<p className="rounded-lg border p-6 text-center text-sm text-muted-foreground">

View file

@ -517,6 +517,7 @@ export interface MCPServerProps {
accessToken: string | null;
userRole: string | null;
userID: string | null;
isViewOnly?: boolean;
}
export interface MCPToolsetTool {
@ -587,3 +588,23 @@ export interface MCPGatewaySessionsResponse {
by_user: MCPGatewaySessionGroupCount[];
sessions: MCPGatewaySession[];
}
export interface MCPGatewaySessionsTerminateResponse {
worker_pid: number;
terminated_sessions: number;
sessions: MCPGatewaySession[];
}
export type MCPGatewaySessionSelector =
| { session_id_prefix: string; user_id?: undefined }
| { user_id: string; session_id_prefix?: undefined };
export type MCPServerUserCredentialType = "oauth2" | "byok";
export interface MCPServerUserCredentialListItem {
user_id: string;
credential_type: MCPServerUserCredentialType;
expires_at: string | null;
connected_at: string | null;
updated_at: string;
}

View file

@ -97,7 +97,14 @@ import type { ModelBudgetUsage, ModelMaxBudget } from "./key_team_helpers/ModelM
import type { ObjectPermission } from "./object_permission_types";
import type { components } from "@/lib/http/schema";
import { jsonFields } from "./common_components/check_openapi_schema";
import type { MCPGatewaySessionsResponse, MCPUserEnvVarsStatus } from "./mcp_tools/types";
import type {
MCPGatewaySessionSelector,
MCPGatewaySessionsResponse,
MCPGatewaySessionsTerminateResponse,
MCPServerUserCredentialListItem,
MCPServerUserCredentialType,
MCPUserEnvVarsStatus,
} from "./mcp_tools/types";
import type {
CoordinationRedisSettings,
CoordinationRedisSettingsResponse,
@ -4976,6 +4983,33 @@ export const fetchMCPSubmissions = async (accessToken: string) => {
export const fetchMCPGatewaySessions = async (accessToken: string): Promise<MCPGatewaySessionsResponse> =>
apiClient.get<MCPGatewaySessionsResponse>(`/v1/mcp/sessions`, { accessToken });
export const terminateMCPGatewaySessions = async (
accessToken: string,
selector: MCPGatewaySessionSelector,
): Promise<MCPGatewaySessionsTerminateResponse> =>
apiClient.delete<MCPGatewaySessionsTerminateResponse>(`/v1/mcp/sessions`, { accessToken, query: { ...selector } });
export const fetchMCPServerUserCredentials = async (
accessToken: string,
serverId: string,
): Promise<MCPServerUserCredentialListItem[]> =>
apiClient.get<MCPServerUserCredentialListItem[]>(`/v1/mcp/server/${encodeURIComponent(serverId)}/user-credentials`, {
accessToken,
});
export const revokeMCPServerUserCredential = async (
accessToken: string,
serverId: string,
userId: string,
credentialType: MCPServerUserCredentialType,
): Promise<void> => {
const route = credentialType === "oauth2" ? "oauth-user-credential" : "user-credential";
await apiClient.delete(`/v1/mcp/server/${encodeURIComponent(serverId)}/${route}`, {
accessToken,
query: { user_id: userId },
});
};
export const approveMCPServer = async (accessToken: string, serverId: string) => {
try {
const url = (proxyBaseUrl ? `${proxyBaseUrl}` : "") + `/v1/mcp/server/${encodeURIComponent(serverId)}/approve`;

View file

@ -19037,7 +19037,7 @@ export interface paths {
post: operations["store_mcp_oauth_user_credential_v1_mcp_server__server_id__oauth_user_credential_post"];
/**
* Delete Mcp Oauth User Credential
* @description Revoke the calling user's stored OAuth2 token for an MCP server
* @description Revoke the calling user's stored OAuth2 token for an MCP server. A proxy admin may pass user_id to revoke another user's stored token.
*/
delete: operations["delete_mcp_oauth_user_credential_v1_mcp_server__server_id__oauth_user_credential_delete"];
options?: never;
@ -19101,7 +19101,7 @@ export interface paths {
post: operations["store_mcp_user_credential_v1_mcp_server__server_id__user_credential_post"];
/**
* Delete Mcp User Credential
* @description Delete the calling user's stored API key for a BYOK MCP server
* @description Delete the calling user's stored API key for a BYOK MCP server. A proxy admin may pass user_id to revoke another user's stored key.
*/
delete: operations["delete_mcp_user_credential_v1_mcp_server__server_id__user_credential_delete"];
options?: never;
@ -19109,6 +19109,26 @@ export interface paths {
patch?: never;
trace?: never;
};
"/v1/mcp/server/{server_id}/user-credentials": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/**
* List Mcp Server User Credentials
* @description List every user's stored BYOK or OAuth2 credential for an MCP server (admin only, no secrets)
*/
get: operations["list_mcp_server_user_credentials_v1_mcp_server__server_id__user_credentials_get"];
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/v1/mcp/server/{server_id}/user-env-vars": {
parameters: {
query?: never;
@ -19151,7 +19171,11 @@ export interface paths {
get: operations["get_mcp_gateway_sessions_v1_mcp_sessions_get"];
put?: never;
post?: never;
delete?: never;
/**
* Delete Mcp Gateway Sessions
* @description Force-close live stateful MCP gateway sessions on this proxy worker, selected by session id prefix and/or by the LiteLLM user that opened them (proxy admin only).
*/
delete: operations["delete_mcp_gateway_sessions_v1_mcp_sessions_delete"];
options?: never;
head?: never;
patch?: never;
@ -26692,6 +26716,16 @@ export interface components {
* @description If True, router fallbacks configured in router_settings are only attempted when the calling key (and its team and project) is allowed to call the fallback model; unauthorized fallback targets are skipped and the primary model's error is returned. Default is False.
*/
enforce_fallback_model_access?: boolean | null;
/**
* Failed Login Block Seconds
* @description How long a blocked source address, or source address and username, stays blocked. Every attempt from a blocked key, right or wrong, is refused with 429 before the password is checked; the block is not extended by refused attempts. Set under `general_settings` in config.yaml. Defaults to 300
*/
failed_login_block_seconds?: number | null;
/**
* Failed Login Window Seconds
* @description Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Set under `general_settings` in config.yaml. Defaults to 60
*/
failed_login_window_seconds?: number | null;
/**
* Forward Client Headers To Llm Api
* @description If True, forwards client headers (e.g. Authorization) to the LLM API. Required for Claude Code with Max subscription.
@ -26736,6 +26770,18 @@ export interface components {
* @description max batch input file size in MB for /v1/files uploads with purpose=batch, if a file is larger than this size it will be rejected before being forwarded to the provider
*/
max_batch_file_size_mb?: number | null;
/**
* Max Failed Login Attempts Per Source
* @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded down but at least 1, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10
*/
max_failed_login_attempts_per_source?: number | null;
/**
* Max Failed Login Attempts Per Source Overrides
* @description Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins (between equivalent keys such as '1.2.3.4' and '1.2.3.4/32', an exemption wins, then the higher limit), and the per-username allowance for that address follows as half the override. A value of 0 exempts the address from both limits. Set under `general_settings` in config.yaml
*/
max_failed_login_attempts_per_source_overrides?: {
[key: string]: number;
} | null;
/**
* Max File Size Mb
* @description max file size in MB for /v1/files uploads, for any purpose, if a file is larger than this size it will be rejected before being forwarded to the provider
@ -26916,7 +26962,7 @@ export interface components {
transcribe_media_buckets?: string[] | null;
/**
* Trusted Proxy Ranges
* @description CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler.
* @description CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, or containing an entry that is not an address or CIDR range, the per-source sign-in limit is off.
*/
trusted_proxy_ranges?: string[] | null;
/**
@ -27693,6 +27739,11 @@ export interface components {
};
/** DailySpendMetadata */
DailySpendMetadata: {
/**
* Api Key Limit
* @description When set, api_keys and every api_key_breakdown list at most this many keys, ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.
*/
api_key_limit?: number | null;
/**
* Has More
* @default false
@ -27703,6 +27754,11 @@ export interface components {
* @default 1
*/
page: number;
/**
* Total Api Keys
* @description Distinct API keys matching the filters. When this exceeds api_key_limit, the per-key lists are truncated to the highest-spend keys.
*/
total_api_keys?: number | null;
/**
* Total Api Requests
* @default 0
@ -32621,6 +32677,18 @@ export interface components {
/** Worker Pid */
worker_pid: number;
};
/**
* MCPGatewaySessionsTerminateResponse
* @description Stateful sessions an administrator force-closed on this proxy worker.
*/
MCPGatewaySessionsTerminateResponse: {
/** Sessions */
sessions?: components["schemas"]["MCPGatewaySession"][];
/** Terminated Sessions */
terminated_sessions: number;
/** Worker Pid */
worker_pid: number;
};
/**
* MCPOAuthUserCredentialRequest
* @description Stores a user's OAuth2 token for an OpenAPI MCP server.
@ -32725,6 +32793,25 @@ export interface components {
[key: string]: unknown;
};
};
/**
* MCPServerUserCredentialListItem
* @description One user's stored credential for an MCP server, as an admin sees it. Never carries the secret.
*/
MCPServerUserCredentialListItem: {
/** Connected At */
connected_at?: string | null;
/**
* Credential Type
* @enum {string}
*/
credential_type: "oauth2" | "byok";
/** Expires At */
expires_at?: string | null;
/** Updated At */
updated_at: string;
/** User Id */
user_id: string;
};
/** MCPSubmissionsSummary */
MCPSubmissionsSummary: {
/** Active */
@ -36801,7 +36888,9 @@ export interface components {
/** Type */
type?: string | null;
/** Value */
value: string;
value?: string | null;
} & {
[key: string]: unknown;
};
/** SCIMPatchOp */
SCIMPatchOp: {
@ -65691,7 +65780,9 @@ export interface operations {
};
delete_mcp_oauth_user_credential_v1_mcp_server__server_id__oauth_user_credential_delete: {
parameters: {
query?: never;
query?: {
user_id?: string | null;
};
header?: never;
path: {
server_id: string;
@ -65823,7 +65914,9 @@ export interface operations {
};
delete_mcp_user_credential_v1_mcp_server__server_id__user_credential_delete: {
parameters: {
query?: never;
query?: {
user_id?: string | null;
};
header?: never;
path: {
server_id: string;
@ -65852,6 +65945,37 @@ export interface operations {
};
};
};
list_mcp_server_user_credentials_v1_mcp_server__server_id__user_credentials_get: {
parameters: {
query?: never;
header?: never;
path: {
server_id: string;
};
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["MCPServerUserCredentialListItem"][];
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
get_mcp_user_env_vars_v1_mcp_server__server_id__user_env_vars_get: {
parameters: {
query?: never;
@ -65969,6 +66093,38 @@ export interface operations {
};
};
};
delete_mcp_gateway_sessions_v1_mcp_sessions_delete: {
parameters: {
query?: {
session_id_prefix?: string | null;
user_id?: string | null;
};
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["MCPGatewaySessionsTerminateResponse"];
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
get_mcp_tools_v1_mcp_tools_get: {
parameters: {
query?: never;

388
uv.lock generated
View file

@ -10,7 +10,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-09-14T23:55:55.024292355Z"
exclude-newer = "0001-01-01T00:00:00Z" # This has no effect and is included for backwards compatibility when using relative exclude-newer values.
exclude-newer-span = "P3D"
[manifest]
@ -225,9 +225,9 @@ name = "aiologic"
version = "0.17.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "sniffio", marker = "python_full_version < '3.13'" },
{ name = "typing-extensions", marker = "python_full_version < '3.13'" },
{ name = "wrapt", marker = "python_full_version < '3.13'" },
{ name = "sniffio" },
{ name = "typing-extensions" },
{ name = "wrapt" },
]
sdist = { url = "https://files.pythonhosted.org/packages/53/a7/809482759f40079f4c4328c7318bf569ae25d457f5017aad30a1b9aafedc/aiologic-0.17.0.tar.gz", hash = "sha256:65aa058e858c94cd208badb188e7f00b54dcabb3ba85b34f794db98074d108b9", size = 251625, upload-time = "2026-06-14T12:24:35.367Z" }
wheels = [
@ -315,16 +315,16 @@ vertex = [
[[package]]
name = "anyio"
version = "4.13.0"
version = "4.14.2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "exceptiongroup", marker = "python_full_version < '3.11'" },
{ name = "idna" },
{ name = "typing-extensions", marker = "python_full_version < '3.13'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" }
sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" },
{ url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" },
]
[[package]]
@ -519,14 +519,14 @@ name = "aurelio-sdk"
version = "0.0.19"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "aiofiles", marker = "python_full_version < '3.14'" },
{ name = "aiohttp", marker = "python_full_version < '3.14'" },
{ name = "colorlog", marker = "python_full_version < '3.14'" },
{ name = "pydantic", marker = "python_full_version < '3.14'" },
{ name = "python-dotenv", marker = "python_full_version < '3.14'" },
{ name = "requests", marker = "python_full_version < '3.14'" },
{ name = "requests-toolbelt", marker = "python_full_version < '3.14'" },
{ name = "tornado", marker = "python_full_version < '3.14'" },
{ name = "aiofiles" },
{ name = "aiohttp" },
{ name = "colorlog" },
{ name = "pydantic" },
{ name = "python-dotenv" },
{ name = "requests" },
{ name = "requests-toolbelt" },
{ name = "tornado" },
]
sdist = { url = "https://files.pythonhosted.org/packages/27/0e/c2e369ad173fb3d76448e46d10beb3dcc53388318933ddf8169a3f21a810/aurelio_sdk-0.0.19.tar.gz", hash = "sha256:14107e7440ff2efd0b4a08c52fb595e7680bd4bc973a0ddfb3b64157c6666b91", size = 15258, upload-time = "2025-03-24T14:37:32.203Z" }
wheels = [
@ -538,9 +538,9 @@ name = "aws-sdk-bedrock-runtime"
version = "0.11.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "smithy-aws-core", extra = ["eventstream", "json"], marker = "python_full_version >= '3.12'" },
{ name = "smithy-core", marker = "python_full_version >= '3.12'" },
{ name = "smithy-http", extra = ["aiohttp"], marker = "python_full_version >= '3.12'" },
{ name = "smithy-aws-core", extra = ["eventstream", "json"] },
{ name = "smithy-core" },
{ name = "smithy-http", extra = ["aiohttp"] },
]
sdist = { url = "https://files.pythonhosted.org/packages/8e/b3/9c225cbfe9f17ea2e3d75a0fdd0b325ef79839b9c09a376bda63a7bf3bb3/aws_sdk_bedrock_runtime-0.11.0.tar.gz", hash = "sha256:f2c45d34625bf6a7b56375e29a53a16b376880bda771e4bbf7d84491622eb193", size = 173854, upload-time = "2026-08-24T21:17:16.304Z" }
wheels = [
@ -549,7 +549,7 @@ wheels = [
[package.optional-dependencies]
awscrt = [
{ name = "smithy-http", extra = ["awscrt"], marker = "python_full_version >= '3.12'" },
{ name = "smithy-http", extra = ["awscrt"] },
]
[[package]]
@ -1207,7 +1207,7 @@ name = "colorlog"
version = "6.10.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "colorama", marker = "python_full_version < '3.14' and sys_platform == 'win32'" },
{ name = "colorama", marker = "sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/a2/61/f083b5ac52e505dfc1c624eafbf8c7589a0d7f32daa398d2e7590efa5fda/colorlog-6.10.1.tar.gz", hash = "sha256:eb4ae5cb65fe7fec7773c2306061a8e63e02efc2c72eba9d27b0fa23c94f1321", size = 17162, upload-time = "2025-10-16T16:14:11.978Z" }
wheels = [
@ -1231,7 +1231,7 @@ resolution-markers = [
"python_full_version < '3.11'",
]
dependencies = [
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" } },
]
sdist = { url = "https://files.pythonhosted.org/packages/66/54/eb9bfc647b19f2009dd5c7f5ec51c4e6ca831725f1aea7a993034f483147/contourpy-1.3.2.tar.gz", hash = "sha256:b6945942715a034c671b7fc54f9588126b0b8bf23db2696e3ca8328f3ff0ab54", size = 13466130, upload-time = "2025-04-15T17:47:53.79Z" }
wheels = [
@ -1304,7 +1304,7 @@ resolution-markers = [
"python_full_version == '3.11.*'",
]
dependencies = [
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" },
{ name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/58/01/1253e6698a07380cd31a736d248a3f2a50a7c88779a1813da27503cadc2a/contourpy-1.3.3.tar.gz", hash = "sha256:083e12155b210502d0bca491432bb04d56dc3432f95a979b429f2848c3dbe880", size = 13466174, upload-time = "2025-07-26T12:03:12.549Z" }
@ -1574,8 +1574,8 @@ name = "culsans"
version = "0.11.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "aiologic", marker = "python_full_version < '3.13'" },
{ name = "typing-extensions", marker = "python_full_version < '3.13'" },
{ name = "aiologic" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/d9/e3/49afa1bc180e0d28008ec6bcdf82a4072d1c7a41032b5b759b60814ca4b0/culsans-0.11.0.tar.gz", hash = "sha256:0b43d0d05dce6106293d114c86e3fb4bfc63088cfe8ff08ed3fe36891447fe33", size = 107546, upload-time = "2025-12-31T23:15:38.196Z" }
wheels = [
@ -1829,7 +1829,7 @@ name = "exceptiongroup"
version = "1.3.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions", marker = "python_full_version < '3.11'" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/50/79/66800aadf48771f6b62f7eb014e352e5d06856655206165d775e675a02c9/exceptiongroup-1.3.1.tar.gz", hash = "sha256:8b412432c6055b0b7d14c310000ae93352ed6754f70fa8f7c34141f91c4e3219", size = 30371, upload-time = "2025-11-21T23:01:54.787Z" }
wheels = [
@ -2412,11 +2412,11 @@ resolution-markers = [
"python_full_version >= '3.14'",
]
dependencies = [
{ name = "google-auth", marker = "python_full_version >= '3.14'" },
{ name = "googleapis-common-protos", marker = "python_full_version >= '3.14'" },
{ name = "proto-plus", marker = "python_full_version >= '3.14'" },
{ name = "protobuf", marker = "python_full_version >= '3.14'" },
{ name = "requests", marker = "python_full_version >= '3.14'" },
{ name = "google-auth" },
{ name = "googleapis-common-protos" },
{ name = "proto-plus" },
{ name = "protobuf" },
{ name = "requests" },
]
sdist = { url = "https://files.pythonhosted.org/packages/09/cd/63f1557235c2440fe0577acdbc32577c5c002684c58c7f4d770a92366a24/google_api_core-2.25.2.tar.gz", hash = "sha256:1c63aa6af0d0d5e37966f157a77f9396d820fba59f9e43e9415bc3dc5baff300", size = 166266, upload-time = "2025-10-03T00:07:34.778Z" }
wheels = [
@ -2425,8 +2425,8 @@ wheels = [
[package.optional-dependencies]
grpc = [
{ name = "grpcio", marker = "python_full_version >= '3.14'" },
{ name = "grpcio-status", marker = "python_full_version >= '3.14'" },
{ name = "grpcio" },
{ name = "grpcio-status" },
]
[[package]]
@ -2440,11 +2440,11 @@ resolution-markers = [
"python_full_version < '3.11'",
]
dependencies = [
{ name = "google-auth", marker = "python_full_version < '3.14'" },
{ name = "googleapis-common-protos", marker = "python_full_version < '3.14'" },
{ name = "proto-plus", marker = "python_full_version < '3.14'" },
{ name = "protobuf", marker = "python_full_version < '3.14'" },
{ name = "requests", marker = "python_full_version < '3.14'" },
{ name = "google-auth" },
{ name = "googleapis-common-protos" },
{ name = "proto-plus" },
{ name = "protobuf" },
{ name = "requests" },
]
sdist = { url = "https://files.pythonhosted.org/packages/16/ce/502a57fb0ec752026d24df1280b162294b22a0afb98a326084f9a979138b/google_api_core-2.30.3.tar.gz", hash = "sha256:e601a37f148585319b26db36e219df68c5d07b6382cff2d580e83404e44d641b", size = 177001, upload-time = "2026-04-10T00:41:28.035Z" }
wheels = [
@ -2453,8 +2453,8 @@ wheels = [
[package.optional-dependencies]
grpc = [
{ name = "grpcio", marker = "python_full_version < '3.14'" },
{ name = "grpcio-status", marker = "python_full_version < '3.14'" },
{ name = "grpcio" },
{ name = "grpcio-status" },
]
[[package]]
@ -2623,12 +2623,12 @@ resolution-markers = [
"python_full_version >= '3.14'",
]
dependencies = [
{ name = "google-api-core", version = "2.25.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.14'" },
{ name = "google-auth", marker = "python_full_version >= '3.14'" },
{ name = "google-cloud-core", marker = "python_full_version >= '3.14'" },
{ name = "google-crc32c", marker = "python_full_version >= '3.14'" },
{ name = "google-resumable-media", marker = "python_full_version >= '3.14'" },
{ name = "requests", marker = "python_full_version >= '3.14'" },
{ name = "google-api-core", version = "2.25.2", source = { registry = "https://pypi.org/simple" } },
{ name = "google-auth" },
{ name = "google-cloud-core" },
{ name = "google-crc32c" },
{ name = "google-resumable-media" },
{ name = "requests" },
]
sdist = { url = "https://files.pythonhosted.org/packages/bd/ef/7cefdca67a6c8b3af0ec38612f9e78e5a9f6179dd91352772ae1a9849246/google_cloud_storage-3.4.1.tar.gz", hash = "sha256:6f041a297e23a4b485fad8c305a7a6e6831855c208bcbe74d00332a909f82268", size = 17238203, upload-time = "2025-10-08T18:43:39.665Z" }
wheels = [
@ -2646,12 +2646,12 @@ resolution-markers = [
"python_full_version < '3.11'",
]
dependencies = [
{ name = "google-api-core", version = "2.30.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.14'" },
{ name = "google-auth", marker = "python_full_version < '3.14'" },
{ name = "google-cloud-core", marker = "python_full_version < '3.14'" },
{ name = "google-crc32c", marker = "python_full_version < '3.14'" },
{ name = "google-resumable-media", marker = "python_full_version < '3.14'" },
{ name = "requests", marker = "python_full_version < '3.14'" },
{ name = "google-api-core", version = "2.30.3", source = { registry = "https://pypi.org/simple" } },
{ name = "google-auth" },
{ name = "google-cloud-core" },
{ name = "google-crc32c" },
{ name = "google-resumable-media" },
{ name = "requests" },
]
sdist = { url = "https://files.pythonhosted.org/packages/4c/47/205eb8e9a1739b5345843e5a425775cbdc472cc38e7eda082ba5b8d02450/google_cloud_storage-3.10.1.tar.gz", hash = "sha256:97db9aa4460727982040edd2bd13ff3d5e2260b5331ad22895802da1fc2a5286", size = 17309950, upload-time = "2026-03-23T09:35:23.409Z" }
wheels = [
@ -4081,13 +4081,13 @@ name = "langchain-classic"
version = "1.0.7"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "langchain-core", marker = "python_full_version >= '3.11'" },
{ name = "langchain-text-splitters", marker = "python_full_version >= '3.11'" },
{ name = "langsmith", marker = "python_full_version >= '3.11'" },
{ name = "pydantic", marker = "python_full_version >= '3.11'" },
{ name = "pyyaml", marker = "python_full_version >= '3.11'" },
{ name = "requests", marker = "python_full_version >= '3.11'" },
{ name = "sqlalchemy", marker = "python_full_version >= '3.11'" },
{ name = "langchain-core" },
{ name = "langchain-text-splitters" },
{ name = "langsmith" },
{ name = "pydantic" },
{ name = "pyyaml" },
{ name = "requests" },
{ name = "sqlalchemy" },
]
sdist = { url = "https://files.pythonhosted.org/packages/9b/78/84b5065816f348c39fefa4316f209f0135e8410216340a953bec17d9e4e4/langchain_classic-1.0.7.tar.gz", hash = "sha256:debbec8065e69b95108d2652e8d5c44f4516e19aa8d716c02ed2211c3aee099d", size = 10554118, upload-time = "2026-05-07T15:46:56.8Z" }
wheels = [
@ -4102,18 +4102,18 @@ resolution-markers = [
"python_full_version < '3.11'",
]
dependencies = [
{ name = "aiohttp", marker = "python_full_version < '3.11'" },
{ name = "dataclasses-json", marker = "python_full_version < '3.11'" },
{ name = "httpx-sse", marker = "python_full_version < '3.11'" },
{ name = "langchain", marker = "python_full_version < '3.11'" },
{ name = "langchain-core", marker = "python_full_version < '3.11'" },
{ name = "langsmith", marker = "python_full_version < '3.11'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
{ name = "pydantic-settings", marker = "python_full_version < '3.11'" },
{ name = "pyyaml", marker = "python_full_version < '3.11'" },
{ name = "requests", marker = "python_full_version < '3.11'" },
{ name = "sqlalchemy", marker = "python_full_version < '3.11'" },
{ name = "tenacity", marker = "python_full_version < '3.11'" },
{ name = "aiohttp" },
{ name = "dataclasses-json" },
{ name = "httpx-sse" },
{ name = "langchain" },
{ name = "langchain-core" },
{ name = "langsmith" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" } },
{ name = "pydantic-settings" },
{ name = "pyyaml" },
{ name = "requests" },
{ name = "sqlalchemy" },
{ name = "tenacity" },
]
sdist = { url = "https://files.pythonhosted.org/packages/83/49/2ff5354273809e9811392bc24bcffda545a196070666aef27bc6aacf1c21/langchain_community-0.3.31.tar.gz", hash = "sha256:250e4c1041539130f6d6ac6f9386cb018354eafccd917b01a4cff1950b80fd81", size = 33241237, upload-time = "2025-10-07T20:17:57.857Z" }
wheels = [
@ -4131,19 +4131,19 @@ resolution-markers = [
"python_full_version == '3.11.*'",
]
dependencies = [
{ name = "aiohttp", marker = "python_full_version >= '3.11'" },
{ name = "dataclasses-json", marker = "python_full_version >= '3.11'" },
{ name = "httpx-sse", marker = "python_full_version >= '3.11'" },
{ name = "langchain-classic", marker = "python_full_version >= '3.11'" },
{ name = "langchain-core", marker = "python_full_version >= '3.11'" },
{ name = "langsmith", marker = "python_full_version >= '3.11'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" },
{ name = "aiohttp" },
{ name = "dataclasses-json" },
{ name = "httpx-sse" },
{ name = "langchain-classic" },
{ name = "langchain-core" },
{ name = "langsmith" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" },
{ name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" },
{ name = "pydantic-settings", marker = "python_full_version >= '3.11'" },
{ name = "pyyaml", marker = "python_full_version >= '3.11'" },
{ name = "requests", marker = "python_full_version >= '3.11'" },
{ name = "sqlalchemy", marker = "python_full_version >= '3.11'" },
{ name = "tenacity", marker = "python_full_version >= '3.11'" },
{ name = "pydantic-settings" },
{ name = "pyyaml" },
{ name = "requests" },
{ name = "sqlalchemy" },
{ name = "tenacity" },
]
sdist = { url = "https://files.pythonhosted.org/packages/53/97/a03585d42b9bdb6fbd935282d6e3348b10322a24e6ce12d0c99eb461d9af/langchain_community-0.4.1.tar.gz", hash = "sha256:f3b211832728ee89f169ddce8579b80a085222ddb4f4ed445a46e977d17b1e85", size = 33241144, upload-time = "2025-10-27T15:20:32.504Z" }
wheels = [
@ -4215,7 +4215,7 @@ name = "langchain-text-splitters"
version = "1.1.2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "langchain-core", marker = "python_full_version >= '3.11'" },
{ name = "langchain-core" },
]
sdist = { url = "https://files.pythonhosted.org/packages/26/9f/6c545900fefb7b00ddfa3f16b80d61338a0ec68c31c5451eeeab99082760/langchain_text_splitters-1.1.2.tar.gz", hash = "sha256:782a723db0a4746ac91e251c7c1d57fd23636e4f38ed733074e28d7a86f41627", size = 293580, upload-time = "2026-04-16T14:20:39.162Z" }
wheels = [
@ -4959,16 +4959,16 @@ resolution-markers = [
"python_full_version < '3.11'",
]
dependencies = [
{ name = "aiohttp", marker = "python_full_version < '3.11'" },
{ name = "chevron", marker = "python_full_version < '3.11'" },
{ name = "jsonpickle", marker = "python_full_version < '3.11'" },
{ name = "langchain-community", version = "0.3.31", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
{ name = "packaging", marker = "python_full_version < '3.11'" },
{ name = "pydantic", marker = "python_full_version < '3.11'" },
{ name = "pyhumps", marker = "python_full_version < '3.11'" },
{ name = "requests", marker = "python_full_version < '3.11'" },
{ name = "setuptools", marker = "python_full_version < '3.11'" },
{ name = "tenacity", marker = "python_full_version < '3.11'" },
{ name = "aiohttp" },
{ name = "chevron" },
{ name = "jsonpickle" },
{ name = "langchain-community", version = "0.3.31", source = { registry = "https://pypi.org/simple" } },
{ name = "packaging" },
{ name = "pydantic" },
{ name = "pyhumps" },
{ name = "requests" },
{ name = "setuptools" },
{ name = "tenacity" },
]
sdist = { url = "https://files.pythonhosted.org/packages/4a/6f/9ca1acf766848aaf5f0ac4140c34c91ad0dbfad2654359699644be3352c9/lunary-1.4.36.tar.gz", hash = "sha256:53f002f385c83d9c0e6368e7999923acffbde987f53c5205c2c249c38ee2d75c", size = 20253, upload-time = "2026-02-09T20:49:30.56Z" }
wheels = [
@ -4986,16 +4986,16 @@ resolution-markers = [
"python_full_version == '3.11.*'",
]
dependencies = [
{ name = "aiohttp", marker = "python_full_version >= '3.11'" },
{ name = "chevron", marker = "python_full_version >= '3.11'" },
{ name = "jsonpickle", marker = "python_full_version >= '3.11'" },
{ name = "langchain-community", version = "0.4.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" },
{ name = "packaging", marker = "python_full_version >= '3.11'" },
{ name = "pydantic", marker = "python_full_version >= '3.11'" },
{ name = "pyhumps", marker = "python_full_version >= '3.11'" },
{ name = "requests", marker = "python_full_version >= '3.11'" },
{ name = "setuptools", marker = "python_full_version >= '3.11'" },
{ name = "tenacity", marker = "python_full_version >= '3.11'" },
{ name = "aiohttp" },
{ name = "chevron" },
{ name = "jsonpickle" },
{ name = "langchain-community", version = "0.4.1", source = { registry = "https://pypi.org/simple" } },
{ name = "packaging" },
{ name = "pydantic" },
{ name = "pyhumps" },
{ name = "requests" },
{ name = "setuptools" },
{ name = "tenacity" },
]
sdist = { url = "https://files.pythonhosted.org/packages/37/ef/1acbc6957585cc0110e648d787663871717ced3df27fcd3cb5e18fa418f3/lunary-1.4.37.tar.gz", hash = "sha256:1781091e9dceffcc28ebc4be7e085c9fec4102d98d7ca945ed0021e9ce03c36f", size = 20248, upload-time = "2026-02-12T08:15:02.091Z" }
wheels = [
@ -8787,10 +8787,10 @@ resolution-markers = [
"python_full_version < '3.11'",
]
dependencies = [
{ name = "joblib", marker = "python_full_version < '3.11'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
{ name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
{ name = "threadpoolctl", marker = "python_full_version < '3.11'" },
{ name = "joblib" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" } },
{ name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" } },
{ name = "threadpoolctl" },
]
sdist = { url = "https://files.pythonhosted.org/packages/98/c2/a7855e41c9d285dfe86dc50b250978105dce513d6e459ea66a6aeb0e1e0c/scikit_learn-1.7.2.tar.gz", hash = "sha256:20e9e49ecd130598f1ca38a1d85090e1a600147b9c02fa6f15d69cb53d968fda", size = 7193136, upload-time = "2025-09-09T08:21:29.075Z" }
wheels = [
@ -8837,11 +8837,11 @@ resolution-markers = [
"python_full_version == '3.11.*'",
]
dependencies = [
{ name = "joblib", marker = "python_full_version >= '3.11'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" },
{ name = "joblib" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" },
{ name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" },
{ name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" },
{ name = "threadpoolctl", marker = "python_full_version >= '3.11'" },
{ name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" } },
{ name = "threadpoolctl" },
]
sdist = { url = "https://files.pythonhosted.org/packages/0e/d4/40988bf3b8e34feec1d0e6a051446b1f66225f8529b9309becaeef62b6c4/scikit_learn-1.8.0.tar.gz", hash = "sha256:9bccbb3b40e3de10351f8f5068e105d0f4083b1a65fa07b6634fbc401a6287fd", size = 7335585, upload-time = "2025-12-10T07:08:53.618Z" }
wheels = [
@ -8891,7 +8891,7 @@ resolution-markers = [
"python_full_version < '3.11'",
]
dependencies = [
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" } },
]
sdist = { url = "https://files.pythonhosted.org/packages/0f/37/6964b830433e654ec7485e45a00fc9a27cf868d622838f6b6d9c5ec0d532/scipy-1.15.3.tar.gz", hash = "sha256:eae3cf522bc7df64b42cad3925c876e1b0b6c35c1337c93e12c0f366f55b0eaf", size = 59419214, upload-time = "2025-05-08T16:13:05.955Z" }
wheels = [
@ -8953,7 +8953,7 @@ resolution-markers = [
"python_full_version == '3.11.*'",
]
dependencies = [
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" },
{ name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/7a/97/5a3609c4f8d58b039179648e62dd220f89864f56f7357f5d4f45c29eb2cc/scipy-1.17.1.tar.gz", hash = "sha256:95d8e012d8cb8816c226aef832200b1d45109ed4464303e997c5b13122b297c0", size = 30573822, upload-time = "2026-02-23T00:26:24.851Z" }
@ -9038,20 +9038,20 @@ name = "semantic-router"
version = "0.1.15"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "aiohttp", marker = "python_full_version < '3.14'" },
{ name = "aurelio-sdk", marker = "python_full_version < '3.14'" },
{ name = "colorama", marker = "python_full_version < '3.14'" },
{ name = "colorlog", marker = "python_full_version < '3.14'" },
{ name = "litellm", marker = "python_full_version < '3.14'" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" },
{ name = "aiohttp" },
{ name = "aurelio-sdk" },
{ name = "colorama" },
{ name = "colorlog" },
{ name = "litellm" },
{ name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12' or python_full_version >= '3.14'" },
{ name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12' and python_full_version < '3.14'" },
{ name = "openai", marker = "python_full_version < '3.14'" },
{ name = "pydantic", marker = "python_full_version < '3.14'" },
{ name = "pyyaml", marker = "python_full_version < '3.14'" },
{ name = "regex", marker = "python_full_version < '3.14'" },
{ name = "tiktoken", marker = "python_full_version < '3.14'" },
{ name = "tornado", marker = "python_full_version < '3.14'" },
{ name = "urllib3", marker = "python_full_version < '3.14'" },
{ name = "openai" },
{ name = "pydantic" },
{ name = "pyyaml" },
{ name = "regex" },
{ name = "tiktoken" },
{ name = "tornado" },
{ name = "urllib3" },
]
sdist = { url = "https://files.pythonhosted.org/packages/dc/a9/1a689e916e8b280f1fd8fb335cc059be626a22fe4533baa045d32fcd6de5/semantic_router-0.1.15.tar.gz", hash = "sha256:328256ddc3c2b713101ec69561d6585aecbf1198ea3461e1486289d8c3a35288", size = 95605, upload-time = "2026-05-23T12:58:15.444Z" }
wheels = [
@ -9134,9 +9134,9 @@ name = "smithy-aws-core"
version = "0.11.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "aws-sdk-signers", marker = "python_full_version >= '3.12'" },
{ name = "smithy-core", marker = "python_full_version >= '3.12'" },
{ name = "smithy-http", marker = "python_full_version >= '3.12'" },
{ name = "aws-sdk-signers" },
{ name = "smithy-core" },
{ name = "smithy-http" },
]
sdist = { url = "https://files.pythonhosted.org/packages/7d/d3/501c0023548173416109ac42298ca33b708469dc922005770811a597949f/smithy_aws_core-0.11.0.tar.gz", hash = "sha256:29ee89976a520a87e3db557e03e115fdc21a0a60b81161e95174395a1b064da1", size = 38791, upload-time = "2026-08-24T21:16:59.631Z" }
wheels = [
@ -9145,10 +9145,10 @@ wheels = [
[package.optional-dependencies]
eventstream = [
{ name = "smithy-aws-event-stream", marker = "python_full_version >= '3.12'" },
{ name = "smithy-aws-event-stream" },
]
json = [
{ name = "smithy-json", marker = "python_full_version >= '3.12'" },
{ name = "smithy-json" },
]
[[package]]
@ -9156,7 +9156,7 @@ name = "smithy-aws-event-stream"
version = "0.3.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "smithy-core", marker = "python_full_version >= '3.12'" },
{ name = "smithy-core" },
]
sdist = { url = "https://files.pythonhosted.org/packages/38/0e/6efb3a4ed92c0f1ada6de060ac92e7115a1e34d0ab1fb99a6056734a88ea/smithy_aws_event_stream-0.3.0.tar.gz", hash = "sha256:a0e227367a973144e205a075d0a424f95c92f26656a1018d08900da2ae547c49", size = 12818, upload-time = "2026-05-05T18:04:14.317Z" }
wheels = [
@ -9177,7 +9177,7 @@ name = "smithy-http"
version = "0.5.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "smithy-core", marker = "python_full_version >= '3.12'" },
{ name = "smithy-core" },
]
sdist = { url = "https://files.pythonhosted.org/packages/98/78/b5f3113d6c8f0bc1f9777a7f5ca84b892d29efac05850e14f7d4f7e645b5/smithy_http-0.5.0.tar.gz", hash = "sha256:bb4a19672f7c7eeb872a308f777eb505281a5bafb1ee3d1ea9c760c06c352510", size = 31122, upload-time = "2026-08-24T21:16:56.488Z" }
wheels = [
@ -9186,11 +9186,11 @@ wheels = [
[package.optional-dependencies]
aiohttp = [
{ name = "aiohttp", marker = "python_full_version >= '3.12'" },
{ name = "yarl", marker = "python_full_version >= '3.12'" },
{ name = "aiohttp" },
{ name = "yarl" },
]
awscrt = [
{ name = "awscrt", marker = "python_full_version >= '3.12'" },
{ name = "awscrt" },
]
[[package]]
@ -9198,8 +9198,8 @@ name = "smithy-json"
version = "0.3.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "ijson", marker = "python_full_version >= '3.12'" },
{ name = "smithy-core", marker = "python_full_version >= '3.12'" },
{ name = "ijson" },
{ name = "smithy-core" },
]
sdist = { url = "https://files.pythonhosted.org/packages/c7/ac/04164eefb3da7479f52f6535b4b39cc8384c292cb2bb74279f2acc4f4b4d/smithy_json-0.3.0.tar.gz", hash = "sha256:c81c7034587e01bc64767cbbecb05a7d65ca9070612fd94e8a03e80540290a22", size = 7956, upload-time = "2026-08-20T17:55:32.177Z" }
wheels = [
@ -9277,23 +9277,23 @@ resolution-markers = [
"python_full_version < '3.11'",
]
dependencies = [
{ name = "alabaster", marker = "python_full_version < '3.11'" },
{ name = "babel", marker = "python_full_version < '3.11'" },
{ name = "colorama", marker = "python_full_version < '3.11' and sys_platform == 'win32'" },
{ name = "docutils", version = "0.21.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
{ name = "imagesize", marker = "python_full_version < '3.11'" },
{ name = "jinja2", marker = "python_full_version < '3.11'" },
{ name = "packaging", marker = "python_full_version < '3.11'" },
{ name = "pygments", marker = "python_full_version < '3.11'" },
{ name = "requests", marker = "python_full_version < '3.11'" },
{ name = "snowballstemmer", marker = "python_full_version < '3.11'" },
{ name = "sphinxcontrib-applehelp", marker = "python_full_version < '3.11'" },
{ name = "sphinxcontrib-devhelp", marker = "python_full_version < '3.11'" },
{ name = "sphinxcontrib-htmlhelp", marker = "python_full_version < '3.11'" },
{ name = "sphinxcontrib-jsmath", marker = "python_full_version < '3.11'" },
{ name = "sphinxcontrib-qthelp", marker = "python_full_version < '3.11'" },
{ name = "sphinxcontrib-serializinghtml", marker = "python_full_version < '3.11'" },
{ name = "tomli", marker = "python_full_version < '3.11'" },
{ name = "alabaster" },
{ name = "babel" },
{ name = "colorama", marker = "sys_platform == 'win32'" },
{ name = "docutils", version = "0.21.2", source = { registry = "https://pypi.org/simple" } },
{ name = "imagesize" },
{ name = "jinja2" },
{ name = "packaging" },
{ name = "pygments" },
{ name = "requests" },
{ name = "snowballstemmer" },
{ name = "sphinxcontrib-applehelp" },
{ name = "sphinxcontrib-devhelp" },
{ name = "sphinxcontrib-htmlhelp" },
{ name = "sphinxcontrib-jsmath" },
{ name = "sphinxcontrib-qthelp" },
{ name = "sphinxcontrib-serializinghtml" },
{ name = "tomli" },
]
sdist = { url = "https://files.pythonhosted.org/packages/6f/6d/be0b61178fe2cdcb67e2a92fc9ebb488e3c51c4f74a36a7824c0adf23425/sphinx-8.1.3.tar.gz", hash = "sha256:43c1911eecb0d3e161ad78611bc905d1ad0e523e4ddc202a58a821773dc4c927", size = 8184611, upload-time = "2024-10-13T20:27:13.93Z" }
wheels = [
@ -9308,23 +9308,23 @@ resolution-markers = [
"python_full_version == '3.11.*'",
]
dependencies = [
{ name = "alabaster", marker = "python_full_version == '3.11.*'" },
{ name = "babel", marker = "python_full_version == '3.11.*'" },
{ name = "colorama", marker = "python_full_version == '3.11.*' and sys_platform == 'win32'" },
{ name = "docutils", version = "0.22.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" },
{ name = "imagesize", marker = "python_full_version == '3.11.*'" },
{ name = "jinja2", marker = "python_full_version == '3.11.*'" },
{ name = "packaging", marker = "python_full_version == '3.11.*'" },
{ name = "pygments", marker = "python_full_version == '3.11.*'" },
{ name = "requests", marker = "python_full_version == '3.11.*'" },
{ name = "roman-numerals", marker = "python_full_version == '3.11.*'" },
{ name = "snowballstemmer", marker = "python_full_version == '3.11.*'" },
{ name = "sphinxcontrib-applehelp", marker = "python_full_version == '3.11.*'" },
{ name = "sphinxcontrib-devhelp", marker = "python_full_version == '3.11.*'" },
{ name = "sphinxcontrib-htmlhelp", marker = "python_full_version == '3.11.*'" },
{ name = "sphinxcontrib-jsmath", marker = "python_full_version == '3.11.*'" },
{ name = "sphinxcontrib-qthelp", marker = "python_full_version == '3.11.*'" },
{ name = "sphinxcontrib-serializinghtml", marker = "python_full_version == '3.11.*'" },
{ name = "alabaster" },
{ name = "babel" },
{ name = "colorama", marker = "sys_platform == 'win32'" },
{ name = "docutils", version = "0.22.4", source = { registry = "https://pypi.org/simple" } },
{ name = "imagesize" },
{ name = "jinja2" },
{ name = "packaging" },
{ name = "pygments" },
{ name = "requests" },
{ name = "roman-numerals" },
{ name = "snowballstemmer" },
{ name = "sphinxcontrib-applehelp" },
{ name = "sphinxcontrib-devhelp" },
{ name = "sphinxcontrib-htmlhelp" },
{ name = "sphinxcontrib-jsmath" },
{ name = "sphinxcontrib-qthelp" },
{ name = "sphinxcontrib-serializinghtml" },
]
sdist = { url = "https://files.pythonhosted.org/packages/42/50/a8c6ccc36d5eacdfd7913ddccd15a9cee03ecafc5ee2bc40e1f168d85022/sphinx-9.0.4.tar.gz", hash = "sha256:594ef59d042972abbc581d8baa577404abe4e6c3b04ef61bd7fc2acbd51f3fa3", size = 8710502, upload-time = "2025-12-04T07:45:27.343Z" }
wheels = [
@ -9341,23 +9341,23 @@ resolution-markers = [
"python_full_version == '3.12.*'",
]
dependencies = [
{ name = "alabaster", marker = "python_full_version >= '3.12'" },
{ name = "babel", marker = "python_full_version >= '3.12'" },
{ name = "colorama", marker = "python_full_version >= '3.12' and sys_platform == 'win32'" },
{ name = "docutils", version = "0.22.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" },
{ name = "imagesize", marker = "python_full_version >= '3.12'" },
{ name = "jinja2", marker = "python_full_version >= '3.12'" },
{ name = "packaging", marker = "python_full_version >= '3.12'" },
{ name = "pygments", marker = "python_full_version >= '3.12'" },
{ name = "requests", marker = "python_full_version >= '3.12'" },
{ name = "roman-numerals", marker = "python_full_version >= '3.12'" },
{ name = "snowballstemmer", marker = "python_full_version >= '3.12'" },
{ name = "sphinxcontrib-applehelp", marker = "python_full_version >= '3.12'" },
{ name = "sphinxcontrib-devhelp", marker = "python_full_version >= '3.12'" },
{ name = "sphinxcontrib-htmlhelp", marker = "python_full_version >= '3.12'" },
{ name = "sphinxcontrib-jsmath", marker = "python_full_version >= '3.12'" },
{ name = "sphinxcontrib-qthelp", marker = "python_full_version >= '3.12'" },
{ name = "sphinxcontrib-serializinghtml", marker = "python_full_version >= '3.12'" },
{ name = "alabaster" },
{ name = "babel" },
{ name = "colorama", marker = "sys_platform == 'win32'" },
{ name = "docutils", version = "0.22.4", source = { registry = "https://pypi.org/simple" } },
{ name = "imagesize" },
{ name = "jinja2" },
{ name = "packaging" },
{ name = "pygments" },
{ name = "requests" },
{ name = "roman-numerals" },
{ name = "snowballstemmer" },
{ name = "sphinxcontrib-applehelp" },
{ name = "sphinxcontrib-devhelp" },
{ name = "sphinxcontrib-htmlhelp" },
{ name = "sphinxcontrib-jsmath" },
{ name = "sphinxcontrib-qthelp" },
{ name = "sphinxcontrib-serializinghtml" },
]
sdist = { url = "https://files.pythonhosted.org/packages/cd/bd/f08eb0f4eed5c83f1ba2a3bd18f7745a2b1525fad70660a1c00224ec468a/sphinx-9.1.0.tar.gz", hash = "sha256:7741722357dd75f8190766926071fed3bdc211c74dd2d7d4df5404da95930ddb", size = 8718324, upload-time = "2025-12-31T15:09:27.646Z" }
wheels = [
@ -9505,8 +9505,8 @@ name = "standard-aifc"
version = "3.13.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "audioop-lts", marker = "python_full_version >= '3.13'" },
{ name = "standard-chunk", marker = "python_full_version >= '3.13'" },
{ name = "audioop-lts" },
{ name = "standard-chunk" },
]
sdist = { url = "https://files.pythonhosted.org/packages/c4/53/6050dc3dde1671eb3db592c13b55a8005e5040131f7509cef0215212cb84/standard_aifc-3.13.0.tar.gz", hash = "sha256:64e249c7cb4b3daf2fdba4e95721f811bde8bdfc43ad9f936589b7bb2fae2e43", size = 15240, upload-time = "2024-10-30T16:01:31.772Z" }
wheels = [
@ -9527,7 +9527,7 @@ name = "standard-sunau"
version = "3.13.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "audioop-lts", marker = "python_full_version >= '3.13'" },
{ name = "audioop-lts" },
]
sdist = { url = "https://files.pythonhosted.org/packages/66/e3/ce8d38cb2d70e05ffeddc28bb09bad77cfef979eb0a299c9117f7ed4e6a9/standard_sunau-3.13.0.tar.gz", hash = "sha256:b319a1ac95a09a2378a8442f403c66f4fd4b36616d6df6ae82b8e536ee790908", size = 9368, upload-time = "2024-10-30T16:01:41.626Z" }
wheels = [
@ -9561,8 +9561,8 @@ name = "taskgroup"
version = "0.2.2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "exceptiongroup", marker = "python_full_version < '3.11'" },
{ name = "typing-extensions", marker = "python_full_version < '3.11'" },
{ name = "exceptiongroup" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/f0/8d/e218e0160cc1b692e6e0e5ba34e8865dbb171efeb5fc9a704544b3020605/taskgroup-0.2.2.tar.gz", hash = "sha256:078483ac3e78f2e3f973e2edbf6941374fbea81b9c5d0a96f51d297717f4752d", size = 11504, upload-time = "2025-01-03T09:24:13.761Z" }
wheels = [

3
vscode-extension/.gitignore vendored Normal file
View file

@ -0,0 +1,3 @@
node_modules/
dist/
*.vsix

View file

@ -0,0 +1,10 @@
.gitignore
.vscodeignore
node_modules/**
src/**
test/**
tsconfig.json
package-lock.json
**/*.map
**/*.vsix
vitest.config.mts

21
vscode-extension/LICENSE Normal file
View file

@ -0,0 +1,21 @@
MIT License
Copyright (c) 2023 Berri AI
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.

View file

@ -0,0 +1,31 @@
# LiteLLM for VS Code
Chat in VS Code with every model your [LiteLLM AI Gateway](https://docs.litellm.ai) exposes. The extension registers LiteLLM as a language model provider, so the gateway's models show up in the chat model picker next to the built-in ones, with the price and reasoning effort controls the gateway reports for each of them
## What you get
The model list comes from the gateway's `GET /model_group/info` endpoint, scoped to the virtual key you configure, so the picker shows exactly the chat models that key can use. Each model carries its input and output price per 1M tokens in the picker and in the Language Models editor, and its context limits come from the gateway too, so VS Code sizes prompts correctly. A model whose gateway entry lists `supported_reasoning_efforts` gets a Reasoning Effort submenu in the picker's Configure Model menu, and the chosen effort is sent as `reasoning_effort` on every request to that model. Requests go to `POST /v1/chat/completions` on the gateway as streaming chat completions with tools and images passed through, so routing, fallbacks, guardrails, and spend tracking all apply as usual
## Setup
1. Install the extension
2. Run `Chat: Manage Language Models` from the Command Palette and pick `LiteLLM`
3. Enter a name for the connection, the gateway URL (for example `https://litellm.example.com`), and a LiteLLM virtual key. The key is stored in VS Code's secret storage
4. Open the chat model picker. The gateway's chat models are listed under the name you chose, each with its price
Add the same provider again with another name to reach a second gateway or a second key. Run `LiteLLM: Refresh Models` after the gateway's model list changes. To change the key of an existing connection or to drop it, use the gear on its row in the Language Models editor (`Update API Key`, `Delete`); to change the URL, open its entry with `Open in Language Models (JSON)` from the same menu. If the stored key is ever lost the editor shows a `missing its API key` row for that connection until you update the key
## Requirements
VS Code 1.115 or newer and a LiteLLM AI Gateway the key can reach. The key needs access to at least one model group whose mode is `chat`
## Development
```
npm ci
npm run typecheck
npm test
npm run package
```
`npm run package` writes a `.vsix` you can install with `code --install-extension litellm-vscode-<version>.vsix`

3570
vscode-extension/package-lock.json generated Normal file

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,87 @@
{
"name": "litellm-vscode",
"displayName": "LiteLLM",
"description": "Chat with every model behind your LiteLLM AI Gateway in VS Code, with live pricing and reasoning effort controls in the model picker",
"version": "0.1.0",
"publisher": "litellm",
"license": "MIT",
"repository": {
"type": "git",
"url": "https://github.com/BerriAI/litellm.git",
"directory": "vscode-extension"
},
"homepage": "https://docs.litellm.ai",
"bugs": {
"url": "https://github.com/BerriAI/litellm/issues"
},
"engines": {
"vscode": "^1.115.0"
},
"categories": [
"AI",
"Chat"
],
"keywords": [
"litellm",
"ai gateway",
"llm",
"chat",
"copilot"
],
"main": "./dist/extension.js",
"activationEvents": [],
"contributes": {
"languageModelChatProviders": [
{
"vendor": "litellm",
"displayName": "LiteLLM",
"configuration": {
"type": "object",
"properties": {
"baseUrl": {
"type": "string",
"title": "Gateway URL",
"description": "Base URL of your LiteLLM AI Gateway, for example https://litellm.example.com",
"default": "http://localhost:4000"
},
"apiKey": {
"type": "string",
"title": "API key",
"description": "A LiteLLM virtual key. The models offered are the ones this key can access",
"secret": true
}
},
"required": [
"baseUrl",
"apiKey"
]
}
}
],
"commands": [
{
"command": "litellm.refreshModels",
"title": "Refresh Models",
"category": "LiteLLM"
}
]
},
"scripts": {
"build": "esbuild src/extension.ts --bundle --outfile=dist/extension.js --external:vscode --format=cjs --platform=node --target=node22",
"typecheck": "tsc --noEmit",
"test": "vitest run",
"vscode:prepublish": "npm run typecheck && npm run build",
"package": "vsce package --no-dependencies"
},
"dependencies": {
"openai": "^7.18.0"
},
"devDependencies": {
"@types/node": "^22.20.3",
"@types/vscode": "1.115.0",
"@vscode/vsce": "^4.0.0",
"esbuild": "^0.28.2",
"typescript": "^5.9.3",
"vitest": "^4.1.11"
}
}

View file

@ -0,0 +1,17 @@
import * as vscode from "vscode";
import { createGatewayClient } from "./gateway";
import { LiteLLMChatProvider } from "./provider";
export const VENDOR = "litellm";
export const REFRESH_COMMAND = "litellm.refreshModels";
export function activate(context: vscode.ExtensionContext): void {
const provider = new LiteLLMChatProvider(createGatewayClient());
context.subscriptions.push(
provider,
vscode.lm.registerLanguageModelChatProvider(VENDOR, provider),
vscode.commands.registerCommand(REFRESH_COMMAND, () => provider.refresh()),
);
}
export function deactivate(): void {}

View file

@ -0,0 +1,114 @@
import OpenAI from "openai";
import type { ChatCompletionChunk, ChatCompletionCreateParamsStreaming } from "openai/resources/chat/completions";
import packageJson from "../package.json";
import { parseModelGroups, type ConfigurationValues, type ModelGroupInfo } from "./models";
export interface GatewayConfig {
readonly baseUrl: string;
readonly apiKey: string;
}
export type GatewayConfigResult =
| { readonly kind: "ok"; readonly config: GatewayConfig }
| { readonly kind: "unconfigured" }
| { readonly kind: "missing_fields"; readonly fields: readonly string[] }
| { readonly kind: "invalid_url"; readonly baseUrl: string };
export type ModelGroupsResult =
| { readonly kind: "ok"; readonly groups: readonly ModelGroupInfo[] }
| { readonly kind: "http_error"; readonly status: number; readonly body: string }
| { readonly kind: "invalid_response"; readonly reason: string };
export interface GatewayClient {
listModelGroups(config: GatewayConfig, signal: AbortSignal): Promise<ModelGroupsResult>;
streamChatCompletion(
config: GatewayConfig,
params: ChatCompletionCreateParamsStreaming,
signal: AbortSignal,
): Promise<AsyncIterable<ChatCompletionChunk>>;
}
export const USER_AGENT = `litellm-vscode/${packageJson.version}`;
export const ERROR_SUMMARY_LIMIT = 200;
const GATEWAY_PROTOCOLS: ReadonlySet<string> = new Set(["http:", "https:"]);
const parsesAsHttpUrl = (value: string): boolean => {
try {
return GATEWAY_PROTOCOLS.has(new URL(value).protocol);
} catch {
return false;
}
};
export const gatewayRoot = (baseUrl: string): string | undefined => {
const root = baseUrl.trim().replace(/\/+$/, "").replace(/\/v1$/, "");
return parsesAsHttpUrl(root) ? root : undefined;
};
const nonEmptyString = (value: unknown): string | undefined =>
typeof value === "string" && value.trim() !== "" ? value.trim() : undefined;
export const gatewayConfigFrom = (configuration: ConfigurationValues | undefined): GatewayConfigResult => {
if (configuration === undefined) {
return { kind: "unconfigured" };
}
const baseUrl = nonEmptyString(configuration.baseUrl);
const apiKey = nonEmptyString(configuration.apiKey);
if (baseUrl === undefined || apiKey === undefined) {
const fields = [...(baseUrl === undefined ? ["Gateway URL"] : []), ...(apiKey === undefined ? ["API key"] : [])];
return { kind: "missing_fields", fields };
}
const root = gatewayRoot(baseUrl);
return root === undefined ? { kind: "invalid_url", baseUrl } : { kind: "ok", config: { baseUrl: root, apiKey } };
};
export const modelGroupInfoUrl = (root: string): string => `${root}/model_group/info`;
export const openAiBaseUrl = (root: string): string => `${root}/v1`;
const isRecord = (value: unknown): value is Record<string, unknown> => typeof value === "object" && value !== null;
const errorMessageIn = (body: string): string | undefined => {
try {
const parsed: unknown = JSON.parse(body);
if (!isRecord(parsed)) {
return undefined;
}
if (isRecord(parsed.error) && typeof parsed.error.message === "string") {
return parsed.error.message;
}
return typeof parsed.detail === "string" ? parsed.detail : undefined;
} catch {
return undefined;
}
};
export const summarizeErrorBody = (body: string): string => {
const message = (errorMessageIn(body) ?? body).replace(/\s+/g, " ").trim();
return message.length > ERROR_SUMMARY_LIMIT ? `${message.slice(0, ERROR_SUMMARY_LIMIT)}...` : message;
};
export const createGatewayClient = (fetchImpl: typeof fetch = fetch): GatewayClient => ({
async listModelGroups(config, signal) {
const response = await fetchImpl(modelGroupInfoUrl(config.baseUrl), {
headers: { Authorization: `Bearer ${config.apiKey}`, "User-Agent": USER_AGENT },
signal,
});
if (!response.ok) {
return { kind: "http_error", status: response.status, body: await response.text() };
}
const parsed = parseModelGroups(await response.json());
return parsed.kind === "ok" ? parsed : { kind: "invalid_response", reason: parsed.reason };
},
streamChatCompletion(config, params, signal) {
const client = new OpenAI({
apiKey: config.apiKey,
baseURL: openAiBaseUrl(config.baseUrl),
defaultHeaders: { "User-Agent": USER_AGENT },
fetch: fetchImpl,
maxRetries: 0,
});
return client.chat.completions.create(params, { signal });
},
});

View file

@ -0,0 +1,188 @@
import type * as vscode from "vscode";
import type {
ChatCompletionAssistantMessageParam,
ChatCompletionContentPart,
ChatCompletionCreateParamsStreaming,
ChatCompletionMessageParam,
ChatCompletionMessageToolCall,
ChatCompletionTool,
ChatCompletionToolMessageParam,
} from "openai/resources/chat/completions";
import { estimateTokens } from "./models";
export interface ChatRequestInput {
readonly model: string;
readonly messages: readonly vscode.LanguageModelChatRequestMessage[];
readonly tools: readonly vscode.LanguageModelChatTool[];
readonly requireToolCall: boolean;
readonly reasoningEffort: string | undefined;
readonly modelOptions: { readonly [key: string]: unknown };
}
interface TextPart {
readonly value: string;
}
interface ToolCallPart {
readonly callId: string;
readonly name: string;
readonly input: object;
}
interface ToolResultPart {
readonly callId: string;
readonly content: ReadonlyArray<unknown>;
}
interface DataPart {
readonly mimeType: string;
readonly data: Uint8Array;
}
const USER_ROLE = 1;
const ASSISTANT_ROLE = 2;
const SYSTEM_ROLE = 3;
export const ESTIMATED_TOKENS_PER_IMAGE = 1000;
const isRecord = (value: unknown): value is Record<string, unknown> => typeof value === "object" && value !== null;
const isTextPart = (part: unknown): part is TextPart => isRecord(part) && typeof part.value === "string";
const isToolCallPart = (part: unknown): part is ToolCallPart =>
isRecord(part) && typeof part.callId === "string" && typeof part.name === "string" && isRecord(part.input);
const isToolResultPart = (part: unknown): part is ToolResultPart =>
isRecord(part) && typeof part.callId === "string" && Array.isArray(part.content);
const isDataPart = (part: unknown): part is DataPart =>
isRecord(part) && typeof part.mimeType === "string" && part.data instanceof Uint8Array;
const isImagePart = (part: unknown): part is DataPart => isDataPart(part) && part.mimeType.startsWith("image/");
const dataUrl = (part: DataPart): string => `data:${part.mimeType};base64,${Buffer.from(part.data).toString("base64")}`;
const textOf = (part: unknown): string => {
if (isTextPart(part)) {
return part.value;
}
if (isDataPart(part) && part.mimeType.startsWith("text/")) {
return Buffer.from(part.data).toString("utf8");
}
if (isRecord(part) && "value" in part) {
return JSON.stringify(part.value);
}
return "";
};
const contentParts = (parts: readonly unknown[]): readonly ChatCompletionContentPart[] =>
parts.flatMap((part): readonly ChatCompletionContentPart[] => {
if (isImagePart(part)) {
return [{ type: "image_url", image_url: { url: dataUrl(part) } }];
}
const text = textOf(part);
return text === "" ? [] : [{ type: "text", text }];
});
const toolMessage = (part: ToolResultPart): ChatCompletionToolMessageParam => ({
role: "tool",
tool_call_id: part.callId,
content: part.content.filter((item) => !isImagePart(item)).map(textOf).join(""),
});
const userMessages = (parts: readonly unknown[]): readonly ChatCompletionMessageParam[] => {
const toolResults = parts.filter(isToolResultPart);
const toolResultImages = toolResults.flatMap((result) => result.content.filter(isImagePart));
const remaining = parts.filter((part) => !isToolResultPart(part));
const userContent = contentParts([...remaining, ...toolResultImages]);
const userMessage: readonly ChatCompletionMessageParam[] =
userContent.length === 0 ? [] : [{ role: "user", content: [...userContent] }];
return [...toolResults.map(toolMessage), ...userMessage];
};
const toolCall = (part: ToolCallPart): ChatCompletionMessageToolCall => ({
id: part.callId,
type: "function",
function: { name: part.name, arguments: JSON.stringify(part.input) },
});
const assistantMessages = (parts: readonly unknown[]): readonly ChatCompletionAssistantMessageParam[] => {
const text = parts.filter(isTextPart).map((part) => part.value).join("");
const toolCalls = parts.filter(isToolCallPart).map(toolCall);
if (text === "" && toolCalls.length === 0) {
return [];
}
return [
{
role: "assistant",
content: text === "" ? null : text,
...(toolCalls.length === 0 ? {} : { tool_calls: toolCalls }),
},
];
};
const convertMessage = (message: vscode.LanguageModelChatRequestMessage): readonly ChatCompletionMessageParam[] => {
const role: number = message.role;
switch (role) {
case USER_ROLE:
return userMessages(message.content);
case ASSISTANT_ROLE:
return assistantMessages(message.content);
case SYSTEM_ROLE:
return [{ role: "system", content: message.content.map(textOf).join("") }];
default:
return [];
}
};
export const toChatCompletionMessages = (
messages: readonly vscode.LanguageModelChatRequestMessage[],
): readonly ChatCompletionMessageParam[] => messages.flatMap(convertMessage);
const imagePartsIn = (parts: readonly unknown[]): readonly DataPart[] => [
...parts.filter(isImagePart),
...parts.filter(isToolResultPart).flatMap((result) => result.content.filter(isImagePart)),
];
const withoutImageData = (key: string, value: unknown): unknown => (key === "image_url" ? undefined : value);
export const estimateMessageTokens = (message: vscode.LanguageModelChatRequestMessage): number => {
const converted = convertMessage(message);
if (converted.length === 0) {
return 0;
}
const images = imagePartsIn(message.content).length;
return estimateTokens(JSON.stringify(converted, withoutImageData)) + images * ESTIMATED_TOKENS_PER_IMAGE;
};
const toTool = (tool: vscode.LanguageModelChatTool): ChatCompletionTool => ({
type: "function",
function: {
name: tool.name,
description: tool.description,
...(tool.inputSchema === undefined ? {} : { parameters: tool.inputSchema as Record<string, unknown> }),
},
});
const NUMERIC_OPTIONS = ["temperature", "top_p", "max_tokens", "presence_penalty", "frequency_penalty", "seed"] as const;
const forwardedModelOptions = (modelOptions: { readonly [key: string]: unknown }): Record<string, number> =>
Object.fromEntries(
NUMERIC_OPTIONS.flatMap((key) => {
const value = modelOptions[key];
return typeof value === "number" ? [[key, value] as const] : [];
}),
);
export const buildChatCompletionParams = (input: ChatRequestInput): ChatCompletionCreateParamsStreaming => ({
model: input.model,
messages: [...toChatCompletionMessages(input.messages)],
stream: true,
stream_options: { include_usage: true },
...forwardedModelOptions(input.modelOptions),
...(input.tools.length === 0 ? {} : { tools: input.tools.map(toTool) }),
...(input.tools.length === 0 || !input.requireToolCall ? {} : { tool_choice: "required" }),
...(input.reasoningEffort === undefined
? {}
: { reasoning_effort: input.reasoningEffort as ChatCompletionCreateParamsStreaming["reasoning_effort"] }),
});

View file

@ -0,0 +1,170 @@
export interface ModelGroupInfo {
readonly modelGroup: string;
readonly providers: readonly string[];
readonly mode: string | undefined;
readonly maxInputTokens: number | undefined;
readonly maxOutputTokens: number | undefined;
readonly inputCostPerToken: number | undefined;
readonly outputCostPerToken: number | undefined;
readonly supportsVision: boolean;
readonly supportsFunctionCalling: boolean;
readonly supportedReasoningEfforts: readonly string[];
}
export type ConfigurationValues = { readonly [key: string]: unknown };
export interface ConfigurationSchemaProperty {
readonly type: "string";
readonly title: string;
readonly enum: readonly string[];
readonly enumItemLabels: readonly string[];
readonly default: string;
readonly group: "navigation";
}
export interface ConfigurationSchema {
readonly properties: { readonly [key: string]: ConfigurationSchemaProperty };
}
export interface ModelDescriptor {
readonly id: string;
readonly name: string;
readonly family: string;
readonly version: string;
readonly detail: string;
readonly tooltip: string;
readonly maxInputTokens: number;
readonly maxOutputTokens: number;
readonly imageInput: boolean;
readonly toolCalling: boolean;
readonly configurationSchema: ConfigurationSchema | undefined;
}
export type ModelGroupsParseResult =
| { readonly kind: "ok"; readonly groups: readonly ModelGroupInfo[] }
| { readonly kind: "invalid"; readonly reason: string };
export const REASONING_EFFORT_KEY = "reasoningEffort";
export const GATEWAY_DEFAULT_EFFORT = "default";
export const ASSUMED_MAX_INPUT_TOKENS = 128000;
export const ASSUMED_MAX_OUTPUT_TOKENS = 4096;
export const MARKDOWN_LINE_BREAK = " \n";
const isRecord = (value: unknown): value is Record<string, unknown> => typeof value === "object" && value !== null;
const optionalNumber = (value: unknown): number | undefined =>
typeof value === "number" && Number.isFinite(value) ? value : undefined;
const optionalString = (value: unknown): string | undefined => (typeof value === "string" ? value : undefined);
const stringList = (value: unknown): readonly string[] =>
Array.isArray(value) ? value.filter((item): item is string => typeof item === "string") : [];
const parseGroup = (value: unknown): ModelGroupInfo | undefined => {
if (!isRecord(value) || typeof value.model_group !== "string") {
return undefined;
}
return {
modelGroup: value.model_group,
providers: stringList(value.providers),
mode: optionalString(value.mode),
maxInputTokens: optionalNumber(value.max_input_tokens),
maxOutputTokens: optionalNumber(value.max_output_tokens),
inputCostPerToken: optionalNumber(value.input_cost_per_token),
outputCostPerToken: optionalNumber(value.output_cost_per_token),
supportsVision: value.supports_vision === true,
supportsFunctionCalling: value.supports_function_calling === true,
supportedReasoningEfforts: stringList(value.supported_reasoning_efforts),
};
};
export const parseModelGroups = (body: unknown): ModelGroupsParseResult => {
if (!isRecord(body) || !Array.isArray(body.data)) {
return { kind: "invalid", reason: "response has no data array" };
}
const groups = body.data.map(parseGroup).filter((group): group is ModelGroupInfo => group !== undefined);
return { kind: "ok", groups };
};
const isChatGroup = (group: ModelGroupInfo): boolean => group.mode === undefined || group.mode === "chat";
export const formatUsdPerMillionTokens = (costPerToken: number): string => {
const perMillion = costPerToken * 1_000_000;
const digits = perMillion === 0 || perMillion >= 0.01 ? perMillion.toFixed(2) : perMillion.toPrecision(2);
return `$${digits}`;
};
const priceLine = (label: string, costPerToken: number | undefined): string =>
costPerToken === undefined ? `${label}: no price configured` : `${label}: ${formatUsdPerMillionTokens(costPerToken)} per 1M tokens`;
const pricingDetail = (group: ModelGroupInfo): string => {
if (group.inputCostPerToken === undefined && group.outputCostPerToken === undefined) {
return "No pricing configured";
}
const input = group.inputCostPerToken === undefined ? "n/a" : formatUsdPerMillionTokens(group.inputCostPerToken);
const output = group.outputCostPerToken === undefined ? "n/a" : formatUsdPerMillionTokens(group.outputCostPerToken);
return `${input} in / ${output} out per 1M tokens`;
};
const capitalize = (value: string): string => value.charAt(0).toUpperCase() + value.slice(1);
const effortSchema = (efforts: readonly string[]): ConfigurationSchema | undefined => {
if (efforts.length === 0) {
return undefined;
}
return {
properties: {
[REASONING_EFFORT_KEY]: {
type: "string",
title: "Reasoning Effort",
enum: [GATEWAY_DEFAULT_EFFORT, ...efforts],
enumItemLabels: ["Gateway default", ...efforts.map(capitalize)],
default: GATEWAY_DEFAULT_EFFORT,
group: "navigation",
},
},
};
};
const tooltipFor = (group: ModelGroupInfo): string => {
const providers = group.providers.length === 0 ? "" : ` via ${group.providers.join(", ")}`;
const context =
group.maxInputTokens === undefined || group.maxOutputTokens === undefined
? `Context: unknown, assuming ${ASSUMED_MAX_INPUT_TOKENS} in / ${ASSUMED_MAX_OUTPUT_TOKENS} out tokens`
: `Context: ${group.maxInputTokens} in / ${group.maxOutputTokens} out tokens`;
const efforts =
group.supportedReasoningEfforts.length === 0
? "Reasoning effort: not configurable"
: `Reasoning effort: ${group.supportedReasoningEfforts.join(", ")}`;
return [
`LiteLLM model group ${group.modelGroup}${providers}`,
priceLine("Input", group.inputCostPerToken),
priceLine("Output", group.outputCostPerToken),
context,
efforts,
].join(MARKDOWN_LINE_BREAK);
};
const describeGroup = (group: ModelGroupInfo): ModelDescriptor => ({
id: group.modelGroup,
name: group.modelGroup,
family: group.modelGroup,
version: "1.0",
detail: pricingDetail(group),
tooltip: tooltipFor(group),
maxInputTokens: group.maxInputTokens ?? ASSUMED_MAX_INPUT_TOKENS,
maxOutputTokens: group.maxOutputTokens ?? ASSUMED_MAX_OUTPUT_TOKENS,
imageInput: group.supportsVision,
toolCalling: group.supportsFunctionCalling,
configurationSchema: effortSchema(group.supportedReasoningEfforts),
});
export const describeModels = (groups: readonly ModelGroupInfo[]): readonly ModelDescriptor[] =>
groups.filter(isChatGroup).map(describeGroup);
export const reasoningEffortFrom = (configuration: ConfigurationValues | undefined): string | undefined => {
const effort = configuration?.[REASONING_EFFORT_KEY];
return typeof effort === "string" && effort !== GATEWAY_DEFAULT_EFFORT ? effort : undefined;
};
export const estimateTokens = (text: string): number => Math.ceil(text.length / 4);

View file

@ -0,0 +1,136 @@
import * as vscode from "vscode";
import { gatewayConfigFrom, summarizeErrorBody, type GatewayClient, type GatewayConfig, type GatewayConfigResult, type ModelGroupsResult } from "./gateway";
import { buildChatCompletionParams, estimateMessageTokens } from "./messages";
import { describeModels, estimateTokens, reasoningEffortFrom, type ModelDescriptor } from "./models";
import { responseParts, type ResponsePart } from "./stream";
export interface LiteLLMModel extends vscode.LanguageModelChatInformation {
readonly gateway: GatewayConfig;
}
export const TRUNCATED_MESSAGE = "The model stopped at its output token limit before finishing the response";
const RECONFIGURE_HINT =
'Fix it from the gear on its row in Manage Language Models: "Update API Key" for the key, "Open in Language Models (JSON)" for the URL';
const configurationProblem = (result: Exclude<GatewayConfigResult, { kind: "ok" | "unconfigured" }>): string => {
switch (result.kind) {
case "missing_fields":
return `LiteLLM provider is missing its ${result.fields.join(" and ")}. ${RECONFIGURE_HINT}`;
case "invalid_url":
return `LiteLLM gateway URL "${result.baseUrl}" is not an http or https URL. ${RECONFIGURE_HINT}`;
}
};
const discoveryFailure = (result: Exclude<ModelGroupsResult, { kind: "ok" }>, baseUrl: string): string => {
switch (result.kind) {
case "http_error":
return `LiteLLM gateway at ${baseUrl} answered ${result.status} for /model_group/info: ${summarizeErrorBody(result.body)}`;
case "invalid_response":
return `LiteLLM gateway at ${baseUrl} returned an unexpected /model_group/info payload: ${result.reason}`;
}
};
const toModel = (descriptor: ModelDescriptor, gateway: GatewayConfig): LiteLLMModel => ({
id: descriptor.id,
name: descriptor.name,
family: descriptor.family,
version: descriptor.version,
detail: descriptor.detail,
tooltip: descriptor.tooltip,
maxInputTokens: descriptor.maxInputTokens,
maxOutputTokens: descriptor.maxOutputTokens,
capabilities: { imageInput: descriptor.imageInput, toolCalling: descriptor.toolCalling },
...(descriptor.configurationSchema === undefined ? {} : { configurationSchema: descriptor.configurationSchema }),
gateway,
});
const toVscodePart = (part: ResponsePart): vscode.LanguageModelResponsePart => {
switch (part.kind) {
case "text":
return new vscode.LanguageModelTextPart(part.value);
case "tool_call":
return new vscode.LanguageModelToolCallPart(part.callId, part.name, part.input);
case "invalid_tool_call":
throw new Error(`Model returned invalid JSON arguments for tool ${part.name}: ${part.arguments}`);
case "truncated":
throw new Error(TRUNCATED_MESSAGE);
}
};
const withAbortSignal = async <T>(token: vscode.CancellationToken, run: (signal: AbortSignal) => Promise<T>): Promise<T> => {
const controller = new AbortController();
const subscription = token.onCancellationRequested(() => controller.abort());
try {
return await run(controller.signal);
} finally {
subscription.dispose();
}
};
export class LiteLLMChatProvider implements vscode.LanguageModelChatProvider<LiteLLMModel>, vscode.Disposable {
private readonly changeEmitter = new vscode.EventEmitter<void>();
readonly onDidChangeLanguageModelChatInformation = this.changeEmitter.event;
constructor(private readonly gateway: GatewayClient) {}
refresh(): void {
this.changeEmitter.fire();
}
dispose(): void {
this.changeEmitter.dispose();
}
async provideLanguageModelChatInformation(
options: vscode.PrepareLanguageModelChatModelOptions,
token: vscode.CancellationToken,
): Promise<LiteLLMModel[]> {
const configured = gatewayConfigFrom(options.configuration);
if (configured.kind === "unconfigured") {
return [];
}
if (configured.kind !== "ok") {
throw new Error(configurationProblem(configured));
}
const result = await withAbortSignal(token, (signal) => this.gateway.listModelGroups(configured.config, signal));
if (result.kind !== "ok") {
throw new Error(discoveryFailure(result, configured.config.baseUrl));
}
return describeModels(result.groups).map((descriptor) => toModel(descriptor, configured.config));
}
async provideLanguageModelChatResponse(
model: LiteLLMModel,
messages: readonly vscode.LanguageModelChatRequestMessage[],
options: vscode.ProvideLanguageModelChatResponseOptions,
progress: vscode.Progress<vscode.LanguageModelResponsePart>,
token: vscode.CancellationToken,
): Promise<void> {
const params = buildChatCompletionParams({
model: model.id,
messages,
tools: options.tools ?? [],
requireToolCall: options.toolMode === vscode.LanguageModelChatToolMode.Required,
reasoningEffort: reasoningEffortFrom(options.modelConfiguration),
modelOptions: options.modelOptions ?? {},
});
try {
await withAbortSignal(token, async (signal) => {
const chunks = await this.gateway.streamChatCompletion(model.gateway, params, signal);
for await (const part of responseParts(chunks)) {
progress.report(toVscodePart(part));
}
});
} catch (error) {
if (token.isCancellationRequested) {
return;
}
throw error;
}
}
async provideTokenCount(_model: LiteLLMModel, text: string | vscode.LanguageModelChatRequestMessage): Promise<number> {
return typeof text === "string" ? estimateTokens(text) : estimateMessageTokens(text);
}
}

View file

@ -0,0 +1,90 @@
import type { ChatCompletionChunk } from "openai/resources/chat/completions";
export type ResponsePart =
| { readonly kind: "text"; readonly value: string }
| { readonly kind: "tool_call"; readonly callId: string; readonly name: string; readonly input: object }
| { readonly kind: "invalid_tool_call"; readonly callId: string; readonly name: string; readonly arguments: string }
| { readonly kind: "truncated" };
interface PendingToolCall {
readonly index: number;
readonly callId: string;
readonly name: string;
readonly arguments: string;
}
export type PendingToolCalls = readonly PendingToolCall[];
export interface ChunkOutcome {
readonly pending: PendingToolCalls;
readonly parts: readonly ResponsePart[];
}
export const NO_PENDING_TOOL_CALLS: PendingToolCalls = [];
type ToolCallDelta = NonNullable<ChatCompletionChunk.Choice.Delta["tool_calls"]>[number];
const nonEmpty = (value: string | undefined): string | undefined => (value === undefined || value === "" ? undefined : value);
const targetOf = (pending: PendingToolCalls, delta: ToolCallDelta): PendingToolCall | undefined => {
const id = nonEmpty(delta.id);
if (id !== undefined) {
return pending.find((call) => call.callId === id);
}
const sameIndex = pending.filter((call) => call.index === delta.index);
return sameIndex.at(-1) ?? (delta.index === undefined ? pending.at(-1) : undefined);
};
const mergeToolCallDelta = (pending: PendingToolCalls, delta: ToolCallDelta): PendingToolCalls => {
const target = targetOf(pending, delta);
const base: PendingToolCall = target ?? { index: delta.index ?? pending.length, callId: nonEmpty(delta.id) ?? "", name: "", arguments: "" };
const merged: PendingToolCall = {
...base,
name: nonEmpty(delta.function?.name) ?? base.name,
arguments: base.arguments + (delta.function?.arguments ?? ""),
};
return target === undefined ? [...pending, merged] : pending.map((call) => (call === target ? merged : call));
};
export const applyChunk = (pending: PendingToolCalls, chunk: ChatCompletionChunk): ChunkOutcome => {
const choice = chunk.choices[0];
if (choice === undefined) {
return { pending, parts: [] };
}
const text = typeof choice.delta.content === "string" && choice.delta.content !== "" ? [{ kind: "text", value: choice.delta.content } as const] : [];
const truncated = choice.finish_reason === "length" ? [{ kind: "truncated" } as const] : [];
const nextPending = (choice.delta.tool_calls ?? []).reduce(mergeToolCallDelta, pending);
return { pending: nextPending, parts: [...text, ...truncated] };
};
const parseArguments = (raw: string): object | undefined => {
if (raw.trim() === "") {
return {};
}
try {
const parsed: unknown = JSON.parse(raw);
return typeof parsed === "object" && parsed !== null ? parsed : undefined;
} catch {
return undefined;
}
};
const finishToolCall = (call: PendingToolCall): ResponsePart => {
const input = parseArguments(call.arguments);
return input === undefined
? { kind: "invalid_tool_call", callId: call.callId, name: call.name, arguments: call.arguments }
: { kind: "tool_call", callId: call.callId, name: call.name, input };
};
export const flushToolCalls = (pending: PendingToolCalls): readonly ResponsePart[] =>
[...pending].sort((left, right) => left.index - right.index).map(finishToolCall);
export async function* responseParts(chunks: AsyncIterable<ChatCompletionChunk>): AsyncGenerator<ResponsePart> {
let pending: PendingToolCalls = NO_PENDING_TOOL_CALLS;
for await (const chunk of chunks) {
const outcome = applyChunk(pending, chunk);
pending = outcome.pending;
yield* outcome.parts;
}
yield* flushToolCalls(pending);
}

View file

@ -0,0 +1,15 @@
import type { ConfigurationSchema, ConfigurationValues } from "./models";
declare module "vscode" {
interface LanguageModelChatInformation {
readonly configurationSchema?: ConfigurationSchema;
}
interface PrepareLanguageModelChatModelOptions {
readonly configuration?: ConfigurationValues;
}
interface ProvideLanguageModelChatResponseOptions {
readonly modelConfiguration?: ConfigurationValues;
}
}

View file

@ -0,0 +1,205 @@
import { createServer, type IncomingMessage, type Server, type ServerResponse } from "node:http";
import type { AddressInfo } from "node:net";
import { afterEach, describe, expect, it } from "vitest";
import {
ERROR_SUMMARY_LIMIT,
createGatewayClient,
gatewayConfigFrom,
gatewayRoot,
modelGroupInfoUrl,
openAiBaseUrl,
summarizeErrorBody,
USER_AGENT,
type GatewayConfig,
} from "../src/gateway";
import { buildChatCompletionParams } from "../src/messages";
interface RecordedRequest {
readonly method: string | undefined;
readonly url: string | undefined;
readonly authorization: string | undefined;
readonly userAgent: string | undefined;
readonly body: string;
}
type Handler = (request: RecordedRequest, response: ServerResponse) => void;
const readBody = (request: IncomingMessage): Promise<string> =>
new Promise((resolve) => {
const chunks: Buffer[] = [];
request.on("data", (chunk: Buffer) => chunks.push(chunk));
request.on("end", () => resolve(Buffer.concat(chunks).toString("utf8")));
});
const servers: Server[] = [];
const startGateway = (handler: Handler): Promise<{ readonly url: string; readonly requests: readonly RecordedRequest[] }> =>
new Promise((resolve) => {
const requests: RecordedRequest[] = [];
const server = createServer(async (request, response) => {
const recorded: RecordedRequest = {
method: request.method,
url: request.url,
authorization: request.headers.authorization,
userAgent: request.headers["user-agent"],
body: await readBody(request),
};
requests.push(recorded);
handler(recorded, response);
});
servers.push(server);
server.listen(0, "127.0.0.1", () => {
const { port } = server.address() as AddressInfo;
resolve({ url: `http://127.0.0.1:${port}`, requests });
});
});
afterEach(() => {
servers.splice(0).forEach((server) => server.close());
});
const configFor = (baseUrl: string, apiKey: string): GatewayConfig => {
const result = gatewayConfigFrom({ baseUrl, apiKey });
if (result.kind !== "ok") {
throw new Error(result.kind);
}
return result.config;
};
const sse = (response: ServerResponse, events: readonly object[]): void => {
response.writeHead(200, { "content-type": "text/event-stream" });
events.forEach((event) => response.write(`data: ${JSON.stringify(event)}\n\n`));
response.end("data: [DONE]\n\n");
};
describe("gateway URLs", () => {
it("accepts the gateway root with or without a trailing slash or /v1", () => {
expect(gatewayRoot("https://litellm.example.com/")).toBe("https://litellm.example.com");
expect(gatewayRoot("https://litellm.example.com/v1")).toBe("https://litellm.example.com");
expect(gatewayRoot(" http://localhost:4000 ")).toBe("http://localhost:4000");
expect(modelGroupInfoUrl("https://litellm.example.com")).toBe("https://litellm.example.com/model_group/info");
expect(openAiBaseUrl("https://litellm.example.com")).toBe("https://litellm.example.com/v1");
});
it("rejects anything that is not an http or https URL", () => {
expect(gatewayRoot("litellm.example.com")).toBeUndefined();
expect(gatewayRoot("ftp://litellm.example.com")).toBeUndefined();
expect(gatewayRoot("")).toBeUndefined();
});
});
describe("gatewayConfigFrom", () => {
it("distinguishes the unconfigured probe, a lost secret, and a bad URL from a usable configuration", () => {
expect(gatewayConfigFrom(undefined)).toEqual({ kind: "unconfigured" });
expect(gatewayConfigFrom({ baseUrl: "http://localhost:4000", apiKey: undefined })).toEqual({ kind: "missing_fields", fields: ["API key"] });
expect(gatewayConfigFrom({ baseUrl: " ", apiKey: "" })).toEqual({ kind: "missing_fields", fields: ["Gateway URL", "API key"] });
expect(gatewayConfigFrom({ baseUrl: "localhost:4000", apiKey: "sk" })).toEqual({ kind: "invalid_url", baseUrl: "localhost:4000" });
expect(gatewayConfigFrom({ baseUrl: " http://localhost:4000/v1/ ", apiKey: " sk-test " })).toEqual({
kind: "ok",
config: { baseUrl: "http://localhost:4000", apiKey: "sk-test" },
});
});
});
describe("summarizeErrorBody", () => {
it("prefers the gateway's error message and caps the length", () => {
expect(summarizeErrorBody('{"error":{"message":"invalid key","type":"auth_error","param":"sk-...abcd"}}')).toBe("invalid key");
expect(summarizeErrorBody('{"detail":"Not Found"}')).toBe("Not Found");
expect(summarizeErrorBody("<html>\n 502 Bad Gateway\n</html>")).toBe("<html> 502 Bad Gateway </html>");
const long = summarizeErrorBody("x".repeat(ERROR_SUMMARY_LIMIT + 50));
expect(long).toBe(`${"x".repeat(ERROR_SUMMARY_LIMIT)}...`);
});
});
describe("listModelGroups", () => {
it("calls /model_group/info with the virtual key and this extension's user agent", async () => {
const gateway = await startGateway((_request, response) => {
response.writeHead(200, { "content-type": "application/json" });
response.end(JSON.stringify({ data: [{ model_group: "gpt-5.6", mode: "chat", input_cost_per_token: 4e-6 }] }));
});
const result = await createGatewayClient().listModelGroups(configFor(`${gateway.url}/v1`, "sk-test"), new AbortController().signal);
expect(result).toEqual({
kind: "ok",
groups: [expect.objectContaining({ modelGroup: "gpt-5.6", inputCostPerToken: 4e-6 })],
});
expect(gateway.requests).toEqual([
expect.objectContaining({ method: "GET", url: "/model_group/info", authorization: "Bearer sk-test", userAgent: USER_AGENT }),
]);
});
it("reports the gateway's status and body when the key is rejected", async () => {
const gateway = await startGateway((_request, response) => {
response.writeHead(401, { "content-type": "application/json" });
response.end('{"error":{"message":"invalid key"}}');
});
expect(await createGatewayClient().listModelGroups({ baseUrl: gateway.url, apiKey: "sk-bad" }, new AbortController().signal)).toEqual({
kind: "http_error",
status: 401,
body: '{"error":{"message":"invalid key"}}',
});
});
it("reports a payload that is not a model group listing", async () => {
const gateway = await startGateway((_request, response) => {
response.writeHead(200, { "content-type": "application/json" });
response.end('{"object":"list","models":[]}');
});
expect(await createGatewayClient().listModelGroups({ baseUrl: gateway.url, apiKey: "sk" }, new AbortController().signal)).toEqual({
kind: "invalid_response",
reason: "response has no data array",
});
});
});
describe("streamChatCompletion", () => {
it("streams /v1/chat/completions through the gateway with the chosen reasoning effort", async () => {
const gateway = await startGateway((_request, response) =>
sse(response, [
{ id: "c", object: "chat.completion.chunk", created: 0, model: "gpt-5.6", choices: [{ index: 0, delta: { content: "Hi" }, finish_reason: null }] },
{ id: "c", object: "chat.completion.chunk", created: 0, model: "gpt-5.6", choices: [{ index: 0, delta: {}, finish_reason: "stop" }] },
]),
);
const params = buildChatCompletionParams({
model: "gpt-5.6",
messages: [{ role: 1, content: [{ value: "hello" }], name: undefined }],
tools: [],
requireToolCall: false,
reasoningEffort: "high",
modelOptions: {},
});
const chunks = await createGatewayClient().streamChatCompletion(configFor(`${gateway.url}/`, "sk-test"), params, new AbortController().signal);
const contents: string[] = [];
for await (const chunk of chunks) {
contents.push(chunk.choices[0]?.delta.content ?? "");
}
expect(contents.join("")).toBe("Hi");
const [request] = gateway.requests;
expect(request).toMatchObject({ method: "POST", url: "/v1/chat/completions", authorization: "Bearer sk-test", userAgent: USER_AGENT });
expect(JSON.parse(request?.body ?? "{}")).toMatchObject({
model: "gpt-5.6",
stream: true,
stream_options: { include_usage: true },
reasoning_effort: "high",
messages: [{ role: "user", content: [{ type: "text", text: "hello" }] }],
});
});
it("leaves retries to the gateway instead of resending a failed request", async () => {
const gateway = await startGateway((_request, response) => {
response.writeHead(502, { "content-type": "application/json" });
response.end('{"error":{"message":"upstream unavailable"}}');
});
const params = buildChatCompletionParams({
model: "gpt-5.6",
messages: [{ role: 1, content: [{ value: "hello" }], name: undefined }],
tools: [],
requireToolCall: false,
reasoningEffort: undefined,
modelOptions: {},
});
await expect(
createGatewayClient().streamChatCompletion({ baseUrl: gateway.url, apiKey: "sk-test" }, params, new AbortController().signal),
).rejects.toThrow(/upstream unavailable/);
expect(gateway.requests).toHaveLength(1);
});
});

View file

@ -0,0 +1,168 @@
import { describe, expect, it } from "vitest";
import type * as vscode from "vscode";
import { ESTIMATED_TOKENS_PER_IMAGE, buildChatCompletionParams, estimateMessageTokens, toChatCompletionMessages, type ChatRequestInput } from "../src/messages";
const USER = 1 as vscode.LanguageModelChatMessageRole;
const ASSISTANT = 2 as vscode.LanguageModelChatMessageRole;
const SYSTEM = 3 as vscode.LanguageModelChatMessageRole;
const message = (role: vscode.LanguageModelChatMessageRole, content: readonly unknown[]): vscode.LanguageModelChatRequestMessage => ({
role,
content,
name: undefined,
});
const text = (value: string): unknown => ({ value });
const image = (bytes: readonly number[], mimeType = "image/png"): unknown => ({ mimeType, data: Uint8Array.from(bytes) });
const toolCall = (callId: string, name: string, input: object): unknown => ({ callId, name, input });
const toolResult = (callId: string, content: readonly unknown[]): unknown => ({ callId, content });
const request = (overrides: Partial<ChatRequestInput> = {}): ChatRequestInput => ({
model: "gpt-5.6",
messages: [message(USER, [text("hi")])],
tools: [],
requireToolCall: false,
reasoningEffort: undefined,
modelOptions: {},
...overrides,
});
describe("toChatCompletionMessages", () => {
it("maps system, user, and assistant text", () => {
expect(
toChatCompletionMessages([
message(SYSTEM, [text("be terse")]),
message(USER, [text("hello "), text("there")]),
message(ASSISTANT, [text("hi")]),
]),
).toEqual([
{ role: "system", content: "be terse" },
{ role: "user", content: [{ type: "text", text: "hello " }, { type: "text", text: "there" }] },
{ role: "assistant", content: "hi" },
]);
});
it("sends user images as data URLs", () => {
expect(toChatCompletionMessages([message(USER, [text("what is this"), image([1, 2, 3])])])).toEqual([
{
role: "user",
content: [
{ type: "text", text: "what is this" },
{ type: "image_url", image_url: { url: "data:image/png;base64,AQID" } },
],
},
]);
});
it("round-trips tool calls and puts tool results before the user's follow-up text", () => {
expect(
toChatCompletionMessages([
message(ASSISTANT, [text("checking"), toolCall("call_1", "read_file", { path: "a.ts" })]),
message(USER, [toolResult("call_1", [text("export const a = 1;")]), text("thanks")]),
]),
).toEqual([
{
role: "assistant",
content: "checking",
tool_calls: [{ id: "call_1", type: "function", function: { name: "read_file", arguments: '{"path":"a.ts"}' } }],
},
{ role: "tool", tool_call_id: "call_1", content: "export const a = 1;" },
{ role: "user", content: [{ type: "text", text: "thanks" }] },
]);
});
it("emits a content-less assistant turn that only called tools", () => {
expect(toChatCompletionMessages([message(ASSISTANT, [toolCall("c", "t", {})])])).toEqual([
{ role: "assistant", content: null, tool_calls: [{ id: "c", type: "function", function: { name: "t", arguments: "{}" } }] },
]);
});
it("drops an assistant turn with neither text nor tool calls", () => {
expect(toChatCompletionMessages([message(USER, [text("hi")]), message(ASSISTANT, [text("")]), message(USER, [text("again")])])).toEqual([
{ role: "user", content: [{ type: "text", text: "hi" }] },
{ role: "user", content: [{ type: "text", text: "again" }] },
]);
});
it("hoists images out of tool results into a user message and serializes prompt-tsx values", () => {
expect(
toChatCompletionMessages([
message(USER, [toolResult("call_2", [text("screenshot:"), image([9], "image/jpeg"), { value: { node: 1 } }])]),
]),
).toEqual([
{ role: "tool", tool_call_id: "call_2", content: 'screenshot:{"node":1}' },
{ role: "user", content: [{ type: "image_url", image_url: { url: "data:image/jpeg;base64,CQ==" } }] },
]);
});
it("decodes text data parts and ignores unknown parts", () => {
expect(toChatCompletionMessages([message(USER, [{ mimeType: "text/plain", data: Uint8Array.from([104, 105]) }, 42])])).toEqual([
{ role: "user", content: [{ type: "text", text: "hi" }] },
]);
});
});
describe("estimateMessageTokens", () => {
it("counts what the gateway will receive, tool results and tool calls included", () => {
const plain = estimateMessageTokens(message(USER, [text("ok")]));
const withToolResult = estimateMessageTokens(message(USER, [toolResult("call_1", [text("y".repeat(800))]), text("ok")]));
const withToolCall = estimateMessageTokens(message(ASSISTANT, [toolCall("call_1", "read_file", { path: "z".repeat(800) })]));
expect(plain).toBeGreaterThan(0);
expect(withToolResult).toBeGreaterThanOrEqual(plain + 200);
expect(withToolCall).toBeGreaterThanOrEqual(200);
});
it("charges each image a flat estimate rather than its base64 length", () => {
const withoutImage = estimateMessageTokens(message(USER, [text("see")]));
const withImages = estimateMessageTokens(message(USER, [text("see"), image(new Array(30000).fill(0)), image([1])]));
expect(withImages - withoutImage).toBeGreaterThanOrEqual(2 * ESTIMATED_TOKENS_PER_IMAGE);
expect(withImages - withoutImage).toBeLessThan(2 * ESTIMATED_TOKENS_PER_IMAGE + 20);
});
it("counts nothing for a turn the gateway will never see", () => {
expect(estimateMessageTokens(message(ASSISTANT, []))).toBe(0);
});
});
describe("buildChatCompletionParams", () => {
it("streams with usage and forwards only the chosen extras", () => {
expect(buildChatCompletionParams(request())).toEqual({
model: "gpt-5.6",
messages: [{ role: "user", content: [{ type: "text", text: "hi" }] }],
stream: true,
stream_options: { include_usage: true },
});
});
it("declares tools as functions and requires a call only when VS Code does", () => {
const tools: readonly vscode.LanguageModelChatTool[] = [
{ name: "read_file", description: "Read a file", inputSchema: { type: "object", properties: { path: { type: "string" } } } },
{ name: "noop", description: "No input" },
];
const auto = buildChatCompletionParams(request({ tools }));
expect(auto.tools).toEqual([
{
type: "function",
function: { name: "read_file", description: "Read a file", parameters: { type: "object", properties: { path: { type: "string" } } } },
},
{ type: "function", function: { name: "noop", description: "No input" } },
]);
expect(auto.tool_choice).toBeUndefined();
expect(buildChatCompletionParams(request({ tools, requireToolCall: true })).tool_choice).toBe("required");
expect(buildChatCompletionParams(request({ requireToolCall: true })).tool_choice).toBeUndefined();
});
it("sends reasoning_effort only when the user picked one", () => {
expect(buildChatCompletionParams(request({ reasoningEffort: "xhigh" })).reasoning_effort).toBe("xhigh");
expect(buildChatCompletionParams(request()).reasoning_effort).toBeUndefined();
});
it("forwards numeric sampling options and drops everything else", () => {
const params = buildChatCompletionParams(
request({ modelOptions: { temperature: 0.2, max_tokens: 500, seed: "7", foo: "bar", top_p: 0.9 } }),
);
expect(params).toMatchObject({ temperature: 0.2, max_tokens: 500, top_p: 0.9 });
expect(params).not.toHaveProperty("seed");
expect(params).not.toHaveProperty("foo");
});
});

View file

@ -0,0 +1,185 @@
import { describe, expect, it } from "vitest";
import {
ASSUMED_MAX_INPUT_TOKENS,
ASSUMED_MAX_OUTPUT_TOKENS,
MARKDOWN_LINE_BREAK,
describeModels,
estimateTokens,
formatUsdPerMillionTokens,
parseModelGroups,
reasoningEffortFrom,
type ModelGroupInfo,
} from "../src/models";
const gatewayGroup = (overrides: Partial<Record<string, unknown>> = {}): Record<string, unknown> => ({
model_group: "gpt-5.6",
providers: ["openai"],
max_input_tokens: 922000,
max_output_tokens: 128000,
input_cost_per_token: 4e-6,
output_cost_per_token: 2e-5,
mode: "chat",
supports_vision: true,
supports_function_calling: true,
supports_reasoning: true,
supported_reasoning_efforts: ["none", "low", "medium", "high", "xhigh"],
...overrides,
});
const parsed = (...groups: readonly Record<string, unknown>[]): readonly ModelGroupInfo[] => {
const result = parseModelGroups({ data: groups });
if (result.kind !== "ok") {
throw new Error(result.reason);
}
return result.groups;
};
describe("parseModelGroups", () => {
it("maps the gateway's /model_group/info shape", () => {
expect(parsed(gatewayGroup())).toEqual([
{
modelGroup: "gpt-5.6",
providers: ["openai"],
mode: "chat",
maxInputTokens: 922000,
maxOutputTokens: 128000,
inputCostPerToken: 4e-6,
outputCostPerToken: 2e-5,
supportsVision: true,
supportsFunctionCalling: true,
supportedReasoningEfforts: ["none", "low", "medium", "high", "xhigh"],
},
]);
});
it("treats null limits, prices, and efforts as unknown", () => {
const [group] = parsed(
gatewayGroup({
max_input_tokens: null,
max_output_tokens: null,
input_cost_per_token: null,
output_cost_per_token: null,
supported_reasoning_efforts: null,
supports_vision: null,
}),
);
expect(group).toMatchObject({
maxInputTokens: undefined,
inputCostPerToken: undefined,
supportedReasoningEfforts: [],
supportsVision: false,
});
});
it("drops entries without a model_group and rejects payloads without data", () => {
expect(parsed({ providers: ["openai"] }, gatewayGroup()).map((group) => group.modelGroup)).toEqual(["gpt-5.6"]);
expect(parseModelGroups({ detail: "Unauthorized" })).toEqual({ kind: "invalid", reason: "response has no data array" });
});
});
describe("describeModels", () => {
it("lists chat groups with USD pricing in the detail and a full tooltip", () => {
const [model] = describeModels(parsed(gatewayGroup()));
expect(model).toMatchObject({
id: "gpt-5.6",
name: "gpt-5.6",
family: "gpt-5.6",
detail: "$4.00 in / $20.00 out per 1M tokens",
maxInputTokens: 922000,
maxOutputTokens: 128000,
imageInput: true,
toolCalling: true,
});
expect(model?.tooltip).toBe(
[
"LiteLLM model group gpt-5.6 via openai",
"Input: $4.00 per 1M tokens",
"Output: $20.00 per 1M tokens",
"Context: 922000 in / 128000 out tokens",
"Reasoning effort: none, low, medium, high, xhigh",
].join(MARKDOWN_LINE_BREAK),
);
});
it("offers the gateway's reasoning efforts behind a gateway default entry", () => {
const [model] = describeModels(parsed(gatewayGroup({ supported_reasoning_efforts: ["low", "high"] })));
expect(model?.configurationSchema).toEqual({
properties: {
reasoningEffort: {
type: "string",
title: "Reasoning Effort",
enum: ["default", "low", "high"],
enumItemLabels: ["Gateway default", "Low", "High"],
default: "default",
group: "navigation",
},
},
});
});
it("has no configuration schema when the group lists no reasoning efforts", () => {
const [model] = describeModels(parsed(gatewayGroup({ supported_reasoning_efforts: null })));
expect(model?.configurationSchema).toBeUndefined();
expect(model?.tooltip).toContain("Reasoning effort: not configurable");
});
it("keeps groups without a mode and skips non-chat groups", () => {
const models = describeModels(
parsed(
gatewayGroup({ model_group: "text-embedding-4", mode: "embedding" }),
gatewayGroup({ model_group: "whisper-3", mode: "audio_transcription" }),
gatewayGroup({ model_group: "gpt-image-2", mode: "image_generation" }),
gatewayGroup({ model_group: "unlabeled", mode: null }),
gatewayGroup(),
),
);
expect(models.map((model) => model.id)).toEqual(["unlabeled", "gpt-5.6"]);
});
it("falls back to assumed context limits and says so", () => {
const [model] = describeModels(parsed(gatewayGroup({ max_input_tokens: null, max_output_tokens: null })));
expect(model).toMatchObject({ maxInputTokens: ASSUMED_MAX_INPUT_TOKENS, maxOutputTokens: ASSUMED_MAX_OUTPUT_TOKENS });
expect(model?.tooltip).toContain(`Context: unknown, assuming ${ASSUMED_MAX_INPUT_TOKENS} in / ${ASSUMED_MAX_OUTPUT_TOKENS} out tokens`);
});
it("shows missing prices instead of inventing zeros", () => {
const [both, inputOnly, free] = describeModels(
parsed(
gatewayGroup({ input_cost_per_token: null, output_cost_per_token: null }),
gatewayGroup({ output_cost_per_token: null }),
gatewayGroup({ input_cost_per_token: 0, output_cost_per_token: 0 }),
),
);
expect(both?.detail).toBe("No pricing configured");
expect(both?.tooltip).toContain("Input: no price configured");
expect(inputOnly?.detail).toBe("$4.00 in / n/a out per 1M tokens");
expect(free?.detail).toBe("$0.00 in / $0.00 out per 1M tokens");
});
});
describe("formatUsdPerMillionTokens", () => {
it("renders cents for ordinary prices and two significant digits below a cent", () => {
expect(formatUsdPerMillionTokens(4e-6)).toBe("$4.00");
expect(formatUsdPerMillionTokens(7.5e-7)).toBe("$0.75");
expect(formatUsdPerMillionTokens(2.5e-5)).toBe("$25.00");
expect(formatUsdPerMillionTokens(1e-9)).toBe("$0.0010");
expect(formatUsdPerMillionTokens(0)).toBe("$0.00");
});
});
describe("reasoningEffortFrom", () => {
it("forwards a chosen effort and leaves the gateway default unset", () => {
expect(reasoningEffortFrom({ reasoningEffort: "high" })).toBe("high");
expect(reasoningEffortFrom({ reasoningEffort: "default" })).toBeUndefined();
expect(reasoningEffortFrom({ reasoningEffort: 3 })).toBeUndefined();
expect(reasoningEffortFrom(undefined)).toBeUndefined();
});
});
describe("estimateTokens", () => {
it("rounds four characters per token upward", () => {
expect(estimateTokens("")).toBe(0);
expect(estimateTokens("abcd")).toBe(1);
expect(estimateTokens("abcde")).toBe(2);
});
});

View file

@ -0,0 +1,258 @@
import type { ChatCompletionChunk, ChatCompletionCreateParamsStreaming } from "openai/resources/chat/completions";
import { describe, expect, it } from "vitest";
import type * as vscode from "vscode";
import type { GatewayClient, GatewayConfig, ModelGroupsResult } from "../src/gateway";
import { ESTIMATED_TOKENS_PER_IMAGE } from "../src/messages";
import { LiteLLMChatProvider, TRUNCATED_MESSAGE, type LiteLLMModel } from "../src/provider";
import { CancellationTokenSource, LanguageModelTextPart, LanguageModelToolCallPart } from "./vscode-mock";
interface StreamRequest {
readonly config: GatewayConfig;
readonly params: ChatCompletionCreateParamsStreaming;
}
interface FakeGateway extends GatewayClient {
readonly listCalls: readonly GatewayConfig[];
readonly streamRequests: readonly StreamRequest[];
}
const chunk = (delta: ChatCompletionChunk.Choice.Delta, finishReason: ChatCompletionChunk.Choice["finish_reason"] = null): ChatCompletionChunk => ({
id: "chatcmpl-1",
object: "chat.completion.chunk",
created: 0,
model: "gpt-5.6",
choices: [{ index: 0, delta, finish_reason: finishReason }],
});
const modelGroups: ModelGroupsResult = {
kind: "ok",
groups: [
{
modelGroup: "gpt-5.6",
providers: ["openai"],
mode: "chat",
maxInputTokens: 922000,
maxOutputTokens: 128000,
inputCostPerToken: 4e-6,
outputCostPerToken: 2e-5,
supportsVision: true,
supportsFunctionCalling: true,
supportedReasoningEfforts: ["low", "high"],
},
],
};
const fakeGateway = (
listResult: ModelGroupsResult = modelGroups,
stream: (signal: AbortSignal) => AsyncIterable<ChatCompletionChunk> = () => (async function* () {})(),
): FakeGateway => {
const listCalls: GatewayConfig[] = [];
const streamRequests: StreamRequest[] = [];
return {
listCalls,
streamRequests,
async listModelGroups(config) {
listCalls.push(config);
return listResult;
},
async streamChatCompletion(config, params, signal) {
streamRequests.push({ config, params });
return stream(signal);
},
};
};
const token = (): vscode.CancellationToken => new CancellationTokenSource().token as unknown as vscode.CancellationToken;
const prepare = (configuration: Record<string, unknown> | undefined): vscode.PrepareLanguageModelChatModelOptions =>
({ silent: true, configuration }) as vscode.PrepareLanguageModelChatModelOptions;
const gateway: GatewayConfig = { baseUrl: "http://127.0.0.1:4000", apiKey: "sk-test" };
const model = (overrides: Partial<LiteLLMModel> = {}): LiteLLMModel => ({
id: "gpt-5.6",
name: "gpt-5.6",
family: "gpt-5.6",
version: "1.0",
maxInputTokens: 922000,
maxOutputTokens: 128000,
capabilities: { imageInput: true, toolCalling: true },
gateway,
...overrides,
});
const userMessage = (parts: readonly unknown[]): vscode.LanguageModelChatRequestMessage =>
({ role: 1, content: parts, name: undefined }) as vscode.LanguageModelChatRequestMessage;
const responseOptions = (overrides: Partial<vscode.ProvideLanguageModelChatResponseOptions> = {}): vscode.ProvideLanguageModelChatResponseOptions =>
({ toolMode: 1, ...overrides }) as vscode.ProvideLanguageModelChatResponseOptions;
const collect = (
provider: LiteLLMChatProvider,
cancellation: CancellationTokenSource = new CancellationTokenSource(),
): { readonly parts: readonly vscode.LanguageModelResponsePart[]; readonly run: Promise<void> } => {
const parts: vscode.LanguageModelResponsePart[] = [];
const run = provider.provideLanguageModelChatResponse(
model(),
[userMessage([{ value: "hi" }])],
responseOptions(),
{ report: (part) => parts.push(part) },
cancellation.token as unknown as vscode.CancellationToken,
);
return { parts, run };
};
describe("provideLanguageModelChatInformation", () => {
it("returns nothing for the unconfigured probe without touching the gateway", async () => {
const client = fakeGateway();
expect(await new LiteLLMChatProvider(client).provideLanguageModelChatInformation(prepare(undefined), token())).toEqual([]);
expect(client.listCalls).toEqual([]);
});
it("names the API key when the stored secret is gone instead of listing nothing", async () => {
const client = fakeGateway();
await expect(
new LiteLLMChatProvider(client).provideLanguageModelChatInformation(prepare({ baseUrl: "http://127.0.0.1:4000" }), token()),
).rejects.toThrow(/missing its API key/);
expect(client.listCalls).toEqual([]);
});
it("rejects a gateway URL that is not http or https", async () => {
await expect(
new LiteLLMChatProvider(fakeGateway()).provideLanguageModelChatInformation(prepare({ baseUrl: "litellm.example.com", apiKey: "sk" }), token()),
).rejects.toThrow(/"litellm.example.com" is not an http or https URL/);
});
it("lists the gateway's chat models with pricing, effort choices, and the gateway attached", async () => {
const client = fakeGateway();
const models = await new LiteLLMChatProvider(client).provideLanguageModelChatInformation(
prepare({ baseUrl: "http://127.0.0.1:4000/v1/", apiKey: "sk-test" }),
token(),
);
expect(client.listCalls).toEqual([gateway]);
expect(models).toEqual([
expect.objectContaining({
id: "gpt-5.6",
detail: "$4.00 in / $20.00 out per 1M tokens",
maxInputTokens: 922000,
capabilities: { imageInput: true, toolCalling: true },
configurationSchema: expect.objectContaining({ properties: expect.objectContaining({ reasoningEffort: expect.anything() }) }),
gateway,
}),
]);
});
it("shows the gateway's error message, not its whole JSON body, when discovery fails", async () => {
const body = JSON.stringify({
error: { message: "Authentication Error, Invalid proxy server token passed", type: "auth_error", param: "sk-...abcd", code: "401" },
});
await expect(
new LiteLLMChatProvider(fakeGateway({ kind: "http_error", status: 401, body })).provideLanguageModelChatInformation(
prepare({ baseUrl: "http://127.0.0.1:4000", apiKey: "sk-bad" }),
token(),
),
).rejects.toThrow("LiteLLM gateway at http://127.0.0.1:4000 answered 401 for /model_group/info: Authentication Error, Invalid proxy server token passed");
});
});
describe("provideLanguageModelChatResponse", () => {
it("streams text and tool calls with the picked reasoning effort and a required tool choice", async () => {
const client = fakeGateway(modelGroups, () =>
(async function* () {
yield chunk({ content: "Reading" });
yield chunk({ tool_calls: [{ index: 0, id: "call_1", type: "function", function: { name: "read_file", arguments: '{"path":"a"}' } }] });
yield chunk({}, "tool_calls");
})(),
);
const parts: vscode.LanguageModelResponsePart[] = [];
await new LiteLLMChatProvider(client).provideLanguageModelChatResponse(
model(),
[userMessage([{ value: "read a" }])],
responseOptions({ toolMode: 2, tools: [{ name: "read_file", description: "Read" }], modelConfiguration: { reasoningEffort: "high" } }),
{ report: (part) => parts.push(part) },
token(),
);
expect(parts).toEqual([new LanguageModelTextPart("Reading"), new LanguageModelToolCallPart("call_1", "read_file", { path: "a" })]);
expect(client.streamRequests).toEqual([
{
config: gateway,
params: expect.objectContaining({ model: "gpt-5.6", reasoning_effort: "high", tool_choice: "required", tools: [expect.anything()] }),
},
]);
});
it("finishes quietly when the user cancels mid-stream and drops its cancellation listener", async () => {
const cancellation = new CancellationTokenSource();
const client = fakeGateway(modelGroups, (signal) =>
(async function* () {
yield chunk({ content: "partial" });
await new Promise<void>((resolve) => signal.addEventListener("abort", () => resolve(), { once: true }));
throw new Error("Request was aborted.");
})(),
);
const { parts, run } = collect(new LiteLLMChatProvider(client), cancellation);
await new Promise((resolve) => setTimeout(resolve, 0));
cancellation.cancel();
await expect(run).resolves.toBeUndefined();
expect(parts).toEqual([new LanguageModelTextPart("partial")]);
expect(cancellation.disposedListeners).toBe(1);
});
it("surfaces a gateway failure as an error and still drops its cancellation listener", async () => {
const cancellation = new CancellationTokenSource();
const client = fakeGateway(modelGroups, () =>
(async function* () {
throw new Error("502 Bad Gateway");
})(),
);
const { run } = collect(new LiteLLMChatProvider(client), cancellation);
await expect(run).rejects.toThrow("502 Bad Gateway");
expect(cancellation.disposedListeners).toBe(1);
});
it("reports the text it got and then fails when the model hits its output limit", async () => {
const client = fakeGateway(modelGroups, () =>
(async function* () {
yield chunk({ content: "half an ans" });
yield chunk({}, "length");
})(),
);
const { parts, run } = collect(new LiteLLMChatProvider(client));
await expect(run).rejects.toThrow(TRUNCATED_MESSAGE);
expect(parts).toEqual([new LanguageModelTextPart("half an ans")]);
});
it("fails on tool arguments that are not JSON", async () => {
const client = fakeGateway(modelGroups, () =>
(async function* () {
yield chunk({ tool_calls: [{ index: 0, id: "call_1", type: "function", function: { name: "grep", arguments: "{oops" } }] });
})(),
);
const { run } = collect(new LiteLLMChatProvider(client));
await expect(run).rejects.toThrow("invalid JSON arguments for tool grep");
});
});
describe("provideTokenCount", () => {
const provider = new LiteLLMChatProvider(fakeGateway());
it("estimates plain text at four characters per token", async () => {
expect(await provider.provideTokenCount(model(), "abcdefgh")).toBe(2);
});
it("counts tool results and tool calls, not only text parts", async () => {
const textOnly = await provider.provideTokenCount(model(), userMessage([{ value: "ok" }]));
const withToolResult = await provider.provideTokenCount(
model(),
userMessage([{ callId: "call_1", content: [{ value: "x".repeat(400) }] }, { value: "ok" }]),
);
expect(withToolResult).toBeGreaterThan(textOnly + 100);
});
it("charges a flat estimate per image instead of counting its bytes", async () => {
const withImage = await provider.provideTokenCount(model(), userMessage([{ value: "see" }, { mimeType: "image/png", data: new Uint8Array(50000) }]));
const withoutImage = await provider.provideTokenCount(model(), userMessage([{ value: "see" }]));
expect(withImage - withoutImage).toBeGreaterThanOrEqual(ESTIMATED_TOKENS_PER_IMAGE);
expect(withImage - withoutImage).toBeLessThan(ESTIMATED_TOKENS_PER_IMAGE + 20);
});
});

View file

@ -0,0 +1,93 @@
import { describe, expect, it } from "vitest";
import type { ChatCompletionChunk } from "openai/resources/chat/completions";
import { responseParts, type ResponsePart } from "../src/stream";
type ToolCallDelta = NonNullable<ChatCompletionChunk.Choice.Delta["tool_calls"]>[number];
const chunk = (delta: ChatCompletionChunk.Choice.Delta, finishReason: ChatCompletionChunk.Choice["finish_reason"] = null): ChatCompletionChunk => ({
id: "chatcmpl-1",
object: "chat.completion.chunk",
created: 0,
model: "gpt-5.6",
choices: [{ index: 0, delta, finish_reason: finishReason }],
});
const usageChunk: ChatCompletionChunk = {
id: "chatcmpl-1",
object: "chat.completion.chunk",
created: 0,
model: "gpt-5.6",
choices: [],
usage: { prompt_tokens: 3, completion_tokens: 2, total_tokens: 5 },
};
async function* stream(chunks: readonly ChatCompletionChunk[]): AsyncGenerator<ChatCompletionChunk> {
yield* chunks;
}
const collect = async (chunks: readonly ChatCompletionChunk[]): Promise<readonly ResponsePart[]> => {
const parts: ResponsePart[] = [];
for await (const part of responseParts(stream(chunks))) {
parts.push(part);
}
return parts;
};
describe("responseParts", () => {
it("yields text deltas as they arrive and ignores empty and usage-only chunks", async () => {
expect(await collect([chunk({ role: "assistant", content: "" }), chunk({ content: "Hel" }), chunk({ content: "lo" }), usageChunk])).toEqual([
{ kind: "text", value: "Hel" },
{ kind: "text", value: "lo" },
]);
});
it("assembles tool calls split across chunks and emits them after the text, in index order", async () => {
expect(
await collect([
chunk({ content: "Looking" }),
chunk({ tool_calls: [{ index: 1, id: "call_b", type: "function", function: { name: "grep", arguments: "" } }] }),
chunk({ tool_calls: [{ index: 0, id: "call_a", type: "function", function: { name: "read_file", arguments: '{"pa' } }] }),
chunk({ tool_calls: [{ index: 0, function: { name: "read_file", arguments: 'th":"a"}' } }] }),
chunk({ tool_calls: [{ index: 1, function: { arguments: '{"q":"x"}' } }] }, "tool_calls"),
]),
).toEqual([
{ kind: "text", value: "Looking" },
{ kind: "tool_call", callId: "call_a", name: "read_file", input: { path: "a" } },
{ kind: "tool_call", callId: "call_b", name: "grep", input: { q: "x" } },
]);
});
it("starts a new call when a fresh id reuses an index and appends index-less deltas to the last call", async () => {
expect(
await collect([
chunk({ tool_calls: [{ index: 0, id: "call_a", type: "function", function: { name: "grep", arguments: '{"q":' } }] }),
chunk({ tool_calls: [{ function: { arguments: '"a"}' } } as ToolCallDelta] }),
chunk({ tool_calls: [{ index: 0, id: "call_b", type: "function", function: { name: "grep", arguments: '{"q":"b"}' } }] }),
]),
).toEqual([
{ kind: "tool_call", callId: "call_a", name: "grep", input: { q: "a" } },
{ kind: "tool_call", callId: "call_b", name: "grep", input: { q: "b" } },
]);
});
it("flags a response cut off at the output token limit after the text it did produce", async () => {
expect(await collect([chunk({ content: "half" }), chunk({}, "length"), usageChunk])).toEqual([
{ kind: "text", value: "half" },
{ kind: "truncated" },
]);
});
it("treats empty arguments as an empty object and flags malformed JSON", async () => {
expect(
await collect([
chunk({ tool_calls: [{ index: 0, id: "call_0", type: "function", function: { name: "noop", arguments: "" } }] }),
chunk({ tool_calls: [{ index: 1, id: "call_1", type: "function", function: { name: "bad", arguments: "{oops" } }] }),
chunk({ tool_calls: [{ index: 2, id: "call_2", type: "function", function: { name: "scalar", arguments: "42" } }] }),
]),
).toEqual([
{ kind: "tool_call", callId: "call_0", name: "noop", input: {} },
{ kind: "invalid_tool_call", callId: "call_1", name: "bad", arguments: "{oops" },
{ kind: "invalid_tool_call", callId: "call_2", name: "scalar", arguments: "42" },
]);
});
});

View file

@ -0,0 +1,82 @@
export class LanguageModelTextPart {
constructor(readonly value: string) {}
}
export class LanguageModelToolCallPart {
constructor(
readonly callId: string,
readonly name: string,
readonly input: object,
) {}
}
export const LanguageModelChatToolMode = { Auto: 1, Required: 2 } as const;
type Listener<T> = (value: T) => void;
interface Subscription {
dispose(): void;
}
export class EventEmitter<T> {
private readonly listeners: Listener<T>[] = [];
readonly event = (listener: Listener<T>): Subscription => {
this.listeners.push(listener);
return {
dispose: () => {
const index = this.listeners.indexOf(listener);
if (index >= 0) {
this.listeners.splice(index, 1);
}
},
};
};
fire(value: T): void {
[...this.listeners].forEach((listener) => listener(value));
}
dispose(): void {
this.listeners.splice(0);
}
}
export interface MockCancellationToken {
readonly isCancellationRequested: boolean;
onCancellationRequested(listener: Listener<void>): Subscription;
}
export class CancellationTokenSource {
private readonly emitter = new EventEmitter<void>();
private cancelled = false;
private disposed = 0;
readonly token: MockCancellationToken;
constructor() {
const source = this;
this.token = {
get isCancellationRequested(): boolean {
return source.cancelled;
},
onCancellationRequested: (listener) => {
const subscription = source.emitter.event(listener);
return {
dispose: () => {
source.disposed += 1;
subscription.dispose();
},
};
},
};
}
get disposedListeners(): number {
return this.disposed;
}
cancel(): void {
this.cancelled = true;
this.emitter.fire();
}
}

View file

@ -0,0 +1,19 @@
{
"compilerOptions": {
"target": "ES2022",
"module": "ESNext",
"moduleResolution": "Bundler",
"lib": ["ES2022"],
"types": ["node"],
"strict": true,
"noUncheckedIndexedAccess": true,
"noImplicitOverride": true,
"isolatedModules": true,
"verbatimModuleSyntax": true,
"esModuleInterop": true,
"resolveJsonModule": true,
"skipLibCheck": true,
"noEmit": true
},
"include": ["src", "test"]
}

View file

@ -0,0 +1,13 @@
import { fileURLToPath } from "node:url";
import { defineConfig } from "vitest/config";
export default defineConfig({
resolve: {
alias: {
vscode: fileURLToPath(new URL("./test/vscode-mock.ts", import.meta.url)),
},
},
test: {
include: ["test/**/*.test.ts"],
},
});