fix(relay): preserve intentional request removals

Signed-off-by: Alex Fournier <afournier@nvidia.com>
This commit is contained in:
Alex Fournier 2026-07-23 14:22:49 -07:00
parent 06ed41aa1f
commit a3ef27ab70
2 changed files with 150 additions and 14 deletions

View file

@ -43,6 +43,12 @@ def execute(
relay_request_body = _relay_request_body(request, metadata)
relay_request = runtime.relay.LLMRequest({}, relay_request_body)
codec_baseline_body = _codec_round_trip_request_body(
runtime.relay,
relay_request,
relay_request_body=relay_request_body,
metadata=metadata,
)
raw_response: dict[str, Any] = {}
callback_error: BaseException | None = None
callback_context = contextvars.copy_context()
@ -54,6 +60,7 @@ def execute(
request,
next_request,
relay_request_body=relay_request_body,
codec_baseline_body=codec_baseline_body,
metadata=metadata,
)
raw = callback_context.copy().run(callback, final_request)
@ -121,6 +128,12 @@ async def execute_async(
relay_request_body = _relay_request_body(request, metadata)
relay_request = runtime.relay.LLMRequest({}, relay_request_body)
codec_baseline_body = _codec_round_trip_request_body(
runtime.relay,
relay_request,
relay_request_body=relay_request_body,
metadata=metadata,
)
raw_response: dict[str, Any] = {}
callback_error: BaseException | None = None
callback_context = contextvars.copy_context()
@ -132,6 +145,7 @@ async def execute_async(
request,
next_request,
relay_request_body=relay_request_body,
codec_baseline_body=codec_baseline_body,
metadata=metadata,
)
async def call_provider() -> Any:
@ -354,6 +368,12 @@ class ManagedLlmStream(Iterator[Any]):
parent = self._logical[1]
relay_request_body = _relay_request_body(request, metadata)
relay_request = runtime.relay.LLMRequest({}, relay_request_body)
codec_baseline_body = _codec_round_trip_request_body(
runtime.relay,
relay_request,
relay_request_body=relay_request_body,
metadata=metadata,
)
async def provider_stream(next_request: Any):
raw_stream = None
@ -364,6 +384,7 @@ class ManagedLlmStream(Iterator[Any]):
request,
next_request,
relay_request_body=relay_request_body,
codec_baseline_body=codec_baseline_body,
metadata=metadata,
)
)
@ -810,6 +831,7 @@ def _provider_request(
request: Any,
*,
relay_request_body: dict[str, Any],
codec_baseline_body: dict[str, Any],
metadata: dict[str, Any] | None,
) -> dict[str, Any]:
content = getattr(request, "content", request)
@ -818,16 +840,26 @@ def _provider_request(
if _json_equal(content, relay_request_body):
final = dict(original)
else:
baseline = _provider_request_body(relay_request_body, metadata)
baseline = codec_baseline_body
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)
for key in baseline.keys() | intercepted.keys():
if key not in intercepted:
final.pop(key, None)
elif key not in baseline or not _json_equal(
intercepted[key],
baseline[key],
):
final[key] = intercepted[key]
_restore_provider_message_extensions(
original,
final,
baseline=baseline,
intercepted=intercepted,
)
headers = getattr(request, "headers", None)
if isinstance(headers, dict):
headers = {
@ -889,25 +921,83 @@ def _relay_request_body(
def _restore_provider_message_extensions(
original: dict[str, Any], final: dict[str, Any]
original: dict[str, Any],
final: dict[str, Any],
*,
baseline: dict[str, Any],
intercepted: dict[str, Any],
) -> None:
"""Restore provider wire fields that Relay's typed codec cannot represent."""
original_messages = original.get("messages")
final_messages = final.get("messages")
if not isinstance(original_messages, list) or not isinstance(final_messages, list):
return
if len(original_messages) != len(final_messages):
return
for original_message, final_message in zip(
original_messages, final_messages, strict=True
baseline_messages = baseline.get("messages")
intercepted_messages = intercepted.get("messages")
if not all(
isinstance(messages, list)
for messages in (
original_messages,
final_messages,
baseline_messages,
intercepted_messages,
)
):
if not isinstance(original_message, dict) or not isinstance(final_message, dict):
return
if not (
len(original_messages)
== len(final_messages)
== len(baseline_messages)
== len(intercepted_messages)
):
return
for original_message, final_message, baseline_message, intercepted_message in zip(
original_messages,
final_messages,
baseline_messages,
intercepted_messages,
strict=True,
):
if not all(
isinstance(message, dict)
for message in (
original_message,
final_message,
baseline_message,
intercepted_message,
)
):
continue
for key in _PROVIDER_MESSAGE_EXTENSION_KEYS:
if key in original_message and key not in final_message:
if (
key in original_message
and key not in baseline_message
and key not in intercepted_message
and key not in final_message
):
final_message[key] = original_message[key]
def _codec_round_trip_request_body(
relay: Any,
relay_request: Any,
*,
relay_request_body: dict[str, Any],
metadata: dict[str, Any] | None,
) -> dict[str, Any]:
"""Return the codec-only request shape used to identify real rewrites."""
codec = _codec(relay, metadata)
if codec is None:
return _provider_request_body(relay_request_body, metadata)
try:
annotated = codec.decode(relay_request)
encoded = codec.encode(annotated, relay_request)
content = getattr(encoded, "content", encoded)
if isinstance(content, dict):
return _provider_request_body(content, metadata)
except Exception:
logger.debug("NeMo Relay request codec baseline failed", exc_info=True)
return _provider_request_body(relay_request_body, metadata)
def _provider_request_body(
content: dict[str, Any], metadata: dict[str, Any] | None
) -> dict[str, Any]:

View file

@ -1228,6 +1228,15 @@ def test_request_rewrite_preserves_fields_dropped_by_codec(relay_turn, monkeypat
return callback(relay.LLMRequest(request.headers, content))
monkeypatch.setattr(relay.llm, "execute", lossy_execute)
monkeypatch.setattr(
relay_llm,
"_codec_round_trip_request_body",
lambda *_args, relay_request_body, **_kwargs: {
key: value
for key, value in relay_request_body.items()
if key != "extra_body"
},
)
relay_llm.execute(
{
@ -1254,3 +1263,40 @@ def test_request_rewrite_preserves_fields_dropped_by_codec(relay_turn, monkeypat
"extra_body": vendor_body,
}
]
def test_request_rewrite_can_remove_codec_represented_field(relay_turn, monkeypatch):
relay, _turn = relay_turn
captured_requests = []
async def remove_temperature(_name, request, callback, **_kwargs):
content = dict(request.content)
content.pop("temperature")
return callback(relay.LLMRequest(request.headers, content))
monkeypatch.setattr(relay.llm, "execute", remove_temperature)
relay_llm.execute(
{
"model": "test-model",
"messages": [],
"temperature": 0.25,
"extra_body": {"routing": {"provider": "nim"}},
},
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-remove-field",
},
)
assert captured_requests == [
{
"model": "test-model",
"messages": [],
"extra_body": {"routing": {"provider": "nim"}},
}
]