From b4b4f58a76d01dd635085fa3685d10de567e7ce3 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Mon, 7 Sep 2026 16:44:15 -0700 Subject: [PATCH] fix(auto_router): non-database callers, master-key logging, and worker budget bypass Three fixes from this review round, plus a CI shard registration and a comment-density pass. A caller admitted through JWT or a custom auth path has no DB-backed key hash (UserAPIKeyAuth.api_key is None or some other non-sk- value), which crashed the capability token mint instead of leaving the tool_use untouched. _mint_caller_capability_token now returns None for that case, and _endpoints_for_request treats it the same as an unreachable base URL. A resolved master-key caller carried the real master key as its own api_key, which _worker_text later placed in the outbound request's metadata["user_api_key"] -- reachable by any raw-metadata logging callback. Now substitutes LITELLM_PROXY_MASTER_KEY_ALIAS there, matching what normal master-key auth already does for exactly this reason. The worker call went straight to llm_router.acompletion, skipping every registered rate-limit and budget callback (they run as async_pre_call_hook, which only proxy_logging_obj.pre_call_hook walks). A caller already over budget or rate-limited could keep spending through this endpoint indefinitely. _worker_text now calls pre_call_hook first and lets a block propagate. Registers tests/test_litellm/proxy/shunt_endpoints in test-unit.yml's proxy-endpoints shard; CI's shard-coverage assertion failed without it since the directory held tests but named no owning shard. Trims several docstrings/module comments in the shunt modules down to the non-obvious "why" they were justified by, cutting repetition and one stale field name a rename had left behind. --- .github/workflows/test-unit.yml | 1 + litellm/proxy/guardrails/auto_router_shunt.py | 92 +++++++++---------- .../guardrails/shunt_capability_token.py | 21 ++--- litellm/proxy/guardrails/shunt_rewrite.py | 36 +++----- litellm/proxy/shunt_endpoints/endpoints.py | 71 ++++++++------ .../guardrails/test_auto_router_shunt.py | 29 ++++++ .../proxy/shunt_endpoints/test_endpoints.py | 81 +++++++++++++++- 7 files changed, 209 insertions(+), 122 deletions(-) diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index cc606339a20..c0f2b79eb24 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -146,6 +146,7 @@ jobs: tests/test_litellm/proxy/list_api tests/test_litellm/proxy/memory tests/test_litellm/proxy/guardrails + tests/test_litellm/proxy/shunt_endpoints tests/test_litellm/proxy/management_helpers tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/google_endpoints diff --git a/litellm/proxy/guardrails/auto_router_shunt.py b/litellm/proxy/guardrails/auto_router_shunt.py index 5e72f03ffd6..e5b6c382973 100644 --- a/litellm/proxy/guardrails/auto_router_shunt.py +++ b/litellm/proxy/guardrails/auto_router_shunt.py @@ -1,17 +1,13 @@ """ Server-side port of Spotify's ``shunt`` Claude Code plugin -(https://engineering.atspotify.com/2026/9/portal-by-spotify-cut-my-claude-code-token-usage-by-90), -as an auto-router preset rather than a client-side plugin. +(https://engineering.atspotify.com/2026/9/portal-by-spotify-cut-my-claude-code-token-usage-by-90). -shunt intercepts large file reads and boilerplate generation at the client via Claude Code -``PreToolUse`` hooks and delegates them to a cheap worker model. This module does the same -decision on the proxy instead: it arms via ``auto_router_shunt_min_lines`` / -``auto_router_shunt_bulk_read_model`` / ``auto_router_shunt_code_write_model`` on an -auto-router marker deployment (the same "read `litellm_params` off the resolved deployment, no -``guardrails:`` config entry needed" shape as ``auto_router_compression.py``), then a single -always-on ``ShuntGuardrail`` callback injects ``bulk_read``/``code_write`` tool definitions -pre-call and rewrites large ``Read``/``Bash``/``bulk_read``/``code_write`` tool_use blocks -post-call so the file bytes and generated code never reach the routed model's context. +Arms via ``auto_router_shunt_min_lines`` / ``auto_router_shunt_bulk_read_model`` / +``auto_router_shunt_code_write_model`` on an auto-router marker's ``litellm_params`` (same +resolution shape as ``auto_router_compression.py``). The always-on ``ShuntGuardrail`` injects +``bulk_read``/``code_write`` tool definitions pre-call and rewrites large +``Read``/``Bash``/``bulk_read``/``code_write`` tool_use blocks post-call, so the file bytes and +generated code never reach the routed model's context. """ import json @@ -42,11 +38,7 @@ class ShuntConfig: def _simple_tier_model(litellm_params: Mapping[str, object]) -> str | None: - """The first model in the router's SIMPLE tier, or None. - - The cheapest tier the router already has, which is what the UI and the preset docs promise - an unset worker model falls back to. - """ + """The router's SIMPLE tier model, the fallback the UI promises for an unset worker model.""" config: Final = litellm_params.get("complexity_router_config") if not isinstance(config, Mapping): return None @@ -64,11 +56,10 @@ def _simple_tier_model(litellm_params: Mapping[str, object]) -> str | None: def _config_from_litellm_params( litellm_params: Mapping[str, object], *, default_model: str | None ) -> ShuntConfig | None: - """The marker's shunt config, or None when `auto_router_shunt_min_lines` is absent. + """The marker's shunt config, or None when unarmed or its worker models resolve to nothing. - An unset worker model falls back to the SIMPLE tier first, then the router's default model. - Returns None rather than a config naming no model at all: arming shunt with an empty worker - model would rewrite reads into calls that can only fail, which is worse than not arming. + An unset worker model falls back to the SIMPLE tier, then the router's default model; if + neither resolves, this returns None rather than a config naming no model at all. """ raw_min_lines: Final = litellm_params.get("auto_router_shunt_min_lines") if not isinstance(raw_min_lines, int): @@ -131,6 +122,7 @@ def _config_from_marker(litellm_params: Mapping[str, object]) -> ShuntConfig | N # is no client-side plugin here for a skill file to live in. BULK_READ_TOOL_NAME: Final = "bulk_read" CODE_WRITE_TOOL_NAME: Final = "code_write" +_SHUNT_TOOL_NAMES: Final = frozenset({BULK_READ_TOOL_NAME, CODE_WRITE_TOOL_NAME}) # JSON source, not dict literals: these go straight into the outbound payload, so they must stay # plain JSON-serializable dicts, and parsing keeps one construction site instead of a suppression @@ -246,14 +238,15 @@ class _ShuntEndpoints: capability_token: str -def _mint_caller_capability_token(user_api_key_dict: "UserAPIKeyAuth") -> str: - """Seal a short-lived grant identifying this request's caller. +def _mint_caller_capability_token(user_api_key_dict: "UserAPIKeyAuth") -> str | None: + """Seal a short-lived grant identifying this request's caller, or None if it can't be. - ``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. + `UserAPIKeyAuth.api_key` is already the hashed token for a DB-backed virtual key, so the + common case carries that hash forward. Master-key auth stores a stable alias there instead + (`LITELLM_PROXY_MASTER_KEY_ALIAS`), so that case seals the real master key itself for the + worker endpoint to compare directly. A JWT- or custom-auth-admitted caller can carry + neither; None there rather than raising, since an already-answered request must not fail + over an optimization that couldn't run. """ from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS from litellm.proxy.guardrails.shunt_capability_token import mint_shunt_capability_token @@ -261,6 +254,8 @@ def _mint_caller_capability_token(user_api_key_dict: "UserAPIKeyAuth") -> str: 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) + if not user_api_key_dict.api_key: + return None return mint_shunt_capability_token(key_hash=user_api_key_dict.api_key, master_key=None) @@ -269,20 +264,20 @@ def _endpoints_for_request( ) -> "_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 - 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. + None when the base URL can't be recovered or no capability token can be minted for this + caller, since a generated command could then not authenticate; the tool_use is left + unmodified rather than shipped broken. The token carries the caller's identity by reference + (never the real credential) in the command's `Authorization` header, never a query string. - 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 - to them and the delegated call would 400, even though this rewrite matched that marker. + Tags ride along in the query string because the worker endpoints re-resolve the marker from + scratch; without them a tag-scoped marker would 400 there even though this rewrite matched it. """ base_url: Final = _request_base_url(data) if base_url is None: return None + capability_token: Final = _mint_caller_capability_token(user_api_key_dict) + if capability_token is None: + return None from urllib.parse import urlencode from litellm.router_strategy.tag_based_routing import ( @@ -293,7 +288,7 @@ def _endpoints_for_request( 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), + capability_token=capability_token, ) @@ -490,7 +485,7 @@ def caller_owns_shunt_tool_name(tools: object) -> bool: if not isinstance(tools, Sequence) or isinstance(tools, (str, bytes)): return False declared: Final = frozenset(name for name in (_declared_tool_name(tool) for tool in tools) if name is not None) - return not declared.isdisjoint({BULK_READ_TOOL_NAME, CODE_WRITE_TOOL_NAME}) + return not declared.isdisjoint(_SHUNT_TOOL_NAMES) def _tools_payload( @@ -573,16 +568,14 @@ class ShuntGuardrail(CustomLogger): ) -> "AsyncGenerator[ModelResponseStream, None]": """Buffer the whole stream, rewrite any shunt-shaped tool_use, then replay it. - Buffer-then-replay, not per-fragment rewriting, matching `tool_permission.py`'s own - streaming hook: `input_json_delta` fragments split mid-token (confirmed against a real - Anthropic trace during design), so a tool_use's `input` can only be read once its - `content_block_stop` has arrived. Unlike that guardrail, an unparseable or unassemblable - stream is passed through unmodified rather than raised on: shunt is an optimization, not - a safety control, so a request should never fail because this rewrite couldn't run. + Buffer-then-replay, matching `tool_permission.py`'s own streaming hook: `input_json_delta` + fragments split mid-token, so a tool_use's `input` is only readable once its + `content_block_stop` arrives. Unlike that guardrail, an unparseable stream passes through + untouched rather than raising, since shunt is an optimization, not a safety control. - The declared return type matches the base class and `tool_permission.py`'s own override: - on the Anthropic path this actually yields raw `bytes` SSE frames, not - `ModelResponseStream` objects, the same documented mismatch `tool_permission.py` carries. + Declared return type matches the base class and `tool_permission.py`'s own override; on + the Anthropic path this actually yields raw `bytes` SSE frames, the same mismatch that + guardrail's own override carries. """ from litellm.main import stream_chunk_builder from litellm.proxy.guardrails.anthropic_sse import ( @@ -592,11 +585,10 @@ class ShuntGuardrail(CustomLogger): ) from litellm.types.utils import ModelResponse, TextCompletionResponse - # Declared element type matches `tool_permission.py`'s own override rather than the - # `object` the raw-SSE path really carries: see the docstring's note on that mismatch. + # Typed as ModelResponseStream, though the raw-SSE path really carries bytes here. all_chunks: Final[ list[ModelResponseStream] - ] = [ # mutable-ok: stream_chunk_builder/is_raw_sse_stream both take a concrete list + ] = [ # mutable-ok: stream_chunk_builder/is_raw_sse_stream take a list chunk async for chunk in response ] diff --git a/litellm/proxy/guardrails/shunt_capability_token.py b/litellm/proxy/guardrails/shunt_capability_token.py index 8383aeb0f75..e293b4523ae 100644 --- a/litellm/proxy/guardrails/shunt_capability_token.py +++ b/litellm/proxy/guardrails/shunt_capability_token.py @@ -2,15 +2,9 @@ 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. +Identifies the caller by a reference to their existing key rather than the key itself, sealed +with the proxy's own encryption helper and given a short expiry, carried in an `Authorization` +header rather than 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 @@ -32,12 +26,9 @@ 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. + Exactly one of `key_hash` (the same hash `UserAPIKeyAuth.api_key` already stores for a + DB-backed key) or `master_key` (the real master key, for the one caller with no such row) + is set. """ model_config = ConfigDict(frozen=True, extra="forbid") diff --git a/litellm/proxy/guardrails/shunt_rewrite.py b/litellm/proxy/guardrails/shunt_rewrite.py index 203ec32f9c4..46698552f90 100644 --- a/litellm/proxy/guardrails/shunt_rewrite.py +++ b/litellm/proxy/guardrails/shunt_rewrite.py @@ -1,15 +1,10 @@ """ Decides which tool_use blocks shunt rewrites, and builds the Bash command that replaces them. -Ports two decisions from Spotify's shunt plugin (see auto_router_shunt.py's module docstring): -`check-file-size`'s "is this an untargeted Read on a large file" gate, and -`check-bash-read`'s "is this cat/head/tail/less/more on a bare file path" parser, including its -documented parser bug (an option-value token like the `5` in `head -n 5 file` is mistaken for -the path). shunt's version tolerates that bug because its hook runs on the file's own machine -and checks the guessed path exists before blocking; a misparse then falls through to `allow`, -leaving the original command untouched. This module has no filesystem access to make that same -check, so it must never rewrite a command it isn't sure it parsed correctly — see -`extract_bare_read_path`'s docstring for how the port preserves shunt's fail-safe direction. +Ports two decisions from Spotify's shunt plugin: `check-file-size`'s "is this an untargeted +Read on a large file" gate, and `check-bash-read`'s "is this cat/head/tail/less/more on a bare +file path" parser (see `extract_bare_read_path` for its documented parser bug and how this port +handles it, since this module has no filesystem to fall back on if a rewrite is wrong). """ import re @@ -35,23 +30,14 @@ def is_targeted_read(offset: object, limit: object) -> bool: def extract_bare_read_path(command: str) -> str | None: """The file path a bare `cat`/`head`/`tail`/`less`/`more` command reads, or None. - None means "do not rewrite this command": either it isn't a bulk read (piped, redirected, - not one of the five commands), or no non-flag argument was found. + None means don't rewrite: piped/redirected, not one of the five commands, or no non-flag + argument found. shunt's own parser misreads a flag's value token as the path (the `5` in + `head -n 5 file`), tolerable there since its hook runs on the file's own machine and can + check the guessed path exists before blocking. This port has no filesystem for that check, + so a bare numeric token (every flag value in these commands' option sets happens to be one) + is skipped rather than returned, fixing shunt's documented bug instead of reproducing it. - shunt's own parser walks flags as if none of them take a value, which misreads a flag's - value token as the path — e.g. the `5` in `head -n 5 file`. shunt tolerates this because its - hook runs on the file's own machine and checks the guessed path actually exists before - blocking; a nonexistent guess falls through to `allow`, leaving the original command - untouched. This port has no filesystem to make that same check, so it cannot rely on a - downstream existence check to catch a bad guess — treating a wrong guess as the real path - would rewrite the command against a file the model never named. A bare numeric token (as - every flag value in the five commands' own option sets happens to be — `-n`, `head -c`, - `tail -n +N`) is skipped as a likely flag value rather than returned, which fixes shunt's - documented bug for exactly the case it names rather than reproducing it. A path made only - of digits (or `tail`'s `+N` follow-from-line syntax) is stripped either way, matching - shunt's own scope: neither version claims to handle a bare-numeric filename correctly. - - Like shunt's own bash word-splitting, this does not handle a quoted path containing spaces + Like shunt's own bash word-splitting, a quoted path containing spaces is not handled (`cat "my file.txt"` returns `"my`) — matching parity, not a regression. """ if _PIPE_OR_REDIRECT_PATTERN.search(command): diff --git a/litellm/proxy/shunt_endpoints/endpoints.py b/litellm/proxy/shunt_endpoints/endpoints.py index 48f7f019897..1061357cdff 100644 --- a/litellm/proxy/shunt_endpoints/endpoints.py +++ b/litellm/proxy/shunt_endpoints/endpoints.py @@ -1,18 +1,11 @@ """ -`/v1/bulk_read` and `/v1/code_write`: the worker endpoints shunt's generated Bash commands -call. See `auto_router_shunt.py`'s module docstring for the source this ports. +`/v1/bulk_read` and `/v1/code_write`: the worker endpoints shunt's generated Bash commands call. -Each request names the auto-router marker it belongs to (`router=`, the same -value the client originally called), so the worker model is the one configured on that -marker's `auto_router_shunt_bulk_read_model` / `auto_router_shunt_code_write_model`, not a -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. +Each request names the auto-router marker it belongs to (`router=`), so the +worker model is whatever that marker's `auto_router_shunt_bulk_read_model` / +`auto_router_shunt_code_write_model` configures, not one the caller chooses. Auth is the +short-lived capability token minted at rewrite time (`shunt_capability_token.py`), not a +normal virtual key: only a shunt-generated command should ever call these routes. """ from collections.abc import Sequence @@ -35,7 +28,7 @@ from litellm.proxy.shunt_endpoints.worker import ( build_code_write_message, strip_code_fences, ) -from litellm.types.llms.openai import ChatCompletionSystemMessage, ChatCompletionUserMessage +from litellm.types.llms.openai import AllMessageValues, ChatCompletionSystemMessage, ChatCompletionUserMessage if TYPE_CHECKING: from litellm.router import Router @@ -63,11 +56,16 @@ async def _caller_from_capability_token(authorization: Annotated[str | None, Hea raise HTTPException(status_code=401, detail="Invalid or expired shunt capability token") if grant.master_key is not None: + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS 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) + # The alias substitutes for the real master key here for the same reason normal + # master-key auth substitutes it (user_api_key_auth.py): neither the key nor a + # reversible derivation of it should reach spend logs, Prometheus labels, or any raw- + # metadata logging callback the worker call's own metadata is later forwarded to. + return UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN) if grant.key_hash is None: raise HTTPException(status_code=401, detail="Invalid or expired shunt capability token") @@ -129,27 +127,40 @@ async def _worker_text( ) -> str: """The worker model's reply text, or a 502 if it produced none. - One call site for both endpoints: they differ only in model, system prompt, and message, so - the `acompletion` shape (temperature, non-streaming, caller attribution) lives here once. + One call site for both endpoints, since they differ only in model, system prompt, and + message. Attribution reuses the proxy's own key-metadata builder so the call is billed and + budgeted against the calling key/user/team/org like a normal request. - Attribution reuses the proxy's own key-metadata builder rather than hand-picking a couple of - fields, so the worker call is billed and budgeted against the calling key, user, team, and - org exactly like a normal request instead of only carrying a team id. + `proxy_logging_obj.pre_call_hook` runs first: `llm_router.acompletion` alone skips every + rate-limit and budget callback, since those register as `async_pre_call_hook` and only + `/chat/completions` and friends normally walk that list before routing. Without this call a + caller already over budget or rate-limited could keep spending through this endpoint. """ + from litellm.proxy.proxy_server import proxy_logging_obj + system: Final = ChatCompletionSystemMessage(role="system", content=system_prompt) user: Final = ChatCompletionUserMessage(role="user", content=message) + messages: Final[ + list[AllMessageValues] + ] = [ # mutable-ok: shared between pre_call_hook and acompletion, both take a list + system, + user, + ] key_metadata: Final = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( user_api_key_dict=user_api_key_dict ) + metadata: Final = {**key_metadata, "user_api_key": user_api_key_dict.api_key} # mutable-ok: same + request_data: Final = { # mutable-ok: pre_call_hook's own signature takes a plain dict + "model": model, + "messages": messages, + "metadata": metadata, + } + await proxy_logging_obj.pre_call_hook( + user_api_key_dict=user_api_key_dict, data=request_data, call_type="acompletion" + ) + response: Final = await llm_router.acompletion( - model=model, - messages=[system, user], # mutable-ok: acompletion's own signature takes a concrete list - temperature=WORKER_TEMPERATURE, - stream=False, - metadata={ # mutable-ok: same - **key_metadata, - "user_api_key": user_api_key_dict.api_key, - }, + model=model, messages=messages, temperature=WORKER_TEMPERATURE, stream=False, metadata=metadata ) text: Final = response.choices[0].message.content if not isinstance(text, str): @@ -186,7 +197,7 @@ async def bulk_read( 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(_caller_from_capability_token)], - tags: Annotated[list[str] | None, Query()] = None, + tags: Annotated[list[str] | None, Query()] = None, # mutable-ok: FastAPI requires a list for a repeated query param ) -> str: """Summarize or answer a question about one or more files via a cheap worker model. @@ -225,7 +236,7 @@ async def code_write( spec: Annotated[str, Form()], reference: Annotated[UploadFile, File()], user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(_caller_from_capability_token)], - tags: Annotated[list[str] | None, Query()] = None, + tags: Annotated[list[str] | None, Query()] = None, # mutable-ok: FastAPI requires a list for a repeated query param ) -> 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 24004448f20..e653ac3a602 100644 --- a/tests/test_litellm/proxy/guardrails/test_auto_router_shunt.py +++ b/tests/test_litellm/proxy/guardrails/test_auto_router_shunt.py @@ -540,3 +540,32 @@ class TestAsyncPostCallSuccessHookOpenAIShape: data=self._armed_request_data(), user_api_key_dict=_FAKE_USER_API_KEY_DICT, response=response ) assert result.choices[0].message.content == "hi" + + +# Regression: a caller with no DB-backed key hash (JWT/custom-auth admission, where +# UserAPIKeyAuth.api_key can be None) crashed mint_shunt_capability_token's exactly-one-of +# check instead of leaving the tool_use untouched. +class TestCallerWithNoMintableIdentity: + 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"}, + } + + @pytest.mark.asyncio + async def test_none_api_key_leaves_the_response_untouched_rather_than_raising(self, monkeypatch): + import litellm.proxy.guardrails.auto_router_shunt as mod + + monkeypatch.setattr(mod, "_resolve_shunt_config", lambda data: self._config()) + guardrail = mod.ShuntGuardrail() + jwt_admitted_caller = UserAPIKeyAuth(api_key=None) + 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=jwt_admitted_caller, response=response + ) + assert result["content"][0]["name"] == "Read" diff --git a/tests/test_litellm/proxy/shunt_endpoints/test_endpoints.py b/tests/test_litellm/proxy/shunt_endpoints/test_endpoints.py index 984dc2ec5b8..f1b2bcd06a6 100644 --- a/tests/test_litellm/proxy/shunt_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/shunt_endpoints/test_endpoints.py @@ -1,11 +1,11 @@ -"""Unit tests for litellm.proxy.shunt_endpoints.endpoints's capability-token auth dependency.""" +"""Unit tests for litellm.proxy.shunt_endpoints.endpoints.""" 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 +from litellm.proxy.shunt_endpoints.endpoints import _caller_from_capability_token, _worker_text @pytest.fixture(autouse=True) @@ -83,6 +83,20 @@ class TestMasterKeyGrant: result = await _caller_from_capability_token(authorization=f"Bearer {token}") assert result.user_role == LitellmUserRoles.PROXY_ADMIN + # Regression: a resolved master-key caller carried the raw master key as its own api_key, + # which _worker_text later places in the outbound request's metadata["user_api_key"] -- + # reachable by any raw-metadata logging callback. Normal master-key auth substitutes a + # stable alias there specifically to keep the real key out of that sink; this must match. + @pytest.mark.asyncio + async def test_resolved_caller_never_carries_the_raw_master_key(self, monkeypatch): + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + + 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.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS + assert result.api_key != "sk-the-real-master-key" + @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") @@ -98,3 +112,66 @@ class TestMasterKeyGrant: with pytest.raises(HTTPException) as exc_info: await _caller_from_capability_token(authorization=f"Bearer {token}") assert exc_info.value.status_code == 401 + + +class _FakeRouter: + def __init__(self, response_text: str): + self._response_text = response_text + + async def acompletion(self, **kwargs): + from litellm.types.utils import Choices, Message, ModelResponse + + return ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content=self._response_text))]) + + +class _FakeProxyLogging: + def __init__(self, *, blocks: bool): + self._blocks = blocks + self.calls = [] + + async def pre_call_hook(self, *, user_api_key_dict, data, call_type): + self.calls.append((user_api_key_dict, data, call_type)) + if self._blocks: + raise HTTPException(status_code=429, detail="rate limited") + return data + + +# Regression: the worker call went straight to llm_router.acompletion, skipping every +# registered rate-limit/budget callback (they run as async_pre_call_hook, which only +# proxy_logging_obj.pre_call_hook walks). A caller already over budget or rate-limited could +# keep spending through this endpoint indefinitely. +class TestWorkerTextEnforcesRateLimitsAndBudget: + @pytest.mark.asyncio + async def test_calls_pre_call_hook_before_the_worker_model(self, monkeypatch): + fake_logging = _FakeProxyLogging(blocks=False) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", fake_logging) + holder = UserAPIKeyAuth(api_key="fakehash1234567890") + text = await _worker_text( + _FakeRouter("the worker's answer"), + model="claude-haiku-4-5", + system_prompt="be precise", + message="what does this do", + user_api_key_dict=holder, + label="bulk_read", + ) + assert text == "the worker's answer" + assert len(fake_logging.calls) == 1 + called_key, _, called_type = fake_logging.calls[0] + assert called_key is holder + assert called_type == "acompletion" + + @pytest.mark.asyncio + async def test_a_blocked_pre_call_hook_prevents_the_worker_call(self, monkeypatch): + fake_logging = _FakeProxyLogging(blocks=True) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", fake_logging) + holder = UserAPIKeyAuth(api_key="fakehash1234567890") + with pytest.raises(HTTPException) as exc_info: + await _worker_text( + _FakeRouter("should never be reached"), + model="claude-haiku-4-5", + system_prompt="be precise", + message="what does this do", + user_api_key_dict=holder, + label="bulk_read", + ) + assert exc_info.value.status_code == 429