diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 92b73867e67..e113e65d98f 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -55,6 +55,11 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = ( "/v1/skills", "/v1/a2a/", "/a2a/", + # Shunt worker endpoints: both call a worker model, so they belong on the data plane + "/v1/bulk_read", + "/bulk_read", + "/v1/code_write", + "/code_write", # LiteLLM-native LLM surface "/v1/rerank", "/v2/rerank", diff --git a/litellm/proxy/client/cli/commands/autoroute/settings.py b/litellm/proxy/client/cli/commands/autoroute/settings.py index c1f999eed13..e1a4565ea40 100644 --- a/litellm/proxy/client/cli/commands/autoroute/settings.py +++ b/litellm/proxy/client/cli/commands/autoroute/settings.py @@ -32,7 +32,12 @@ ANTHROPIC_DEFAULT_MODEL_ENV_KEYS: Final = ( # matching, so a rule narrow enough to name "curl" at all cannot also pin the destination the # trailing "*" is free to name any URL. A real boundary needs a PreToolUse hook instead, which # is the docs' own recommendation for exactly this case. +# +# One rule per command shape the rewrite emits, matched on the literal prefix each one starts +# with. The bounded read needs its own rule because it opens with the `wc -l` size check rather +# than with `curl`, so a curl-prefixed rule would never match the main large-file path. SHUNT_BASH_ALLOW_RULES: Final = ( + "Bash(L=$(wc -l < *)", "Bash(curl -sS -F question=*)", "Bash(curl -sS -F spec=*)", ) diff --git a/litellm/proxy/guardrails/auto_router_shunt.py b/litellm/proxy/guardrails/auto_router_shunt.py index a6fa659db05..356bb64f98a 100644 --- a/litellm/proxy/guardrails/auto_router_shunt.py +++ b/litellm/proxy/guardrails/auto_router_shunt.py @@ -41,19 +41,49 @@ class ShuntConfig: code_write_model: str +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. + """ + config: Final = litellm_params.get("complexity_router_config") + if not isinstance(config, Mapping): + return None + tiers: Final = config.get("tiers") + if not isinstance(tiers, Mapping): + return None + simple: Final = tiers.get("SIMPLE") + if isinstance(simple, str) and simple: + return simple + if not isinstance(simple, Sequence): + return None + return next((model for model in simple if isinstance(model, str) and model), 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. + + 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. + """ raw_min_lines: Final = litellm_params.get("auto_router_shunt_min_lines") if not isinstance(raw_min_lines, int): return None - fallback_model: Final = default_model or "" + fallback_model: Final = _simple_tier_model(litellm_params) or default_model raw_bulk_read: Final = litellm_params.get("auto_router_shunt_bulk_read_model") raw_code_write: Final = litellm_params.get("auto_router_shunt_code_write_model") + bulk_read_model: Final = raw_bulk_read if isinstance(raw_bulk_read, str) and raw_bulk_read else fallback_model + code_write_model: Final = raw_code_write if isinstance(raw_code_write, str) and raw_code_write else fallback_model + if not bulk_read_model or not code_write_model: + return None return ShuntConfig( min_lines=raw_min_lines, - bulk_read_model=raw_bulk_read if isinstance(raw_bulk_read, str) and raw_bulk_read else fallback_model, - code_write_model=raw_code_write if isinstance(raw_code_write, str) and raw_code_write else fallback_model, + bulk_read_model=bulk_read_model, + code_write_model=code_write_model, ) @@ -97,16 +127,14 @@ def _config_from_marker(litellm_params: Mapping[str, object]) -> ShuntConfig | N ) -# Tool descriptions carry shunt's own SKILL.md guidance: this is shunt's skills layer, -# delivered as tool metadata instead of a bundled markdown file, since there is no client-side -# plugin here for a skill file to live in. +# Tool descriptions carry shunt's own SKILL.md guidance, delivered as tool metadata since there +# 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" -# Held as JSON source, not dict literals: these are JSON Schema documents that go straight into -# the outbound provider payload, so they must stay plain JSON-serializable dicts (a -# MappingProxyType raises in the JSON encoder). Parsing the schema from the notation it is -# written in keeps one construction site instead of a suppression on every nested literal. +# 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 +# on every nested literal. _TOOL_DEFINITIONS_JSON: Final = """ { "bulk_read": { @@ -139,9 +167,7 @@ _TOOL_DEFINITIONS_JSON: Final = """ def _anthropic_tool(name: str) -> Mapping[str, object]: """The Anthropic-shape tool definition for `name`, as a fresh plain dict. - Fresh per call, never a shared module-level dict: these go into the outbound payload, and - handing every request the same mutable object would let one request's downstream mutation - (a provider transform normalizing a schema in place, say) leak into every later request. + Fresh per call so one request's downstream mutation cannot leak into every later request. """ definition: Final = json.loads(_TOOL_DEFINITIONS_JSON)[name] return {"name": name, **definition} # mutable-ok: goes into the outbound provider payload as plain JSON @@ -184,23 +210,6 @@ def _request_base_url(data: Mapping[str, object]) -> str | None: return f"{parts.scheme}://{parts.netloc}" -def _request_auth_header(data: Mapping[str, object]) -> str | None: - """The client's own `authorization` header, or None. - - Forwarded rather than a stored key: the generated curl authenticates to this proxy as the - same caller who sent the request, so it is billed and rate-limited the same way, and this - module never needs to hold or mint a credential of its own. - """ - secret_fields: Final = data.get("secret_fields") - if not isinstance(secret_fields, Mapping): - return None - raw_headers: Final = secret_fields.get("raw_headers") - if not isinstance(raw_headers, Mapping): - return None - header: Final = raw_headers.get("authorization") - return header if isinstance(header, str) and header else None - - def _resolve_shunt_config(data: Mapping[str, object]) -> ShuntConfig | None: """The armed `ShuntConfig` for this request's resolved model, or None. @@ -234,27 +243,33 @@ DEFAULT_BULK_READ_QUESTION: Final = "Summarize this file's exports and overall s class _ShuntEndpoints: bulk_read_url: str code_write_url: str - auth_header: str def _endpoints_for_request(data: Mapping[str, object], model_alias: str) -> "_ShuntEndpoints | None": """Where this request's generated curl commands should point, or None if unreachable. - None when the base URL or the caller's own auth header can't be recovered: without both, - a generated command could not reach this proxy as this caller, so the tool_use is left - unmodified rather than shipped with a broken command. + 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. + + 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. """ base_url: Final = _request_base_url(data) - auth_header: Final = _request_auth_header(data) - if base_url is None or auth_header is None: + if base_url is None: return None - from urllib.parse import quote + from urllib.parse import urlencode - router_query: Final = f"router={quote(model_alias)}" + from litellm.router_strategy.tag_based_routing import ( + _get_tags_from_request_kwargs, # pyright: ignore[reportPrivateUsage] # same helper router.py and auto_router_compression.py already use + ) + + query: Final = urlencode((("router", model_alias), *(("tags", tag) for tag in _get_tags_from_request_kwargs(data)))) return _ShuntEndpoints( - bulk_read_url=f"{base_url}/v1/bulk_read?{router_query}", - code_write_url=f"{base_url}/v1/code_write?{router_query}", - auth_header=auth_header, + bulk_read_url=f"{base_url}/v1/bulk_read?{query}", + code_write_url=f"{base_url}/v1/code_write?{query}", ) @@ -291,7 +306,6 @@ def _bash_replacement_for_tool_use( question=question, paths=paths, bulk_read_endpoint=endpoints.bulk_read_url, - auth_header=endpoints.auth_header, ) if name == CODE_WRITE_TOOL_NAME: @@ -305,7 +319,6 @@ 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, - auth_header=endpoints.auth_header, ) if name == "Read": @@ -319,7 +332,6 @@ def _bash_replacement_for_tool_use( question=DEFAULT_BULK_READ_QUESTION, min_lines=config.min_lines, bulk_read_endpoint=endpoints.bulk_read_url, - auth_header=endpoints.auth_header, ) if name == "Bash": @@ -334,7 +346,6 @@ def _bash_replacement_for_tool_use( question=DEFAULT_BULK_READ_QUESTION, min_lines=config.min_lines, bulk_read_endpoint=endpoints.bulk_read_url, - auth_header=endpoints.auth_header, ) return None @@ -421,6 +432,39 @@ def _rewrite_anthropic_response_in_place( return any(results) +# Stamped on the request by the pre-call hook when it declined to inject, and read back by the +# post-call hook. Post-call cannot re-inspect `data["tools"]` to make the same call: by then the +# list holds whatever the pre-call hook added, so the name would always look taken. +_CALLER_OWNS_TOOL_NAME_KEY: Final = "_shunt_caller_owns_tool_name" + + +def _declared_tool_name(tool: object) -> str | None: + """A tool definition's name, in either the Anthropic or the OpenAI shape.""" + if not isinstance(tool, Mapping): + return None + name: Final = tool.get("name") + if isinstance(name, str): + return name + function: Final = tool.get("function") + if not isinstance(function, Mapping): + return None + function_name: Final = function.get("name") + return function_name if isinstance(function_name, str) else None + + +def caller_owns_shunt_tool_name(tools: object) -> bool: + """Whether the caller already declared a tool named `bulk_read` or `code_write`. + + When it has, shunt stays out of the way entirely: injecting would hand the model two + different definitions of one name, and rewriting would silently turn the caller's own tool + call into shunt's unrelated curl. Leaving the request alone is the only safe read. + """ + 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}) + + def _tools_payload( existing: Sequence[object], added: Sequence[Mapping[str, object]] ) -> list[object]: # mutable-ok: outbound provider payload @@ -447,6 +491,10 @@ class ShuntGuardrail(CustomLogger): if config is None: return None + if caller_owns_shunt_tool_name(data.get("tools")): + data[_CALLER_OWNS_TOOL_NAME_KEY] = True # rebind-ok: read back by the post-call hook + return None + use_anthropic_format: Final = call_type in ("anthropic_messages", "aanthropic_messages") anthropic_tools: Final = (_anthropic_tool(BULK_READ_TOOL_NAME), _anthropic_tool(CODE_WRITE_TOOL_NAME)) new_tools: Final = ( @@ -469,7 +517,7 @@ class ShuntGuardrail(CustomLogger): from litellm.types.utils import ModelResponse config: Final = _resolve_shunt_config(data) - if config is None: + if config is None or data.get(_CALLER_OWNS_TOOL_NAME_KEY): return response model: Final = data.get("model") @@ -522,7 +570,7 @@ class ShuntGuardrail(CustomLogger): chunk async for chunk in response ] - config: Final = _resolve_shunt_config(request_data) + 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) diff --git a/litellm/proxy/guardrails/shunt_rewrite.py b/litellm/proxy/guardrails/shunt_rewrite.py index 0f391de770b..58b37a75977 100644 --- a/litellm/proxy/guardrails/shunt_rewrite.py +++ b/litellm/proxy/guardrails/shunt_rewrite.py @@ -13,6 +13,7 @@ check, so it must never rewrite a command it isn't sure it parsed correctly — """ import re +import shlex from collections.abc import Sequence from dataclasses import dataclass from typing import Final @@ -74,12 +75,18 @@ class ShuntBashRewrite: note: str -def _quote(value: str) -> str: - return value.replace('"', '\\"') +# 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 build_bounded_read_command( - *, path: str, question: str, min_lines: int, bulk_read_endpoint: str, auth_header: str + *, path: str, question: str, min_lines: int, bulk_read_endpoint: str ) -> ShuntBashRewrite: """The shunt conditional: read small files directly, delegate large ones. @@ -87,16 +94,21 @@ def build_bounded_read_command( `cat`/`head`/`tail`/`less`/`more` command) since both made the identical decision in shunt. `-sS` on curl mirrors shunt's own `-s` plus fails loudly on a server error rather than silently returning an HTML error page as if it were the answer. + + Every interpolated value goes through `shlex.quote`, and the path is reported with `printf` + rather than inside a double-quoted `echo`, because these values come from the model and the + command runs a shell on the developer's own machine: inside double quotes a `$(...)` in a + path would still be command-substituted. `min_lines` is an int, so it needs no quoting. """ - quoted_path: Final = _quote(path) - quoted_question: Final = _quote(question) + quoted_path: Final = shlex.quote(path) command: Final = ( - f'L=$(wc -l < "{quoted_path}" 2>/dev/null || echo 0); ' + f"L=$(wc -l < {quoted_path} 2>/dev/null || echo 0); " f'if [ "$L" -gt {min_lines} ]; then ' - f'echo "[shunt] {quoted_path}: $L lines, bounded read delegated" >&2; ' - f'curl -sS -F question="{quoted_question}" -F "paths[]=@{quoted_path}" ' - f'-H "Authorization: {auth_header}" "{bulk_read_endpoint}"; ' - f'else cat "{quoted_path}"; fi' + 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"else cat {quoted_path}; fi" ) return ShuntBashRewrite( command=command, @@ -104,25 +116,21 @@ def build_bounded_read_command( ) -def build_bulk_read_command( - *, question: str, paths: Sequence[str], bulk_read_endpoint: str, auth_header: str -) -> ShuntBashRewrite: +def build_bulk_read_command(*, question: str, paths: Sequence[str], bulk_read_endpoint: 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 `build_bounded_read_command` applies to a plain `Read`/`cat`/`head`/`tail` call. """ - quoted_question: Final = _quote(question) - path_flags: Final = " ".join(f'-F "paths[]=@{_quote(path)}"' for path in paths) + path_flags: Final = " ".join(f"-F {shlex.quote(f'paths=@{path}')}" for path in paths) command: Final = ( - f'curl -sS -F question="{quoted_question}" {path_flags} ' - f'-H "Authorization: {auth_header}" "{bulk_read_endpoint}"' + f"curl -sS -F {shlex.quote(f'question={question}')} {path_flags} {_AUTH_FLAG} {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, auth_header: str + *, spec: str, reference: str, target: str | None, code_write_endpoint: str ) -> ShuntBashRewrite: """The curl a model's own explicit `code_write(spec, reference, target)` tool call becomes. @@ -131,14 +139,12 @@ def build_code_write_command( itself when `target` is given. Either way the generated code enters the client's context only as a file write, never as text the routed model has to hold or repeat. """ - quoted_spec: Final = _quote(spec) - quoted_reference: Final = _quote(reference) request: Final = ( - f'curl -sS -F spec="{quoted_spec}" -F "reference=@{quoted_reference}" ' - f'-H "Authorization: {auth_header}" "{code_write_endpoint}"' + f"curl -sS -F {shlex.quote(f'spec={spec}')} " + f"-F {shlex.quote(f'reference=@{reference}')} " + f"{_AUTH_FLAG} {shlex.quote(code_write_endpoint)}" ) if target is None: return ShuntBashRewrite(command=request, note="Delegated to a cheaper model via code_write.") - quoted_target: Final = _quote(target) - command: Final = f'{request} > "{quoted_target}"' + command: Final = f"{request} > {shlex.quote(target)}" return ShuntBashRewrite(command=command, note=f"Delegated to a cheaper model via code_write, written to {target}.") diff --git a/litellm/proxy/shunt_endpoints/endpoints.py b/litellm/proxy/shunt_endpoints/endpoints.py index e6c6a0c0093..f6941465fe8 100644 --- a/litellm/proxy/shunt_endpoints/endpoints.py +++ b/litellm/proxy/shunt_endpoints/endpoints.py @@ -10,6 +10,7 @@ model chosen by the caller. The call goes through `llm_router.acompletion`, not caller's key/team exactly like any other request. """ +from collections.abc import Sequence from enum import Enum from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Final @@ -20,6 +21,7 @@ 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.guardrails.auto_router_shunt import ShuntConfig, shunt_config_for_model +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.shunt_endpoints.worker import ( BULK_READ_SYSTEM_PROMPT, CODE_WRITE_SYSTEM_PROMPT, @@ -41,18 +43,22 @@ _AUTH_DEPENDENCIES: Final = [Depends(user_api_key_auth)] # mutable-ok: FastAPI' _SHUNT_TAGS: Final[list[str | Enum]] = ["shunt"] # mutable-ok: FastAPI's `tags=` takes an invariant list -def _worker_config(model_alias: str, user_api_key_dict: UserAPIKeyAuth) -> tuple["Router", ShuntConfig]: +def _worker_config( + model_alias: str, user_api_key_dict: UserAPIKeyAuth, request_tags: Sequence[str] +) -> tuple["Router", ShuntConfig]: from litellm.proxy.proxy_server import llm_router if llm_router is None: raise ProxyException( message="LLM Router not found", type=ProxyErrorTypes.internal_server_error, param=None, code=500 ) + # `request_tags` comes from the query string the rewrite generated, so a marker armed only + # under a tag resolves here the same way it did when the original request was rewritten. config: Final = shunt_config_for_model( llm_router=llm_router, model_alias=model_alias, team_id=user_api_key_dict.team_id, - request_tags=(), + request_tags=request_tags, ) if config is None: raise ProxyException( @@ -65,21 +71,37 @@ def _worker_config(model_alias: str, user_api_key_dict: UserAPIKeyAuth) -> tuple async def _worker_text( - llm_router: "Router", *, model: str, system_prompt: str, message: str, team_id: str | None, label: str + llm_router: "Router", + *, + model: str, + system_prompt: str, + message: str, + user_api_key_dict: UserAPIKeyAuth, + label: str, ) -> 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, team metadata) lives here once. + the `acompletion` shape (temperature, non-streaming, caller attribution) lives here once. + + 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. """ system: Final = ChatCompletionSystemMessage(role="system", content=system_prompt) user: Final = ChatCompletionUserMessage(role="user", content=message) + key_metadata: Final = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( + user_api_key_dict=user_api_key_dict + ) 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={"user_api_key_team_id": team_id}, # mutable-ok: same + metadata={ # mutable-ok: same + **key_metadata, + "user_api_key": user_api_key_dict.api_key, + }, ) text: Final = response.choices[0].message.content if not isinstance(text, str): @@ -116,13 +138,14 @@ 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(user_api_key_auth)], + tags: Annotated[list[str] | None, Query()] = None, ) -> str: """Summarize or answer a question about one or more files via a cheap worker model. The paths shunt's generated command uploads are read here as plain UTF-8 text and never written to disk; only the text and the question reach the worker model. """ - llm_router, config = _worker_config(router_name, user_api_key_dict) + llm_router, config = _worker_config(router_name, user_api_key_dict, tags or ()) files: Final = MappingProxyType({upload.filename or "unnamed": await _read_upload_text(upload) for upload in paths}) message: Final = build_bulk_read_message(question=question, files=files) @@ -133,7 +156,7 @@ async def bulk_read( model=config.bulk_read_model, system_prompt=BULK_READ_SYSTEM_PROMPT, message=message, - team_id=user_api_key_dict.team_id, + user_api_key_dict=user_api_key_dict, label="bulk_read", ) @@ -154,6 +177,7 @@ async def code_write( spec: Annotated[str, Form()], reference: Annotated[UploadFile, File()], user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + tags: Annotated[list[str] | None, Query()] = None, ) -> str: """Generate boilerplate code matching a reference file's patterns, via a cheap worker model. @@ -161,7 +185,7 @@ async def code_write( it runs on the file's own machine; this has no local filesystem, so the client writes the returned text itself (the `Bash` rewrite this backs redirects the curl output to `target`). """ - llm_router, config = _worker_config(router_name, user_api_key_dict) + llm_router, config = _worker_config(router_name, user_api_key_dict, tags or ()) reference_content: Final = await _read_upload_text(reference) message: Final = build_code_write_message( @@ -175,7 +199,7 @@ async def code_write( model=config.code_write_model, system_prompt=CODE_WRITE_SYSTEM_PROMPT, message=message, - team_id=user_api_key_dict.team_id, + user_api_key_dict=user_api_key_dict, label="code_write", ) ) 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 c2d2ed35886..e3575256364 100644 --- a/tests/test_litellm/proxy/guardrails/test_auto_router_shunt.py +++ b/tests/test_litellm/proxy/guardrails/test_auto_router_shunt.py @@ -38,6 +38,44 @@ def _marker(shunt_fields: dict[str, Any], tags: list[str] | None = None) -> dict } +def _tiered_marker(shunt_fields: dict[str, Any], simple: list[str] | None) -> dict[str, Any]: + """A marker with a SIMPLE tier and no default model, to isolate the tier fallback.""" + return { + "model_name": "shunt", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": {"tiers": {"SIMPLE": simple} if simple is not None else {}}, + **shunt_fields, + }, + } + + +# Regression: an unset worker model fell through to "" when the marker had no default model, so +# the delegated call ran against an empty model name. The UI and preset docs promise SIMPLE. +class TestWorkerModelFallback: + def test_unset_worker_models_fall_back_to_the_simple_tier(self): + router = _FakeRouter([_tiered_marker({"auto_router_shunt_min_lines": 350}, ["claude-haiku-4-5"])]) + config = shunt_config_for_model(llm_router=router, model_alias="shunt", team_id=None, request_tags=()) + assert config.bulk_read_model == "claude-haiku-4-5" + assert config.code_write_model == "claude-haiku-4-5" + + def test_simple_tier_wins_over_the_default_model(self): + marker = _marker({"auto_router_shunt_min_lines": 350}) + marker["litellm_params"]["complexity_router_config"] = {"tiers": {"SIMPLE": ["gpt-5.6-luna"]}} + config = shunt_config_for_model( + llm_router=_FakeRouter([marker]), model_alias="shunt", team_id=None, request_tags=() + ) + assert config.bulk_read_model == "gpt-5.6-luna" + + def test_no_tier_and_no_default_model_is_unarmed_rather_than_empty(self): + router = _FakeRouter([_tiered_marker({"auto_router_shunt_min_lines": 350}, None)]) + assert shunt_config_for_model(llm_router=router, model_alias="shunt", team_id=None, request_tags=()) is None + + def test_empty_tier_list_and_no_default_model_is_unarmed(self): + router = _FakeRouter([_tiered_marker({"auto_router_shunt_min_lines": 350}, [])]) + assert shunt_config_for_model(llm_router=router, model_alias="shunt", team_id=None, request_tags=()) is None + + class TestShuntConfigForModel: def test_no_router_returns_none(self): assert shunt_config_for_model(llm_router=None, model_alias="shunt", team_id=None, request_tags=()) is None @@ -153,6 +191,68 @@ class TestAsyncPreCallHook: assert BULK_READ_TOOL_NAME in tool_names +# Regression: a caller that already had its own `bulk_read`/`code_write` tool got a duplicate +# definition injected, and its own tool calls silently rewritten into shunt's curl. +class TestCallerOwnedToolNamesAreLeftAlone: + @pytest.mark.asyncio + async def test_pre_call_declines_to_inject_over_an_anthropic_shape_collision(self, monkeypatch): + import litellm.proxy.guardrails.auto_router_shunt as mod + + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", _FakeRouter([_marker({"auto_router_shunt_min_lines": 350})])) + caller_tool = {"name": BULK_READ_TOOL_NAME, "description": "the caller's own", "input_schema": {}} + data = {"model": "shunt", "tools": [caller_tool]} + result = await mod.ShuntGuardrail().async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + assert result is None + assert data["tools"] == [caller_tool] + + @pytest.mark.asyncio + async def test_pre_call_declines_to_inject_over_an_openai_shape_collision(self, monkeypatch): + import litellm.proxy.guardrails.auto_router_shunt as mod + + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", _FakeRouter([_marker({"auto_router_shunt_min_lines": 350})])) + caller_tool = {"type": "function", "function": {"name": CODE_WRITE_TOOL_NAME, "parameters": {}}} + data = {"model": "shunt", "tools": [caller_tool]} + result = await mod.ShuntGuardrail().async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="acompletion" + ) + assert result is None + assert data["tools"] == [caller_tool] + + @pytest.mark.asyncio + async def test_post_call_leaves_the_callers_own_tool_call_unrewritten(self, monkeypatch): + import litellm.proxy.guardrails.auto_router_shunt as mod + + monkeypatch.setattr(mod, "_resolve_shunt_config", lambda data: self._config()) + data = { + "model": "shunt", + "proxy_server_request": {"url": "http://localhost:4000/v1/messages"}, + "tools": [{"name": BULK_READ_TOOL_NAME, "input_schema": {}}], + } + guardrail = mod.ShuntGuardrail() + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + response = { + "content": [ + { + "type": "tool_use", + "id": "t1", + "name": BULK_READ_TOOL_NAME, + "input": {"question": "q", "paths": ["a.py"]}, + } + ] + } + result = await guardrail.async_post_call_success_hook( + data=data, user_api_key_dict=None, response=response + ) + assert result["content"][0]["name"] == BULK_READ_TOOL_NAME + + def _config(self) -> ShuntConfig: + return ShuntConfig(min_lines=350, bulk_read_model="claude-haiku-4-5", code_write_model="claude-haiku-4-5") + + class TestAsyncPostCallSuccessHookAnthropicShape: def _config(self) -> ShuntConfig: return ShuntConfig(min_lines=350, bulk_read_model="claude-haiku-4-5", code_write_model="claude-haiku-4-5") @@ -242,8 +342,8 @@ class TestAsyncPostCallSuccessHookAnthropicShape: ) block = result["content"][0] assert block["name"] == "Bash" - assert "paths[]=@a.py" in block["input"]["command"] - assert "paths[]=@b.py" in block["input"]["command"] + assert "paths=@a.py" in block["input"]["command"] + assert "paths=@b.py" in block["input"]["command"] @pytest.mark.asyncio async def test_rewrites_explicit_code_write_call_with_target(self, monkeypatch): @@ -266,7 +366,7 @@ class TestAsyncPostCallSuccessHookAnthropicShape: ) block = result["content"][0] assert block["name"] == "Bash" - assert '> "tests/x_test.py"' in block["input"]["command"] + assert block["input"]["command"].endswith("> tests/x_test.py") @pytest.mark.asyncio async def test_rewrites_bare_bash_read_command(self, monkeypatch): diff --git a/tests/test_litellm/proxy/guardrails/test_shunt_rewrite.py b/tests/test_litellm/proxy/guardrails/test_shunt_rewrite.py index df448ba5958..6d4dcf9f3cf 100644 --- a/tests/test_litellm/proxy/guardrails/test_shunt_rewrite.py +++ b/tests/test_litellm/proxy/guardrails/test_shunt_rewrite.py @@ -1,6 +1,7 @@ """Unit tests for litellm.proxy.guardrails.shunt_rewrite.""" import subprocess +from pathlib import Path import pytest @@ -76,14 +77,12 @@ def _bounded_read( question: str = "Summarize this file's structure.", min_lines: int = 350, bulk_read_endpoint: str = "http://localhost:4000/v1/bulk_read", - auth_header: str = "Bearer sk-1234", ) -> ShuntBashRewrite: return build_bounded_read_command( path=path, question=question, min_lines=min_lines, bulk_read_endpoint=bulk_read_endpoint, - auth_header=auth_header, ) @@ -107,7 +106,7 @@ class TestBuildBoundedReadCommand: assert endpoint in _bounded_read(bulk_read_endpoint=endpoint).command def test_command_falls_back_to_cat_for_small_files(self): - assert 'else cat "litellm/router.py"; fi' in _bounded_read().command + assert "else cat litellm/router.py; fi" in _bounded_read().command def test_command_escapes_embedded_double_quotes_in_path(self): _assert_valid_bash(_bounded_read(path='weird"file.py').command) @@ -116,13 +115,170 @@ class TestBuildBoundedReadCommand: assert "200" in _bounded_read(min_lines=200).note +# Every value in a generated command comes from the model and is run by a shell on the +# developer's own machine, so a shell metacharacter in any of them must stay inert data. +# Regression: the values were interpolated inside double quotes with only `"` escaped, so a +# path of `$(cmd)` was command-substituted and ran `cmd` locally. +_INJECTIONS = [ + "$(touch {marker})", + "`touch {marker}`", + "; touch {marker}", + "&& touch {marker}", + "| touch {marker}", + "$(touch {marker})'; touch {marker}; '", + 'x" ; touch {marker} ; "', + "\n touch {marker} \n", +] + + +def _assert_runs_without_side_effect(command: str, marker: Path) -> None: + """Run `command` in a real bash and assert the injected marker file was never created.""" + subprocess.run(["bash", "-c", command], capture_output=True, text=True, check=False, timeout=30) + assert not marker.exists(), f"injection executed, marker created by: {command}" + + +@pytest.mark.parametrize("payload", _INJECTIONS) +class TestGeneratedCommandsResistShellInjection: + def test_bounded_read_path_is_inert(self, payload: str, tmp_path: Path): + marker = tmp_path / "pwned_path" + rewrite = _bounded_read(path=payload.format(marker=marker)) + _assert_runs_without_side_effect(rewrite.command, marker) + + def test_bounded_read_question_is_inert(self, payload: str, tmp_path: Path): + marker = tmp_path / "pwned_question" + rewrite = _bounded_read(question=payload.format(marker=marker)) + _assert_runs_without_side_effect(rewrite.command, marker) + + def test_bounded_read_endpoint_is_inert(self, payload: str, tmp_path: Path): + marker = tmp_path / "pwned_endpoint" + rewrite = _bounded_read(bulk_read_endpoint=payload.format(marker=marker)) + _assert_runs_without_side_effect(rewrite.command, marker) + + def test_bulk_read_paths_are_inert(self, payload: str, tmp_path: Path): + marker = tmp_path / "pwned_bulk" + rewrite = build_bulk_read_command( + question="q", + paths=["ok.py", payload.format(marker=marker)], + bulk_read_endpoint="http://127.0.0.1:9/v1/bulk_read", + ) + _assert_runs_without_side_effect(rewrite.command, marker) + + def test_code_write_target_is_inert(self, payload: str, tmp_path: Path): + marker = tmp_path / "pwned_target" + rewrite = build_code_write_command( + spec="s", + reference="r.py", + target=payload.format(marker=marker), + code_write_endpoint="http://127.0.0.1:9/v1/code_write", + ) + _assert_runs_without_side_effect(rewrite.command, marker) + + def test_code_write_spec_is_inert(self, payload: str, tmp_path: Path): + marker = tmp_path / "pwned_spec" + rewrite = build_code_write_command( + spec=payload.format(marker=marker), + reference="r.py", + target=None, + code_write_endpoint="http://127.0.0.1:9/v1/code_write", + ) + _assert_runs_without_side_effect(rewrite.command, marker) + + +# 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 + assert "sk-" not in rewrite.command + + def test_code_write_references_the_env_var_instead_of_a_secret(self): + rewrite = build_code_write_command( + spec="s", reference="r.py", target=None, code_write_endpoint="http://localhost:4000/v1/code_write" + ) + assert "ANTHROPIC_AUTH_TOKEN" 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.""" + 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" + ) + 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_API_KEY": "sk-from-api-key", "PATH": "/usr/bin:/bin"}, + ) + assert out.read_text() == "Authorization: Bearer sk-from-api-key" + + +# Regression: the commands uploaded files as `paths[]`, but the endpoint binds them under +# `paths`. FastAPI matches the form name exactly, so every delegated read 422'd. +class TestUploadFieldNameMatchesTheEndpoint: + def test_bounded_read_uses_the_bare_paths_field_name(self): + command = _bounded_read().command + assert "paths=@" in command + assert "paths[]=@" not in command + + 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" + ) + assert "paths[]=@" not in rewrite.command + for path in ("a.py", "b.py"): + assert f"paths=@{path}" in rewrite.command + + +def test_bounded_read_still_reads_a_small_file_verbatim(tmp_path: Path): + """The quoting must not break the real path: a small file is still cat'd through.""" + target = tmp_path / "small.py" + target.write_text("line one\nline two\n") + rewrite = _bounded_read(path=str(target), min_lines=350) + result = subprocess.run(["bash", "-c", rewrite.command], capture_output=True, text=True, check=False) + assert result.stdout == "line one\nline two\n" + + +def test_bounded_read_reports_a_path_containing_a_dollar_sign_literally(tmp_path: Path): + target = tmp_path / "odd$name.py" + target.write_text("x\n") + rewrite = _bounded_read(path=str(target), min_lines=350) + result = subprocess.run(["bash", "-c", rewrite.command], capture_output=True, text=True, check=False) + assert result.stdout == "x\n" + + class TestBuildBulkReadCommand: def test_unconditional_no_size_check(self): rewrite = build_bulk_read_command( question="what does this do", paths=["a.py", "b.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read", - auth_header="Bearer sk-1234", ) assert "wc -l" not in rewrite.command assert "if [" not in rewrite.command @@ -132,14 +288,13 @@ class TestBuildBulkReadCommand: question="q", paths=["a.py", "b.py", "c.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read", - auth_header="Bearer sk-1234", ) for path in ("a.py", "b.py", "c.py"): - assert f"paths[]=@{path}" in rewrite.command + 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", auth_header="x" + question="q", paths=["a.py"], bulk_read_endpoint="http://localhost:4000/v1/bulk_read" ) _assert_valid_bash(rewrite.command) @@ -151,7 +306,6 @@ class TestBuildCodeWriteCommand: reference="tests/y_test.py", target=None, code_write_endpoint="http://localhost:4000/v1/code_write", - auth_header="Bearer sk-1234", ) assert ">" not in rewrite.command @@ -161,13 +315,12 @@ class TestBuildCodeWriteCommand: reference="tests/y_test.py", target="tests/x_test.py", code_write_endpoint="http://localhost:4000/v1/code_write", - auth_header="Bearer sk-1234", ) - assert '> "tests/x_test.py"' in rewrite.command + 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", auth_header="x" + spec="s", reference="r.py", target=target, code_write_endpoint="http://x/v1/code_write" ) _assert_valid_bash(rewrite.command) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c0d3a242f74..329ba0c10b9 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -42615,6 +42615,7 @@ export interface operations { parameters: { query: { router: string; + tags?: string[] | null; }; header?: never; path?: never; @@ -43508,6 +43509,7 @@ export interface operations { parameters: { query: { router: string; + tags?: string[] | null; }; header?: never; path?: never; @@ -61521,6 +61523,7 @@ export interface operations { parameters: { query: { router: string; + tags?: string[] | null; }; header?: never; path?: never; @@ -61755,6 +61758,7 @@ export interface operations { parameters: { query: { router: string; + tags?: string[] | null; }; header?: never; path?: never;