From 2edfe2dc1029aa471c867ce2ffbde7ef6b128c1a Mon Sep 17 00:00:00 2001 From: moe-berri Date: Mon, 7 Sep 2026 16:07:11 -0700 Subject: [PATCH] feat(auto_router): mint a server-issued capability token for shunt worker auth Replaces the ANTHROPIC_AUTH_TOKEN/ANTHROPIC_API_KEY environment-variable read with a short-lived token the proxy mints itself, so the generated command authenticates to /v1/bulk_read and /v1/code_write without depending on the calling client's shell holding either variable. That dependency only ever held for Claude Code; any other client (Cursor, a custom agent) would have sent an empty bearer and 401'd. The token is a sealed grant (litellm/proxy/guardrails/shunt_capability_token.py) built on the proxy's own encrypt_value_helper, the same primitive the gateway's OAuth flow already seals values with. It carries a reference to the caller's key hash rather than the key itself, and a two-minute expiry rather than a single-use guard: an agent retrying a timed-out Bash command must still authenticate, and a single-use claim would turn that ordinary retry into a permanent 401. The worst a replay inside the window can do is spend the caller's own already- budgeted quota on a request they already made. Carried in the generated command's Authorization header, never a URL query string: every other sealed token in this proxy already avoids query strings, since they routinely end up in access logs. /v1/bulk_read and /v1/code_write now authenticate exclusively via this token instead of the normal user_api_key_auth path, since nothing but a shunt-generated command should ever call them. Master-key callers (UserAPIKeyAuth.api_key holds a stable alias rather than a DB-backed hash for that case) carry the real master key in the grant instead, compared directly at the endpoint. --- litellm/proxy/guardrails/auto_router_shunt.py | 40 ++++++- .../guardrails/shunt_capability_token.py | 81 ++++++++++++++ litellm/proxy/guardrails/shunt_rewrite.py | 32 +++--- litellm/proxy/shunt_endpoints/endpoints.py | 60 +++++++++- .../guardrails/test_auto_router_shunt.py | 105 ++++++++++++++++-- .../guardrails/test_shunt_capability_token.py | 89 +++++++++++++++ .../proxy/guardrails/test_shunt_rewrite.py | 77 +++++++------ .../proxy/shunt_endpoints/test_endpoints.py | 100 +++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 16 ++- 9 files changed, 525 insertions(+), 75 deletions(-) create mode 100644 litellm/proxy/guardrails/shunt_capability_token.py create mode 100644 tests/test_litellm/proxy/guardrails/test_shunt_capability_token.py create mode 100644 tests/test_litellm/proxy/shunt_endpoints/test_endpoints.py diff --git a/litellm/proxy/guardrails/auto_router_shunt.py b/litellm/proxy/guardrails/auto_router_shunt.py index 356bb64f98a..5e72f03ffd6 100644 --- a/litellm/proxy/guardrails/auto_router_shunt.py +++ b/litellm/proxy/guardrails/auto_router_shunt.py @@ -243,15 +243,38 @@ DEFAULT_BULK_READ_QUESTION: Final = "Summarize this file's exports and overall s class _ShuntEndpoints: bulk_read_url: str code_write_url: str + capability_token: str -def _endpoints_for_request(data: Mapping[str, object], model_alias: str) -> "_ShuntEndpoints | None": +def _mint_caller_capability_token(user_api_key_dict: "UserAPIKeyAuth") -> str: + """Seal a short-lived grant identifying this request's caller. + + ``UserAPIKeyAuth.api_key`` is already the hashed token for a DB-backed virtual key + (`_safe_hash_litellm_api_key` on the model itself), so the common case just carries that + hash forward. Master-key auth is the one caller with no such row: it stores a stable alias + there instead (`LITELLM_PROXY_MASTER_KEY_ALIAS`), so that case carries the real master key, + itself sealed rather than embedded in the clear, for the worker endpoint to compare directly. + """ + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + from litellm.proxy.guardrails.shunt_capability_token import mint_shunt_capability_token + from litellm.proxy.proxy_server import master_key + + if user_api_key_dict.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS and master_key is not None: + return mint_shunt_capability_token(key_hash=None, master_key=master_key) + return mint_shunt_capability_token(key_hash=user_api_key_dict.api_key, master_key=None) + + +def _endpoints_for_request( + data: Mapping[str, object], model_alias: str, user_api_key_dict: "UserAPIKeyAuth" +) -> "_ShuntEndpoints | None": """Where this request's generated curl commands should point, or None if unreachable. None when the base URL can't be recovered, since a generated command could then not reach this proxy at all; the tool_use is left unmodified rather than shipped broken. The caller's - credential is deliberately not read here: the generated command picks it up from the - client's own environment at run time instead, so it never enters the model's response. + own credential never appears in the generated command: instead a short-lived capability + token identifying the caller is minted here and carried in the command's `Authorization` + header, so the worker call authenticates without the real key ever entering the model's + response, the conversation history, or (unlike a query-string token) an access log line. The request's own tags ride along in the query string, because the worker endpoints resolve the marker again from scratch: a marker armed only under a tag would otherwise be invisible @@ -270,6 +293,7 @@ def _endpoints_for_request(data: Mapping[str, object], model_alias: str) -> "_Sh return _ShuntEndpoints( bulk_read_url=f"{base_url}/v1/bulk_read?{query}", code_write_url=f"{base_url}/v1/code_write?{query}", + capability_token=_mint_caller_capability_token(user_api_key_dict), ) @@ -306,6 +330,7 @@ def _bash_replacement_for_tool_use( question=question, paths=paths, bulk_read_endpoint=endpoints.bulk_read_url, + capability_token=endpoints.capability_token, ) if name == CODE_WRITE_TOOL_NAME: @@ -319,6 +344,7 @@ def _bash_replacement_for_tool_use( reference=reference, target=target if isinstance(target, str) and target else None, code_write_endpoint=endpoints.code_write_url, + capability_token=endpoints.capability_token, ) if name == "Read": @@ -332,6 +358,7 @@ def _bash_replacement_for_tool_use( question=DEFAULT_BULK_READ_QUESTION, min_lines=config.min_lines, bulk_read_endpoint=endpoints.bulk_read_url, + capability_token=endpoints.capability_token, ) if name == "Bash": @@ -346,6 +373,7 @@ def _bash_replacement_for_tool_use( question=DEFAULT_BULK_READ_QUESTION, min_lines=config.min_lines, bulk_read_endpoint=endpoints.bulk_read_url, + capability_token=endpoints.capability_token, ) return None @@ -521,7 +549,9 @@ class ShuntGuardrail(CustomLogger): return response model: Final = data.get("model") - endpoints: Final = _endpoints_for_request(data, model) if isinstance(model, str) and model else None + endpoints: Final = ( + _endpoints_for_request(data, model, user_api_key_dict) if isinstance(model, str) and model else None + ) if endpoints is None: return response @@ -573,7 +603,7 @@ class ShuntGuardrail(CustomLogger): config: Final = None if request_data.get(_CALLER_OWNS_TOOL_NAME_KEY) else _resolve_shunt_config(request_data) model: Final = request_data.get("model") endpoints: Final = ( - _endpoints_for_request(request_data, model) + _endpoints_for_request(request_data, model, user_api_key_dict) if config is not None and isinstance(model, str) and model else None ) diff --git a/litellm/proxy/guardrails/shunt_capability_token.py b/litellm/proxy/guardrails/shunt_capability_token.py new file mode 100644 index 00000000000..8383aeb0f75 --- /dev/null +++ b/litellm/proxy/guardrails/shunt_capability_token.py @@ -0,0 +1,81 @@ +""" +Mints and opens the short-lived credential a shunt-generated command carries to authenticate +its own call back into `/v1/bulk_read` / `/v1/code_write`. + +Embedding the caller's real API key in the generated command (an earlier version of this) put +that key in the model's response and the conversation history. Reading an Anthropic-flavored +env var off the client's shell (a later version) only ever worked for Claude Code, since no +other client has a reason to export `ANTHROPIC_AUTH_TOKEN`/`ANTHROPIC_API_KEY`. This module +mints, at rewrite time, a token that identifies the caller by a reference to their existing key +rather than the key itself, sealed with the proxy's own encryption helper (the same primitive +`gateway_dcr_flow.py`'s `_seal`/`_open_sealed` use) and given a short expiry. The generated +command carries it in an `Authorization` header, the same transport every other sealed token in +this proxy already uses -- never a URL query string, which routinely ends up in access logs. + +Deliberately not single-use: an agent retries a timed-out or interrupted Bash command, and a +single-use guard would turn that ordinary retry into a permanent 401. The token's only defense +is its short TTL; the worst a replay within that window can do is spend the caller's own budget +on a request they already made. +""" + +import time +from typing import Final + +from pydantic import BaseModel, ConfigDict, Field, ValidationError + +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper + +_TOKEN_PREFIX: Final = "shunt_cap_v1:" +TOKEN_TTL_SECONDS: Final = 120 + + +class ShuntCapabilityGrant(BaseModel): + """What the sealed token attests: the caller to bill this call to, and until when it's valid. + + ``key_hash`` is the same hashed value already stored as ``UserAPIKeyAuth.api_key`` for a + DB-backed key, so opening the token is a plain ``get_key_object`` lookup, not a new identity + scheme. ``is_master_key`` covers the one caller with no such row: master-key auth stores a + stable alias there instead (see ``LITELLM_PROXY_MASTER_KEY_ALIAS``), so the grant carries the + real master key itself to hand back to the worker call, sealed the same as everything else + here rather than embedded in the clear. + """ + + model_config = ConfigDict(frozen=True, extra="forbid") + key_hash: str | None = None + master_key: str | None = None + exp: int = Field(gt=0) + + +def mint_shunt_capability_token(*, key_hash: str | None, master_key: str | None, now: float | None = None) -> str: + """Seal a grant for the caller identified by exactly one of `key_hash` or `master_key`.""" + if (key_hash is None) == (master_key is None): + raise ValueError("mint_shunt_capability_token requires exactly one of key_hash or master_key") + grant: Final = ShuntCapabilityGrant( + key_hash=key_hash, master_key=master_key, exp=int((now if now is not None else time.time()) + TOKEN_TTL_SECONDS) + ) + return _TOKEN_PREFIX + encrypt_value_helper(grant.model_dump_json(exclude_none=True)) + + +def open_shunt_capability_token(token: str, *, now: float | None = None) -> ShuntCapabilityGrant | None: + """The grant a token carries, or None if it is malformed, unparseable, or expired. + + Total: every failure mode (wrong prefix, decrypt failure, schema mismatch, past expiry) + returns None rather than raising, so the caller has one branch to handle -- reject the + request -- instead of distinguishing why the token didn't validate. + """ + if not token.startswith(_TOKEN_PREFIX): + return None + decrypted: Final = decrypt_value_helper( + token[len(_TOKEN_PREFIX) :], "shunt_capability_token", return_original_value=False + ) + if not isinstance(decrypted, str): + return None + try: + grant: Final = ShuntCapabilityGrant.model_validate_json(decrypted) + except ValidationError: + return None + if (grant.key_hash is None) == (grant.master_key is None): + return None + if grant.exp < (now if now is not None else time.time()): + return None + return grant diff --git a/litellm/proxy/guardrails/shunt_rewrite.py b/litellm/proxy/guardrails/shunt_rewrite.py index 58b37a75977..203ec32f9c4 100644 --- a/litellm/proxy/guardrails/shunt_rewrite.py +++ b/litellm/proxy/guardrails/shunt_rewrite.py @@ -75,18 +75,19 @@ class ShuntBashRewrite: note: str -# The generated command reads the client's own credential out of its environment at run time -# instead of carrying it. Embedding the value would copy the caller's key into the model's -# response, the conversation history, and the next upstream turn, which is exactly what keeping -# it in `secret_fields` is meant to prevent. `${VAR:-$OTHER}` also covers both header styles: -# Claude Code sets ANTHROPIC_AUTH_TOKEN, while an x-api-key client sets ANTHROPIC_API_KEY. -_AUTH_ENV_EXPR: Final = "${ANTHROPIC_AUTH_TOKEN:-$ANTHROPIC_API_KEY}" -# Deliberately double-quoted, not shlex.quote'd: this is shell syntax to evaluate, not data. -_AUTH_FLAG: Final = f'-H "Authorization: Bearer {_AUTH_ENV_EXPR}"' +def _auth_flag(capability_token: str) -> str: + """The `-H` flag carrying the caller's short-lived capability token. + + Never the caller's real key: that would copy it into the model's response and the + conversation history, exactly what keeping it in `secret_fields` is meant to prevent. The + token is minted per request (see `auto_router_shunt.py`'s `_mint_caller_capability_token`) + and expires in minutes, so a copy left in a stale transcript is worthless shortly after. + """ + return f"-H {shlex.quote(f'Authorization: Bearer {capability_token}')}" def build_bounded_read_command( - *, path: str, question: str, min_lines: int, bulk_read_endpoint: str + *, path: str, question: str, min_lines: int, bulk_read_endpoint: str, capability_token: str ) -> ShuntBashRewrite: """The shunt conditional: read small files directly, delegate large ones. @@ -107,7 +108,7 @@ def build_bounded_read_command( f"printf '[shunt] %s: %s lines, bounded read delegated\\n' {quoted_path} \"$L\" >&2; " f"curl -sS -F {shlex.quote(f'question={question}')} " f"-F {shlex.quote(f'paths=@{path}')} " - f"{_AUTH_FLAG} {shlex.quote(bulk_read_endpoint)}; " + f"{_auth_flag(capability_token)} {shlex.quote(bulk_read_endpoint)}; " f"else cat {quoted_path}; fi" ) return ShuntBashRewrite( @@ -116,7 +117,9 @@ def build_bounded_read_command( ) -def build_bulk_read_command(*, question: str, paths: Sequence[str], bulk_read_endpoint: str) -> ShuntBashRewrite: +def build_bulk_read_command( + *, question: str, paths: Sequence[str], bulk_read_endpoint: str, capability_token: str +) -> ShuntBashRewrite: """The curl a model's own explicit `bulk_read(question, paths)` tool call becomes. Unconditional (no size check): the model chose to delegate, unlike the automatic bounding @@ -124,13 +127,14 @@ def build_bulk_read_command(*, question: str, paths: Sequence[str], bulk_read_en """ path_flags: Final = " ".join(f"-F {shlex.quote(f'paths=@{path}')}" for path in paths) command: Final = ( - f"curl -sS -F {shlex.quote(f'question={question}')} {path_flags} {_AUTH_FLAG} {shlex.quote(bulk_read_endpoint)}" + f"curl -sS -F {shlex.quote(f'question={question}')} {path_flags} " + f"{_auth_flag(capability_token)} {shlex.quote(bulk_read_endpoint)}" ) return ShuntBashRewrite(command=command, note="Delegated to a cheaper model via bulk_read.") def build_code_write_command( - *, spec: str, reference: str, target: str | None, code_write_endpoint: str + *, spec: str, reference: str, target: str | None, code_write_endpoint: str, capability_token: str ) -> ShuntBashRewrite: """The curl a model's own explicit `code_write(spec, reference, target)` tool call becomes. @@ -142,7 +146,7 @@ def build_code_write_command( request: Final = ( f"curl -sS -F {shlex.quote(f'spec={spec}')} " f"-F {shlex.quote(f'reference=@{reference}')} " - f"{_AUTH_FLAG} {shlex.quote(code_write_endpoint)}" + f"{_auth_flag(capability_token)} {shlex.quote(code_write_endpoint)}" ) if target is None: return ShuntBashRewrite(command=request, note="Delegated to a cheaper model via code_write.") diff --git a/litellm/proxy/shunt_endpoints/endpoints.py b/litellm/proxy/shunt_endpoints/endpoints.py index f6941465fe8..48f7f019897 100644 --- a/litellm/proxy/shunt_endpoints/endpoints.py +++ b/litellm/proxy/shunt_endpoints/endpoints.py @@ -8,6 +8,11 @@ marker's `auto_router_shunt_bulk_read_model` / `auto_router_shunt_code_write_mod model chosen by the caller. The call goes through `llm_router.acompletion`, not `litellm.acompletion` directly, so worker-model spend is tracked and budgeted against the caller's key/team exactly like any other request. + +Auth is the short-lived capability token minted at rewrite time +(`shunt_capability_token.py`), not a normal virtual key: these routes exist only to be hit by +a shunt-generated command, never called directly, so `user_api_key_auth`'s full DB-backed path +is the wrong tool here and the token is the only credential accepted. """ from collections.abc import Sequence @@ -15,12 +20,12 @@ from enum import Enum from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Final -from fastapi import APIRouter, Depends, File, Form, HTTPException, Query, Request, UploadFile +from fastapi import APIRouter, Depends, File, Form, Header, HTTPException, Query, Request, UploadFile from litellm._logging import verbose_proxy_logger -from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy._types import LitellmUserRoles, ProxyErrorTypes, ProxyException, UserAPIKeyAuth from litellm.proxy.guardrails.auto_router_shunt import ShuntConfig, shunt_config_for_model +from litellm.proxy.guardrails.shunt_capability_token import open_shunt_capability_token from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.shunt_endpoints.worker import ( BULK_READ_SYSTEM_PROMPT, @@ -39,9 +44,52 @@ router: Final = APIRouter() # Both endpoints are registered twice, at the `/v1`-prefixed path the generated commands call # and at the bare path, matching how the rest of the proxy exposes its native routes. -_AUTH_DEPENDENCIES: Final = [Depends(user_api_key_auth)] # mutable-ok: FastAPI's `dependencies=` takes a list _SHUNT_TAGS: Final[list[str | Enum]] = ["shunt"] # mutable-ok: FastAPI's `tags=` takes an invariant list +_BEARER_PREFIX: Final = "Bearer " + + +async def _caller_from_capability_token(authorization: Annotated[str | None, Header()] = None) -> UserAPIKeyAuth: + """Resolve the request's caller from its shunt capability token, or reject the request. + + Never falls through to the proxy's own key/DB lookup: a request that reaches these routes + without a valid token is rejected outright, since a shunt-generated command is the only + thing that should ever call them. + """ + if authorization is None or not authorization.startswith(_BEARER_PREFIX): + raise HTTPException(status_code=401, detail="Missing or malformed Authorization header") + grant: Final = open_shunt_capability_token(authorization[len(_BEARER_PREFIX) :]) + if grant is None: + raise HTTPException(status_code=401, detail="Invalid or expired shunt capability token") + + if grant.master_key is not None: + from litellm.proxy.proxy_server import master_key + + if master_key is None or grant.master_key != master_key: + raise HTTPException(status_code=401, detail="Invalid or expired shunt capability token") + return UserAPIKeyAuth(api_key=grant.master_key, user_role=LitellmUserRoles.PROXY_ADMIN) + + if grant.key_hash is None: + raise HTTPException(status_code=401, detail="Invalid or expired shunt capability token") + + from litellm.proxy.auth.auth_checks import get_key_object + from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache + + try: + return await get_key_object( + hashed_token=grant.key_hash, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + raise HTTPException(status_code=401, detail="Invalid or expired shunt capability token") from e + + +_AUTH_DEPENDENCIES: Final = [ + Depends(_caller_from_capability_token) +] # mutable-ok: FastAPI's `dependencies=` takes a list + def _worker_config( model_alias: str, user_api_key_dict: UserAPIKeyAuth, request_tags: Sequence[str] @@ -137,7 +185,7 @@ async def bulk_read( router_name: Annotated[str, Query(alias="router")], question: Annotated[str, Form()], paths: Annotated[list[UploadFile], File()], # mutable-ok: FastAPI requires a list for a repeated file field - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(_caller_from_capability_token)], tags: Annotated[list[str] | None, Query()] = None, ) -> str: """Summarize or answer a question about one or more files via a cheap worker model. @@ -176,7 +224,7 @@ async def code_write( router_name: Annotated[str, Query(alias="router")], spec: Annotated[str, Form()], reference: Annotated[UploadFile, File()], - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(_caller_from_capability_token)], tags: Annotated[list[str] | None, Query()] = None, ) -> str: """Generate boilerplate code matching a reference file's patterns, via a cheap worker model. diff --git a/tests/test_litellm/proxy/guardrails/test_auto_router_shunt.py b/tests/test_litellm/proxy/guardrails/test_auto_router_shunt.py index e3575256364..24004448f20 100644 --- a/tests/test_litellm/proxy/guardrails/test_auto_router_shunt.py +++ b/tests/test_litellm/proxy/guardrails/test_auto_router_shunt.py @@ -12,8 +12,21 @@ from litellm.proxy.guardrails.auto_router_shunt import ( ShuntGuardrail, shunt_config_for_model, ) +from litellm.proxy._types import UserAPIKeyAuth from litellm.types.utils import ChatCompletionMessageToolCall, Choices, Function, Message, ModelResponse +# A real UserAPIKeyAuth is required once a request actually reaches the rewrite path: it mints a +# capability token identifying the caller (auto_router_shunt.py's _mint_caller_capability_token), +# which needs a real api_key hash to seal. `None` still works for every test that stays on the +# unarmed/unchanged path, since that path returns before ever touching user_api_key_dict. +_FAKE_USER_API_KEY_DICT = UserAPIKeyAuth(api_key="fakehash1234567890") + + +@pytest.fixture(autouse=True) +def _salt_key(monkeypatch): + """Minting a capability token needs a signing key; see shunt_capability_token.py.""" + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234-test-salt-key") + class _FakeRouter: """Minimal stand-in for litellm.Router.get_model_list, mirroring @@ -283,7 +296,7 @@ class TestAsyncPostCallSuccessHookAnthropicShape: "content": [{"type": "tool_use", "id": "t1", "name": "Read", "input": {"file_path": "litellm/router.py"}}] } result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) block = result["content"][0] assert block["name"] == "Bash" @@ -302,7 +315,7 @@ class TestAsyncPostCallSuccessHookAnthropicShape: ] } result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) assert result["content"][0]["name"] == "Read" @@ -338,7 +351,7 @@ class TestAsyncPostCallSuccessHookAnthropicShape: ] } result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) block = result["content"][0] assert block["name"] == "Bash" @@ -362,7 +375,7 @@ class TestAsyncPostCallSuccessHookAnthropicShape: ] } result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) block = result["content"][0] assert block["name"] == "Bash" @@ -380,7 +393,7 @@ class TestAsyncPostCallSuccessHookAnthropicShape: ] } result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) assert "wc -l" in result["content"][0]["input"]["command"] @@ -393,7 +406,7 @@ class TestAsyncPostCallSuccessHookAnthropicShape: original_command = "cat litellm/router.py | grep foo" response = {"content": [{"type": "tool_use", "id": "t1", "name": "Bash", "input": {"command": original_command}}]} result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) assert result["content"][0]["input"]["command"] == original_command @@ -405,12 +418,86 @@ class TestAsyncPostCallSuccessHookAnthropicShape: guardrail = mod.ShuntGuardrail() response = {"content": [{"type": "text", "text": "hello"}]} result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) assert result["content"][0]["type"] == "text" assert result["content"][0]["text"] == "hello" +# Regression: the generated command used to embed the caller's raw Authorization header, which +# put the real key in the model's response and conversation history. It now carries a sealed, +# short-lived capability token that identifies the caller by reference instead. +class TestRewriteNeverCarriesTheCallersRealCredential: + def _config(self) -> ShuntConfig: + return ShuntConfig(min_lines=350, bulk_read_model="claude-haiku-4-5", code_write_model="claude-haiku-4-5") + + def _armed_request_data(self) -> dict: + return { + "model": "shunt", + "proxy_server_request": {"url": "http://localhost:4000/v1/messages"}, + # A real caller's Authorization header may still be present on the request (secret_ + # fields is populated regardless of shunt), but the rewrite must never read it now. + "secret_fields": {"raw_headers": {"authorization": "Bearer sk-the-callers-real-key"}}, + } + + @pytest.mark.asyncio + async def test_rewritten_command_never_contains_the_callers_real_key(self, monkeypatch): + import litellm.proxy.guardrails.auto_router_shunt as mod + + monkeypatch.setattr(mod, "_resolve_shunt_config", lambda data: self._config()) + guardrail = mod.ShuntGuardrail() + real_key_holder = UserAPIKeyAuth(api_key="fakehash1234567890") + response = { + "content": [{"type": "tool_use", "id": "t1", "name": "Read", "input": {"file_path": "litellm/router.py"}}] + } + result = await guardrail.async_post_call_success_hook( + data=self._armed_request_data(), user_api_key_dict=real_key_holder, response=response + ) + command = result["content"][0]["input"]["command"] + assert "sk-the-callers-real-key" not in command + assert "shunt_cap_v1:" in command + + @pytest.mark.asyncio + async def test_master_key_caller_gets_a_token_too(self, monkeypatch): + """A master-key caller has no DB-backed key hash (LITELLM_PROXY_MASTER_KEY_ALIAS instead + of a real hash), so the mint path must handle it without raising.""" + import litellm.proxy.guardrails.auto_router_shunt as mod + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + + monkeypatch.setattr(mod, "_resolve_shunt_config", lambda data: self._config()) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-the-real-master-key") + guardrail = mod.ShuntGuardrail() + master_key_holder = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS) + response = { + "content": [{"type": "tool_use", "id": "t1", "name": "Read", "input": {"file_path": "litellm/router.py"}}] + } + result = await guardrail.async_post_call_success_hook( + data=self._armed_request_data(), user_api_key_dict=master_key_holder, response=response + ) + command = result["content"][0]["input"]["command"] + assert "sk-the-real-master-key" not in command + assert "shunt_cap_v1:" in command + + @pytest.mark.asyncio + async def test_token_is_carried_in_the_authorization_header_not_a_query_string(self, monkeypatch): + import litellm.proxy.guardrails.auto_router_shunt as mod + + monkeypatch.setattr(mod, "_resolve_shunt_config", lambda data: self._config()) + guardrail = mod.ShuntGuardrail() + holder = UserAPIKeyAuth(api_key="fakehash1234567890") + response = { + "content": [{"type": "tool_use", "id": "t1", "name": "Read", "input": {"file_path": "litellm/router.py"}}] + } + result = await guardrail.async_post_call_success_hook( + data=self._armed_request_data(), user_api_key_dict=holder, response=response + ) + command = result["content"][0]["input"]["command"] + assert "-H " in command + before_header, _, after_header = command.partition("-H ") + assert "shunt_cap_v1:" not in before_header + assert "shunt_cap_v1:" in after_header + + class TestAsyncPostCallSuccessHookOpenAIShape: def _config(self) -> ShuntConfig: return ShuntConfig(min_lines=350, bulk_read_model="claude-haiku-4-5", code_write_model="claude-haiku-4-5") @@ -435,7 +522,7 @@ class TestAsyncPostCallSuccessHookOpenAIShape: choices=[Choices(index=0, message=Message(role="assistant", tool_calls=[tool_call]))] ) result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) rewritten = result.choices[0].message.tool_calls[0] assert rewritten.function.name == "Bash" @@ -450,6 +537,6 @@ class TestAsyncPostCallSuccessHookOpenAIShape: guardrail = mod.ShuntGuardrail() response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="hi"))]) result = await guardrail.async_post_call_success_hook( - data=self._armed_request_data(), user_api_key_dict=None, response=response + data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) assert result.choices[0].message.content == "hi" diff --git a/tests/test_litellm/proxy/guardrails/test_shunt_capability_token.py b/tests/test_litellm/proxy/guardrails/test_shunt_capability_token.py new file mode 100644 index 00000000000..74e593869c5 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/test_shunt_capability_token.py @@ -0,0 +1,89 @@ +"""Unit tests for litellm.proxy.guardrails.shunt_capability_token.""" + +import pytest + +from litellm.proxy.guardrails.shunt_capability_token import ( + TOKEN_TTL_SECONDS, + mint_shunt_capability_token, + open_shunt_capability_token, +) + + +@pytest.fixture(autouse=True) +def _salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234-test-salt-key") + + +class TestMintRequiresExactlyOneIdentity: + def test_neither_field_raises(self): + with pytest.raises(ValueError, match="exactly one of key_hash or master_key"): + mint_shunt_capability_token(key_hash=None, master_key=None) + + def test_both_fields_raises(self): + with pytest.raises(ValueError, match="exactly one of key_hash or master_key"): + mint_shunt_capability_token(key_hash="abc", master_key="sk-1234") + + +class TestRoundTrip: + def test_key_hash_round_trips(self): + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None) + grant = open_shunt_capability_token(token) + assert grant is not None + assert grant.key_hash == "deadbeef" + assert grant.master_key is None + + def test_master_key_round_trips(self): + token = mint_shunt_capability_token(key_hash=None, master_key="sk-real-master-key") + grant = open_shunt_capability_token(token) + assert grant is not None + assert grant.master_key == "sk-real-master-key" + assert grant.key_hash is None + + def test_token_carries_the_shunt_prefix(self): + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None) + assert token.startswith("shunt_cap_v1:") + + def test_token_never_contains_the_raw_master_key_in_plaintext(self): + token = mint_shunt_capability_token(key_hash=None, master_key="sk-real-master-key") + assert "sk-real-master-key" not in token + + +class TestExpiry: + def test_fresh_token_is_valid(self): + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None, now=1_000_000) + assert open_shunt_capability_token(token, now=1_000_000) is not None + + def test_token_valid_just_before_expiry(self): + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None, now=1_000_000) + assert open_shunt_capability_token(token, now=1_000_000 + TOKEN_TTL_SECONDS - 1) is not None + + def test_token_expired_after_ttl(self): + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None, now=1_000_000) + assert open_shunt_capability_token(token, now=1_000_000 + TOKEN_TTL_SECONDS + 1) is None + + def test_replay_within_ttl_still_opens(self): + """Deliberately not single-use: a retried Bash command must still authenticate.""" + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None, now=1_000_000) + first = open_shunt_capability_token(token, now=1_000_005) + second = open_shunt_capability_token(token, now=1_000_010) + assert first is not None + assert second is not None + assert first.key_hash == second.key_hash + + +class TestMalformedInput: + def test_wrong_prefix_returns_none(self): + assert open_shunt_capability_token("not-a-shunt-token") is None + + def test_empty_string_returns_none(self): + assert open_shunt_capability_token("") is None + + def test_prefix_with_garbage_payload_returns_none(self): + assert open_shunt_capability_token("shunt_cap_v1:not-valid-ciphertext") is None + + def test_a_sealed_but_differently_shaped_payload_returns_none(self): + """Cross-type confusion: another sealed value's ciphertext must not parse as a grant.""" + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + foreign = "shunt_cap_v1:" + encrypt_value_helper('{"totally": "unrelated"}') + assert open_shunt_capability_token(foreign) is None diff --git a/tests/test_litellm/proxy/guardrails/test_shunt_rewrite.py b/tests/test_litellm/proxy/guardrails/test_shunt_rewrite.py index 6d4dcf9f3cf..460b3ef50ab 100644 --- a/tests/test_litellm/proxy/guardrails/test_shunt_rewrite.py +++ b/tests/test_litellm/proxy/guardrails/test_shunt_rewrite.py @@ -77,12 +77,14 @@ def _bounded_read( question: str = "Summarize this file's structure.", min_lines: int = 350, bulk_read_endpoint: str = "http://localhost:4000/v1/bulk_read", + capability_token: str = "shunt_cap_v1:test-token", ) -> ShuntBashRewrite: return build_bounded_read_command( path=path, question=question, min_lines=min_lines, bulk_read_endpoint=bulk_read_endpoint, + capability_token=capability_token, ) @@ -160,6 +162,7 @@ class TestGeneratedCommandsResistShellInjection: question="q", paths=["ok.py", payload.format(marker=marker)], bulk_read_endpoint="http://127.0.0.1:9/v1/bulk_read", + capability_token="shunt_cap_v1:test-token", ) _assert_runs_without_side_effect(rewrite.command, marker) @@ -170,6 +173,7 @@ class TestGeneratedCommandsResistShellInjection: reference="r.py", target=payload.format(marker=marker), code_write_endpoint="http://127.0.0.1:9/v1/code_write", + capability_token="shunt_cap_v1:test-token", ) _assert_runs_without_side_effect(rewrite.command, marker) @@ -180,6 +184,7 @@ class TestGeneratedCommandsResistShellInjection: reference="r.py", target=None, code_write_endpoint="http://127.0.0.1:9/v1/code_write", + capability_token="shunt_cap_v1:test-token", ) _assert_runs_without_side_effect(rewrite.command, marker) @@ -187,46 +192,40 @@ class TestGeneratedCommandsResistShellInjection: # Regression: the caller's key was interpolated straight into the generated command, so it # landed in the model's response, the conversation history, and the next upstream turn. class TestGeneratedCommandsNeverCarryTheCallersCredential: - def test_bounded_read_references_the_env_var_instead_of_a_secret(self): - command = _bounded_read().command - assert "ANTHROPIC_AUTH_TOKEN" in command - assert "sk-" not in command - - def test_bulk_read_references_the_env_var_instead_of_a_secret(self): - rewrite = build_bulk_read_command( - question="q", paths=["a.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read" - ) - assert "ANTHROPIC_AUTH_TOKEN" in rewrite.command + def test_bounded_read_carries_the_token_not_a_real_key(self): + rewrite = _bounded_read(capability_token="shunt_cap_v1:abc123") + assert "shunt_cap_v1:abc123" in rewrite.command assert "sk-" not in rewrite.command - def test_code_write_references_the_env_var_instead_of_a_secret(self): + def test_bulk_read_carries_the_token_not_a_real_key(self): + rewrite = build_bulk_read_command( + question="q", + paths=["a.py"], + bulk_read_endpoint="http://localhost:4000/v1/bulk_read", + capability_token="shunt_cap_v1:abc123", + ) + assert "shunt_cap_v1:abc123" in rewrite.command + assert "sk-" not in rewrite.command + + def test_code_write_carries_the_token_not_a_real_key(self): rewrite = build_code_write_command( - spec="s", reference="r.py", target=None, code_write_endpoint="http://localhost:4000/v1/code_write" + spec="s", + reference="r.py", + target=None, + code_write_endpoint="http://localhost:4000/v1/code_write", + capability_token="shunt_cap_v1:abc123", ) - assert "ANTHROPIC_AUTH_TOKEN" in rewrite.command + assert "shunt_cap_v1:abc123" in rewrite.command assert "sk-" not in rewrite.command - def test_the_env_var_expands_at_run_time(self, tmp_path: Path): - """The header must carry the client's real token once bash evaluates the command.""" + def test_the_token_is_carried_as_a_bearer_authorization_header(self, tmp_path: Path): + """A generated command must send the token in the header, never a URL query string.""" out = tmp_path / "seen_header.txt" rewrite = build_bulk_read_command( - question="q", paths=["a.py"], bulk_read_endpoint="http://127.0.0.1:9/v1/bulk_read" - ) - # Echo the expanded header rather than sending it, so the assertion needs no server. - header_only = rewrite.command.split(" -H ", 1)[1].rsplit(" ", 1)[0] - subprocess.run( - ["bash", "-c", f"printf '%s' {header_only} > {out}"], - capture_output=True, - text=True, - check=False, - env={"ANTHROPIC_AUTH_TOKEN": "sk-real-token", "PATH": "/usr/bin:/bin"}, - ) - assert out.read_text() == "Authorization: Bearer sk-real-token" - - def test_falls_back_to_the_api_key_env_var_for_x_api_key_clients(self, tmp_path: Path): - out = tmp_path / "seen_header.txt" - rewrite = build_bulk_read_command( - question="q", paths=["a.py"], bulk_read_endpoint="http://127.0.0.1:9/v1/bulk_read" + question="q", + paths=["a.py"], + bulk_read_endpoint="http://127.0.0.1:9/v1/bulk_read", + capability_token="shunt_cap_v1:abc123", ) header_only = rewrite.command.split(" -H ", 1)[1].rsplit(" ", 1)[0] subprocess.run( @@ -234,9 +233,9 @@ class TestGeneratedCommandsNeverCarryTheCallersCredential: capture_output=True, text=True, check=False, - env={"ANTHROPIC_API_KEY": "sk-from-api-key", "PATH": "/usr/bin:/bin"}, + env={"PATH": "/usr/bin:/bin"}, ) - assert out.read_text() == "Authorization: Bearer sk-from-api-key" + assert out.read_text() == "Authorization: Bearer shunt_cap_v1:abc123" # Regression: the commands uploaded files as `paths[]`, but the endpoint binds them under @@ -249,7 +248,7 @@ class TestUploadFieldNameMatchesTheEndpoint: def test_bulk_read_uses_the_bare_paths_field_name(self): rewrite = build_bulk_read_command( - question="q", paths=["a.py", "b.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read" + question="q", paths=["a.py", "b.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read", capability_token="shunt_cap_v1:test-token" ) assert "paths[]=@" not in rewrite.command for path in ("a.py", "b.py"): @@ -279,6 +278,7 @@ class TestBuildBulkReadCommand: question="what does this do", paths=["a.py", "b.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read", + capability_token="shunt_cap_v1:test-token", ) assert "wc -l" not in rewrite.command assert "if [" not in rewrite.command @@ -288,13 +288,14 @@ class TestBuildBulkReadCommand: question="q", paths=["a.py", "b.py", "c.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read", + capability_token="shunt_cap_v1:test-token", ) for path in ("a.py", "b.py", "c.py"): assert f"paths=@{path}" in rewrite.command def test_command_is_valid_bash(self): rewrite = build_bulk_read_command( - question="q", paths=["a.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read" + question="q", paths=["a.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read", capability_token="shunt_cap_v1:test-token" ) _assert_valid_bash(rewrite.command) @@ -306,6 +307,7 @@ class TestBuildCodeWriteCommand: reference="tests/y_test.py", target=None, code_write_endpoint="http://localhost:4000/v1/code_write", + capability_token="shunt_cap_v1:test-token", ) assert ">" not in rewrite.command @@ -315,12 +317,13 @@ class TestBuildCodeWriteCommand: reference="tests/y_test.py", target="tests/x_test.py", code_write_endpoint="http://localhost:4000/v1/code_write", + capability_token="shunt_cap_v1:test-token", ) assert rewrite.command.endswith("> tests/x_test.py") @pytest.mark.parametrize("target", [None, "tests/x_test.py"]) def test_command_is_valid_bash_with_and_without_target(self, target: str | None): rewrite = build_code_write_command( - spec="s", reference="r.py", target=target, code_write_endpoint="http://x/v1/code_write" + spec="s", reference="r.py", target=target, code_write_endpoint="http://x/v1/code_write", capability_token="shunt_cap_v1:test-token" ) _assert_valid_bash(rewrite.command) diff --git a/tests/test_litellm/proxy/shunt_endpoints/test_endpoints.py b/tests/test_litellm/proxy/shunt_endpoints/test_endpoints.py new file mode 100644 index 00000000000..984dc2ec5b8 --- /dev/null +++ b/tests/test_litellm/proxy/shunt_endpoints/test_endpoints.py @@ -0,0 +1,100 @@ +"""Unit tests for litellm.proxy.shunt_endpoints.endpoints's capability-token auth dependency.""" + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.guardrails.shunt_capability_token import mint_shunt_capability_token +from litellm.proxy.shunt_endpoints.endpoints import _caller_from_capability_token + + +@pytest.fixture(autouse=True) +def _salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234-test-salt-key") + + +class TestMissingOrMalformedHeader: + @pytest.mark.asyncio + async def test_no_header_is_rejected(self): + with pytest.raises(HTTPException) as exc_info: + await _caller_from_capability_token(authorization=None) + assert exc_info.value.status_code == 401 + + @pytest.mark.asyncio + async def test_non_bearer_header_is_rejected(self): + with pytest.raises(HTTPException) as exc_info: + await _caller_from_capability_token(authorization="Basic abc123") + assert exc_info.value.status_code == 401 + + @pytest.mark.asyncio + async def test_malformed_token_is_rejected(self): + with pytest.raises(HTTPException) as exc_info: + await _caller_from_capability_token(authorization="Bearer not-a-real-token") + assert exc_info.value.status_code == 401 + + @pytest.mark.asyncio + async def test_expired_token_is_rejected(self, monkeypatch): + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None, now=1_000_000) + # 121s after mint, one second past the 120s TTL. + import litellm.proxy.guardrails.shunt_capability_token as token_mod + + monkeypatch.setattr(token_mod.time, "time", lambda: 1_000_121) + with pytest.raises(HTTPException) as exc_info: + await _caller_from_capability_token(authorization=f"Bearer {token}") + assert exc_info.value.status_code == 401 + + +class TestKeyHashGrant: + @pytest.mark.asyncio + async def test_resolves_the_key_object_for_the_grants_hash(self, monkeypatch): + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None) + resolved = UserAPIKeyAuth(api_key="deadbeef", team_id="team-1") + + async def _fake_get_key_object(**kwargs): + assert kwargs["hashed_token"] == "deadbeef" + return resolved + + import litellm.proxy.auth.auth_checks as auth_checks + + monkeypatch.setattr(auth_checks, "get_key_object", _fake_get_key_object) + result = await _caller_from_capability_token(authorization=f"Bearer {token}") + assert result is resolved + + @pytest.mark.asyncio + async def test_a_lookup_failure_is_rejected_not_propagated(self, monkeypatch): + token = mint_shunt_capability_token(key_hash="deadbeef", master_key=None) + + async def _raising_get_key_object(**kwargs): + raise Exception("key not found") + + import litellm.proxy.auth.auth_checks as auth_checks + + monkeypatch.setattr(auth_checks, "get_key_object", _raising_get_key_object) + with pytest.raises(HTTPException) as exc_info: + await _caller_from_capability_token(authorization=f"Bearer {token}") + assert exc_info.value.status_code == 401 + + +class TestMasterKeyGrant: + @pytest.mark.asyncio + async def test_matching_master_key_resolves_as_proxy_admin(self, monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-the-real-master-key") + token = mint_shunt_capability_token(key_hash=None, master_key="sk-the-real-master-key") + result = await _caller_from_capability_token(authorization=f"Bearer {token}") + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + + @pytest.mark.asyncio + async def test_master_key_mismatch_is_rejected(self, monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-the-current-master-key") + token = mint_shunt_capability_token(key_hash=None, master_key="sk-a-stale-master-key") + with pytest.raises(HTTPException) as exc_info: + await _caller_from_capability_token(authorization=f"Bearer {token}") + assert exc_info.value.status_code == 401 + + @pytest.mark.asyncio + async def test_no_configured_master_key_rejects_a_master_key_grant(self, monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + token = mint_shunt_capability_token(key_hash=None, master_key="sk-anything") + with pytest.raises(HTTPException) as exc_info: + await _caller_from_capability_token(authorization=f"Bearer {token}") + assert exc_info.value.status_code == 401 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 329ba0c10b9..b642d1f901a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -42617,7 +42617,9 @@ export interface operations { router: string; tags?: string[] | null; }; - header?: never; + header?: { + authorization?: string | null; + }; path?: never; cookie?: never; }; @@ -43511,7 +43513,9 @@ export interface operations { router: string; tags?: string[] | null; }; - header?: never; + header?: { + authorization?: string | null; + }; path?: never; cookie?: never; }; @@ -61525,7 +61529,9 @@ export interface operations { router: string; tags?: string[] | null; }; - header?: never; + header?: { + authorization?: string | null; + }; path?: never; cookie?: never; }; @@ -61760,7 +61766,9 @@ export interface operations { router: string; tags?: string[] | null; }; - header?: never; + header?: { + authorization?: string | null; + }; path?: never; cookie?: never; };