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

# Conflicts:
#	tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py
This commit is contained in:
mateo-berri 2026-08-15 15:56:11 -07:00
commit bb1c3366cf
52 changed files with 2305 additions and 132 deletions

View file

@ -4,12 +4,12 @@
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
info lint lint-dev lint-checks format \
info lint lint-inner lint-dev lint-checks format \
lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
install-dev install-proxy-dev install-test-deps install-hooks \
install-helm-unittest check-circular-imports check-import-safety check pre-commit \
lint-install lint-fetch-base bootstrap
install-helm-unittest check-circular-imports check-import-safety check check-inner pre-commit \
lint-install lint-fetch-base bootstrap bootstrap-inner
# Default target
help:
@ -52,10 +52,17 @@ help:
@echo " make test-proxy-unit-b - Run proxy_unit_tests (p-z, ~28 files)"
@echo " make test-integration - Run integration tests"
@echo " make test-unit-helm - Run helm unit tests"
@echo ""
@echo "Heavy targets (check, bootstrap, lint) queue for LITELLM_GATE_SLOTS machine-wide"
@echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine."
UV := uv
UV_RUN := $(UV) run --no-sync
# Machine-wide slot queue for the heavy targets below; python3 + stdlib only, so
# it runs before any venv exists. See scripts/gate_slot_lock.py.
GATE_SLOT_LOCK := python3 scripts/gate_slot_lock.py
LINT_DEP_INSTALL ?= install-dev
LINT_E2E_DEP_INSTALL ?= lint-install
LINT_DEP_BASE ?= lint-fetch-base
@ -74,6 +81,9 @@ install-dev:
$(UV) sync --inexact --frozen
bootstrap:
@$(GATE_SLOT_LOCK) $(MAKE) bootstrap-inner
bootstrap-inner:
$(UV) sync --inexact --frozen --extra proxy --group proxy-dev --group e2e-dev
$(UV_RUN) python scripts/prisma_generate_if_needed.py
cd ui/litellm-dashboard && ../../scripts/with_dashboard_node.sh npm install --no-audit --no-fund
@ -229,7 +239,10 @@ check-import-safety: $(LINT_DEP_INSTALL)
# does (merge-base with origin/litellm_internal_staging). Setup (env sync, Prisma client,
# base fetch) runs once up front; the checks themselves are independent, so a sub-make
# fans them out with -j and the fast ones finish under basedpyright's shadow.
lint: lint-install lint-fetch-base
lint:
@$(GATE_SLOT_LOCK) $(MAKE) lint-inner
lint-inner: lint-install lint-fetch-base
$(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_E2E_DEP_INSTALL= LINT_DEP_BASE= lint-checks
lint-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright lint-e2e-basedpyright check-circular-imports check-import-safety
@ -244,7 +257,10 @@ lint-dev: lint-format-changed check-circular-imports check-import-safety
# test-linting.yml (Python), test-litellm-ui-build.yml's frontend-lint (dashboard), and
# check-ui-api-types.yml (API-type drift), skipping any whose files aren't in scope.
# Not auto-installed as a git hook so it never slows an unrelated human commit.
check: bootstrap
check:
@$(GATE_SLOT_LOCK) $(MAKE) check-inner
check-inner: bootstrap
./scripts/pre_commit_lint.sh
pre-commit:

View file

@ -3,6 +3,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t
"""
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Dict, Final, List, Optional, Tuple
from litellm._logging import verbose_proxy_logger
@ -564,6 +565,7 @@ class CheckBatchCost:
credentials = self.llm_router.get_deployment_credentials_with_provider(model_id) or {}
_file_content = await afile_content(
file_id=raw_output_file_id,
_litellm_internal_model_credentials=MappingProxyType(dict(credentials)),
**credentials,
)

View file

@ -5,6 +5,7 @@ from typing import Any, Final, Literal
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
from litellm.types.llms.openai import Batch
from litellm.types.utils import CallTypes, ModelInfo, Usage
@ -295,7 +296,7 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
if litellm_params:
# List of credential keys that should be passed to file operations
credential_keys: Final = [
credential_keys: Final = (
"api_key",
"api_base",
"api_version",
@ -309,7 +310,9 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
"bucket_name",
"timeout",
"max_retries",
]
"_litellm_internal_model_credentials",
*AWS_CREDENTIAL_KWARGS_KEYS,
)
for key in credential_keys:
if key in litellm_params:
credentials[key] = litellm_params[key]

View file

@ -22,6 +22,7 @@ from openai.types.batch import BatchRequestCounts
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_litellm_params import add_trusted_model_credentials_to_litellm_params
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.anthropic.batches.handler import AnthropicBatchesHandler
from litellm.llms.azure.batches.handler import AzureBatchesAPI
@ -527,6 +528,7 @@ def retrieve_batch(
custom_llm_provider=custom_llm_provider,
**kwargs,
)
add_trusted_model_credentials_to_litellm_params(litellm_params, kwargs)
if litellm_logging_obj is not None:
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,

View file

@ -11,7 +11,6 @@ import time
import uuid as uuid_module
from collections.abc import Coroutine
from functools import partial
from types import MappingProxyType
from typing import Any, Final, Literal, cast
import httpx
@ -34,6 +33,7 @@ import litellm
from litellm import get_secret_str
from litellm.files.streaming import FileContentStreamingResponse
from litellm.files.types import FileContentProvider, FileContentStreamingResult
from litellm.litellm_core_utils.get_litellm_params import add_trusted_model_credentials_to_litellm_params
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.azure.common_utils import get_azure_credentials
@ -85,14 +85,6 @@ bedrock_files_instance: Final = BedrockFilesHandler()
#################################################
def _add_trusted_model_credentials_to_litellm_params(
litellm_params_dict: dict[str, Any], kwargs: dict[str, Any]
) -> None:
trusted_model_credentials: Final = kwargs.get("_litellm_internal_model_credentials")
if isinstance(trusted_model_credentials, type(MappingProxyType({}))):
litellm_params_dict["_litellm_internal_model_credentials"] = trusted_model_credentials
@client
async def acreate_file(
file: FileTypes,
@ -372,7 +364,7 @@ def file_retrieve(
)
if provider_config is not None:
litellm_params_dict: Final = get_litellm_params(**kwargs)
_add_trusted_model_credentials_to_litellm_params(
add_trusted_model_credentials_to_litellm_params(
litellm_params_dict=litellm_params_dict,
kwargs=kwargs,
)
@ -494,7 +486,7 @@ def file_delete(
pass
optional_params: Final = GenericLiteLLMParams(**kwargs)
litellm_params_dict: Final = get_litellm_params(**kwargs)
_add_trusted_model_credentials_to_litellm_params(
add_trusted_model_credentials_to_litellm_params(
litellm_params_dict=litellm_params_dict,
kwargs=kwargs,
)
@ -834,7 +826,7 @@ def file_content(
try:
optional_params: Final = GenericLiteLLMParams(**kwargs)
litellm_params_dict: Final = get_litellm_params(**kwargs)
_add_trusted_model_credentials_to_litellm_params(
add_trusted_model_credentials_to_litellm_params(
litellm_params_dict=litellm_params_dict,
kwargs=kwargs,
)

View file

@ -1,3 +1,5 @@
from collections.abc import Mapping, MutableMapping
from types import MappingProxyType
from typing import Final
from litellm.llms.openai.data_residency import infer_openai_data_residency
@ -184,3 +186,19 @@ def get_litellm_params(
litellm_params[key] = kwargs[key]
return litellm_params
def add_trusted_model_credentials_to_litellm_params(
litellm_params_dict: MutableMapping[str, object], kwargs: Mapping[str, object]
) -> None:
"""
Carry the immutable server-side credential snapshot into litellm_params.
get_litellm_params has a fixed signature, so callers that need the snapshot to
survive into the logging object and the downstream file read have to re-add it. Only
a MappingProxyType is accepted, since providers resolve trusted configuration such
as a Bedrock file bucket from it and must not read a request-supplied mapping.
"""
trusted_model_credentials: Final = kwargs.get("_litellm_internal_model_credentials")
if isinstance(trusted_model_credentials, MappingProxyType):
litellm_params_dict["_litellm_internal_model_credentials"] = trusted_model_credentials

View file

@ -39,7 +39,11 @@ from ...openai.chat.gpt_transformation import (
OpenAIChatCompletionStreamingHandler,
OpenAIGPTConfig,
)
from ..common_utils import FireworksAIException, FireworksAIMixin
from ..common_utils import (
FireworksAIException,
FireworksAIMixin,
resolve_fireworks_resource_name,
)
def _extract_fireworks_hidden_params(payload: dict) -> dict:
@ -627,12 +631,10 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
litellm_params: dict,
headers: dict,
) -> dict:
if not model.startswith("accounts/") and "#" not in model:
if model.endswith("-fast"):
model = f"accounts/fireworks/routers/{model}"
else:
model = f"accounts/fireworks/models/{model}"
messages = self._transform_messages_helper(messages=messages, model=model, litellm_params=litellm_params)
resolved_model: Final = resolve_fireworks_resource_name(model)
messages = self._transform_messages_helper(
messages=messages, model=resolved_model, litellm_params=litellm_params
)
if "tools" in optional_params and optional_params["tools"] is not None:
tools: Final = self._transform_tools(tools=optional_params["tools"])
optional_params["tools"] = tools
@ -646,7 +648,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
"include_usage": True,
}
return super().transform_request(
model=model,
model=resolved_model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,

View file

@ -29,6 +29,17 @@ def get_fireworks_session_id(litellm_params: dict) -> str | None:
return None
def resolve_fireworks_resource_name(model: str) -> str:
stripped: Final = model.removeprefix("fireworks_ai/")
if stripped.startswith("accounts/") or "#" in stripped:
return stripped
if stripped.startswith(("routers/", "models/")):
return f"accounts/fireworks/{stripped}"
if stripped.endswith("-fast"):
return f"accounts/fireworks/routers/{stripped}"
return f"accounts/fireworks/models/{stripped}"
class FireworksAIMixin:
"""
Common Base Config functions across Fireworks AI Endpoints

View file

@ -13,7 +13,7 @@ from ..chat.transformation import (
FireworksAIConfig,
effort_from_chat_template_kwargs,
)
from ..common_utils import FireworksAIMixin
from ..common_utils import FireworksAIMixin, resolve_fireworks_resource_name
_TEXT_COMPLETION_STRIP_PARAMS: Final = (
frozenset({"truncate_prompt_tokens", "prompt_truncate_len"}) | NIM_VLLM_STRIP_PARAMS
@ -167,11 +167,8 @@ class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig
translated_params: Final = self.map_extra_body_params(optional_params=optional_params, model=model)
prompt: Final = _transform_prompt(messages=messages)
if not model.startswith("accounts/") and "#" not in model:
model = f"accounts/fireworks/models/{model}"
data: Final = {
"model": model,
"model": resolve_fireworks_resource_name(model),
"prompt": prompt,
**translated_params,
}

View file

@ -18,6 +18,7 @@ _PASS_THROUGH_PROTECTED_HEADERS: Final[frozenset] = frozenset(
"x-goog-api-key",
"host",
"content-length",
"accept-encoding",
}
)
@ -69,6 +70,9 @@ class BasePassthroughUtils:
# Header We Should NOT forward
request_headers.pop("content-length", None)
request_headers.pop("host", None)
# accept-encoding must stay client-negotiated: forwarding e.g. "br" when
# the brotli package is absent relays undecodable bytes to the caller
request_headers.pop("accept-encoding", None)
custom_header_names: Final = {header_name.lower() for header_name in headers}
for header_name in list(request_headers.keys()):

View file

@ -83,7 +83,7 @@ class UnloadableEntitlementError(Exception):
def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: list[str] | None = None) -> list[str] | None:
"""Resolve the single MCP server name a cold-start passthrough bypass may
target. Delegates parsing to
:meth:`MCPRequestHandler._extract_target_server_names_from_path` so the
:meth:`MCPRequestHandler.extract_target_server_names_from_path` so the
names used here always match the names downstream routing uses; returns
``None`` whenever the bypass must not activate (aggregate ``/mcp``,
multi-server CSV paths, or any other unrecognized path).
@ -94,7 +94,7 @@ def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: list[str] |
header/path mismatch here is a sign of a confused or hostile caller —
refuse the cold-start bypass rather than admit anonymously based on the
path while the header advertises a stricter, non-passthrough target."""
servers: Final = MCPRequestHandler._extract_target_server_names_from_path(path)
servers: Final = MCPRequestHandler.extract_target_server_names_from_path(path)
if len(servers) != 1:
verbose_logger.debug(
"MCP cold-start: path %r resolved to %r; passthrough 401 bypass "
@ -215,7 +215,7 @@ def _is_gateway_dcr_challenge_scope(
return False
if _has_client_supplied_mcp_auth(mcp_auth_header, mcp_server_auth_headers):
return False
if len(MCPRequestHandler._extract_target_server_names_from_path(route)) == 0:
if len(MCPRequestHandler.extract_target_server_names_from_path(route)) == 0:
return True
return _gateway_dcr_challenge_target(route, mcp_servers, client_ip) is not None
@ -579,7 +579,7 @@ class MCPRequestHandler:
return oauth2_headers, raw_headers, mcp_auth_header, mcp_server_auth_headers
@staticmethod
def _extract_target_server_names_from_path(path: str) -> list[str]:
def extract_target_server_names_from_path(path: str) -> list[str]:
"""
Extract the target MCP server name(s) from the standard MCP transport
URL patterns: ``/mcp/{server_name_or_csv}[/...]`` and
@ -836,6 +836,7 @@ class MCPRequestHandler:
case SessionBearerAdmitted():
try:
admitted: Final = await MCPRequestHandler._reload_admitted_user(result.principal.user_id)
admitted.mcp_session_resource_server_id = result.principal.resource_server_id
await MCPRequestHandler._enforce_admitted_live_policy(
admitted=admitted, request=request, route=route
)
@ -1168,7 +1169,7 @@ class MCPRequestHandler:
(header/path TOCTOU). For non-``/mcp/...`` paths (where the path
does not encode targets), fall back to the header.
"""
path_targets: Final = MCPRequestHandler._extract_target_server_names_from_path(path)
path_targets: Final = MCPRequestHandler.extract_target_server_names_from_path(path)
if path_targets:
return path_targets
# Path did not resolve to /mcp/... targets — trust the header

View file

@ -1655,6 +1655,7 @@ async def authorize(
code_challenge_method: str | None = None,
response_type: str | None = None,
scope: str | None = None,
resource: str | None = None,
):
# Redirect to real OAuth provider with PKCE support
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -1671,6 +1672,7 @@ async def authorize(
code_challenge_method=code_challenge_method,
response_type=response_type,
session_user_id=_session_cookie_user_id(request),
resource=resource,
)
lookup_name: Final[str | None] = mcp_server_name or client_id
@ -1721,6 +1723,7 @@ async def token_endpoint(
code_verifier: str = Form(None),
refresh_token: str | None = Form(None),
scope: str | None = Form(None),
resource: str | None = Form(None),
mcp_server_name: str | None = None,
):
"""
@ -1753,6 +1756,7 @@ async def token_endpoint(
master_key=master_key,
reload_user=_reload_active_user_by_id,
cache=user_api_key_cache,
resource=resource,
)
lookup_name: Final = mcp_server_name or client_id

View file

@ -56,6 +56,8 @@ from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
canonical_resource_uri,
canonicalize_url_identity,
get_request_base_url,
is_loopback_redirect_host,
validate_redirect_uri_shape,
@ -77,6 +79,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
GATEWAY_DCR_CLIENT_ID_PREFIX: Final = "llm_dcrc_"
"""Marker prefix on every gateway-issued DCR client_id so the root authorize/token
@ -169,6 +172,7 @@ class _ConnectFlow(BaseModel):
code_challenge: str = Field(min_length=1)
jti: str = Field(min_length=1)
exp: int
resource_server_id: str | None = None
class _GatewayAuthCode(BaseModel):
@ -185,6 +189,7 @@ class _GatewayAuthCode(BaseModel):
jti: str = Field(min_length=1)
iat: int
exp: int
resource_server_id: str | None = None
def is_gateway_dcr_client_id(client_id: str | None) -> bool:
@ -204,7 +209,13 @@ def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse
def _seal(prefix: str, payload: BaseModel) -> str:
return prefix + encrypt_value_helper(payload.model_dump_json())
"""Serialized ``exclude_none`` for the same reason session JWTs are minted that way: an
optional claim that is unset never reaches the wire, so during a rolling deploy a blob
sealed by a new pod without the new claim set stays byte-compatible with predating pods
whose strict models forbid unknown keys. This holds for every sealed artifact and every
future optional claim by construction; it requires each optional field to default to
``None`` so reopening restores exactly what was sealed."""
return prefix + encrypt_value_helper(payload.model_dump_json(exclude_none=True))
_SealedModelT = TypeVar("_SealedModelT", bound=BaseModel)
@ -320,6 +331,44 @@ def relative_request_url(request: Request) -> str:
return f"{path}?{request.url.query}" if request.url.query else path
def resolve_scoped_resource_server(request: Request, resource: str | None) -> MCPServer | None:
"""Resolve an RFC 8707 ``resource`` value to the single gateway-managed oauth2 server it
names, or ``None`` for every other shape: absent, the aggregate resource, a foreign
host, an unparseable value, a multi-server path, an unknown name, or any server mode the
keyless gateway flow does not serve (whose protected-resource metadata never directs a
client here). ``None`` means the flow stays unscoped and byte-identical to today, so a
hostile or confused ``resource`` can never widen anything; a resolved server only ever
NARROWS the session via the sealed scope.
Resolution is an IDENTITY question, deliberately free of the per-IP visibility filter:
access is enforced where it belongs (grant intersection at admission, IP checks on the
MCP routes), while filtering here would mint an entitlement-wide UNSCOPED bearer exactly
when the caller asked to narrow, and would let authorize-time vs token-time IP drift
turn a matching redemption into a spurious ``invalid_target``."""
if resource is None:
return None
canonical: Final = canonical_resource_uri(resource)
if canonical is None:
return None
base: Final = canonicalize_url_identity(get_request_base_url(request))
if canonical == f"{base}/mcp" or not canonical.startswith(f"{base}/"):
return None
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: PLC0415 # proxy import cycle
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # proxy import cycle
global_mcp_server_manager,
)
names: Final = MCPRequestHandler.extract_target_server_names_from_path(canonical[len(base) :])
if len(names) != 1:
return None
server: Final = global_mcp_server_manager.get_mcp_server_by_name(names[0])
if server is None or not server.is_gateway_managed_oauth2:
return None
return server
def aggregate_authorize(
request: Request,
client_id: str,
@ -329,11 +378,16 @@ def aggregate_authorize(
code_challenge_method: str | None,
response_type: str | None,
session_user_id: str | None,
resource: str | None = None,
) -> Response:
"""The aggregate authorize verb: validate the client, require S256 PKCE, interpose
LiteLLM sign-in, and hand the browser to the connect page with the flow sealed into a
per-flow cookie.
A per-server RFC 8707 ``resource`` naming a gateway-managed oauth2 server scopes the
flow to that one server: the scope is sealed into the flow, carried into the code, and
bound into the session token, while the connect page interlude runs exactly as before.
Validation failures respond directly with 400 and never redirect: per RFC 6749
section 4.1.2.1 an unvalidated redirect URI must not receive an error redirect, and
once the client is at fault there is no trusted place to send the browser.
@ -358,6 +412,7 @@ def aggregate_authorize(
login_url: Final = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}"
return RedirectResponse(login_url, status_code=303)
now: Final = datetime.now(timezone.utc)
scoped_server: Final = resolve_scoped_resource_server(request, resource)
handle: Final = secrets.token_urlsafe(24)
flow: Final = _ConnectFlow(
user_id=session_user_id,
@ -367,6 +422,7 @@ def aggregate_authorize(
code_challenge=code_challenge,
jti=secrets.token_urlsafe(24),
exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS,
resource_server_id=scoped_server.server_id if scoped_server is not None else None,
)
connect_url: Final = _append_query_params(
f"{base_url}/ui/connect",
@ -455,6 +511,7 @@ async def complete_connect_flow(
jti=secrets.token_urlsafe(24),
iat=int(now.timestamp()),
exp=int(now.timestamp()) + code_ttl,
resource_server_id=flow.resource_server_id,
),
)
params: Final = {"code": code, **({"state": flow.state} if flow.state else {})}
@ -587,6 +644,20 @@ def _reload_failure_response(failure: ReloadUserFailure) -> Response:
assert_never(failure)
def _resource_conflicts_with_scope(
request: Request, resource: str | None, sealed_resource_server_id: str | None
) -> bool:
"""True when a scoped grant is being redeemed for a DIFFERENT resource than the one
sealed into it (RFC 8707 section 2.2: reject with ``invalid_target``). An absent
``resource`` never conflicts (the sealed scope still binds the minted session), and an
unscoped grant ignores the parameter entirely, exactly as the endpoint always has, so
no pre-existing client breaks."""
if sealed_resource_server_id is None or resource is None:
return False
resolved: Final = resolve_scoped_resource_server(request, resource)
return resolved is None or resolved.server_id != sealed_resource_server_id
async def aggregate_token(
request: Request,
grant_type: str,
@ -598,6 +669,7 @@ async def aggregate_token(
master_key: str | None,
reload_user: ReloadUser,
cache: DualCache,
resource: str | None = None,
) -> Response:
"""The aggregate token verb: authorization_code and refresh_token grants for the
identity-only session pair. Every path re-validates the litellm user live before
@ -609,10 +681,12 @@ async def aggregate_token(
now: Final = datetime.now(timezone.utc)
if grant_type == "authorization_code":
return await _authorization_code_grant(
request=request,
code=code,
redirect_uri=redirect_uri,
client_id=client_id,
code_verifier=code_verifier,
resource=resource,
keys=keys,
now=now,
reload_user=reload_user,
@ -620,8 +694,10 @@ async def aggregate_token(
)
if grant_type == "refresh_token":
return await _refresh_token_grant(
request=request,
refresh_token=refresh_token,
client_id=client_id,
resource=resource,
keys=keys,
now=now,
reload_user=reload_user,
@ -631,10 +707,12 @@ async def aggregate_token(
async def _authorization_code_grant(
request: Request,
code: str | None,
redirect_uri: str | None,
client_id: str,
code_verifier: str | None,
resource: str | None,
keys: SessionKeys,
now: datetime,
reload_user: ReloadUser,
@ -651,6 +729,8 @@ async def _authorization_code_grant(
return _oauth_error(400, "invalid_grant", "the authorization code has expired")
if client_id != parsed.client_id or redirect_uri != parsed.redirect_uri:
return _oauth_error(400, "invalid_grant", "the authorization code was issued to a different client")
if _resource_conflicts_with_scope(request, resource, parsed.resource_server_id):
return _oauth_error(400, "invalid_target", "resource does not match the scope this code was issued for")
if not _pkce_verifier_matches(code_verifier, parsed.code_challenge):
return _oauth_error(400, "invalid_grant", "PKCE verification failed")
# Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable
@ -666,12 +746,18 @@ async def _authorization_code_grant(
parsed.exp - int(now.timestamp()) + _CLAIM_TTL_BUFFER_SECONDS,
):
return _oauth_error(400, "invalid_grant", "the authorization code was already used")
return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now)
return _session_token_pair(
SessionPrincipal(user_id=parsed.user_id, client_id=client_id, resource_server_id=parsed.resource_server_id),
keys,
now,
)
async def _refresh_token_grant(
request: Request,
refresh_token: str | None,
client_id: str,
resource: str | None,
keys: SessionKeys,
now: datetime,
reload_user: ReloadUser,
@ -682,6 +768,8 @@ async def _refresh_token_grant(
opened: Final = open_session_refresh_bearer(refresh_token, keys, now, expected_client_id=client_id)
if not isinstance(opened, SessionRefreshOpened):
return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client")
if _resource_conflicts_with_scope(request, resource, opened.principal.resource_server_id):
return _oauth_error(400, "invalid_target", "resource does not match the scope this token was issued for")
failure: Final = await reload_user(opened.principal.user_id)
if failure is not None:
return _reload_failure_response(failure)

View file

@ -2491,6 +2491,18 @@ class MCPServerManager:
open_ids.update(submitted_server_ids)
return open_ids
@staticmethod
def _admitted_session_resource_scope(user_api_key_auth: UserAPIKeyAuth | None) -> str | None:
"""The single server an admitted session subject's bearer was scoped to at authorize
time (RFC 8707 resource), or None for every other principal shape and for unscoped
sessions. Read at every return path of :meth:`get_allowed_mcp_servers`, including
the exception fallback, and applied AFTER every union (grants, operator-open,
submitted) because the scope is a ceiling over the whole reachable set; a resolver
fault therefore never widens a scoped bearer to the allow-all set."""
if user_api_key_auth is None or not _is_mcp_admitted_user_subject(user_api_key_auth):
return None
return user_api_key_auth.mcp_session_resource_server_id
async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth | None = None) -> list[str]:
"""
Get the allowed MCP Servers for the user.
@ -2600,13 +2612,19 @@ class MCPServerManager:
if len(combined_servers) == 0:
verbose_logger.debug("No allowed MCP Servers found for user api key auth.")
return list(combined_servers)
scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth)
return [server_id for server_id in combined_servers if scope is None or server_id == scope]
except Exception: # noqa: BLE001
verbose_logger.exception(
"Failed to get allowed MCP servers; team-level object_permission "
"grants may be dropped. Falling back to global and submitted servers."
)
return list(dict.fromkeys(allow_all_server_ids + submitted_server_ids))
scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth)
return [
server_id
for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids)
if scope is None or server_id == scope
]
async def resolve_toolset_tool_permissions(
self,

View file

@ -633,7 +633,7 @@ def canonicalize_url_identity(url: str) -> str:
return urlunparse((scheme, netloc, parsed.path.rstrip("/"), "", "", ""))
def _canonical_resource_uri(url: str) -> str | None:
def canonical_resource_uri(url: str) -> str | None:
"""Canonicalize an upstream MCP server URL into an RFC 8707 resource identifier.
Keeps only the scheme, host, port and path, which is the shape the MCP authorization spec's
@ -693,7 +693,7 @@ def resolve_upstream_resource(mcp_server: "MCPServer") -> str | None:
mcp_server.server_id,
)
return None
canonical: Final = _canonical_resource_uri(mcp_server.url)
canonical: Final = canonical_resource_uri(mcp_server.url)
if canonical is None:
verbose_logger.warning(
"MCP server %s sets upstream_resource=auto but its url is not an absolute URI, so no "

View file

@ -85,11 +85,18 @@ class SessionPrincipal(BaseModel):
enforced at use time rather than frozen at mint time. ``client_id`` is the (stateless,
gateway-sealed) DCR client identifier the token was issued to; the token endpoint
requires it to match on the refresh grant.
``resource_server_id`` is the single MCP server this session was authorized for when
the client requested a per-server RFC 8707 resource at authorize time, or ``None`` for
the aggregate scope. It is a RESTRICTION carried for admission to intersect against
the live grant resolution, never a grant by itself; the refresh grant re-mints from
this principal so the restriction survives rotation.
"""
model_config = ConfigDict(frozen=True)
user_id: str = Field(min_length=1)
client_id: str = Field(min_length=1)
resource_server_id: str | None = None
class SessionKeys(BaseModel):
@ -186,6 +193,7 @@ class _SessionClaims(BaseModel):
kind: SessionTokenKind
user_id: str = Field(min_length=1)
client_id: str = Field(min_length=1)
resource_server_id: str | None = None
def is_session_token(candidate: str) -> bool:
@ -286,9 +294,10 @@ def _mint(
kind=kind,
user_id=principal.user_id,
client_id=principal.client_id,
resource_server_id=principal.resource_server_id,
)
token: Final = prefix + jwt.encode(
claims.model_dump(), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM
claims.model_dump(exclude_none=True), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM
)
size_bytes: Final = len(token.encode("utf-8"))
if size_bytes > MAX_SESSION_TOKEN_BYTES:
@ -323,7 +332,10 @@ def _open(
if now.timestamp() >= claims.exp:
return SessionExpired()
return OpenedSessionToken(
principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id), jti=claims.jti
principal=SessionPrincipal(
user_id=claims.user_id, client_id=claims.client_id, resource_server_id=claims.resource_server_id
),
jti=claims.jti,
)

View file

@ -2762,6 +2762,13 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# key off. Server-only and stripped from validated input for the same reason as the marker
# above: a forged entry would let a caller pick which team's rpm bucket it is charged against.
mcp_source_team_rpm_limits: dict[str, dict[str, int]] | None = Field(default=None, exclude=True)
# The single MCP server_id a gateway session bearer was scoped to at authorize time (RFC 8707
# resource), or None for an aggregate-scope session. A RESTRICTION intersected against the live
# grant resolution, never a grant. Server-only, set exclusively by the MCP gateway admission
# path via post-construction assignment and stripped from validated input like the markers
# above; a forged value could at most narrow, but the stripping keeps the field's provenance
# single-owner so its meaning stays trustworthy.
mcp_session_resource_server_id: str | None = Field(default=None, exclude=True)
via_virtual_key: bool = Field(
default=False,
exclude=True,
@ -2798,6 +2805,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data.
values.pop("mcp_admitted_user_subject", None)
values.pop("mcp_source_team_rpm_limits", None)
values.pop("mcp_session_resource_server_id", None)
values.pop("via_virtual_key", None)
if values.get("api_key") is not None:
values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))})

View file

@ -23,6 +23,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
add_internal_model_credentials,
apply_team_provider_credentials,
batch_cost_poller_is_active,
decode_model_from_file_id,
@ -546,6 +547,13 @@ async def retrieve_batch(
detail={"error": "LLM Router not initialized. Ensure models added to proxy."},
)
if unified_batch_id:
add_internal_model_credentials(
data=data,
llm_router=llm_router,
model_id=get_model_id_from_unified_batch_id(unified_batch_id),
)
response = await llm_router.aretrieve_batch(**data)
response._hidden_params["unified_batch_id"] = unified_batch_id
if unified_batch_id:

View file

@ -71,6 +71,14 @@ class PanwPrismaAirsHandler(CustomGuardrail):
_PROVIDER_NAME = "panw_prisma_airs"
#: AIRS fields withheld from the client-visible error detail.
#: ``response_masked_data`` is the model's own generation. The block branch that builds
#: this detail is only reached when ``mask_response_content`` is False, so echoing it
#: back would hand the caller exactly the text the operator declined to deliver.
#: ``prompt_masked_data`` is deliberately NOT withheld: it is the caller's own input,
#: and it is one of the fields the ticket asks for.
_CLIENT_HIDDEN_SCAN_FIELDS: Final = frozenset({"response_masked_data"})
def __init__(
self,
guardrail_name: str,
@ -632,12 +640,21 @@ class PanwPrismaAirsHandler(CustomGuardrail):
choice.message.function_call.arguments = masked_text
def _build_error_detail(
self, scan_result: Mapping[str, object], is_response: bool = False
self,
scan_result: Mapping[str, object],
is_response: bool = False,
also_hide: str | None = None,
) -> Mapping[str, Mapping[str, object]]:
"""Build enhanced error detail with scan information."""
"""Build enhanced error detail with scan information.
``also_hide`` names one more scan field to withhold, for the caller that knows
its AIRS verdict carries model-generated content under a key that is normally
caller input.
"""
action_type: Final = "Response" if is_response else "Prompt"
code_suffix: Final = "_response_blocked" if is_response else "_blocked"
detection_key: Final = "response_detected" if is_response else "prompt_detected"
hidden_fields: Final = self._CLIENT_HIDDEN_SCAN_FIELDS.union(() if also_hide is None else (also_hide,))
category: Final = scan_result.get("category", "unknown")
default_msg: Final = f"{action_type} blocked by PANW Prisma AI Security policy (Category: {category})"
@ -653,8 +670,13 @@ class PanwPrismaAirsHandler(CustomGuardrail):
},
)
error_detail: Final[dict[str, dict[str, object]]] = {
return {
"error": {
**{
key: value
for key, value in scan_result.items()
if not key.startswith("_") and key not in hidden_fields
},
"message": error_msg,
"type": "guardrail_violation",
"code": f"panw_prisma_airs{code_suffix}",
@ -663,24 +685,6 @@ class PanwPrismaAirsHandler(CustomGuardrail):
}
}
# Add optional fields if present
optional_fields: Final = [
"scan_id",
"report_id",
"profile_name",
"profile_id",
"tr_id",
]
for field in optional_fields:
if scan_result.get(field):
error_detail["error"][field] = scan_result[field]
# Add detection details
if scan_result.get(detection_key):
error_detail["error"][detection_key] = scan_result[detection_key]
return error_detail
def _record_scan_id(self, request_data: dict[str, Any], scan_result: Mapping[str, object]) -> None:
"""Surface the AIRS scan id on the response, so allowed calls are auditable too."""
scan_id: Final = scan_result.get("scan_id")
@ -1481,7 +1485,17 @@ class PanwPrismaAirsHandler(CustomGuardrail):
):
self._set_tool_call_arguments(tool_call, masked_text)
else:
error_detail = self._build_error_detail(scan_result, is_response=is_response)
# tool_event scans are request-side in the AIRS schema, so AIRS returns
# the model's own tool arguments under prompt_masked_data. On a
# response-side block that is generated content, not caller input, and
# the class-level default only withholds response_masked_data — which is
# empty on this path. Withhold it explicitly so the 400 does not become
# the content channel this branch declined to deliver.
error_detail = self._build_error_detail(
scan_result,
is_response=is_response,
also_hide="prompt_masked_data" if is_response else None,
)
raise HTTPException(status_code=400, detail=error_detail)
@staticmethod

View file

@ -465,6 +465,31 @@ def apply_team_provider_credentials(
prepare_data_with_credentials(data=data, credentials=credentials)
def add_internal_model_credentials(
data: dict,
llm_router: "Router",
model_id: str | None,
) -> None:
"""
Attach the deployment's immutable server-side credential snapshot to a router-routed
batch call (in-place).
Cost accounting for a completed batch reads the batch's output file, and the Bedrock
file config resolves its bucket only from this snapshot, never from a request param,
because the bucket is what managed file ids are validated against. Without it that
read fails and the batch's cost is never recorded.
"""
if model_id is None:
return
try:
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
except Exception: # noqa: BLE001 # the snapshot only enables cost accounting; a batch whose deployment no longer resolves must still be retrievable
return
if credentials is None:
return
data["_litellm_internal_model_credentials"] = MappingProxyType(dict(credentials))
def prepare_data_with_credentials(
data: dict,
credentials: dict,

View file

@ -43,6 +43,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
add_internal_model_credentials,
apply_team_provider_credentials,
encode_file_id_with_model,
extract_file_creation_params,
@ -706,6 +707,7 @@ async def get_file_content(
model: Final = cast(str | None, data.get("model"))
if model:
add_internal_model_credentials(data=data, llm_router=llm_router, model_id=model)
response = await llm_router.afile_content(
**{
"model": model,

View file

@ -266,6 +266,8 @@ class CredentialLiteLLMParams(BaseModel):
s3_region_name: str | None = None
s3_encryption_key_id: str | None = None
aws_batch_role_arn: str | None = None
s3_output_bucket_name: str | None = None
bedrock_tags: list | None = None
## IBM WATSONX ##
watsonx_region_name: str | None = None

View file

@ -3398,9 +3398,21 @@ agentic_loop_internal_litellm_params: Final = [
# the provider.
TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars"
# Bedrock managed-batch deployment config, read from litellm_params by the batch and
# files transformations. Listed for the same reason as the fields above: these sit on
# a deployment that also serves chat, so leaking them into extra_body makes Bedrock
# reject every non-batch request to that deployment.
bedrock_batch_litellm_params: Final = (
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",
"s3_output_bucket_name",
"bedrock_tags",
)
all_litellm_params = (
agentic_loop_internal_litellm_params
+ [TRUSTED_CALLBACK_VARS_FIELD]
+ [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params]
+ [
"metadata",
"litellm_metadata",

173
scripts/gate_slot_lock.py Normal file
View file

@ -0,0 +1,173 @@
#!/usr/bin/env python3
"""Machine-wide slot lock for this repo's heavy entrypoints.
`make check`, `make bootstrap`, `make lint`, and the standalone budget gates
(scripts/ruff_strict_gate.py, scripts/type_discipline_gate.py,
scripts/type_check_gate.py) each hold one of N machine-wide slots while they
run, so however many sessions and worktrees share one machine, at most N of
them execute a basedpyright/pytest/prettier storm at a time instead of all
thrashing it at once. Slots are fcntl.flock files (macOS ships no flock(1)
binary, hence python3 + stdlib only, runnable before any venv exists) under a
per-user cache directory shared by every worktree and session:
~/.cache/litellm/gate-slots by default, $LITELLM_GATE_SLOT_DIR to override.
A holder's lock dies with its process, so a crash leaves nothing to clean up.
$LITELLM_GATE_SLOTS sets the slot count (default 2); 0 disables locking.
Waiting is a blocking flock on a turnstile file plus a slow poll of the slots,
so contenders queue roughly first-come-first-served without busy-spinning.
A process that acquired (or deliberately skipped) a slot exports
LITELLM_GATE_SLOT_HELD, and nested acquisitions under that marker are no-ops,
so `make check` invoking the gates internally can never deadlock against
itself. Any filesystem error fails open and the command runs unlocked: the
lock is a courtesy to the machine, never a gate that may break a build (CI
runs one job per machine, so there it only ever takes the instant path).
CLI: python3 scripts/gate_slot_lock.py <command> [args...]
"""
from __future__ import annotations
import contextlib
import fcntl
import os
import subprocess
import sys
import time
from pathlib import Path
from typing import IO, TYPE_CHECKING, Final
if TYPE_CHECKING:
from collections.abc import Iterator
HELD_MARKER_ENV: Final = "LITELLM_GATE_SLOT_HELD"
SLOT_COUNT_ENV: Final = "LITELLM_GATE_SLOTS"
SLOT_DIR_ENV: Final = "LITELLM_GATE_SLOT_DIR"
DEFAULT_SLOT_COUNT: Final = 2
POLL_SECONDS: Final = 2.0
def _slot_dir() -> Path:
override: Final = os.environ.get(SLOT_DIR_ENV)
return Path(override) if override else Path.home() / ".cache" / "litellm" / "gate-slots"
def _slot_count() -> int:
raw: Final = os.environ.get(SLOT_COUNT_ENV)
if not raw:
return DEFAULT_SLOT_COUNT
try:
return int(raw)
except ValueError:
print(
f"gate_slot_lock: ignoring non-integer {SLOT_COUNT_ENV}={raw!r}; "
f"using {DEFAULT_SLOT_COUNT} slots",
file=sys.stderr,
)
return DEFAULT_SLOT_COUNT
def _try_slot(directory: Path, index: int) -> IO[bytes] | None:
handle: Final = (directory / f"slot-{index}.lock").open("wb")
try:
fcntl.flock(handle, fcntl.LOCK_EX | fcntl.LOCK_NB)
except BlockingIOError:
handle.close()
return None
except OSError:
handle.close()
raise
return handle
def _wait_for_slot(directory: Path, count: int) -> IO[bytes]:
print(
f"gate_slot_lock: all {count} machine-wide slots are busy; queueing "
f"(set {SLOT_COUNT_ENV}=0 to disable)",
file=sys.stderr,
flush=True,
)
with (directory / "turnstile.lock").open("wb") as turnstile:
fcntl.flock(turnstile, fcntl.LOCK_EX)
while True:
for index in range(count):
held = _try_slot(directory, index)
if held is not None:
return held
time.sleep(POLL_SECONDS)
def _locked_handle(count: int) -> IO[bytes]:
directory: Final = _slot_dir()
directory.mkdir(parents=True, exist_ok=True)
for index in range(count):
immediate = _try_slot(directory, index)
if immediate is not None:
return immediate
return _wait_for_slot(directory, count)
def acquire_slot() -> IO[bytes] | None:
"""Hold a machine-wide slot for the life of the returned handle.
The caller must keep the handle referenced until the process exits;
dropping it closes the file and releases the slot. Returns None without
locking when this process already runs under a held slot, when locking is
disabled, or when the filesystem refuses to cooperate."""
if os.environ.get(HELD_MARKER_ENV):
return None
count: Final = _slot_count()
if count <= 0:
os.environ[HELD_MARKER_ENV] = "1"
return None
try:
handle: Final = _locked_handle(count)
except (OSError, RuntimeError) as error:
print(f"gate_slot_lock: locking unavailable ({error}); running unlocked", file=sys.stderr)
os.environ[HELD_MARKER_ENV] = "1"
return None
os.environ[HELD_MARKER_ENV] = "1"
return handle
@contextlib.contextmanager
def held_slot() -> Iterator[None]:
"""Run the with-block while holding a machine-wide slot (or its no-op forms)."""
prior_marker: Final = os.environ.get(HELD_MARKER_ENV)
handle: Final = acquire_slot()
try:
yield
finally:
if handle is not None:
handle.close()
if not prior_marker:
os.environ.pop(HELD_MARKER_ENV, None)
def _wait_ignoring_interrupts(process: subprocess.Popen[bytes]) -> int:
while True:
try:
return process.wait()
except KeyboardInterrupt:
continue
def main() -> int:
if len(sys.argv) < 2:
print("usage: gate_slot_lock.py <command> [args...]", file=sys.stderr)
return 2
try:
held: Final = acquire_slot()
except KeyboardInterrupt:
return 130
try:
code: Final = _wait_ignoring_interrupts(subprocess.Popen(sys.argv[1:]))
except FileNotFoundError as error:
print(f"gate_slot_lock: {error}", file=sys.stderr)
return 127
if held is not None:
held.close()
return code if code >= 0 else 128 - code
if __name__ == "__main__":
sys.exit(main())

View file

@ -24,6 +24,16 @@
set -eu
# Queue for one of the machine-wide heavy-work slots (see scripts/gate_slot_lock.py)
# before anything else, so N parallel `make check` runs across worktrees execute two
# at a time instead of thrashing the machine. The wrapper exports
# LITELLM_GATE_SLOT_HELD, so this re-exec happens exactly once and everything this
# script spawns (make lint, the budget gates) skips its own acquisition.
if [ -z "${LITELLM_GATE_SLOT_HELD:-}" ]; then
script_dir=$(python3 -c 'import os, sys; print(os.path.dirname(os.path.realpath(sys.argv[1])))' "$0")
exec python3 "$script_dir/gate_slot_lock.py" "$0" "$@"
fi
if [ -z "${PRE_COMMIT_LINT_INNER:-}" ]; then
log_file=$(git rev-parse --path-format=absolute --git-path pre_commit_lint.log)
if : > "$log_file" 2>/dev/null; then

View file

@ -215,7 +215,10 @@ def main() -> None:
parser.add_argument("--base", default=DEFAULT_BASE)
parser.add_argument("--update", action="store_true")
args = parser.parse_args()
cmd_update(args.base) if args.update else cmd_check(args.base)
from gate_slot_lock import held_slot
with held_slot():
cmd_update(args.base) if args.update else cmd_check(args.base)
if __name__ == "__main__":

View file

@ -670,16 +670,19 @@ def main() -> None:
parser.add_argument("--update", action="store_true")
parser.add_argument("--emit-counts-dir", type=Path)
args = parser.parse_args()
ensure_typecheck_env()
head = count_basedpyright(run_basedpyright())
if args.emit_counts_dir is not None:
cmd_emit_counts(
head, args.emit_counts_dir, _run(["git", "rev-parse", "HEAD"]).strip()
)
elif args.update:
cmd_update(head, args.base)
else:
cmd_check(head, args.base)
from gate_slot_lock import held_slot
with held_slot():
ensure_typecheck_env()
head = count_basedpyright(run_basedpyright())
if args.emit_counts_dir is not None:
cmd_emit_counts(
head, args.emit_counts_dir, _run(["git", "rev-parse", "HEAD"]).strip()
)
elif args.update:
cmd_update(head, args.base)
else:
cmd_check(head, args.base)
if __name__ == "__main__":

View file

@ -267,7 +267,10 @@ def main() -> None:
parser.add_argument("--base", default=DEFAULT_BASE)
parser.add_argument("--update", action="store_true")
args = parser.parse_args()
cmd_update(args.base) if args.update else cmd_check(args.base)
from gate_slot_lock import held_slot
with held_slot():
cmd_update(args.base) if args.update else cmd_check(args.base)
if __name__ == "__main__":

View file

@ -353,6 +353,102 @@ class TestCheckBatchCost:
), "update() must NOT include batch_processed when column is absent"
assert update_data["status"] == "complete"
@pytest.mark.asyncio
async def test_output_fetch_passes_deployment_credentials_as_trusted_snapshot(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""Bedrock resolves the output bucket ONLY from the immutable snapshot kwarg.
Spreading the credentials as plain kwargs is not enough: get_litellm_params drops
s3_bucket_name, so without _litellm_internal_model_credentials the cost poller
cannot read the output file and every completed Bedrock batch stays unbilled.
"""
from types import MappingProxyType
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_job = MagicMock()
mock_job.id = "job-bedrock-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "file-output-123"
mock_response.model_dump_json.return_value = '{"id":"batch-1","status":"completed"}'
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={
"custom_llm_provider": "bedrock",
"s3_bucket_name": "configured-batch-bucket",
"aws_region_name": "us-east-1",
}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "bedrock"
mock_deployment.litellm_params.model = "bedrock/anthropic.claude-haiku-4-5-20251001-v1:0"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
mock_file_content = MagicMock()
mock_file_content.content = b'{"recordId":"req-1"}'
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
) as mock_afile_content,
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"recordId": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(0.01, {"prompt_tokens": 10, "completion_tokens": 5}, ["claude-haiku-4-5"]),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("anthropic.claude-haiku-4-5-20251001-v1:0", "bedrock", None, None),
),
patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await check_batch_cost_instance.check_batch_cost()
mock_afile_content.assert_awaited()
passed_kwargs = mock_afile_content.await_args[1]
snapshot = passed_kwargs.get("_litellm_internal_model_credentials")
assert snapshot is not None, "cost poller must pass the trusted credential snapshot"
assert isinstance(
snapshot, MappingProxyType
), "snapshot must be a MappingProxyType; a plain dict is rejected by get_configured_s3_bucket_name"
assert snapshot["s3_bucket_name"] == "configured-batch-bucket"
@pytest.mark.asyncio
async def test_primary_path_completion_update_includes_batch_processed(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router

View file

@ -17,6 +17,7 @@ deterministic stand-ins so the arithmetic under test is the only variable.
import json
import os
import sys
from types import MappingProxyType
import httpx
import pytest
@ -1229,3 +1230,72 @@ async def test_calculate_batch_cost_and_usage_anthropic_end_to_end():
assert cost == pytest.approx(1000 * 3e-6 / 2 + 8000 * 3e-7 / 2 + 2000 * 3.75e-6 / 2 + 200 * 15e-6 / 2)
assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (11000, 200, 11200)
assert models == ["claude-sonnet-4-5"]
def test_extract_credentials_forwards_the_trusted_model_credential_snapshot():
"""Bedrock resolves a batch's output bucket only from the immutable server-side
snapshot, never from a request param, so cost accounting on the retrieve path cannot
read the output file unless this key is forwarded. Without it the accounting raises
"S3 bucket_name is required" for a bucket the deployment has configured, and the
batch's cost is never recorded."""
snapshot = MappingProxyType({"s3_bucket_name": "configured-bucket", "aws_region_name": "us-east-1"})
credentials = bu._extract_file_access_credentials({"_litellm_internal_model_credentials": snapshot})
assert credentials["_litellm_internal_model_credentials"] is snapshot
def test_extract_credentials_forwards_the_deployment_aws_credentials():
"""The retrieve path's logging object carries the deployment's AWS keys in its
litellm_params, and the S3 read of the output file signs with whatever afile_content
receives. Dropping them here sent the read to the ambient credential chain, so a
deployment whose only AWS credentials live in its litellm_params never recorded
batch cost on retrieve even once the bucket resolved."""
params = {
"aws_access_key_id": "AKIA-deployment",
"aws_secret_access_key": "secret-deployment",
"aws_session_token": "token-deployment",
"aws_region_name": "us-west-2",
"aws_role_name": "arn:aws:iam::123456789012:role/batch-reader",
"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
}
credentials = bu._extract_file_access_credentials(params)
assert credentials == {key: value for key, value in params.items() if key != "model"}
@pytest.mark.asyncio
async def test_output_file_content_bedrock_reads_with_deployment_aws_credentials(monkeypatch):
import litellm.files.main as files_main
captured: dict = {}
async def fake_afile_content(**kw):
captured.update(kw)
return type("R", (), {"content": b""})()
monkeypatch.setattr(files_main, "afile_content", fake_afile_content)
snapshot = MappingProxyType({"s3_bucket_name": "configured-bucket", "aws_region_name": "us-west-2"})
await bu._fetch_batch_output_file_content(
_batch("s3://configured-bucket/litellm-batch-outputs/job-1/out.jsonl.out"),
custom_llm_provider="bedrock",
litellm_params={
"aws_access_key_id": "AKIA-deployment",
"aws_secret_access_key": "secret-deployment",
"aws_session_token": "token-deployment",
"aws_region_name": "us-west-2",
"_litellm_internal_model_credentials": snapshot,
"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
},
)
assert captured["file_id"] == "s3://configured-bucket/litellm-batch-outputs/job-1/out.jsonl.out"
assert captured["custom_llm_provider"] == "bedrock"
assert captured["aws_access_key_id"] == "AKIA-deployment"
assert captured["aws_secret_access_key"] == "secret-deployment"
assert captured["aws_session_token"] == "token-deployment"
assert captured["aws_region_name"] == "us-west-2"
assert captured["_litellm_internal_model_credentials"] is snapshot
assert "model" not in captured

View file

@ -28,6 +28,7 @@ import sys
from contextlib import ExitStack
from dataclasses import dataclass
from typing import Any, Dict
from types import MappingProxyType
from unittest.mock import MagicMock, patch
import pytest
@ -742,3 +743,33 @@ def test_resolve_timeout__httpx_timeout_returns_float_read():
resolved = bm._resolve_timeout(_params(timeout=t), {}, "openai")
assert isinstance(resolved, float)
assert resolved == 99.0
def test_retrieve__forwards_trusted_model_credentials_into_litellm_params(seams):
"""The batch's cost is computed by reading its output file after the retrieve, and
Bedrock resolves that bucket only from this immutable snapshot. get_litellm_params has
a fixed signature that drops it, so without re-adding it here the snapshot never
reaches the logging object and cost accounting fails on a bucket that is configured."""
snapshot = MappingProxyType({"s3_bucket_name": "configured-bucket"})
logging_obj = MagicMock()
bm.retrieve_batch(
batch_id="batch-1",
custom_llm_provider="openai",
litellm_logging_obj=logging_obj,
_litellm_internal_model_credentials=snapshot,
)
litellm_params = logging_obj.update_from_kwargs.call_args.kwargs["litellm_params"]
assert litellm_params["_litellm_internal_model_credentials"] is snapshot
def test_retrieve__omits_trusted_model_credentials_when_not_supplied(seams):
"""A retrieve with no snapshot must not invent an empty one, which would read as a
configured bucket of nothing."""
logging_obj = MagicMock()
bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="openai", litellm_logging_obj=logging_obj)
litellm_params = logging_obj.update_from_kwargs.call_args.kwargs["litellm_params"]
assert "_litellm_internal_model_credentials" not in litellm_params

View file

@ -1295,6 +1295,49 @@ def test_streaming_surfaces_fireworks_response_fields():
assert surfaced["fireworks_prompt_token_ids"] == [1, 2, 3]
def test_transform_request_routes_router_slug():
config = FireworksAIConfig()
data = config.transform_request(
model="routers/glm-latest",
messages=[{"role": "user", "content": "hi"}],
optional_params={},
litellm_params={},
headers={},
)
assert data["model"] == "accounts/fireworks/routers/glm-latest"
def test_transform_request_bare_slug_stays_model():
config = FireworksAIConfig()
data = config.transform_request(
model="glm-4p6",
messages=[{"role": "user", "content": "hi"}],
optional_params={},
litellm_params={},
headers={},
)
assert data["model"] == "accounts/fireworks/models/glm-4p6"
def test_transform_request_direct_route_passthrough():
config = FireworksAIConfig()
model = "accounts/fireworks/models/qwen2p5-coder-7b#accounts/gitlab/deployments/2fb7764c"
data = config.transform_request(
model=model,
messages=[{"role": "user", "content": "hi"}],
optional_params={},
litellm_params={},
headers={},
)
assert data["model"] == model
def test_map_extra_body_params_translates_truncate_prompt_tokens():
config = FireworksAIConfig()
result = config.map_extra_body_params(

View file

@ -0,0 +1,34 @@
import os
import sys
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.fireworks_ai.completion.transformation import (
FireworksAITextCompletionConfig,
)
def test_transform_text_completion_request_routes_router_slug():
config = FireworksAITextCompletionConfig()
data = config.transform_text_completion_request(
model="routers/glm-latest",
messages=[{"role": "user", "content": "hi"}],
optional_params={},
headers={},
)
assert data["model"] == "accounts/fireworks/routers/glm-latest"
def test_transform_text_completion_request_bare_slug_stays_model():
config = FireworksAITextCompletionConfig()
data = config.transform_text_completion_request(
model="glm-4p6",
messages=[{"role": "user", "content": "hi"}],
optional_params={},
headers={},
)
assert data["model"] == "accounts/fireworks/models/glm-4p6"

View file

@ -0,0 +1,45 @@
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
from litellm.llms.fireworks_ai.common_utils import resolve_fireworks_resource_name
@pytest.mark.parametrize(
"model, expected",
[
("routers/glm-latest", "accounts/fireworks/routers/glm-latest"),
("routers/firerouter", "accounts/fireworks/routers/firerouter"),
("fireworks_ai/routers/glm-latest", "accounts/fireworks/routers/glm-latest"),
("models/glm-4p6", "accounts/fireworks/models/glm-4p6"),
("fireworks_ai/models/glm-4p6", "accounts/fireworks/models/glm-4p6"),
("glm-4p6", "accounts/fireworks/models/glm-4p6"),
("fireworks_ai/glm-4p6", "accounts/fireworks/models/glm-4p6"),
("kimi-k2p6-fast", "accounts/fireworks/routers/kimi-k2p6-fast"),
(
"accounts/fireworks/routers/glm-latest",
"accounts/fireworks/routers/glm-latest",
),
(
"accounts/fireworks/models/glm-4p6",
"accounts/fireworks/models/glm-4p6",
),
(
"fireworks_ai/accounts/fireworks/routers/glm-latest",
"accounts/fireworks/routers/glm-latest",
),
(
"accounts/fireworks/models/qwen2p5-coder-7b#accounts/gitlab/deployments/2fb7764c",
"accounts/fireworks/models/qwen2p5-coder-7b#accounts/gitlab/deployments/2fb7764c",
),
(
"glm-4p6#accounts/gitlab/deployments/2fb7764c",
"glm-4p6#accounts/gitlab/deployments/2fb7764c",
),
],
)
def test_resolve_fireworks_resource_name(model, expected):
assert resolve_fireworks_resource_name(model) == expected

View file

@ -2646,7 +2646,7 @@ class TestMCPDelegateAuthToUpstream:
def test_extract_target_server_names_matches_routing_parser(self):
"""
Regression: _extract_target_server_names_from_path must match the
Regression: extract_target_server_names_from_path must match the
downstream regex parser in server.py::_get_mcp_servers_in_path.
Previously, a request to ``/mcp/<delegated>/garbage`` was parsed as
@ -2682,7 +2682,7 @@ class TestMCPDelegateAuthToUpstream:
("/", []),
]
for path_input, expected in cases:
assert MCPRequestHandler._extract_target_server_names_from_path(path_input) == expected, (
assert MCPRequestHandler.extract_target_server_names_from_path(path_input) == expected, (
f"path={path_input!r} → expected {expected!r}"
)
assert (_get_mcp_servers_in_path(path_input) or []) == expected, (
@ -8365,3 +8365,64 @@ class TestEntitlementFaultSemantics:
):
allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert set(allowed) == {"srv1"}
@pytest.mark.asyncio
class TestScopedSessionAdmission:
"""LIT-4917: a session bearer sealed to one server (RFC 8707 resource at authorize)
carries that scope onto the admitted auth object, where the grant resolution intersects
it fail closed; an unscoped bearer carries None and is byte-identical to before."""
_MASTER_KEY = "sk-scoped-session-admission-master-key"
def _bearer(self, resource_server_id):
from datetime import datetime, timezone
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
session_keys_from_master_key,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
SessionPrincipal,
mint_session_token,
)
keys = session_keys_from_master_key(self._MASTER_KEY)
principal = SessionPrincipal(
user_id="scoped-user", client_id="llm_dcrc_abc", resource_server_id=resource_server_id
)
return mint_session_token(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value()
@pytest.mark.parametrize("scope", ["github-server-id", None])
async def test_admission_carries_sealed_resource_scope(self, scope):
token = self._bearer(scope)
scope_dict = {
"type": "http",
"method": "POST",
"path": "/mcp/github",
"headers": [(b"host", b"testserver"), (b"authorization", f"Bearer {token}".encode())],
}
get_user_object = AsyncMock(
return_value=MagicMock(
user_id="scoped-user",
organization_id=None,
metadata={"scim_active": True},
user_role=None,
object_permission=None,
object_permission_id=None,
tpm_limit=None,
rpm_limit=None,
)
)
with (
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
):
auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope_dict)
assert auth_result.mcp_admitted_user_subject is True
assert auth_result.mcp_session_resource_server_id == scope
def test_scope_field_cannot_be_forged_through_construction(self):
forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server")
assert forged.mcp_session_resource_server_id is None

View file

@ -795,3 +795,247 @@ async def test_manual_delivery_page_renders_the_url_as_data_never_as_a_shell_com
assert 'curl "' not in body
assert "curl '" not in body
assert 'value="' in body
def _scoped_mcp_server(name="github", **kw):
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
return MCPServer(
server_id=f"{name}-id",
name=name,
server_name=name,
alias=name,
url="https://upstream.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
**kw,
)
SCOPED_RESOURCE = "https://llm.example.com/mcp/github"
_MANAGER_PATCH = "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
def _scoped_authorize(client_id, resource, session_user_id="u1"):
return aggregate_authorize(
request=_request(query=f"client_id={client_id}"),
client_id=client_id,
redirect_uri=REDIRECT_URI,
state="client-state-123",
code_challenge=CODE_CHALLENGE,
code_challenge_method="S256",
response_type="code",
session_user_id=session_user_id,
resource=resource,
)
async def _redeem(code, client_id, cache=None, **overrides):
arguments = {
"request": _request("/token", method="POST"),
"grant_type": "authorization_code",
"code": code,
"redirect_uri": REDIRECT_URI,
"client_id": client_id,
"code_verifier": CODE_VERIFIER,
"refresh_token": None,
"master_key": MASTER_KEY,
"reload_user": _reload_user_active,
"cache": cache or DualCache(),
}
return await aggregate_token(**{**arguments, **overrides})
def _opened_principal(payload):
keys = session_keys_from_master_key(MASTER_KEY)
admitted = resolve_session_bearer(f"Bearer {payload['access_token']}", keys, datetime.now(timezone.utc))
assert isinstance(admitted, SessionBearerAdmitted)
return admitted.principal
async def _finish_connect_page(response):
handle, cookies = _flow_cookie_from(response)
completed = await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
cache=DualCache(),
)
return parse_qs(urlparse(completed.headers["location"]).query)["code"][0]
def _sealed_wire_json(sealed, prefix, debug_key):
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
raw = decrypt_value_helper(sealed.removeprefix(prefix), debug_key, return_original_value=False)
assert isinstance(raw, str)
return json.loads(raw)
@pytest.mark.asyncio
async def test_scoped_authorize_runs_connect_page_with_sealed_scope():
"""LIT-4917: a per-server RFC 8707 resource naming a gateway-managed oauth2 server
seals that server into the flow. The connect page interlude runs exactly as before
(the scope restricts, it never skips consent), and the code minted at the finish step
and the session pair it redeems for are both scoped."""
from unittest.mock import patch
client_id = (await _register([REDIRECT_URI]))["client_id"]
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = _scoped_mcp_server()
response = _scoped_authorize(client_id, SCOPED_RESOURCE)
assert response.status_code == 303
assert "/ui/connect" in response.headers["location"]
_, cookies = _flow_cookie_from(response)
assert _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow")["resource_server_id"] == "github-id"
code = await _finish_connect_page(response)
assert _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, "gateway_authorization_code")["resource_server_id"] == "github-id"
token_response = await _redeem(code, client_id)
assert token_response.status_code == 200
principal = _opened_principal(json.loads(token_response.body))
assert principal.resource_server_id == "github-id"
assert principal.user_id == "u1"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"resource, resolves",
[
(None, False),
("https://llm.example.com/mcp", False),
("https://other.example.com/mcp/github", False),
("https://llm.example.com/mcp/github,linear", False),
("https://llm.example.com/mcp/unknown", None),
("not a url", False),
],
)
async def test_unscoped_resources_leave_flow_and_token_byte_identical(resource, resolves):
"""Every resource shape outside 'exactly one gateway-managed server' keeps today's flow:
connect page interlude, and NONE of the minted artifacts carry the scope key on the
wire, not the flow cookie, not the code, not the session JWT, so an unscoped flow
started on a new pod completes on a pod whose strict models predate the claim."""
import base64
from unittest.mock import patch
client_id = (await _register([REDIRECT_URI]))["client_id"]
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = None if resolves is None else _scoped_mcp_server()
response = _scoped_authorize(client_id, resource)
assert response.status_code == 303
assert "/ui/connect" in response.headers["location"]
_, cookies = _flow_cookie_from(response)
assert "resource_server_id" not in _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow")
code = await _finish_connect_page(response)
assert "resource_server_id" not in _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, "gateway_authorization_code")
token_response = await _redeem(code, client_id)
payload = json.loads(token_response.body)
assert _opened_principal(payload).resource_server_id is None
jwt_payload_segment = payload["access_token"].removeprefix("llm_session_").split(".")[1]
claims = json.loads(base64.urlsafe_b64decode(jwt_payload_segment + "=" * (-len(jwt_payload_segment) % 4)))
assert "resource_server_id" not in claims
@pytest.mark.asyncio
async def test_scoped_authorize_delegate_server_stays_unscoped():
"""A delegate-auth oauth2 server is outside the gateway-managed set (its keyless flow is
upstream PKCE via the relay), so a resource naming it never scopes the gateway flow."""
from unittest.mock import patch
client_id = (await _register([REDIRECT_URI]))["client_id"]
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = _scoped_mcp_server(delegate_auth_to_upstream=True)
response = _scoped_authorize(client_id, SCOPED_RESOURCE)
assert "/ui/connect" in response.headers["location"]
code = await _finish_connect_page(response)
token_response = await _redeem(code, client_id)
assert _opened_principal(json.loads(token_response.body)).resource_server_id is None
@pytest.mark.asyncio
async def test_token_rejects_resource_conflicting_with_sealed_scope():
"""RFC 8707 section 2.2: redeeming a scoped code (or rotating a scoped refresh token)
for a DIFFERENT resource fails with invalid_target; an absent resource redeems fine and
the sealed scope still binds the minted pair, surviving refresh rotation."""
from unittest.mock import patch
client_id = (await _register([REDIRECT_URI]))["client_id"]
github = _scoped_mcp_server()
linear = _scoped_mcp_server(name="linear")
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = github
response = _scoped_authorize(client_id, SCOPED_RESOURCE)
code = await _finish_connect_page(response)
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = linear
mismatched = await _redeem(code, client_id, resource="https://llm.example.com/mcp/linear")
assert json.loads(mismatched.body)["error"] == "invalid_target"
cache = DualCache()
token_response = await _redeem(code, client_id, cache=cache)
payload = json.loads(token_response.body)
assert _opened_principal(payload).resource_server_id == "github-id"
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = linear
refresh_mismatch = await _redeem(
None,
client_id,
cache=cache,
grant_type="refresh_token",
refresh_token=payload["refresh_token"],
resource="https://llm.example.com/mcp/linear",
)
assert json.loads(refresh_mismatch.body)["error"] == "invalid_target"
rotated = await _redeem(
None, client_id, cache=cache, grant_type="refresh_token", refresh_token=payload["refresh_token"]
)
assert rotated.status_code == 200
assert _opened_principal(json.loads(rotated.body)).resource_server_id == "github-id"
@pytest.mark.asyncio
async def test_resolve_scoped_resource_server_matrix():
"""Unit pin of the resource resolver: both per-server URL spellings resolve; the
aggregate resource, foreign hosts, CSV paths, unknown names, and non-gateway-managed
modes all return None so nothing outside the served set can enter the scoped flow."""
from unittest.mock import patch
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import resolve_scoped_resource_server
request = _request()
github = _scoped_mcp_server()
for resource, resolved_server, expected in [
("https://llm.example.com/mcp/github", github, "github-id"),
("https://llm.example.com/github/mcp", github, "github-id"),
("https://LLM.example.com/mcp/github/", github, "github-id"),
("https://llm.example.com/mcp", github, None),
("https://other.example.com/mcp/github", github, None),
("https://llm.example.com/mcp/a,b", github, None),
("https://llm.example.com/mcp/github", None, None),
("https://llm.example.com/mcp/github", _scoped_mcp_server(delegate_auth_to_upstream=True), None),
(None, github, None),
]:
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = resolved_server
result = resolve_scoped_resource_server(request, resource)
assert (result.server_id if result is not None else None) == expected, resource
@pytest.mark.asyncio
async def test_resource_resolution_is_identity_not_ip_filtered_access():
"""The resolver decides which server a resource NAMES; per-IP visibility filtering
belongs to the MCP routes and grant intersection. Filtering here would mint an
entitlement-wide unscoped bearer exactly when the caller asked to narrow, and IP drift
between authorize and token would turn a matching redemption into invalid_target."""
from unittest.mock import patch
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import resolve_scoped_resource_server
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = _scoped_mcp_server()
result = resolve_scoped_resource_server(_request(), SCOPED_RESOURCE)
assert result is not None
manager.get_mcp_server_by_name.assert_called_once_with("github")

View file

@ -137,7 +137,7 @@ def test_is_mcp_passthrough_cold_start_false_for_empty_servers():
[
("/mcp/sample_docs", ["sample_docs"]),
# Server names may contain at most one slash (mirrors
# ``_extract_target_server_names_from_path``), so when more than two
# ``extract_target_server_names_from_path``), so when more than two
# segments follow ``/mcp/`` the first two are treated as the name.
("/mcp/sample_docs/tools/list", ["sample_docs/tools"]),
("/mcp/custom_solutions/user_123", ["custom_solutions/user_123"]),

View file

@ -9859,3 +9859,66 @@ class TestToolAuthorizationIsNotConditionalOnLogging:
)
upstream.assert_awaited_once()
class TestSessionResourceScopeIntersect:
"""LIT-4917: the sealed session scope intersects the admitted subject's resolved server
set at the single convergence point every fan-out and tool call reads, covering the
exception fallback so a resolver fault never widens a scoped bearer."""
def _admitted_auth(self, scope):
from litellm.proxy._types import UserAPIKeyAuth
auth = UserAPIKeyAuth(user_id="scoped-user")
auth.mcp_admitted_user_subject = True
auth.mcp_session_resource_server_id = scope
return auth
def test_scope_reader_is_none_for_keys_and_unscoped_subjects(self):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy._types import UserAPIKeyAuth
assert MCPServerManager._admitted_session_resource_scope(None) is None
assert MCPServerManager._admitted_session_resource_scope(UserAPIKeyAuth(user_id="u")) is None
assert MCPServerManager._admitted_session_resource_scope(self._admitted_auth(None)) is None
def test_scope_reader_returns_sealed_scope_for_admitted_subjects(self):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
assert MCPServerManager._admitted_session_resource_scope(self._admitted_auth("b")) == "b"
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_scopes_past_operator_open_union(self):
"""The intersect applies AFTER the operator-open (allow_all_keys) union, so a scoped
bearer cannot reach an allow-all server outside its scope, and applies on the
exception fallback so a resolver fault yields the scoped subset of allow-all rather
than the whole set."""
from unittest.mock import AsyncMock, patch
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager = MCPServerManager()
auth = self._admitted_auth("granted-id")
with (
patch.object(MCPServerManager, "get_allow_all_keys_server_ids", return_value=["open-id", "granted-id"]),
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=["granted-id", "other-id"],
),
patch.object(MCPServerManager, "_get_active_submitted_mcp_server_ids_for_user", new_callable=AsyncMock, return_value=[]),
):
allowed = await manager.get_allowed_mcp_servers(auth)
assert allowed == ["granted-id"]
with (
patch.object(MCPServerManager, "get_allow_all_keys_server_ids", return_value=["open-id", "granted-id"]),
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers",
new_callable=AsyncMock,
side_effect=RuntimeError("resolver down"),
),
patch.object(MCPServerManager, "_get_active_submitted_mcp_server_ids_for_user", new_callable=AsyncMock, return_value=[]),
):
fallback = await manager.get_allowed_mcp_servers(auth)
assert fallback == ["granted-id"]

View file

@ -1138,7 +1138,11 @@ async def test_retrieve__unified_batch_id_routes_to_router(retrieve_harness):
# DISPATCH - router fired, direct litellm did not.
assert retrieve_harness.router_aretrieve.call_count == 1
retrieve_harness.litellm_aretrieve.assert_not_called()
retrieve_harness.creds_resolver.assert_not_called()
# Credentials are resolved for the deployment behind the unified id so the batch's
# output file can be read for cost accounting. This id resolves to nothing here, and
# the retrieve must still serve the batch rather than fail on the lookup.
retrieve_harness.creds_resolver.assert_called_once_with(model_id="gpt-4o-mini")
# router receives the (still-encoded) batch id verbatim - this layer does
# not decode it for the unified path.

View file

@ -5647,6 +5647,235 @@ class TestPanwAirsScanIdExposure:
assert "guardrail_scan_ids" in _UNTRUSTED_METADATA_CONTROL_FIELDS
assert "guardrail_scan_ids" in _UNTRUSTED_ROOT_CONTROL_FIELDS
class TestPanwAirsBlockedErrorDetailPassthrough:
"""Regression tests for the full AIRS scan response on blocks.
Before the fix, the error detail was built from a hardcoded allowlist
(scan_id, report_id, profile_name, profile_id, tr_id, prompt/response_detected),
so audit-relevant fields such as prompt_detection_details, prompt_masked_data,
source, transaction_id and session_id never reached the client.
"""
_FULL_BLOCK_RESPONSE = {
"action": "block",
"category": "malicious",
"scan_id": "b2f0a4be-1f6f-4f9a-9f3d-4b6a9d8b1c0e",
"report_id": "R0000000000000000000",
"tr_id": "test-call-id",
"profile_id": "6f5c9f6e-2d0b-4d3f-8a1e-9b7c5d4e3f2a",
"profile_name": "test_profile",
"source": "prisma_airs",
"transaction_id": "4b8c1e2f-5a6d-4c3b-9e8f-1a2b3c4d5e6f",
"session_id": "3a2b1c0d-9e8f-4a7b-8c6d-5e4f3a2b1c0d",
"timeout": False,
"errors": [],
"prompt_detected": {"dlp": True, "injection": False, "url_cats": False},
"prompt_detection_details": {
"dlp_report": {
"dlp_report_id": "1234567890",
"dlp_profile_name": "Sensitive Content",
"data_pattern_rule1_verdict": "MATCHED",
}
},
"prompt_masked_data": {"data": "my ssn is XXX-XX-XXXX"},
"response_detected": {"dlp": False, "url_cats": False},
"response_detection_details": {},
"response_masked_data": {},
}
@pytest.mark.asyncio
@pytest.mark.parametrize("is_response", [False, True])
async def test_block_returns_every_airs_field(
self, base_handler, user_api_key_dict, safe_prompt_data, is_response
):
response = ModelResponse(
id="test_id",
choices=[
Choices(index=0, message=Message(role="assistant", content="Test response")),
],
model="gpt-3.5-turbo",
)
with patch.object(
base_handler, "_call_panw_api", return_value=copy.deepcopy(self._FULL_BLOCK_RESPONSE)
):
with pytest.raises(HTTPException) as exc_info:
if is_response:
await base_handler.async_post_call_success_hook(
data=safe_prompt_data,
user_api_key_dict=user_api_key_dict,
response=response,
)
else:
await base_handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=None,
data=safe_prompt_data,
call_type="completion",
)
error = exc_info.value.detail["error"]
for field, value in self._FULL_BLOCK_RESPONSE.items():
if field == "category":
continue
if field in PanwPrismaAirsHandler._CLIENT_HIDDEN_SCAN_FIELDS:
# Withheld on purpose, covered by TestPanwAirsErrorDetailWithheldFields
continue
assert error[field] == value, f"{field} missing or altered in blocked-request error"
assert error["category"] == "malicious"
assert error["type"] == "guardrail_violation"
assert error["guardrail"] == "test_panw_airs"
assert error["code"] == ("panw_prisma_airs_response_blocked" if is_response else "panw_prisma_airs_blocked")
assert "PANW Prisma AI Security policy" in error["message"]
def test_internal_control_flags_are_not_leaked(self, base_handler):
detail = base_handler._build_error_detail(
{
"action": "block",
"category": "malicious",
"scan_id": "scan-1",
"_always_block": True,
"_is_transient": True,
}
)
assert "_always_block" not in detail["error"]
assert "_is_transient" not in detail["error"]
assert detail["error"]["scan_id"] == "scan-1"
class TestPanwAirsErrorDetailWithheldFields:
"""The blocked-request passthrough must not become a content channel.
``response_masked_data`` is the model's own generation. The block branch is only
reached when ``mask_response_content`` is False, so echoing it back would hand the
caller exactly the text the operator declined to deliver. ``error`` is AIRS's own
message about the operator's Strata Cloud Manager profile configuration.
``prompt_masked_data`` is deliberately NOT withheld by default: it is the caller's
own input, and it is one of the fields LIT-5638 asks for. The one exception is the
response-side tool-call path, covered by
``TestPanwAirsToolCallBlockWithholdsGeneratedArgs`` below — tool_event scans are
request-side in the AIRS schema, so there the key holds model output instead.
"""
@pytest.mark.parametrize("is_response", [False, True])
def test_response_masked_data_never_reaches_client(self, base_handler, is_response):
detail = base_handler._build_error_detail(
{
"action": "block",
"category": "sensitive_data",
"scan_id": "scan-1",
"response_detected": {"dlp": True},
"response_masked_data": {"data": "routing number XXXXXXXXXX"},
"prompt_masked_data": {"data": "my ssn is XXX-XX-XXXX"},
"prompt_detection_details": {"dlp_report": {"dlp_report_id": "1"}},
},
is_response=is_response,
)
error = detail["error"]
assert "response_masked_data" not in error
assert "routing number" not in str(error)
# The audit fields LIT-5638 asks for still come through untouched.
assert error["scan_id"] == "scan-1"
assert error["response_detected"] == {"dlp": True}
assert error["prompt_masked_data"] == {"data": "my ssn is XXX-XX-XXXX"}
assert error["prompt_detection_details"] == {"dlp_report": {"dlp_report_id": "1"}}
def test_upstream_airs_error_field_still_passes_through(self, base_handler):
"""A 2xx AIRS body can carry its own ``error`` (see _call_panw_api's
profile-misconfiguration branch, which only logs and then blocks). It is
diagnostic rather than content, so it stays in the passthrough."""
detail = base_handler._build_error_detail(
{
"action": "block",
"category": "malicious",
"scan_id": "scan-2",
"error": "profile not found",
}
)
assert detail["error"]["error"] == "profile not found"
assert detail["error"]["scan_id"] == "scan-2"
class TestPanwAirsToolCallBlockWithholdsGeneratedArgs:
"""A response-side tool-call block must not ship the model's tool arguments.
``_scan_tool_calls_for_guardrail`` calls AIRS with ``is_response=False`` because
tool_event is request-side in the AIRS schema, so AIRS returns the scanned tool
arguments under ``prompt_masked_data``. When the tool calls being scanned are the
model's own output, that key holds generated content, and the class-level
``_CLIENT_HIDDEN_SCAN_FIELDS`` default (``response_masked_data``, empty on this
path) does not cover it.
"""
MASKED_ARGS = '{"to_account": "XXXXXXXXXX", "amount": 5000}'
SCAN_RESULT = {
"action": "block",
"category": "sensitive_data",
"scan_id": "scan-tool-1",
"prompt_detected": {"dlp": True},
"prompt_masked_data": {"data": MASKED_ARGS},
"response_masked_data": {},
}
@staticmethod
def _tool_call():
return ChatCompletionMessageToolCall(
id="call_1",
type="function",
function=Function(
name="transfer_funds",
arguments='{"to_account": "ACME-VENDOR-001", "amount": 5000}',
),
)
async def _block(self, handler, is_response):
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
mock_api.return_value = dict(self.SCAN_RESULT)
with pytest.raises(HTTPException) as exc_info:
await handler._scan_tool_calls_for_guardrail(
tool_calls=[self._tool_call()],
is_response=is_response,
metadata={},
call_id="test-call-id",
request_data={"metadata": {}},
start_time=datetime.now(),
)
return exc_info.value
@pytest.mark.asyncio
async def test_response_side_block_withholds_generated_tool_args(self):
handler = make_handler(mask_response_content=False)
# The block branch is only reached with masking off; guard the premise.
assert handler.mask_response_content is False
exc = await self._block(handler, is_response=True)
error = exc.detail["error"]
assert exc.status_code == 400
assert "prompt_masked_data" not in error
assert self.MASKED_ARGS not in str(error)
# The audit fields LIT-5638 asks for are unaffected.
assert error["scan_id"] == "scan-tool-1"
assert error["prompt_detected"] == {"dlp": True}
@pytest.mark.asyncio
async def test_request_side_block_still_returns_masked_tool_args(self):
"""Caller-supplied tool arguments stay in the verdict — that is the ticket's ask."""
handler = make_handler(mask_request_content=False)
exc = await self._block(handler, is_response=False)
error = exc.detail["error"]
assert error["prompt_masked_data"] == {"data": self.MASKED_ARGS}
assert error["scan_id"] == "scan-tool-1"
if __name__ == "__main__":

View file

@ -256,7 +256,6 @@ def test_batch_cost_poller_is_active_is_false_when_get_job_raises(monkeypatch):
assert batch_cost_poller_is_active() is False
@pytest.mark.asyncio
async def test_retrieving_a_batch_whose_status_is_unchanged_writes_nothing(monkeypatch):
import litellm.proxy.openai_files_endpoints.common_utils as cu
@ -364,3 +363,71 @@ async def test_a_caller_that_handed_off_accounting_still_leaves_the_marker_alone
data = update_mock.await_args.kwargs["data"]
assert "batch_processed" not in data
assert data["status"] == "complete"
# =========================================================================== #
# add_internal_model_credentials - the snapshot that lets a completed
# batch's output file be read, and therefore its cost be recorded
# =========================================================================== #
def test_add_internal_model_credentials_attaches_an_immutable_snapshot():
"""Cost accounting for a completed batch reads its output file, and Bedrock resolves
that bucket only from this snapshot. It must be immutable so nothing downstream can
redirect the bucket that managed file ids are validated against."""
from litellm.proxy.openai_files_endpoints.common_utils import (
add_internal_model_credentials,
)
router = MagicMock()
router.get_deployment_credentials_with_provider = MagicMock(
return_value={"s3_bucket_name": "configured-bucket", "aws_region_name": "us-east-1"}
)
data = {"batch_id": "unified-batch-id"}
add_internal_model_credentials(data=data, llm_router=router, model_id="deployment-1")
snapshot = data["_litellm_internal_model_credentials"]
assert snapshot["s3_bucket_name"] == "configured-bucket"
assert isinstance(snapshot, MappingProxyType)
with pytest.raises(TypeError):
snapshot["s3_bucket_name"] = "attacker-bucket"
router.get_deployment_credentials_with_provider.assert_called_once_with(model_id="deployment-1")
@pytest.mark.parametrize(
"model_id, credentials",
[(None, {"s3_bucket_name": "b"}), ("deployment-1", None)],
ids=["no-model-id", "deployment-has-no-credentials"],
)
def test_add_internal_model_credentials_is_a_noop_without_a_resolvable_deployment(model_id, credentials):
"""An unroutable batch must be left alone rather than given an empty snapshot, which
would look like a configured bucket of nothing."""
from litellm.proxy.openai_files_endpoints.common_utils import (
add_internal_model_credentials,
)
router = MagicMock()
router.get_deployment_credentials_with_provider = MagicMock(return_value=credentials)
data = {"batch_id": "unified-batch-id"}
add_internal_model_credentials(data=data, llm_router=router, model_id=model_id)
assert "_litellm_internal_model_credentials" not in data
def test_add_internal_model_credentials_survives_a_failing_deployment_lookup():
"""The snapshot only enables cost accounting, so a batch whose deployment no longer
resolves, which happens when a model group is removed while batches are in flight,
must still be retrievable rather than failing the request on the lookup."""
from litellm.proxy.openai_files_endpoints.common_utils import (
add_internal_model_credentials,
)
router = MagicMock()
router.get_deployment_credentials_with_provider = MagicMock(side_effect=KeyError("deployment-gone"))
data = {"batch_id": "unified-batch-id"}
add_internal_model_credentials(data=data, llm_router=router, model_id="deployment-gone")
assert data == {"batch_id": "unified-batch-id"}

View file

@ -3149,6 +3149,99 @@ def test_require_managed_files_rejects_raw_provider_file_id(
mock_call.assert_not_called()
def test_get_file_content_model_routed_attaches_trusted_model_credentials(monkeypatch):
"""A managed batch output id routes by model, and that branch must build the snapshot.
The managed-files pre-call hook sets data["model"] for any id carrying
llm_output_file_id, so batch output retrieval always takes the model-routed branch
and never reaches managed_files_obj.afile_content. Bedrock resolves its output
bucket only from _litellm_internal_model_credentials, so without the snapshot every
Bedrock batch output retrieval fails with "S3 bucket_name is required".
"""
import base64
from types import MappingProxyType
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
from litellm.types.utils import SpecialEnums
router = Router(
model_list=[
{
"model_name": "anthropic.batch.claude-4.5-haiku",
"litellm_params": {
"model": "bedrock/anthropic.claude-haiku-4-5-20251001-v1:0",
"aws_region_name": "us-east-1",
"s3_bucket_name": "configured-batch-bucket",
},
"model_info": {"id": "bedrock-batch-deployment-id"},
}
]
)
from unittest.mock import MagicMock
managed_file_row = MagicMock()
managed_file_row.created_by = "test-user"
managed_file_row.team_id = None
managed_file_row.storage_backend = None
managed_file_row.storage_url = None
prisma_stub = MagicMock()
prisma_stub.db.litellm_managedfiletable.find_first = AsyncMock(return_value=managed_file_row)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_stub)
setup_proxy_logging_object(monkeypatch, router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router)
# One frozen snapshot per call rather than one dict merged across calls, so a second
# invocation is visible instead of silently overwriting the first.
calls: list[MappingProxyType] = []
async def _mock_router_afile_content(**kwargs):
calls.append(MappingProxyType(dict(kwargs)))
return HttpxBinaryResponseContent(
response=httpx.Response(
status_code=200,
content=b'{"recordId":"req-1"}',
headers={"content-type": "application/octet-stream"},
)
)
monkeypatch.setattr(router, "afile_content", _mock_router_afile_content)
unified_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
"application/jsonl",
"unified-output-id",
"anthropic.batch.claude-4.5-haiku",
"llm_output_file_id,s3://configured-batch-bucket/out/batch.jsonl",
"bedrock-batch-deployment-id",
)
encoded_id = base64.urlsafe_b64encode(unified_id.encode()).decode().rstrip("=")
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="test-key",
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="test-user",
)
try:
response = client.get(
f"/v1/files/{encoded_id}/content",
headers={"Authorization": "Bearer test-key", "custom-llm-provider": "bedrock"},
)
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert response.status_code == 200, response.text
assert len(calls) == 1, f"expected exactly one routed retrieval, got {len(calls)}"
snapshot = calls[0].get("_litellm_internal_model_credentials")
assert snapshot is not None, "model-routed branch must attach the trusted credential snapshot"
assert isinstance(
snapshot, MappingProxyType
), "snapshot must be a MappingProxyType; a plain dict is rejected by get_configured_s3_bucket_name"
assert snapshot["s3_bucket_name"] == "configured-batch-bucket"
def _unified_managed_file_id() -> str:
import base64

View file

@ -568,6 +568,32 @@ def test_forward_headers_custom_wins_case_insensitive_over_request_authorization
assert result["x-request-id"] == "req-123"
def test_forward_headers_never_forwards_client_accept_encoding():
"""
The client's Accept-Encoding must not reach the upstream provider: the proxy's
HTTP client decodes the upstream body and advertises only encodings it can
decode. Forwarding e.g. "br" on an install without the brotli package makes
the proxy relay raw compressed bytes with the content-encoding header stripped
(garbled JSON for /v1/models and count_tokens through the Anthropic passthrough).
"""
from litellm.passthrough.utils import BasePassthroughUtils
request_headers = {
"accept-encoding": "gzip, deflate, br, zstd",
"x-pass-accept-encoding": "br",
"x-request-id": "req-123",
}
result = BasePassthroughUtils.forward_headers_from_request(
request_headers=request_headers,
headers={},
forward_headers=True,
)
assert "accept-encoding" not in result
assert result["x-request-id"] == "req-123"
@pytest.mark.asyncio
async def test_vertex_passthrough_custom_model_name_replaced_in_url():
"""

View file

@ -0,0 +1,333 @@
import fcntl
import importlib.util
import os
import signal
import subprocess
import sys
import time
from collections.abc import Callable, Sequence
from contextlib import suppress
from pathlib import Path
import pytest
ROOT = Path(__file__).resolve().parents[2]
HELPER = ROOT / "scripts" / "gate_slot_lock.py"
_spec = importlib.util.spec_from_file_location("gate_slot_lock", HELPER)
assert _spec is not None and _spec.loader is not None
gate_slot_lock = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(gate_slot_lock)
START_THEN_WAIT_FOR = (
"import pathlib, sys, time\n"
"pathlib.Path(sys.argv[1]).touch()\n"
"deadline = time.monotonic() + 20\n"
"while not pathlib.Path(sys.argv[2]).exists():\n"
" if time.monotonic() > deadline:\n"
" sys.exit(3)\n"
" time.sleep(0.05)\n"
)
TOUCH_TARGET = "import pathlib, sys\npathlib.Path(sys.argv[1]).touch()\n"
RECORD_INTERVAL = (
"import sys, time\n"
"with open(sys.argv[1], 'a') as events:\n"
" events.write(f'start {time.monotonic()}\\n')\n"
" events.flush()\n"
" time.sleep(0.6)\n"
" events.write(f'end {time.monotonic()}\\n')\n"
" events.flush()\n"
)
def _env(lock_dir: Path, slots: str) -> dict[str, str]:
return {
"PATH": os.environ["PATH"],
"HOME": str(lock_dir.parent),
"LITELLM_GATE_SLOT_DIR": str(lock_dir),
"LITELLM_GATE_SLOTS": slots,
}
def _wrapped(payload: Sequence[str]) -> list[str]:
return [sys.executable, str(HELPER), sys.executable, "-c", *payload]
def _wait_until(predicate: Callable[[], bool], timeout_seconds: float) -> bool:
deadline = time.monotonic() + timeout_seconds
while time.monotonic() < deadline:
if predicate():
return True
time.sleep(0.05)
return predicate()
def _terminate_group(process: subprocess.Popen[bytes]) -> None:
with suppress(ProcessLookupError, PermissionError):
os.killpg(process.pid, signal.SIGKILL)
def _reap(process: subprocess.Popen[bytes]) -> None:
with suppress(subprocess.TimeoutExpired):
process.wait(timeout=10)
if process.poll() is None:
process.kill()
process.wait(timeout=10)
def test_six_contenders_never_exceed_two_slots_and_all_complete(tmp_path: Path) -> None:
lock_dir = tmp_path / "locks"
events_file = tmp_path / "events.log"
env = _env(lock_dir, "2")
procs = [
subprocess.Popen(
_wrapped([RECORD_INTERVAL, str(events_file)]),
env=env,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
for _ in range(6)
]
try:
assert [proc.wait(timeout=60) for proc in procs] == [0] * 6
finally:
for proc in procs:
if proc.poll() is None:
proc.kill()
proc.wait(timeout=10)
events = sorted(
(float(stamp), 1 if kind == "start" else -1)
for kind, stamp in (line.split() for line in events_file.read_text().splitlines())
)
assert len(events) == 12
concurrency_peaks = []
running = 0
for _, delta in events:
running += delta
concurrency_peaks.append(running)
assert max(concurrency_peaks) <= 2
def test_two_slots_admit_two_holders_at_once(tmp_path: Path) -> None:
lock_dir = tmp_path / "locks"
first_started = tmp_path / "first.started"
second_started = tmp_path / "second.started"
env = _env(lock_dir, "2")
first = subprocess.Popen(_wrapped([START_THEN_WAIT_FOR, str(first_started), str(second_started)]), env=env)
second = subprocess.Popen(_wrapped([START_THEN_WAIT_FOR, str(second_started), str(first_started)]), env=env)
assert first.wait(timeout=30) == 0
assert second.wait(timeout=30) == 0
def test_contender_beyond_capacity_queues_until_the_slot_frees(tmp_path: Path) -> None:
lock_dir = tmp_path / "locks"
holder_started = tmp_path / "holder.started"
release = tmp_path / "release"
done = tmp_path / "done"
env = _env(lock_dir, "1")
holder = subprocess.Popen(_wrapped([START_THEN_WAIT_FOR, str(holder_started), str(release)]), env=env)
try:
assert _wait_until(holder_started.exists, 10)
contender = subprocess.Popen(
_wrapped([TOUCH_TARGET, str(done)]),
env=env,
stderr=subprocess.PIPE,
)
try:
time.sleep(1.5)
assert not done.exists()
release.touch()
assert holder.wait(timeout=10) == 0
assert contender.wait(timeout=30) == 0
assert done.exists()
assert contender.stderr is not None
assert b"queueing" in contender.stderr.read()
finally:
release.touch()
_reap(contender)
finally:
release.touch()
_reap(holder)
def test_nested_wrapping_reenters_instead_of_deadlocking(tmp_path: Path) -> None:
lock_dir = tmp_path / "locks"
nested = [
sys.executable,
str(HELPER),
sys.executable,
str(HELPER),
sys.executable,
"-c",
"print('nested ok')",
]
proc = subprocess.Popen(
nested,
env=_env(lock_dir, "1"),
stdout=subprocess.PIPE,
start_new_session=True,
)
try:
stdout, _ = proc.communicate(timeout=20)
except subprocess.TimeoutExpired:
_terminate_group(proc)
pytest.fail("nested gate_slot_lock invocations deadlocked")
assert proc.returncode == 0
assert b"nested ok" in stdout
def test_wrapped_command_exit_code_is_propagated(tmp_path: Path) -> None:
proc = subprocess.run(
[sys.executable, str(HELPER), sys.executable, "-c", "raise SystemExit(7)"],
env=_env(tmp_path / "locks", "2"),
)
assert proc.returncode == 7
def test_missing_command_exits_127_and_no_command_exits_2(tmp_path: Path) -> None:
env = _env(tmp_path / "locks", "2")
missing = subprocess.run(
[sys.executable, str(HELPER), str(tmp_path / "no-such-binary")],
env=env,
capture_output=True,
)
assert missing.returncode == 127
bare = subprocess.run([sys.executable, str(HELPER)], env=env, capture_output=True)
assert bare.returncode == 2
def test_wrapped_command_killed_by_signal_maps_to_128_plus_signal(tmp_path: Path) -> None:
proc = subprocess.run(
_wrapped(["import os, signal\nos.kill(os.getpid(), signal.SIGTERM)\n"]),
env=_env(tmp_path / "locks", "2"),
)
assert proc.returncode == 128 + signal.SIGTERM
def test_unusable_lock_dir_fails_open_and_still_runs_the_command(tmp_path: Path) -> None:
blocker = tmp_path / "blocker"
blocker.write_text("")
done = tmp_path / "done"
proc = subprocess.run(
_wrapped([TOUCH_TARGET, str(done)]),
env=_env(blocker / "locks", "2"),
capture_output=True,
)
assert proc.returncode == 0
assert done.exists()
assert b"running unlocked" in proc.stderr
def test_zero_slots_disables_locking_entirely(tmp_path: Path) -> None:
lock_dir = tmp_path / "locks"
done = tmp_path / "done"
proc = subprocess.run(
_wrapped([TOUCH_TARGET, str(done)]),
env=_env(lock_dir, "0"),
)
assert proc.returncode == 0
assert done.exists()
assert not lock_dir.exists()
def test_non_integer_slot_count_warns_and_falls_back_to_default(tmp_path: Path) -> None:
proc = subprocess.run(
[sys.executable, str(HELPER), sys.executable, "-c", "print('ran')"],
env=_env(tmp_path / "locks", "lots"),
capture_output=True,
)
assert proc.returncode == 0
assert b"ran" in proc.stdout
assert b"LITELLM_GATE_SLOTS" in proc.stderr
def test_killed_holder_releases_its_slot_for_the_next_contender(tmp_path: Path) -> None:
lock_dir = tmp_path / "locks"
holder_started = tmp_path / "holder.started"
never = tmp_path / "never"
env = _env(lock_dir, "1")
holder = subprocess.Popen(
_wrapped([START_THEN_WAIT_FOR, str(holder_started), str(never)]),
env=env,
start_new_session=True,
)
try:
assert _wait_until(holder_started.exists, 10)
finally:
_terminate_group(holder)
holder.wait(timeout=10)
after = subprocess.run(
[sys.executable, str(HELPER), sys.executable, "-c", "print('freed')"],
env=env,
capture_output=True,
timeout=20,
)
assert after.returncode == 0
assert b"freed" in after.stdout
def test_acquire_slot_holds_marks_and_releases_in_process(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
lock_dir = tmp_path / "locks"
monkeypatch.setenv("LITELLM_GATE_SLOT_HELD", "")
monkeypatch.setenv("LITELLM_GATE_SLOT_DIR", str(lock_dir))
monkeypatch.setenv("LITELLM_GATE_SLOTS", "1")
handle = gate_slot_lock.acquire_slot()
assert handle is not None
assert os.environ["LITELLM_GATE_SLOT_HELD"] == "1"
assert gate_slot_lock.acquire_slot() is None
with (lock_dir / "slot-0.lock").open("wb") as probe:
with pytest.raises(BlockingIOError):
fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB)
handle.close()
fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB)
fcntl.flock(probe, fcntl.LOCK_UN)
def test_held_slot_context_manager_releases_on_exit(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
lock_dir = tmp_path / "locks"
monkeypatch.setenv("LITELLM_GATE_SLOT_HELD", "")
monkeypatch.setenv("LITELLM_GATE_SLOT_DIR", str(lock_dir))
monkeypatch.setenv("LITELLM_GATE_SLOTS", "1")
with gate_slot_lock.held_slot():
assert os.environ["LITELLM_GATE_SLOT_HELD"] == "1"
with (lock_dir / "slot-0.lock").open("wb") as probe:
with pytest.raises(BlockingIOError):
fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB)
assert not os.environ.get("LITELLM_GATE_SLOT_HELD")
with (lock_dir / "slot-0.lock").open("wb") as probe:
fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB)
fcntl.flock(probe, fcntl.LOCK_UN)
def _make_rule(target: str) -> tuple[list[str], list[str]]:
database = subprocess.run(
["make", "--dry-run", "--print-data-base", "info"],
cwd=ROOT,
capture_output=True,
text=True,
check=True,
).stdout
lines = database.splitlines()
for index, line in enumerate(lines):
if line != f"{target}:" and not line.startswith(f"{target}: "):
continue
recipe: list[str] = []
for follower in lines[index + 1 :]:
if follower.startswith("#"):
continue
if not follower.startswith("\t"):
break
recipe.append(follower.strip())
return line.split(":", 1)[1].split(), recipe
raise AssertionError(f"target {target} not found in make database")
def test_direct_make_lint_takes_a_slot_before_any_setup() -> None:
lint_prerequisites, lint_recipe = _make_rule("lint")
assert lint_prerequisites == []
assert any("$(GATE_SLOT_LOCK)" in line for line in lint_recipe)
inner_prerequisites, _ = _make_rule("lint-inner")
assert "lint-install" in inner_prerequisites
assert "lint-fetch-base" in inner_prerequisites

View file

@ -1,4 +1,5 @@
import os
import shutil
import signal
import subprocess
import time
@ -420,6 +421,46 @@ def test_staged_files_matching_no_check_print_an_explicit_noop_note_and_nonempty
assert "skipped: Python lint (make lint) (no litellm/ Python files in scope)" in log
def test_run_queues_through_the_machine_wide_gate_slot_lock(tmp_path: Path) -> None:
repo, bin_dir = _sandbox(tmp_path)
lock_dir = tmp_path / "gate-locks"
proc = _run(repo, bin_dir, {"LITELLM_GATE_SLOT_DIR": str(lock_dir)})
assert proc.returncode == 0, proc.stdout + proc.stderr
assert (lock_dir / "slot-0.lock").exists()
def test_run_under_a_held_slot_skips_reacquiring_the_gate_lock(tmp_path: Path) -> None:
repo, bin_dir = _sandbox(tmp_path)
lock_dir = tmp_path / "gate-locks"
proc = _run(
repo,
bin_dir,
{"LITELLM_GATE_SLOT_DIR": str(lock_dir), "LITELLM_GATE_SLOT_HELD": "1"},
)
assert proc.returncode == 0, proc.stdout + proc.stderr
assert not lock_dir.exists()
def test_hook_symlink_install_still_resolves_the_slot_lock_helper(tmp_path: Path) -> None:
repo, bin_dir = _sandbox(tmp_path)
scripts_dir = repo / "scripts"
scripts_dir.mkdir()
shutil.copy(SCRIPT, scripts_dir / "pre_commit_lint.sh")
shutil.copy(SCRIPT.parent / "gate_slot_lock.py", scripts_dir / "gate_slot_lock.py")
(repo / ".git" / "hooks" / "pre-commit").symlink_to(Path("../../scripts/pre_commit_lint.sh"))
lock_dir = tmp_path / "gate-locks"
proc = subprocess.run(
["git", "-c", "user.email=t@t", "-c", "user.name=t", "commit", "-qm", "hooked"],
cwd=repo,
capture_output=True,
text=True,
env=_env(repo, bin_dir, {"LITELLM_GATE_SLOT_DIR": str(lock_dir)}),
timeout=120,
)
assert proc.returncode == 0, proc.stdout + proc.stderr
assert (lock_dir / "slot-0.lock").exists()
def test_failing_run_ends_with_a_fail_verdict(tmp_path: Path) -> None:
repo, bin_dir = _sandbox(tmp_path)
proc = _run(repo, bin_dir, {"STUB_FAIL": "make-lint"})

View file

@ -29,7 +29,8 @@ from litellm.types.utils import (
StreamingChoices,
Usage,
)
from litellm.types.utils import all_litellm_params
from litellm.types.utils import all_litellm_params, bedrock_batch_litellm_params
from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams
from litellm.utils import (
ProviderConfigManager,
TextCompletionStreamWrapper,
@ -4752,3 +4753,45 @@ def test_websearch_interception_control_fields_never_reach_the_provider():
f"{sorted(set(non_default) - {'a_real_provider_specific_param'})}"
)
assert set(WEBSEARCH_INTERNAL_CONTROL_FIELDS) <= set(all_litellm_params)
def test_bedrock_batch_params_never_reach_the_provider():
"""A Bedrock managed-batch deployment carries aws_batch_role_arn / s3_* /
bedrock_tags in its litellm_params, and the same deployment also serves chat.
Anything the param builder does not recognize is swept into extra_body, so
Bedrock rejects the whole call: `aws_batch_role_arn: Extra inputs are not
permitted` (Anthropic models) or `extraneous key [aws_batch_role_arn] is not
permitted` (Nova/Llama/Titan), turning every non-batch request to that
deployment into a 400.
The batch path is unaffected by registering them, because GenericLiteLLMParams
is extra="allow" and preserves them into litellm_params for the batch and files
transformations that read them.
"""
configured = {
field: ([{"key": "team", "value": "configured-value"}] if field == "bedrock_tags" else "configured-value")
for field in bedrock_batch_litellm_params
}
kwargs = {"a_real_provider_specific_param": 1, **configured}
non_default = get_non_default_completion_params(dict(kwargs))
assert non_default == {"a_real_provider_specific_param": 1}, (
"bedrock batch params leaked into the provider params: "
f"{sorted(set(non_default) - {'a_real_provider_specific_param'})}"
)
assert set(bedrock_batch_litellm_params) <= set(all_litellm_params)
batch_params = dict(GenericLiteLLMParams(**kwargs))
assert all(batch_params.get(field) == configured[field] for field in bedrock_batch_litellm_params), (
"registering these must not strip them from the batch path: "
f"{sorted(f for f in bedrock_batch_litellm_params if batch_params.get(f) != configured[f])}"
)
normalized = CredentialLiteLLMParams.model_validate(
GenericLiteLLMParams(**kwargs).model_dump(exclude_none=True)
).model_dump(exclude_none=True)
assert all(normalized.get(field) == configured[field] for field in bedrock_batch_litellm_params), (
"credential normalization dropped batch params before the transformation: "
f"{sorted(f for f in bedrock_batch_litellm_params if normalized.get(f) != configured[f])}"
)

View file

@ -30,7 +30,7 @@
"limit": 16715
},
"LIT011": {
"limit": 5596
"limit": 5593
},
"LIT012": {
"limit": 4519

View file

@ -46,6 +46,7 @@ vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({
{ model_name: "gpt-auto", litellm_params: { model: "auto_router/gpt-auto" } },
],
})),
usePlainModelGroups: vi.fn(() => new Set(["prod-claude"])),
}));
vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({
@ -72,6 +73,8 @@ const job = (overrides: Partial<ShadowEvalJob> = {}): ShadowEvalJob => ({
job_id: "job-1",
status: "running",
router_name: "claude-auto",
direction: "forward",
baseline_model: null,
judge_model: "anthropic/claude-sonnet-5",
shadow_percentage: 10,
max_turns: 200,
@ -367,6 +370,58 @@ describe("ShadowEvalSection", () => {
expect(start.mutate).toHaveBeenCalledWith(expectedBody);
});
it("requires a baseline model in reverse mode and submits it, while forward mode never shows the picker", async () => {
const user = userEvent.setup();
const { start } = mockHooks({});
render(<ShadowEvalSection />);
expect(screen.queryByPlaceholderText("Select a baseline model")).not.toBeInTheDocument();
await user.click(screen.getByText("Adoption check: key's traffic vs the router"));
await user.click(await screen.findByText("Regression check: router's picks vs a baseline"));
await user.click(screen.getByPlaceholderText("Search keys by alias"));
await user.click(await screen.findByText("prod-alpha"));
await user.click(screen.getByPlaceholderText("Select an auto-router"));
await user.click(await screen.findByText("gpt-auto"));
await user.click(screen.getByPlaceholderText("Select a judge model"));
await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ }));
expect(screen.getByText("Start shadow eval")).toBeDisabled();
await user.click(screen.getByPlaceholderText("Select a baseline model"));
expect(await screen.findByRole("option", { name: /openai\/gpt-4o/ })).toBeInTheDocument();
await user.click(screen.getByRole("option", { name: /prod-claude/ }));
await user.click(screen.getByText("Start shadow eval"));
const expectedBody = {
api_key_id: "hash-alpha",
router_name: "gpt-auto",
direction: "reverse",
baseline_model: "prod-claude",
shadow_percentage: 10,
duration_days: 7,
max_turns: 200,
judge_model: "anthropic/claude-sonnet-5",
};
expect(start.mutate).toHaveBeenCalledWith(expectedBody);
});
it("flips the arm labels and headline for a reverse job's results", () => {
const j = job({ direction: "reverse", baseline_model: "openai/gpt-4o" });
mockHooks({ jobs: [j], detailsById: { "job-1": j } });
render(<ShadowEvalSection />);
expect(screen.getByText(/on 10% of its traffic/)).toBeInTheDocument();
expect(screen.getByText("Router matched or beat the baseline")).toBeInTheDocument();
expect(screen.getByText("52.0%")).toBeInTheDocument();
expect(screen.getByText(/Router won 30.0%/)).toBeInTheDocument();
expect(screen.getByText(/Baseline won 48.0%/)).toBeInTheDocument();
expect(screen.getAllByText("Baseline wins")).toHaveLength(2);
expect(screen.getByText("Router pick")).toBeInTheDocument();
expect(screen.queryByText(/Current model/)).not.toBeInTheDocument();
expect(screen.queryByText("Compared against")).not.toBeInTheDocument();
});
it("keeps an older job's verdicts reachable through the previous evaluations list", async () => {
const user = userEvent.setup();
const emptyOverrides: Partial<ShadowEvalJob> = {

View file

@ -5,7 +5,7 @@ import React, { useMemo, useState } from "react";
import { useInfiniteKeys } from "@/app/(dashboard)/hooks/keys/useKeys";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap";
import { useAutoRouters } from "@/app/(dashboard)/hooks/models/useModels";
import { useAutoRouters, usePlainModelGroups } from "@/app/(dashboard)/hooks/models/useModels";
import { PaginatedSearchSelect } from "@/components/shared/PaginatedSearchSelect";
import { SearchSelect, type SearchSelectOption } from "@/components/shared/SearchSelect";
import { Badge } from "@/components/ui/badge";
@ -31,6 +31,37 @@ const pct = (value: number): string => `${value.toFixed(1)}%`;
const MIN_TURNS_FOR_CONFIDENCE = 30;
type ShadowEvalDirection = ShadowEvalJob["direction"];
const otherArmLabel = (direction: ShadowEvalDirection): string =>
direction === "reverse" ? "Baseline" : "Current model";
const routerWinRate = (direction: ShadowEvalDirection, slice: ShadowEvalSlice): number =>
direction === "reverse" ? slice.real_win_rate_pct : slice.shadow_win_rate_pct;
const otherArmWinRate = (direction: ShadowEvalDirection, slice: ShadowEvalSlice): number =>
direction === "reverse" ? slice.shadow_win_rate_pct : slice.real_win_rate_pct;
const routerMatchedOrBeatPct = (
direction: ShadowEvalDirection,
results: NonNullable<ShadowEvalJob["results"]>,
): number =>
direction === "reverse"
? 100 - results.overall_shadow_win_rate_pct
: results.overall_shadow_win_rate_pct + results.overall_tie_rate_pct;
const jobHeadline = (job: ShadowEvalJob): React.ReactNode =>
job.direction === "reverse" ? (
<>
Comparing <span className="font-mono text-xs">{job.router_name}</span> to{" "}
<span className="font-mono text-xs">{job.baseline_model}</span> on {job.shadow_percentage}% of its traffic
</>
) : (
<>
Shadowing {job.shadow_percentage}% via <span className="font-mono text-xs">{job.router_name}</span>
</>
);
const isActive = (job: ShadowEvalJob): boolean => job.status === "running";
const endsIn = (endsAt: string | null | undefined): string | null => {
@ -54,16 +85,22 @@ const StatusBadge: React.FC<{ status: string }> = ({ status }) => (
</Badge>
);
const SliceTable: React.FC<{ groupHeader: string; slices: readonly ShadowEvalSlice[] }> = ({ groupHeader, slices }) => (
const SliceTable: React.FC<{
groupHeader: string;
direction: ShadowEvalDirection;
slices: readonly ShadowEvalSlice[];
}> = ({ groupHeader, direction, slices }) => (
<Table>
<TableHeader>
<TableRow>
<TableHead>{groupHeader}</TableHead>
{["Judged turns", "Router wins", "Current model wins", "Ties", "Judge confidence"].map((label) => (
<TableHead key={label} className="text-right">
{label}
</TableHead>
))}
{["Judged turns", "Router wins", `${otherArmLabel(direction)} wins`, "Ties", "Judge confidence"].map(
(label) => (
<TableHead key={label} className="text-right">
{label}
</TableHead>
),
)}
</TableRow>
</TableHeader>
<TableBody>
@ -77,9 +114,9 @@ const SliceTable: React.FC<{ groupHeader: string; slices: readonly ShadowEvalSli
</TableCell>
<TableCell className="text-right tabular-nums">{slice.turn_count.toLocaleString()}</TableCell>
<TableCell className="text-right font-medium tabular-nums text-foreground">
{pct(slice.shadow_win_rate_pct)}
{pct(routerWinRate(direction, slice))}
</TableCell>
<TableCell className="text-right tabular-nums">{pct(slice.real_win_rate_pct)}</TableCell>
<TableCell className="text-right tabular-nums">{pct(otherArmWinRate(direction, slice))}</TableCell>
<TableCell className="text-right tabular-nums">{pct(slice.tie_rate_pct)}</TableCell>
<TableCell className="text-right tabular-nums">{slice.avg_judge_confidence.toFixed(2)}</TableCell>
</TableRow>
@ -88,13 +125,23 @@ const SliceTable: React.FC<{ groupHeader: string; slices: readonly ShadowEvalSli
</Table>
);
const VerdictBar: React.FC<{ results: NonNullable<ShadowEvalJob["results"]> }> = ({ results }) => {
const routerWins = results.overall_shadow_win_rate_pct;
const VerdictBar: React.FC<{ direction: ShadowEvalDirection; results: NonNullable<ShadowEvalJob["results"]> }> = ({
direction,
results,
}) => {
const ties = results.overall_tie_rate_pct;
const routerWins =
direction === "reverse"
? Math.max(0, 100 - results.overall_shadow_win_rate_pct - ties)
: results.overall_shadow_win_rate_pct;
const segments = [
{ label: "Router won", value: routerWins, fill: "bg-emerald-500" },
{ label: "Tie", value: ties, fill: "bg-emerald-200" },
{ label: "Current model won", value: Math.max(0, 100 - routerWins - ties), fill: "bg-muted-foreground/30" },
{
label: `${otherArmLabel(direction)} won`,
value: Math.max(0, 100 - routerWins - ties),
fill: "bg-muted-foreground/30",
},
];
return (
<div className="space-y-2 border-b px-6 py-4">
@ -133,20 +180,22 @@ const ResultsBody: React.FC<{ job: ShadowEvalJob; resultsError?: boolean }> = ({
<>
<div className="flex flex-col gap-1 border-b px-6 py-4">
<p className="text-[11px] uppercase tracking-wide text-muted-foreground">
Router matched or beat your current model
</p>
<p className="text-3xl font-semibold text-foreground">
{pct(results.overall_shadow_win_rate_pct + results.overall_tie_rate_pct)}
Router matched or beat {job.direction === "reverse" ? "the baseline" : "your current model"}
</p>
<p className="text-3xl font-semibold text-foreground">{pct(routerMatchedOrBeatPct(job.direction, results))}</p>
<p className="text-xs text-muted-foreground">of {(job.judged_count ?? 0).toLocaleString()} judged responses</p>
</div>
<VerdictBar results={results} />
<VerdictBar direction={job.direction} results={results} />
{results.by_current_model.length > 0 && (
<SliceTable groupHeader="Compared against" slices={results.by_current_model} />
<SliceTable
groupHeader={job.direction === "reverse" ? "Router pick" : "Compared against"}
direction={job.direction}
slices={results.by_current_model}
/>
)}
{results.by_tier.length > 0 && (
<div className={results.by_current_model.length > 0 ? "border-t" : ""}>
<SliceTable groupHeader="Prompt difficulty" slices={results.by_tier} />
<SliceTable groupHeader="Prompt difficulty" direction={job.direction} slices={results.by_tier} />
</div>
)}
</>
@ -168,9 +217,7 @@ const JobResults: React.FC<{
<div className="flex items-center gap-3">
<StatusBadge status={job.status} />
<div>
<p className="text-sm font-medium text-foreground">
Shadowing {job.shadow_percentage}% via <span className="font-mono text-xs">{job.router_name}</span>
</p>
<p className="text-sm font-medium text-foreground">{jobHeadline(job)}</p>
<p className="text-xs text-muted-foreground">
{(job.judged_count ?? 0).toLocaleString()} of {job.max_turns.toLocaleString()} turns judged ·{" "}
{(job.error_count ?? 0).toLocaleString()} errored · {usd(job.judge_spend ?? 0)} judge spend
@ -201,25 +248,55 @@ interface CostMapEntry {
mode?: string;
}
const useJudgeModelOptions = (): SearchSelectOption[] => {
const useChatModelNames = (): string[] => {
const { data: costMap } = useModelCostMap();
return useMemo(() => {
if (!costMap) return [];
const chatModels = Object.entries(costMap as Record<string, CostMapEntry>)
.filter(([, value]) => value?.mode === "chat" && value?.litellm_provider)
.map(([key, value]) => (key.startsWith(`${value.litellm_provider}/`) ? key : `${value.litellm_provider}/${key}`));
return [...new Set(chatModels)].toSorted((a, b) => a.localeCompare(b));
}, [costMap]);
};
const useJudgeModelOptions = (): SearchSelectOption[] => {
const chatModels = useChatModelNames();
return useMemo(() => {
const pinned: SearchSelectOption[] = RECOMMENDED_JUDGE_MODELS.map((model) => ({
label: model,
value: model,
sublabel: "Recommended",
}));
if (!costMap) return pinned;
const pinnedNames = new Set<string>(RECOMMENDED_JUDGE_MODELS);
const chatModels = Object.entries(costMap as Record<string, CostMapEntry>)
.filter(([, value]) => value?.mode === "chat" && value?.litellm_provider)
.map(([key, value]) => (key.startsWith(`${value.litellm_provider}/`) ? key : `${value.litellm_provider}/${key}`));
const rest = [...new Set(chatModels)]
.filter((model) => !pinnedNames.has(model))
.toSorted((a, b) => a.localeCompare(b))
.map((model) => ({ label: model, value: model }));
const rest = chatModels.filter((model) => !pinnedNames.has(model)).map((model) => ({ label: model, value: model }));
return [...pinned, ...rest];
}, [costMap]);
}, [chatModels]);
};
const useBaselineModelOptions = (): SearchSelectOption[] => {
const configuredGroups = usePlainModelGroups();
const chatModels = useChatModelNames();
return useMemo(() => {
const configured = [...configuredGroups]
.toSorted((a, b) => a.localeCompare(b))
.map((model) => ({ label: model, value: model, sublabel: "Configured on this gateway" }));
const rest = chatModels
.filter((model) => !configuredGroups.has(model))
.map((model) => ({ label: model, value: model }));
return [...configured, ...rest];
}, [configuredGroups, chatModels]);
};
const DIRECTION_OPTIONS: readonly { value: ShadowEvalDirection; label: string }[] = [
{ value: "forward", label: "Adoption check: key's traffic vs the router" },
{ value: "reverse", label: "Regression check: router's picks vs a baseline" },
] as const;
const START_FORM_DESCRIPTION: Record<ShadowEvalDirection, string> = {
forward:
"Duplicates a sampled slice of the key's traffic through the auto-router and has an LLM judge compare both answers blind. The router's answers are never served to users; judge calls bill to the shadowed key.",
reverse:
"Duplicates a sampled slice of the traffic the auto-router already serves against a fixed baseline model and has an LLM judge compare both answers blind. The baseline's answers are never served to users; judge calls bill to the shadowed key.",
};
const DURATION_OPTIONS = [
@ -282,12 +359,15 @@ const StartForm: React.FC = () => {
const { accessToken } = useAuthorized();
const [apiKeyId, setApiKeyId] = useState("");
const [routerName, setRouterName] = useState("");
const [direction, setDirection] = useState<ShadowEvalDirection>("forward");
const [baselineModel, setBaselineModel] = useState("");
const [percentage, setPercentage] = useState("10");
const [durationDays, setDurationDays] = useState("7");
const [judgeModel, setJudgeModel] = useState("");
const [maxTurns, setMaxTurns] = useState("200");
const { data: autoRouters } = useAutoRouters();
const judgeModelOptions = useJudgeModelOptions();
const baselineModelOptions = useBaselineModelOptions();
const start = useStartShadowEval();
const routerOptions = useMemo<SearchSelectOption[]>(() => {
@ -301,14 +381,17 @@ const StartForm: React.FC = () => {
const percentageValid = parsedPct >= 0.1 && parsedPct <= 100;
const parsedMaxTurns = Number.parseInt(maxTurns, 10);
const maxTurnsValid = parsedMaxTurns >= 1 && parsedMaxTurns <= 2000;
const filled = [apiKeyId, routerName, judgeModel].every((field) => field !== "");
const filled =
[apiKeyId, routerName, judgeModel].every((field) => field !== "") &&
(direction === "forward" || baselineModel !== "");
const boundsValid = percentageValid && maxTurnsValid;
const valid = Boolean(accessToken) && filled && boundsValid;
const handleStart = () => {
const startBody = {
api_key_id: apiKeyId,
router_name: routerName,
direction: "forward" as const,
direction,
...(direction === "reverse" ? { baseline_model: baselineModel } : {}),
shadow_percentage: parsedPct,
duration_days: Number.parseInt(durationDays, 10),
max_turns: parsedMaxTurns,
@ -321,13 +404,27 @@ const StartForm: React.FC = () => {
<Card size="sm">
<CardHeader>
<CardTitle className="text-sm font-medium text-foreground">Start a shadow eval</CardTitle>
<p className="text-xs text-muted-foreground">
Duplicates a sampled slice of the key&apos;s traffic through the auto-router and has an LLM judge compare both
answers blind. The router&apos;s answers are never served to users; judge calls bill to the shadowed key.
</p>
<p className="text-xs text-muted-foreground">{START_FORM_DESCRIPTION[direction]}</p>
</CardHeader>
<CardContent className="space-y-3">
<div className="grid gap-3 sm:grid-cols-3">
<Field label="Direction">
<Select
value={direction}
onValueChange={(v: string | null) => setDirection(v === "reverse" ? "reverse" : "forward")}
>
<SelectTrigger className="w-full">
<SelectValue>{DIRECTION_OPTIONS.find((o) => o.value === direction)?.label}</SelectValue>
</SelectTrigger>
<SelectContent>
{DIRECTION_OPTIONS.map((option) => (
<SelectItem key={option.value} value={option.value}>
{option.label}
</SelectItem>
))}
</SelectContent>
</Select>
</Field>
<Field label="Key to shadow" htmlFor="shadow-eval-key">
<KeySelect value={apiKeyId} onChange={setApiKeyId} />
</Field>
@ -390,6 +487,17 @@ const StartForm: React.FC = () => {
<p className="text-xs text-destructive">Enter a value from 1 to 2000</p>
)}
</Field>
{direction === "reverse" && (
<Field label="Baseline model">
<SearchSelect
options={baselineModelOptions}
value={baselineModel}
onValueChange={setBaselineModel}
placeholder="Select a baseline model"
emptyText="No chat models available"
/>
</Field>
)}
<Field label="Judge model" className="sm:col-span-2">
<SearchSelect
options={judgeModelOptions}
@ -410,7 +518,7 @@ const StartForm: React.FC = () => {
const previousSummary = (job: ShadowEvalJob): string => {
const results = job.results;
if (results) return pct(results.overall_shadow_win_rate_pct + results.overall_tie_rate_pct);
if (results) return pct(routerMatchedOrBeatPct(job.direction, results));
return job.judged_count === 0 ? "no verdicts" : "view results";
};
@ -429,9 +537,7 @@ const PreviousJob: React.FC<{ job: ShadowEvalJob }> = ({ job }) => {
<div className="flex items-center gap-3">
<StatusBadge status={shown.status} />
<div>
<p className="text-sm font-medium text-foreground">
{shown.shadow_percentage}% via <span className="font-mono text-xs">{shown.router_name}</span>
</p>
<p className="text-sm font-medium text-foreground">{jobHeadline(shown)}</p>
<p className="text-xs text-muted-foreground">
{shown.judged_count != null &&
`${shown.judged_count.toLocaleString()} judged · ${(shown.error_count ?? 0).toLocaleString()} errored · ${usd(shown.judge_spend ?? 0)} judge spend · `}
@ -507,8 +613,8 @@ const ShadowEvalSection: React.FC = () => {
<div className="flex flex-wrap items-baseline gap-2">
<h2 className="text-xl font-semibold text-foreground">Shadow eval</h2>
<p className="text-sm text-muted-foreground">
Would the auto-router have answered as well as the models you use today? Find out on your real traffic, before
switching anything.
Blind-judge the auto-router on your real traffic: against the models a key uses today before switching, or
against a fixed baseline after it has switched.
</p>
</div>

View file

@ -5,6 +5,7 @@ import { beforeEach, describe, expect, it, vi } from "vitest";
import {
isAutoRouterDeployment,
selectAutoRouterModelGroups,
selectPlainModelGroups,
useAllProxyModels,
useAutoRouterModelGroups,
useAutoRouters,
@ -977,6 +978,32 @@ describe("selectAutoRouterModelGroups", () => {
});
});
describe("selectPlainModelGroups", () => {
it("keeps only non-auto-router model groups", () => {
const deployments: AutoRouterCandidateDeployment[] = [
{ model_name: "smart-router", litellm_params: { model: "auto_router/complexity_router" } },
{ model_name: "claude-haiku", litellm_params: { model: "anthropic/claude-haiku-4-5" } },
{ model_name: "claude-sonnet", litellm_params: { model: "anthropic/claude-sonnet-4-5" } },
{ model_name: "cheap-router", litellm_params: { model: "auto_router/adaptive_router" } },
];
expect(selectPlainModelGroups(deployments)).toEqual(new Set(["claude-haiku", "claude-sonnet"]));
});
it("drops a group name that also fronts an auto-router deployment", () => {
const deployments: AutoRouterCandidateDeployment[] = [
{ model_name: "shared-name", litellm_params: { model: "auto_router/complexity_router" } },
{ model_name: "shared-name", litellm_params: { model: "anthropic/claude-sonnet-4-5" } },
];
expect(selectPlainModelGroups(deployments)).toEqual(new Set());
});
it("drops deployments that have no public model_name", () => {
expect(selectPlainModelGroups([{ model_name: "", litellm_params: { model: "openai/gpt-4o" } }])).toEqual(new Set());
});
});
describe("useAutoRouterModelGroups", () => {
let queryClient: QueryClient;

View file

@ -123,6 +123,16 @@ export const selectAutoRouterModelGroups = (deployments: AutoRouterCandidateDepl
export const selectAutoRouterDeployments = (deployments: AutoRouterDeployment[]): AutoRouterDeployment[] =>
deployments.filter(isAutoRouterDeployment);
export const selectPlainModelGroups = (deployments: AutoRouterCandidateDeployment[]): ReadonlySet<string> => {
const autoRouterGroups = selectAutoRouterModelGroups(deployments);
return new Set(
deployments
.map((deployment) => deployment.model_name)
.filter((modelName): modelName is string => Boolean(modelName))
.filter((modelName) => !autoRouterGroups.has(modelName)),
);
};
export const fetchAllModelDeployments = async (
accessToken: string,
userId: string,
@ -172,6 +182,17 @@ export const useAutoRouterModelGroups = (): ReadonlySet<string> => {
return data ?? NO_AUTO_ROUTERS;
};
export const usePlainModelGroups = (): ReadonlySet<string> => {
const { accessToken, userId, userRole } = useAuthorized();
const { data } = useQuery<AutoRouterDeployment[], Error, ReadonlySet<string>>({
queryKey: autoRouterListKey(userId, userRole),
queryFn: async () => await fetchAllModelDeployments(accessToken!, userId!, userRole!),
enabled: Boolean(accessToken && userId && userRole),
select: selectPlainModelGroups,
});
return data ?? NO_AUTO_ROUTERS;
};
export const useAutoRouters = (): UseQueryResult<AutoRouterDeployment[], Error> => {
const { accessToken, userId, userRole } = useAuthorized();
return useQuery<AutoRouterDeployment[], Error, AutoRouterDeployment[]>({

View file

@ -27010,6 +27010,8 @@ export interface components {
aws_web_identity_token?: string | null;
/** Azure Ad Token */
azure_ad_token?: string | null;
/** Bedrock Tags */
bedrock_tags?: unknown[] | null;
/** Budget Duration */
budget_duration?: string | null;
/** Cache Creation Input Audio Token Cost */
@ -27235,6 +27237,8 @@ export interface components {
s3_bucket_name?: string | null;
/** S3 Encryption Key Id */
s3_encryption_key_id?: string | null;
/** S3 Output Bucket Name */
s3_output_bucket_name?: string | null;
/** S3 Region Name */
s3_region_name?: string | null;
/** Search Context Cost Per Query */
@ -35919,6 +35923,8 @@ export interface components {
aws_web_identity_token?: string | null;
/** Azure Ad Token */
azure_ad_token?: string | null;
/** Bedrock Tags */
bedrock_tags?: unknown[] | null;
/** Budget Duration */
budget_duration?: string | null;
/** Cache Creation Input Audio Token Cost */
@ -36144,6 +36150,8 @@ export interface components {
s3_bucket_name?: string | null;
/** S3 Encryption Key Id */
s3_encryption_key_id?: string | null;
/** S3 Output Bucket Name */
s3_output_bucket_name?: string | null;
/** S3 Region Name */
s3_region_name?: string | null;
/** Search Context Cost Per Query */