From fca1b0bc4c7a7db9899ab7ce9b6bad7f780d7684 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 21:13:25 -0700 Subject: [PATCH] refactor(agent/relay_llm): accumulator/translator micro-collapses --- agent/relay_llm.py | 26 ++++++++++---------------- 1 file changed, 10 insertions(+), 16 deletions(-) diff --git a/agent/relay_llm.py b/agent/relay_llm.py index 729dff8fcc..b599d3cfee 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -567,11 +567,10 @@ class AnthropicStreamAccumulator: def observe(self, event: Any) -> None: payload = _jsonable(event) - if not isinstance(payload, dict): - return - handler = self._EVENT_HANDLERS.get(payload.get("type")) - if handler is not None: - handler(self, payload) + if isinstance(payload, dict): + handler = self._EVENT_HANDLERS.get(payload.get("type")) + if handler is not None: + handler(self, payload) def _on_message_start(self, payload: dict[str, Any]) -> None: message = payload.get("message") @@ -579,14 +578,12 @@ class AnthropicStreamAccumulator: self._message.update({k: message[k] for k in ("id", "type", "role", "model", "usage") if k in message}) def _on_content_block_start(self, payload: dict[str, Any]) -> None: - index = payload.get("index") - block = payload.get("content_block") + index, block = payload.get("index"), payload.get("content_block") if isinstance(index, int) and isinstance(block, dict): self._blocks[index] = dict(block) def _on_content_block_delta(self, payload: dict[str, Any]) -> None: - index = payload.get("index") - delta = payload.get("delta") + index, delta = payload.get("index"), payload.get("delta") if not isinstance(index, int) or not isinstance(delta, dict): return block = self._blocks.setdefault(index, {}) @@ -604,8 +601,7 @@ class AnthropicStreamAccumulator: if isinstance(delta, dict): self._message.update({k: delta[k] for k in ("stop_reason", "stop_sequence") if k in delta}) if "usage" in payload: - usage = payload["usage"] - current_usage = self._message.get("usage") + usage, current_usage = payload["usage"], self._message.get("usage") if isinstance(current_usage, dict) and isinstance(usage, dict): usage = {**current_usage, **usage} self._message["usage"] = usage @@ -797,9 +793,8 @@ def _restore_provider_message_extensions( continue original_message, final_message, baseline_message, intercepted_message = messages for key in _PROVIDER_MESSAGE_EXTENSION_KEYS: - if ( - key in original_message and key not in baseline_message - and key not in intercepted_message and key not in final_message + if key in original_message and not any( + key in m for m in (baseline_message, intercepted_message, final_message) ): final_message[key] = original_message[key] @@ -812,8 +807,7 @@ def _codec_round_trip_request_body( if codec is None: return _provider_request_body(relay_request_body, metadata) try: - annotated = codec.decode(relay_request) - encoded = codec.encode(annotated, relay_request) + encoded = codec.encode(codec.decode(relay_request), relay_request) content = getattr(encoded, "content", encoded) if isinstance(content, dict): return _provider_request_body(content, metadata)