fix(auto_router): close shunt command-injection, key-leak, and delegation bugs

Addresses the review findings on this PR.

The generated commands escaped only double quotes, so a model-supplied path, question, spec,
reference, or target containing $(...) or backticks was command-substituted and ran on the
developer's machine. Every interpolated value now goes through shlex.quote, and the size-check
line reports the path with printf instead of a double-quoted echo. Confirmed against a real
bash: the payload used to create its marker file, and no longer does.

The commands also carried the caller's Authorization header verbatim, which put the key in the
model's response, the conversation history, and the next upstream turn. They now reference
${ANTHROPIC_AUTH_TOKEN:-$ANTHROPIC_API_KEY} and the shell resolves it locally, so the secret
never leaves the client. That also fixes clients authenticating with x-api-key, which were
skipped entirely because only the authorization header was read.

Files were uploaded as paths[] while the endpoint binds them under paths. FastAPI matches the
form name exactly, so every delegated read returned 422 and the feature never actually worked.

The worker endpoints resolved the marker with no tags, so a marker armed only under a tag
returned 400 even though the rewrite had matched it; the tags now ride along in the query
string. They also sent only a team id to the router, so worker spend was not attributed to the
calling key. They now reuse the proxy's own key-metadata builder.

An unset worker model could fall through to an empty string. It now falls back to the SIMPLE
tier, then the default model, and refuses to arm rather than naming no model at all.

A caller that already defines its own bulk_read or code_write got a duplicate definition
injected and its own tool calls rewritten into shunt's curl. Shunt now leaves those requests
alone in all three hooks.

The allow rules only matched commands starting with curl, but the bounded read starts with the
wc -l size check, so the main large-file path still prompted. Added a rule for that shape.

Registers the new routes in the gateway allowlist, which the component-coverage test requires.
This commit is contained in:
moe-berri 2026-09-07 14:26:03 -07:00
parent 56009d32dd
commit a71dd731f2
8 changed files with 439 additions and 94 deletions

View file

@ -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",

View file

@ -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=*)",
)

View file

@ -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)

View file

@ -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}.")

View file

@ -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",
)
)

View file

@ -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):

View file

@ -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)

View file

@ -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;