mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
bb1c3366cf
52 changed files with 2305 additions and 132 deletions
26
Makefile
26
Makefile
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"))})
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
173
scripts/gate_slot_lock.py
Normal 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())
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"]),
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
333
tests/test_litellm/test_gate_slot_lock.py
Normal file
333
tests/test_litellm/test_gate_slot_lock.py
Normal 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
|
||||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -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])}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@
|
|||
"limit": 16715
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5596
|
||||
"limit": 5593
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4519
|
||||
|
|
|
|||
|
|
@ -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> = {
|
||||
|
|
|
|||
|
|
@ -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'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.
|
||||
</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>
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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[]>({
|
||||
|
|
|
|||
8
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
8
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue