fix(relay): preserve provider request extensions

Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
Alex Fournier 2026-07-23 13:24:19 -07:00
parent a6fbff1d7e
commit 9dfe690125
2 changed files with 54 additions and 11 deletions

View file

@ -745,17 +745,14 @@ def _provider_request(
if _json_equal(content, relay_request_body):
final = dict(original)
else:
final = _provider_request_body(content, metadata)
# Codec-only normalization must not silently change the provider wire
# request when an unrelated interceptor edits another field.
for key, value in original.items():
if key not in relay_request_body and value is None:
final.setdefault(key, value)
elif (
key in relay_request_body
and key in final
and _json_equal(final[key], relay_request_body[key])
):
baseline = _provider_request_body(relay_request_body, metadata)
intercepted = _provider_request_body(content, metadata)
final = dict(original)
# Typed codecs may not represent provider-specific fields. Overlay only
# values that changed from the codec-facing baseline so unrelated
# intercepts cannot delete or normalize unknown provider arguments.
for key, value in intercepted.items():
if key not in baseline or not _json_equal(value, baseline[key]):
final[key] = value
_restore_provider_message_extensions(original, final)
headers = getattr(request, "headers", None)

View file

@ -881,3 +881,49 @@ def test_request_rewrite_preserves_unmodified_provider_objects(relay_turn):
assert captured_requests[0]["timeout"] is timeout
assert captured_requests[0]["temperature"] == 0.25
def test_request_rewrite_preserves_fields_dropped_by_codec(relay_turn, monkeypatch):
relay, _turn = relay_turn
captured_requests = []
vendor_body = {
"routing": {"provider": "nim", "region": "us-west-2"},
"trace_vendor_request": False,
}
async def lossy_execute(_name, request, callback, **_kwargs):
content = {
key: value
for key, value in request.content.items()
if key != "extra_body"
}
content["temperature"] = 0.25
return callback(relay.LLMRequest(request.headers, content))
monkeypatch.setattr(relay.llm, "execute", lossy_execute)
relay_llm.execute(
{
"model": "test-model",
"messages": [],
"temperature": 0.0,
"extra_body": vendor_body,
},
lambda request: captured_requests.append(request) or {"content": "ok"},
session_id="session-1",
name="test-provider",
model_name="test-model",
metadata={
"api_mode": "chat_completions",
"api_request_id": "request-lossy-codec",
},
)
assert captured_requests == [
{
"model": "test-model",
"messages": [],
"temperature": 0.25,
"extra_body": vendor_body,
}
]