diff --git a/Makefile b/Makefile index 94d8c875af5..bdb643e3ec9 100644 --- a/Makefile +++ b/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: diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 7c88b78ad4e..25b00597355 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -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, ) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index baf522c0bb1..9681d64f656 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -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] diff --git a/litellm/batches/main.py b/litellm/batches/main.py index ce52c12818e..20d38bbb77f 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -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, diff --git a/litellm/files/main.py b/litellm/files/main.py index 34421d13761..9a64c78552b 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -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, ) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index f251ab4d74a..3eb8c163d5c 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -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 diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 10aea51c833..e64237da978 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -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, diff --git a/litellm/llms/fireworks_ai/common_utils.py b/litellm/llms/fireworks_ai/common_utils.py index 143dd151027..e07e7a26f9e 100644 --- a/litellm/llms/fireworks_ai/common_utils.py +++ b/litellm/llms/fireworks_ai/common_utils.py @@ -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 diff --git a/litellm/llms/fireworks_ai/completion/transformation.py b/litellm/llms/fireworks_ai/completion/transformation.py index f03baaddaf6..4f0e302003a 100644 --- a/litellm/llms/fireworks_ai/completion/transformation.py +++ b/litellm/llms/fireworks_ai/completion/transformation.py @@ -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, } diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index e419322dca6..df39b8fad48 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -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()): diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 95a3806e8ad..d13b39661ad 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 693e3f8e47d..86e97b55a8e 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index 4c1b78c754a..85885fc75f5 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4f94f94acb5..c782f0dfa09 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 84b40b72258..a30b5ee9e49 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -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 " diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py index 3ef7327cda3..15f5f82c4b6 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py @@ -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, ) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 934ac9ac3d7..a566d491597 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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"))}) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 554a286555b..6952c0c6f89 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index 3765771247d..7e641814deb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -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 diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 551ba63058d..b9af01e9aea 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -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, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index d2432ea3729..361b5b920e2 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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, diff --git a/litellm/types/router.py b/litellm/types/router.py index 217364c48b7..f3f9276e6ba 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 9baadc36f6b..272fbabf807 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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", diff --git a/scripts/gate_slot_lock.py b/scripts/gate_slot_lock.py new file mode 100644 index 00000000000..e7bd945ad65 --- /dev/null +++ b/scripts/gate_slot_lock.py @@ -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 [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 [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()) diff --git a/scripts/pre_commit_lint.sh b/scripts/pre_commit_lint.sh index 82498ec10cd..0861172056e 100755 --- a/scripts/pre_commit_lint.sh +++ b/scripts/pre_commit_lint.sh @@ -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 diff --git a/scripts/ruff_strict_gate.py b/scripts/ruff_strict_gate.py index 507077ddf25..bf070beeb0f 100644 --- a/scripts/ruff_strict_gate.py +++ b/scripts/ruff_strict_gate.py @@ -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__": diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py index c9f774c6113..763835e6d2e 100644 --- a/scripts/type_check_gate.py +++ b/scripts/type_check_gate.py @@ -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__": diff --git a/scripts/type_discipline_gate.py b/scripts/type_discipline_gate.py index f937283d972..5f6474f20bc 100644 --- a/scripts/type_discipline_gate.py +++ b/scripts/type_discipline_gate.py @@ -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__": diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index cf9c7c5c0c5..72f8b87dd16 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -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 diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 523b512e4cf..d2074853f2b 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -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 diff --git a/tests/test_litellm/batches/test_main.py b/tests/test_litellm/batches/test_main.py index 1f7a91a5511..17e9ee29d4d 100644 --- a/tests/test_litellm/batches/test_main.py +++ b/tests/test_litellm/batches/test_main.py @@ -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 diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 354f4656d6e..b46ba081f6f 100644 --- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -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( diff --git a/tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py b/tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py new file mode 100644 index 00000000000..996f1fd975b --- /dev/null +++ b/tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py @@ -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" diff --git a/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py b/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py new file mode 100644 index 00000000000..4af395baf41 --- /dev/null +++ b/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 52dc91ce24d..0209abee510 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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//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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index d8eab3ecb3b..cc65970a180 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -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") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py index 3e934577a66..f25d3baea0a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py @@ -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"]), diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 55cc8565a71..99181f0f087 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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"] diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 80d52de0012..824654dcf66 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -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. diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py index 17d4a3e304a..fc1465c7c14 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py @@ -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__": diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index 590f74d946f..6ffb7daaa2d 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -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"} diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index f27c8dfd2f4..e363a266688 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -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 diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index aaf1dad4910..8e973fc3771 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -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(): """ diff --git a/tests/test_litellm/test_gate_slot_lock.py b/tests/test_litellm/test_gate_slot_lock.py new file mode 100644 index 00000000000..17fa8547ce7 --- /dev/null +++ b/tests/test_litellm/test_gate_slot_lock.py @@ -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 diff --git a/tests/test_litellm/test_pre_commit_lint.py b/tests/test_litellm/test_pre_commit_lint.py index 33baf7474ce..5ea0e79a196 100644 --- a/tests/test_litellm/test_pre_commit_lint.py +++ b/tests/test_litellm/test_pre_commit_lint.py @@ -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"}) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 01e7e5c7ffd..661b6ed7244 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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])}" + ) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 7651b0125e7..8fbf9b291d5 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -30,7 +30,7 @@ "limit": 16715 }, "LIT011": { - "limit": 5596 + "limit": 5593 }, "LIT012": { "limit": 4519 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx index d4d26650086..e342ee33f25 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx @@ -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 => ({ 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(); + + 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(); + + 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 = { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx index 711fc1af539..df2989d3990 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx @@ -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, +): 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 {job.router_name} to{" "} + {job.baseline_model} on {job.shadow_percentage}% of its traffic + + ) : ( + <> + Shadowing {job.shadow_percentage}% via {job.router_name} + + ); + 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 }) => ( ); -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 }) => ( {groupHeader} - {["Judged turns", "Router wins", "Current model wins", "Ties", "Judge confidence"].map((label) => ( - - {label} - - ))} + {["Judged turns", "Router wins", `${otherArmLabel(direction)} wins`, "Ties", "Judge confidence"].map( + (label) => ( + + {label} + + ), + )} @@ -77,9 +114,9 @@ const SliceTable: React.FC<{ groupHeader: string; slices: readonly ShadowEvalSli {slice.turn_count.toLocaleString()} - {pct(slice.shadow_win_rate_pct)} + {pct(routerWinRate(direction, slice))} - {pct(slice.real_win_rate_pct)} + {pct(otherArmWinRate(direction, slice))} {pct(slice.tie_rate_pct)} {slice.avg_judge_confidence.toFixed(2)} @@ -88,13 +125,23 @@ const SliceTable: React.FC<{ groupHeader: string; slices: readonly ShadowEvalSli
); -const VerdictBar: React.FC<{ results: NonNullable }> = ({ results }) => { - const routerWins = results.overall_shadow_win_rate_pct; +const VerdictBar: React.FC<{ direction: ShadowEvalDirection; results: NonNullable }> = ({ + 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 (
@@ -133,20 +180,22 @@ const ResultsBody: React.FC<{ job: ShadowEvalJob; resultsError?: boolean }> = ({ <>

- Router matched or beat your current model -

-

- {pct(results.overall_shadow_win_rate_pct + results.overall_tie_rate_pct)} + Router matched or beat {job.direction === "reverse" ? "the baseline" : "your current model"}

+

{pct(routerMatchedOrBeatPct(job.direction, results))}

of {(job.judged_count ?? 0).toLocaleString()} judged responses

- + {results.by_current_model.length > 0 && ( - + )} {results.by_tier.length > 0 && (
0 ? "border-t" : ""}> - +
)} @@ -168,9 +217,7 @@ const JobResults: React.FC<{
-

- Shadowing {job.shadow_percentage}% via {job.router_name} -

+

{jobHeadline(job)}

{(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) + .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(RECOMMENDED_JUDGE_MODELS); - const chatModels = Object.entries(costMap as Record) - .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 = { + 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("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(() => { @@ -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 = () => { Start a shadow eval -

- 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. -

+

{START_FORM_DESCRIPTION[direction]}

+ + + @@ -390,6 +487,17 @@ const StartForm: React.FC = () => {

Enter a value from 1 to 2000

)} + {direction === "reverse" && ( + + + + )} { 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 }) => {
-

- {shown.shadow_percentage}% via {shown.router_name} -

+

{jobHeadline(shown)}

{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 = () => {

Shadow eval

- 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.

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts index f04c4b7bfcd..411e8402e11 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts @@ -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; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts index 68044a45edc..a5fbc433ea3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -123,6 +123,16 @@ export const selectAutoRouterModelGroups = (deployments: AutoRouterCandidateDepl export const selectAutoRouterDeployments = (deployments: AutoRouterDeployment[]): AutoRouterDeployment[] => deployments.filter(isAutoRouterDeployment); +export const selectPlainModelGroups = (deployments: AutoRouterCandidateDeployment[]): ReadonlySet => { + 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 => { return data ?? NO_AUTO_ROUTERS; }; +export const usePlainModelGroups = (): ReadonlySet => { + const { accessToken, userId, userRole } = useAuthorized(); + const { data } = useQuery>({ + 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 => { const { accessToken, userId, userRole } = useAuthorized(); return useQuery({ diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 0c128d9815b..cf37709c377 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */