Merge pull request #681 from Routstr/clean-up-forwareded-fields

allow only anthropic fields
This commit is contained in:
9qeklajc
2026-08-22 22:02:41 +02:00
committed by GitHub
2 changed files with 96 additions and 6 deletions
+32 -6
View File
@@ -40,6 +40,10 @@ logger = get_logger(__name__)
# unsupported params; these newer/extension fields get passed through
# verbatim and the upstream rejects them with a 400. Pop them here so the
# request reaches the upstream cleanly.
#
# Note: ``dispatch_anthropic_messages`` additionally enforces
# ``ALLOWED_MESSAGES_REQUEST_FIELDS``, which already excludes all of these.
# This tuple remains for ``gemini_messages``, which pops them explicitly.
ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = (
"thinking",
"cache_control",
@@ -51,6 +55,25 @@ ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = (
"anthropic_beta",
)
# Anthropic Messages API request fields forwarded from the client body
# into the upstream call. Only these are passed on; anything else is
# dropped so the forwarded request is deterministic and limited to the
# documented Messages surface.
ALLOWED_MESSAGES_REQUEST_FIELDS: frozenset[str] = frozenset(
{
"messages",
"max_tokens",
"system",
"temperature",
"top_p",
"top_k",
"stop_sequences",
"tools",
"tool_choice",
"metadata",
}
)
def coerce_litellm_payload(payload: object) -> dict:
"""Convert a litellm event into a plain dict.
@@ -466,15 +489,18 @@ async def dispatch_anthropic_messages(
client_stream = bool(body.pop("stream", False))
upstream_stream = True
dropped: dict[str, Any] = {}
for field in ANTHROPIC_ONLY_FIELDS:
if field in body:
dropped[field] = body.pop(field)
# Forward only allowlisted Anthropic Messages request fields. Any
# other client-supplied key is dropped so it cannot leak into the
# upstream request. See ALLOWED_MESSAGES_REQUEST_FIELDS.
dropped = sorted(set(body) - ALLOWED_MESSAGES_REQUEST_FIELDS)
if dropped:
logger.debug(
"Dropped anthropic-only fields before litellm dispatch",
extra={"dropped_keys": sorted(dropped.keys())},
"Dropped non-forwardable fields before litellm dispatch",
extra={"dropped_keys": dropped},
)
body = {
k: v for k, v in body.items() if k in ALLOWED_MESSAGES_REQUEST_FIELDS
}
# Convention: `model.id` is the canonical upstream model name;
# `forwarded_model_id` is the public alias the internal API exposes
@@ -305,6 +305,10 @@ async def test_dispatch_strips_anthropic_only_fields_before_litellm() -> None:
"service_tier": "auto",
"anthropic_beta": "abc",
"anthropic_version": "2023-06-01",
"api_base": "https://attacker.invalid",
"api_key": "client-controlled-key",
"custom_llm_provider": "client-controlled-provider",
"unexpected_field": "must-not-leak",
}
).encode()
@@ -342,6 +346,8 @@ async def test_dispatch_strips_anthropic_only_fields_before_litellm() -> None:
"service_tier",
"anthropic_beta",
"anthropic_version",
"custom_llm_provider",
"unexpected_field",
):
assert stripped not in forwarded, (
f"Anthropic-only field {stripped!r} leaked through to litellm"
@@ -349,6 +355,64 @@ async def test_dispatch_strips_anthropic_only_fields_before_litellm() -> None:
# Core fields preserved
assert forwarded["max_tokens"] == 64
assert forwarded["messages"] == [{"role": "user", "content": "hi"}]
# Dispatch-controlled values cannot be overridden by the client body.
assert forwarded["model"] == "openai/openai/gpt-4o-mini"
assert forwarded["api_base"] == "http://test"
assert forwarded["api_key"] == "upstream-key"
assert forwarded["stream"] is True
@pytest.mark.asyncio
async def test_dispatch_forwards_all_allowlisted_messages_fields() -> None:
provider = _make_provider()
key = _make_key()
model = _make_model()
session = _make_session()
allowed = {
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 64,
"system": "Be concise",
"temperature": 0.2,
"top_p": 0.9,
"top_k": 20,
"stop_sequences": ["STOP"],
"tools": [
{
"name": "lookup",
"description": "Look something up",
"input_schema": {"type": "object", "properties": {}},
}
],
"tool_choice": {"type": "tool", "name": "lookup"},
"metadata": {"user_id": "test-user"},
}
body = json.dumps({"model": model.id, **allowed}).encode()
captured: dict[str, Any] = {}
async def fake_acreate(**kwargs: Any) -> dict:
captured["kwargs"] = kwargs
return _anthropic_non_stream_response()
with (
patch(
"litellm.anthropic.messages.acreate",
new=AsyncMock(side_effect=fake_acreate),
),
patch(
"routstr.upstream.base.adjust_payment_for_tokens",
new=AsyncMock(return_value={"total_msats": 0, "total_usd": 0.0}),
),
):
await provider._forward_messages_via_litellm(
request_body=body,
key=key,
session=session,
max_cost_for_model=10_000,
model_obj=model,
)
forwarded = captured["kwargs"]
assert {field: forwarded[field] for field in allowed} == allowed
@pytest.mark.asyncio