From bfae9c2572923eebbf4a8c2895035a279fc56d15 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Mon, 27 Jul 2026 08:17:21 -0700 Subject: [PATCH] fix(relay): ignore rewrites after codec baseline failure Signed-off-by: Alex Fournier --- agent/relay_llm.py | 18 +++++++--- tests/agent/test_relay_llm.py | 66 +++++++++++++++++++++++++++++++++++ 2 files changed, 79 insertions(+), 5 deletions(-) diff --git a/agent/relay_llm.py b/agent/relay_llm.py index e139c23baca..c6c6affa78c 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -855,13 +855,13 @@ def _provider_request( request: Any, *, relay_request_body: dict[str, Any], - codec_baseline_body: dict[str, Any], + codec_baseline_body: dict[str, Any] | None, metadata: dict[str, Any] | None, ) -> dict[str, Any]: content = getattr(request, "content", request) if not isinstance(content, dict): content = relay_request_body - if _json_equal(content, relay_request_body): + if codec_baseline_body is None or _json_equal(content, relay_request_body): final = dict(original) else: baseline = codec_baseline_body @@ -1006,7 +1006,7 @@ def _codec_round_trip_request_body( *, relay_request_body: dict[str, Any], metadata: dict[str, Any] | None, -) -> dict[str, Any]: +) -> dict[str, Any] | None: """Return the codec-only request shape used to identify real rewrites.""" codec = _codec(relay, metadata) if codec is None: @@ -1018,8 +1018,16 @@ def _codec_round_trip_request_body( 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) + logger.warning( + "NeMo Relay request codec baseline failed; ignoring request rewrites", + exc_info=True, + ) + return None + logger.warning( + "NeMo Relay request codec returned an unsupported baseline; " + "ignoring request rewrites" + ) + return None def _provider_request_body( diff --git a/tests/agent/test_relay_llm.py b/tests/agent/test_relay_llm.py index d286f8af1a3..32e4788a974 100644 --- a/tests/agent/test_relay_llm.py +++ b/tests/agent/test_relay_llm.py @@ -1349,6 +1349,72 @@ def test_request_rewrite_preserves_fields_dropped_by_codec(relay_turn, monkeypat ] +def test_request_rewrite_is_ignored_when_codec_baseline_fails( + relay_turn, monkeypatch +): + relay, _turn = relay_turn + captured_requests = [] + original = { + "model": "test-model", + "messages": [], + "temperature": 0.0, + "extra_body": {"routing": {"provider": "nim"}}, + } + + async def lossy_execute(_name, request, callback, **_kwargs): + rewritten = { + key: value + for key, value in request.content.items() + if key != "extra_body" + } + rewritten["temperature"] = 0.25 + return callback(relay.LLMRequest(request.headers, rewritten)) + + monkeypatch.setattr(relay.llm, "execute", lossy_execute) + monkeypatch.setattr( + relay_llm, + "_codec_round_trip_request_body", + lambda *_args, **_kwargs: None, + ) + + relay_llm.execute( + original, + 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-codec-failure", + }, + ) + + assert captured_requests == [original] + + +def test_codec_baseline_failure_is_explicit(relay_turn, monkeypatch, caplog): + relay, _turn = relay_turn + request_body = {"model": "test-model", "messages": []} + request = relay.LLMRequest({}, request_body) + + class FailingCodec: + def decode(self, _request): + raise RuntimeError("simulated codec failure") + + monkeypatch.setattr(relay_llm, "_codec", lambda *_args, **_kwargs: FailingCodec()) + + with caplog.at_level("WARNING", logger="agent.relay_llm"): + baseline = relay_llm._codec_round_trip_request_body( + relay, + request, + relay_request_body=request_body, + metadata={"api_mode": "chat_completions"}, + ) + + assert baseline is None + assert "ignoring request rewrites" in caplog.text + + def test_request_rewrite_can_remove_codec_represented_field(relay_turn, monkeypatch): relay, _turn = relay_turn captured_requests = []