mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge pull request #681 from Routstr/clean-up-forwareded-fields
allow only anthropic fields
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user