From 4b1245e3fd7b0cb35b6c222e01c56d0afb871eaa Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 22 Aug 2026 21:55:07 +0200 Subject: [PATCH] allow only anthropic fields --- routstr/upstream/messages_dispatch.py | 38 ++++++++++-- tests/unit/test_messages_litellm_dispatch.py | 64 ++++++++++++++++++++ 2 files changed, 96 insertions(+), 6 deletions(-) diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index 3d689922..9ef7f257 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -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 diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index cb3fe9d7..ec93d555 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -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