diff --git a/agent/relay_llm.py b/agent/relay_llm.py index b49d527ee5..6f6a3358f4 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -3,12 +3,13 @@ from __future__ import annotations import asyncio +import contextlib import contextvars import inspect import json import logging from collections.abc import Callable, Iterator -from dataclasses import dataclass +from functools import partial from types import SimpleNamespace from typing import Any @@ -22,16 +23,11 @@ _RELAY_INTERNAL_PROVIDER_HEADERS = frozenset({"x-dynamo-parent-session-id", "x-d _LogicalCall = tuple[relay_runtime.RelayTurnContext, Any, str] -@dataclass(frozen=True, slots=True) -class _RelayProtocol: - operation: str - codec_class: str - - +# api_mode -> (Relay operation name, codec class name on ``relay.codecs``) _RELAY_PROTOCOL_BY_API_MODE = { - "chat_completions": _RelayProtocol("openai.chat_completions", "OpenAIChatCodec"), - "codex_responses": _RelayProtocol("openai.responses", "OpenAIResponsesCodec"), - "anthropic_messages": _RelayProtocol("anthropic.messages", "AnthropicMessagesCodec"), + "chat_completions": ("openai.chat_completions", "OpenAIChatCodec"), + "codex_responses": ("openai.responses", "OpenAIResponsesCodec"), + "anthropic_messages": ("anthropic.messages", "AnthropicMessagesCodec"), } @@ -39,16 +35,10 @@ def _api_mode(metadata: dict[str, Any] | None) -> str: return str((metadata or {}).get("api_mode") or "") -def _relay_protocol(metadata: dict[str, Any] | None) -> _RelayProtocol | None: - """Return Relay's operation and codec descriptor for an API mode.""" - api_mode = (metadata or {}).get("api_mode") - return _RELAY_PROTOCOL_BY_API_MODE.get(api_mode) if isinstance(api_mode, str) else None - - def _relay_operation_name(provider_name: str, metadata: dict[str, Any] | None) -> str: """Return Relay's canonical operation name when Hermes knows the API mode.""" - protocol = _relay_protocol(metadata) - return protocol.operation if protocol is not None else provider_name + protocol = _RELAY_PROTOCOL_BY_API_MODE.get(_api_mode(metadata)) + return protocol[0] if protocol is not None else provider_name def _relay_metadata(provider_name: str, metadata: dict[str, Any] | None) -> dict[str, Any]: @@ -63,30 +53,20 @@ class _ManagedAttempt: @classmethod def resolve( - cls, - session_id: str, - request: dict[str, Any], - metadata: dict[str, Any] | None, - *, - name: str, - model_name: str, + cls, session_id: str | None, request: dict[str, Any], metadata: dict[str, Any] | None, + *, name: str, model_name: str, ) -> "_ManagedAttempt | None": """Return the managed attempt for ``session_id``, or None to run unmanaged.""" + if not session_id: + return None runtime, session, parent = relay_runtime.resolve_execution_context(session_id) if runtime is None or session is None or not runtime.managed_execution_enabled(): return None return cls(runtime, session, parent, request, metadata, name=name, model_name=model_name) def __init__( - self, - runtime: relay_runtime.RelayRuntime, - session: Any, - parent: Any, - request: dict[str, Any], - metadata: dict[str, Any] | None, - *, - name: str, - model_name: str, + self, runtime: relay_runtime.RelayRuntime, session: Any, parent: Any, + request: dict[str, Any], metadata: dict[str, Any] | None, *, name: str, model_name: str, ) -> None: self.runtime = runtime self.session = session @@ -101,11 +81,8 @@ class _ManagedAttempt: ) self.operation = _relay_operation_name(name, metadata) self.relay_kwargs = { - "handle": self.parent, - "metadata": _relay_metadata(name, metadata), - "model_name": model_name, - "codec": _codec(runtime.relay, metadata), - "response_codec": _codec(runtime.relay, metadata), + "handle": self.parent, "metadata": _relay_metadata(name, metadata), "model_name": model_name, + "codec": _codec(runtime.relay, metadata), "response_codec": _codec(runtime.relay, metadata), } # Provider callback bookkeeping: "value"/"json" once it returned, "error" if it raised. self.raw_response: dict[str, Any] = {} @@ -113,149 +90,79 @@ class _ManagedAttempt: def provider_request(self, next_request: Any) -> dict[str, Any]: return _provider_request( - self.request, - next_request, - relay_request_body=self.body, - codec_baseline_body=self.codec_baseline, - metadata=self.metadata, + self.request, next_request, relay_request_body=self.body, + codec_baseline_body=self.codec_baseline, metadata=self.metadata, ) def run_callback(self, callback: Callable[..., Any], *args: Any) -> Any: """Run a Hermes callback in a fresh copy of the captured context. - - Relay can invoke callbacks while another one still owns the captured - Context, hence the copy. Nested relay calls inside a managed provider - callback must run unmanaged — see relay_runtime.managed_callback_guard. - """ + Relay can invoke callbacks while another still owns the captured Context (hence the + copy); nested relay calls run unmanaged — see relay_runtime.managed_callback_guard.""" def guarded() -> Any: with relay_runtime.managed_callback_guard(): return callback(*args) return self.context.copy().run(guarded) - def record(self, raw: Any) -> Any: + def _record(self, raw: Any) -> Any: self.raw_response["value"] = raw self.raw_response["json"] = _jsonable(raw) return self.raw_response["json"] - def fail(self, exc: BaseException) -> None: - self.raw_response["error"] = exc + def invoke(self, callback: Callable[..., Any], next_request: Any) -> Any: + """Provider callback handed to Relay: run ``callback`` on Relay's (possibly rewritten) request.""" + try: + raw = self.run_callback(callback, self.provider_request(next_request)) + except BaseException as exc: + self.raw_response["error"] = exc + raise + return self._record(raw) + + async def invoke_async(self, callback: Callable[..., Any], next_request: Any) -> Any: + try: + final_request = self.provider_request(next_request) + + async def call_provider() -> Any: + # Nested relay calls run unmanaged — see relay_runtime.managed_callback_guard. + with relay_runtime.managed_callback_guard(): + return await callback(final_request) + + raw = await self.context.copy().run(asyncio.create_task, call_provider()) + except BaseException as exc: + self.raw_response["error"] = exc + raise + return self._record(raw) def run_managed(self, relay_call: Callable[..., Any], *callbacks: Any) -> Any: """Return the awaitable running ``relay_call`` inside the session context.""" return self.runtime.run_in_session_async( - self.session, - relay_call, - self.operation, - self.relay_request, - *callbacks, - **self.relay_kwargs, + self.session, relay_call, self.operation, self.relay_request, *callbacks, **self.relay_kwargs, ) def resolve_failure(self, exc: BaseException, defer_logical_completion: bool) -> Any: """Re-raise the provider's own error, or recover a completed provider result. - - Must be called from the ``except`` handling ``exc`` (bare ``raise``). - """ + Must be called from the ``except`` handling ``exc`` (bare ``raise``).""" callback_error = self.raw_response.get("error") - if ( - callback_error is not None - and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error) - ): + if callback_error is not None and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error): raise callback_error - if ( - not isinstance(exc, Exception) - or callback_error is not None - or "value" not in self.raw_response - ): + if not isinstance(exc, Exception) or callback_error is not None or "value" not in self.raw_response: raise logger.warning( - "NeMo Relay LLM post-processing failed after provider success; " - "returning the provider response", + "NeMo Relay LLM post-processing failed after provider success; returning the provider response", exc_info=True, ) - if not defer_logical_completion: - _complete_logical(self.logical, outcome="success") + self._complete(defer_logical_completion) return self.raw_response["value"] def result(self, managed: Any, defer_logical_completion: bool) -> Any: - if not defer_logical_completion: - _complete_logical(self.logical, outcome="success") + self._complete(defer_logical_completion) if "value" in self.raw_response and _json_equal(managed, self.raw_response["json"]): return self.raw_response["value"] return _namespace(managed) - -def execute( - request: dict[str, Any], - callback: Callable[[dict[str, Any]], Any], - *, - session_id: str, - name: str, - model_name: str, - metadata: dict[str, Any] | None = None, - defer_logical_completion: bool = False, -) -> Any: - """Run one non-streaming physical provider attempt through Relay.""" - attempt = _ManagedAttempt.resolve( - session_id, request, metadata, name=name, model_name=model_name - ) - if attempt is None: - return callback(request) - - def invoke(next_request: Any) -> Any: - try: - raw = attempt.run_callback(callback, attempt.provider_request(next_request)) - except BaseException as exc: - attempt.fail(exc) - raise - return attempt.record(raw) - - try: - managed = _run_awaitable(attempt.run_managed(attempt.runtime.relay.llm.execute, invoke)) - except BaseException as exc: - return attempt.resolve_failure(exc, defer_logical_completion) - return attempt.result(managed, defer_logical_completion) - - -async def execute_async( - request: dict[str, Any], - callback: Callable[[dict[str, Any]], Any], - *, - session_id: str, - name: str, - model_name: str, - metadata: dict[str, Any] | None = None, - defer_logical_completion: bool = False, -) -> Any: - """Run one asynchronous physical provider attempt through Relay.""" - attempt = _ManagedAttempt.resolve( - session_id, request, metadata, name=name, model_name=model_name - ) - if attempt is None: - return await callback(request) - - async def invoke(next_request: Any) -> Any: - try: - final_request = attempt.provider_request(next_request) - - async def call_provider() -> Any: - # Nested relay calls inside a managed provider callback must - # run unmanaged — see relay_runtime.managed_callback_guard. - with relay_runtime.managed_callback_guard(): - return await callback(final_request) - - raw = await attempt.context.copy().run(asyncio.create_task, call_provider()) - except BaseException as exc: - attempt.fail(exc) - raise - return attempt.record(raw) - - try: - managed = await attempt.run_managed(attempt.runtime.relay.llm.execute, invoke) - except BaseException as exc: - return attempt.resolve_failure(exc, defer_logical_completion) - return attempt.result(managed, defer_logical_completion) + def _complete(self, defer_logical_completion: bool) -> None: + if not defer_logical_completion: + _complete_logical(self.logical, outcome="success") def _current_session_id() -> str | None: @@ -264,52 +171,46 @@ def _current_session_id() -> str | None: return None if turn is None else turn.lease.session_id -def execute_current( - request: dict[str, Any], - callback: Callable[[dict[str, Any]], Any], - *, - name: str, - model_name: str, - metadata: dict[str, Any] | None = None, - defer_logical_completion: bool = False, +def execute( + request: dict[str, Any], callback: Callable[[dict[str, Any]], Any], *, name: str, model_name: str, + session_id: str | None = None, metadata: dict[str, Any] | None = None, defer_logical_completion: bool = False, ) -> Any: - """Run a provider attempt under the inherited Hermes turn when present.""" - session_id = _current_session_id() + """Run one non-streaming physical provider attempt through Relay. + ``session_id`` defaults to the inherited Hermes turn's session (unmanaged when there is none).""" if session_id is None: + session_id = _current_session_id() + attempt = _ManagedAttempt.resolve(session_id, request, metadata, name=name, model_name=model_name) + if attempt is None: return callback(request) - return execute( - request, - callback, - session_id=session_id, - name=name, - model_name=model_name, - metadata=metadata, - defer_logical_completion=defer_logical_completion, - ) + try: + managed = _run_awaitable(attempt.run_managed( + attempt.runtime.relay.llm.execute, partial(attempt.invoke, callback) + )) + except BaseException as exc: + return attempt.resolve_failure(exc, defer_logical_completion) + return attempt.result(managed, defer_logical_completion) -async def execute_current_async( - request: dict[str, Any], - callback: Callable[[dict[str, Any]], Any], - *, - name: str, - model_name: str, - metadata: dict[str, Any] | None = None, - defer_logical_completion: bool = False, +async def execute_async( + request: dict[str, Any], callback: Callable[[dict[str, Any]], Any], *, name: str, model_name: str, + session_id: str | None = None, metadata: dict[str, Any] | None = None, defer_logical_completion: bool = False, ) -> Any: - """Run an async provider attempt under the inherited turn when present.""" - session_id = _current_session_id() + """Async ``execute``.""" if session_id is None: + session_id = _current_session_id() + attempt = _ManagedAttempt.resolve(session_id, request, metadata, name=name, model_name=model_name) + if attempt is None: return await callback(request) - return await execute_async( - request, - callback, - session_id=session_id, - name=name, - model_name=model_name, - metadata=metadata, - defer_logical_completion=defer_logical_completion, - ) + try: + managed = await attempt.run_managed(attempt.runtime.relay.llm.execute, partial(attempt.invoke_async, callback)) + except BaseException as exc: + return attempt.resolve_failure(exc, defer_logical_completion) + return attempt.result(managed, defer_logical_completion) + + +# Run under the inherited Hermes turn when present (callers that do not know a session id). +execute_current = execute +execute_current_async = execute_async def _has_running_event_loop() -> bool: @@ -321,90 +222,36 @@ def _has_running_event_loop() -> bool: def stream_current( - request: dict[str, Any], - stream_factory: Callable[[dict[str, Any]], Any], - *, - name: str, - model_name: str, - finalizer: Callable[[], Any], - metadata: dict[str, Any] | None = None, - defer_logical_completion: bool = False, - completed_response_predicate: Callable[[Any], bool] | None = None, + request: dict[str, Any], stream_factory: Callable[[dict[str, Any]], Any], *, name: str, model_name: str, + finalizer: Callable[[], Any], metadata: dict[str, Any] | None = None, + defer_logical_completion: bool = False, completed_response_predicate: Callable[[Any], bool] | None = None, ) -> Any: """Run a provider stream under the inherited Hermes turn when present. - - With ``completed_response_predicate`` set, a factory that ignores - ``stream=True`` and returns a complete response is unwrapped and returned - directly (pre-Relay ``call_llm(stream=True)`` behavior); otherwise it - would stay trapped as ``final_response`` on the inner ManagedLlmStream. - Detecting that shape starts the lazy managed pipeline: a genuine first - chunk is buffered, but provider latency and pre-first-yield errors may - surface before this function returns. - """ + With ``completed_response_predicate`` set, a factory that ignores ``stream=True`` and + returns a complete response is unwrapped and returned directly (pre-Relay behavior) + instead of staying trapped as ``final_response``. Detecting that primes the lazy + pipeline: a genuine first chunk is buffered, but provider latency and pre-first-yield + errors may surface before this returns.""" session_id = _current_session_id() - if session_id is None: - return stream_factory(request) - if _has_running_event_loop(): - # Managed provider callbacks run on the Relay session's event loop; a - # nested ManagedLlmStream would be iterated synchronously on that same - # loop thread, which asyncio forbids. The outer managed stream already - # tracks the enclosing attempt and traps a completed response itself. + # On the Relay session's loop (inside a managed callback) a nested ManagedLlmStream would + # be iterated synchronously on that loop, which asyncio forbids; the outer managed stream + # already tracks this attempt. + if session_id is None or _has_running_event_loop(): return stream_factory(request) managed = stream( - request, - stream_factory, - session_id=session_id, - name=name, - model_name=model_name, - finalizer=finalizer, - metadata=metadata, - defer_logical_completion=defer_logical_completion, + request, stream_factory, session_id=session_id, name=name, model_name=model_name, + finalizer=finalizer, metadata=metadata, defer_logical_completion=defer_logical_completion, completed_response_predicate=completed_response_predicate, ) if completed_response_predicate is not None: - # Relay may defer the provider callback until the first pull; prime - # once so a completed response surfaces. A real first chunk is buffered. + # Relay may defer the provider callback until the first pull; prime once so a + # completed response surfaces (a real first chunk is buffered). managed._prime_completed_response() - completed = getattr(managed, "final_response", None) - if completed is not None: - return completed + if managed.final_response is not None: + return managed.final_response return managed -def stream( - request: dict[str, Any], - stream_factory: Callable[[dict[str, Any]], Any], - *, - session_id: str, - name: str, - model_name: str, - finalizer: Callable[[], Any], - on_stream_created: Callable[[Any], None] | None = None, - on_chunk: Callable[[Any], None] | None = None, - chunk_adapter: Callable[[Any], Any] | None = None, - accept_chunk: Callable[[Any], bool] | None = None, - completed_response_predicate: Callable[[Any], bool] | None = None, - metadata: dict[str, Any] | None = None, - defer_logical_completion: bool = False, -) -> "ManagedLlmStream": - """Return a synchronous view of one Relay-managed provider stream.""" - return ManagedLlmStream( - request, - stream_factory, - session_id=session_id, - name=name, - model_name=model_name, - finalizer=finalizer, - on_stream_created=on_stream_created, - on_chunk=on_chunk, - chunk_adapter=chunk_adapter, - accept_chunk=accept_chunk, - completed_response_predicate=completed_response_predicate, - metadata=metadata, - defer_logical_completion=defer_logical_completion, - ) - - def _aclose_on_loop(loop: asyncio.AbstractEventLoop, stream: Any) -> None: """Await ``stream.aclose()`` on ``loop`` when the stream exposes one.""" close = getattr(stream, "aclose", None) @@ -418,172 +265,135 @@ def _aclose_on_loop(loop: asyncio.AbstractEventLoop, stream: Any) -> None: class ManagedLlmStream(Iterator[Any]): - """Drive Relay's async stream from Hermes's provider worker thread.""" + """Synchronous view of one Relay-managed provider stream, driven from the worker thread.""" + + final_response: Any = None + output_modified = _closed = _provider_completed = False + _loop: asyncio.AbstractEventLoop | None = None + _stream = _raw_stream_resource = None + _runtime_lease: relay_runtime.RelayOperationLease | None = None + _close_error = _callback_error = None # BaseException | None + _logical: _LogicalCall | None = None + _logical_response_model_name: str | None = None def __init__( - self, - request: dict[str, Any], - stream_factory: Callable[[dict[str, Any]], Any], - *, - session_id: str, - name: str, - model_name: str, - finalizer: Callable[[], Any], - on_stream_created: Callable[[Any], None] | None, - on_chunk: Callable[[Any], None] | None, - chunk_adapter: Callable[[Any], Any] | None, - accept_chunk: Callable[[Any], bool] | None, - completed_response_predicate: Callable[[Any], bool] | None, - metadata: dict[str, Any] | None, - defer_logical_completion: bool, + self, request: dict[str, Any], stream_factory: Callable[[dict[str, Any]], Any], *, session_id: str, + name: str, model_name: str, finalizer: Callable[[], Any], + on_stream_created: Callable[[Any], None] | None = None, on_chunk: Callable[[Any], None] | None = None, + chunk_adapter: Callable[[Any], Any] | None = None, accept_chunk: Callable[[Any], bool] | None = None, + completed_response_predicate: Callable[[Any], bool] | None = None, + metadata: dict[str, Any] | None = None, defer_logical_completion: bool = False, ) -> None: - self.final_response: Any = None - self._loop: asyncio.AbstractEventLoop | None = None - self._stream: Any = None - self._raw_stream_resource: Any = None - self._closed = False - self._runtime_lease: relay_runtime.RelayOperationLease | None = None - self._close_error: BaseException | None = None - self._callback_error: BaseException | None = None - self._logical: _LogicalCall | None = None self._defer_logical_completion = defer_logical_completion # Only auxiliary calls report model/provider on their logical scope. auxiliary = str((metadata or {}).get("call_role") or "").startswith("auxiliary:") self._logical_model_name: str | None = model_name if auxiliary else None self._logical_provider_name: str | None = name if auxiliary else None - self._logical_response_model_name: str | None = None self._on_chunk = on_chunk self._chunk_adapter = chunk_adapter or _namespace self._accept_chunk = accept_chunk - self._relay_observes_chunks = False - self._provider_completed = False self._raw_chunks: list[tuple[Any, Any]] = [] self._prefetched_chunks: list[Any] = [] - self.output_modified = False - attempt = _ManagedAttempt.resolve( - session_id, request, metadata, name=name, model_name=model_name - ) + self._stream_factory = stream_factory + self._on_stream_created = on_stream_created + self._completed_response_predicate = completed_response_predicate + self._finalizer = finalizer + attempt = _ManagedAttempt.resolve(session_id, request, metadata, name=name, model_name=model_name) if attempt is None: - self._start_unmanaged( - request, stream_factory, on_stream_created, completed_response_predicate - ) + self._start_unmanaged(request) return self._logical = attempt.logical - self._start_managed( - attempt, stream_factory, on_stream_created, completed_response_predicate, finalizer - ) + self._start_managed(attempt) - def _start_unmanaged( - self, - request: dict[str, Any], - stream_factory: Callable[[dict[str, Any]], Any], - on_stream_created: Callable[[Any], None] | None, - completed_response_predicate: Callable[[Any], bool] | None, - ) -> None: - raw_stream = stream_factory(request) - if completed_response_predicate is not None and completed_response_predicate(raw_stream): + def _start_unmanaged(self, request: dict[str, Any]) -> None: + raw_stream = self._stream_factory(request) + predicate = self._completed_response_predicate + if predicate is not None and predicate(raw_stream): self.final_response = raw_stream self._stream = iter(()) return self._raw_stream_resource = raw_stream - if on_stream_created is not None: - on_stream_created(raw_stream) + if self._on_stream_created is not None: + self._on_stream_created(raw_stream) self._stream = iter(raw_stream) - def _start_managed( - self, - attempt: _ManagedAttempt, - stream_factory: Callable[[dict[str, Any]], Any], - on_stream_created: Callable[[Any], None] | None, - completed_response_predicate: Callable[[Any], bool] | None, - finalizer: Callable[[], Any], - ) -> None: - """Open Relay's stream on a private event loop owned by this iterator.""" + async def _provider_stream(self, attempt: _ManagedAttempt, next_request: Any): + """Relay's provider callback: run the factory and yield JSON-encoded chunks.""" run_callback = attempt.run_callback - - async def provider_stream(next_request: Any): - raw_stream = None - try: - raw_stream = run_callback(stream_factory, attempt.provider_request(next_request)) - if completed_response_predicate is not None and run_callback( - completed_response_predicate, raw_stream - ): - self.final_response = raw_stream - self._provider_completed = True - return - if on_stream_created is not None: - run_callback(on_stream_created, raw_stream) - raw_iterator = run_callback(iter, raw_stream) - while True: - try: - chunk = run_callback(next, raw_iterator) - except StopIteration: - break - if self._accept_chunk is not None and not run_callback( - self._accept_chunk, chunk - ): - break - encoded_chunk = _jsonable(chunk) - self._raw_chunks.append((encoded_chunk, chunk)) - yield encoded_chunk + raw_stream = None + try: + raw_stream = run_callback(self._stream_factory, attempt.provider_request(next_request)) + predicate = self._completed_response_predicate + if predicate is not None and run_callback(predicate, raw_stream): + self.final_response = raw_stream self._provider_completed = True - except BaseException as exc: - self._callback_error = exc - raise - finally: - close = getattr(raw_stream, "close", None) - if callable(close): - try: - run_callback(close) - except BaseException as exc: - self._close_error = exc - raise + return + if self._on_stream_created is not None: + run_callback(self._on_stream_created, raw_stream) + raw_iterator = run_callback(iter, raw_stream) + while True: + try: + chunk = run_callback(next, raw_iterator) + except StopIteration: + break + if self._accept_chunk is not None and not run_callback(self._accept_chunk, chunk): + break + encoded_chunk = _jsonable(chunk) + self._raw_chunks.append((encoded_chunk, chunk)) + yield encoded_chunk + self._provider_completed = True + except BaseException as exc: + self._callback_error = exc + raise + finally: + close = getattr(raw_stream, "close", None) + if callable(close): + try: + run_callback(close) + except BaseException as exc: + self._close_error = exc + raise + + def _relay_finalizer(self, attempt: _ManagedAttempt) -> Any: + # Relay may call this while unwinding a provider-stream failure; keep the + # original error instead of a secondary "missing terminal response". + if self._callback_error is not None: + return None + try: + response = self.final_response + if response is None: + response = attempt.run_callback(self._finalizer) + if self._logical_model_name is not None: + self._logical_response_model_name = _response_model_name(response) + return _jsonable(response) + except BaseException as exc: + self._callback_error = exc + raise + + def _start_managed(self, attempt: _ManagedAttempt) -> None: + """Open Relay's stream on a private event loop owned by this iterator.""" def observe_chunk(chunk: Any) -> None: if self._on_chunk is not None: - run_callback(self._on_chunk, _jsonable(chunk)) - - def relay_finalizer() -> Any: - # Relay can invoke the finalizer while unwinding a provider-stream - # failure; keep that original error instead of a secondary - # "missing terminal response" error. - if self._callback_error is not None: - return None - try: - response = self.final_response - if response is None: - response = run_callback(finalizer) - if self._logical_model_name is not None: - self._logical_response_model_name = _response_model_name(response) - return _jsonable(response) - except BaseException as exc: - self._callback_error = exc - raise + attempt.run_callback(self._on_chunk, _jsonable(chunk)) self._runtime_lease = attempt.runtime.acquire_operation_lease() try: - loop = asyncio.new_event_loop() - except BaseException: - self._release_runtime_lease() - raise - self._loop = loop - self._relay_observes_chunks = True - try: + self._loop = loop = asyncio.new_event_loop() self._stream = loop.run_until_complete( attempt.run_managed( - attempt.runtime.relay.llm.stream_execute, - provider_stream, - observe_chunk, - relay_finalizer, + attempt.runtime.relay.llm.stream_execute, partial(self._provider_stream, attempt), + observe_chunk, partial(self._relay_finalizer, attempt), ) ) except BaseException as exc: - if self._recoverable_relay_failure(exc): + if self._loop is not None and self._recoverable_relay_failure(exc): self._preserve_pending_provider_chunks() return - self._finish_logical("cancelled" if _is_cancellation(exc) else "failed") try: - loop.close() + if self._loop is not None: + self._finish_logical("cancelled" if _is_cancellation(exc) else "failed") + self._loop.close() finally: self._loop = None self._release_runtime_lease() @@ -594,36 +404,27 @@ class ManagedLlmStream(Iterator[Any]): def _prime_completed_response(self) -> None: """Advance once while preserving a genuine first chunk.""" - if self._closed or self._prefetched_chunks: - return - try: - self._prefetched_chunks.append(next(self)) - except StopIteration: - pass + if not self._closed and not self._prefetched_chunks: + with contextlib.suppress(StopIteration): + self._prefetched_chunks.append(next(self)) def _recoverable_relay_failure(self, exc: BaseException) -> bool: """Relay post-processing failed after the provider already succeeded.""" - if ( - isinstance(exc, Exception) and self._provider_completed and self._callback_error is None - ): + recoverable = (isinstance(exc, Exception) and self._provider_completed and self._callback_error is None) + if recoverable: logger.warning( - "NeMo Relay stream post-processing failed after provider success; " - "preserving the provider result", + "NeMo Relay stream post-processing failed after provider success; preserving the provider result", exc_info=True, ) - return True - return False + return recoverable def _finish_logical(self, outcome: str) -> None: """Complete the logical LLM scope unless the caller deferred it.""" if self._defer_logical_completion: return _complete_logical( - self._logical, - outcome=outcome, - model_name=self._logical_model_name, - provider_name=self._logical_provider_name, - response_model_name=self._logical_response_model_name, + self._logical, outcome=outcome, model_name=self._logical_model_name, + provider_name=self._logical_provider_name, response_model_name=self._logical_response_model_name, operation_lease=self._runtime_lease, ) self._logical = None @@ -657,10 +458,7 @@ class ManagedLlmStream(Iterator[Any]): raise StopIteration from None except BaseException as exc: callback_error = self._callback_error - if ( - callback_error is not None - and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error) - ): + if callback_error is not None and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error): self._close(logical_outcome="failed") raise callback_error if self._recoverable_relay_failure(exc): @@ -668,8 +466,6 @@ class ManagedLlmStream(Iterator[Any]): return next(self) self._close(logical_outcome="cancelled" if _is_cancellation(exc) else "failed") raise - if not self._relay_observes_chunks and self._on_chunk is not None: - self._on_chunk(chunk) for index, (encoded, raw) in enumerate(self._raw_chunks): if _json_equal(chunk, encoded): if index > 0: @@ -682,8 +478,7 @@ class ManagedLlmStream(Iterator[Any]): def close(self) -> None: """Close an explicitly abandoned stream and cancel its logical call.""" self._close(logical_outcome="cancelled") - close_error = self._close_error - self._close_error = None + close_error, self._close_error = self._close_error, None if close_error is not None: raise close_error @@ -691,43 +486,36 @@ class ManagedLlmStream(Iterator[Any]): """Switch a failed Relay stream to its undelivered provider chunks.""" pending = [raw for _encoded, raw in self._raw_chunks] self._raw_chunks.clear() - loop = self._loop - relay_stream = self._stream - self._loop = None - self._stream = iter(pending) - self._raw_stream_resource = None - self._accept_chunk = None + loop, relay_stream = self._loop, self._stream + self._loop, self._stream, self._raw_stream_resource, self._accept_chunk = None, iter(pending), None, None try: if loop is not None: try: _aclose_on_loop(loop, relay_stream) except Exception: - logger.debug( - "Relay stream cleanup failed during provider fallback", exc_info=True - ) + logger.debug("Relay stream cleanup failed during provider fallback", exc_info=True) loop.close() self._finish_logical("success") finally: self._release_runtime_lease() + def _keep_first_close_error(self, exc: BaseException) -> None: + if self._close_error is None: + self._close_error = exc + def _close_provider_resources(self) -> None: """Close the unmanaged provider stream/resource once each (they may be the same object).""" - resources = (self._stream, self._raw_stream_resource) + resources = {id(r): r for r in (self._stream, self._raw_stream_resource) if r is not None} self._stream = None self._raw_stream_resource = None - closed_ids: set[int] = set() - for resource in resources: - if resource is None or id(resource) in closed_ids: - continue - closed_ids.add(id(resource)) + for resource in resources.values(): close = getattr(resource, "close", None) if not callable(close): continue try: close() except Exception as exc: - if self._close_error is None: - self._close_error = exc + self._keep_first_close_error(exc) logger.debug("Provider stream cleanup failed", exc_info=True) def _close(self, *, logical_outcome: str) -> None: @@ -745,16 +533,14 @@ class ManagedLlmStream(Iterator[Any]): try: _aclose_on_loop(loop, self._stream) except Exception as exc: - if self._close_error is None: - self._close_error = exc + self._keep_first_close_error(exc) self._finish_logical(logical_outcome) loop.close() finally: self._release_runtime_lease() def _release_runtime_lease(self) -> None: - lease = self._runtime_lease - self._runtime_lease = None + lease, self._runtime_lease = self._runtime_lease, None if lease is not None: lease.release() @@ -762,9 +548,10 @@ class ManagedLlmStream(Iterator[Any]): self._close(logical_outcome="cancelled") -_ANTHROPIC_APPEND_DELTAS = { - "text_delta": "text", "thinking_delta": "thinking", "signature_delta": "signature" -} +stream = ManagedLlmStream + + +_ANTHROPIC_APPEND_DELTAS = {"text_delta": "text", "thinking_delta": "thinking", "signature_delta": "signature"} class AnthropicStreamAccumulator: @@ -776,28 +563,23 @@ 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") if isinstance(message, dict): - for key in ("id", "type", "role", "model", "usage"): - if key in message: - self._message[key] = message[key] + 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, {}) @@ -806,51 +588,41 @@ class AnthropicStreamAccumulator: if field is not None: block[field] = str(block.get(field) or "") + str(delta.get(field) or "") elif delta_type == "input_json_delta": - block["_partial_json"] = str(block.pop("_partial_json", "")) + str( - delta.get("partial_json") or "" - ) + block["_partial_json"] = str(block.pop("_partial_json", "")) + str(delta.get("partial_json") or "") elif delta_type == "citations_delta" and "citation" in delta: block.setdefault("citations", []).append(delta["citation"]) def _on_message_delta(self, payload: dict[str, Any]) -> None: delta = payload.get("delta") if isinstance(delta, dict): - for key in ("stop_reason", "stop_sequence"): - if key in delta: - self._message[key] = delta[key] + 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 _EVENT_HANDLERS = { - "message_start": _on_message_start, - "content_block_start": _on_content_block_start, - "content_block_delta": _on_content_block_delta, - "message_delta": _on_message_delta, + "message_start": _on_message_start, "content_block_start": _on_content_block_start, + "content_block_delta": _on_content_block_delta, "message_delta": _on_message_delta, } def finalize(self) -> dict[str, Any]: - blocks = [] - for index in sorted(self._blocks): - block = dict(self._blocks[index]) + blocks = [dict(self._blocks[index]) for index in sorted(self._blocks)] + for block in blocks: partial = block.pop("_partial_json", None) if partial is not None: try: block["input"] = json.loads(partial) except (TypeError, ValueError): block["input"] = partial - blocks.append(block) return {**self._message, "content": blocks} def response(self, base: Any = None) -> Any: """Return the attribute-shaped response consumed by Hermes.""" assembled = self.finalize() - base_payload = _jsonable_dict(base) content = assembled.pop("content", []) - merged = {**base_payload, **assembled} + merged = {**_jsonable_dict(base), **assembled} if content or "content" not in merged: merged["content"] = content return _namespace(merged) @@ -870,30 +642,18 @@ def _logical_parent( with turn.logical_llm_lock: handle = turn.logical_llm_calls.get(request_id) if handle is None: - handle = runtime.run_in_session( - session, - runtime.relay.scope.push, - relay_runtime.LOGICAL_LLM_SCOPE, - runtime.relay.ScopeType.Function, - handle=parent, - input={}, - metadata=relay_runtime.runtime_metadata( - runtime.runtime_id, - **{"hermes.call_role": str((metadata or {}).get("call_role") or "primary")}, - ), + call_role = str((metadata or {}).get("call_role") or "primary") + handle = turn.logical_llm_calls[request_id] = runtime.run_in_session( + session, runtime.relay.scope.push, relay_runtime.LOGICAL_LLM_SCOPE, + runtime.relay.ScopeType.Function, handle=parent, input={}, + metadata=relay_runtime.runtime_metadata(runtime.runtime_id, **{"hermes.call_role": call_role}), ) - turn.logical_llm_calls[request_id] = handle return turn, handle, request_id def _complete_logical( - logical: _LogicalCall | None, - *, - outcome: str, - model_name: str | None = None, - provider_name: str | None = None, - response_model_name: str | None = None, - operation_lease: relay_runtime.RelayOperationLease | None = None, + logical: _LogicalCall | None, *, outcome: str, model_name: str | None = None, provider_name: str | None = None, + response_model_name: str | None = None, operation_lease: relay_runtime.RelayOperationLease | None = None, ) -> None: if logical is None: return @@ -901,6 +661,11 @@ def _complete_logical( lease = turn.lease if not isinstance(lease.host, relay_runtime.RelayRuntime): return + output = {"outcome": outcome} + if model_name is not None and provider_name is not None: + output.update({"model": model_name, "provider": provider_name}) + if response_model_name is not None: + output["response_model"] = response_model_name with turn.finalize_lock: with turn.logical_llm_lock: if turn.logical_llm_calls.get(request_id) is not handle: @@ -908,30 +673,17 @@ def _complete_logical( if lease.session is None: return try: - output = {"outcome": outcome} - if model_name is not None and provider_name is not None: - output.update({"model": model_name, "provider": provider_name}) - if response_model_name is not None: - output["response_model"] = response_model_name - callback = lease.host.run_in_session - if operation_lease is not None: - callback = operation_lease.run_in_session - callback( - lease.session, - relay_runtime.pop_relay_scope, - lease.host.relay, - handle, - output=output, - metadata=relay_runtime.runtime_metadata(lease.host.runtime_id), + (operation_lease or lease.host).run_in_session( + lease.session, relay_runtime.pop_relay_scope, lease.host.relay, handle, + output=output, metadata=relay_runtime.runtime_metadata(lease.host.runtime_id), ) except Exception: - # The provider result is authoritative. Retain the handle so turn - # finalization can retry cleanup without changing that result. + # Provider result is authoritative; retain the handle so turn finalization can retry. logger.warning("Hermes Relay logical LLM finalization failed", exc_info=True) return with turn.logical_llm_lock: if turn.logical_llm_calls.get(request_id) is handle: - turn.logical_llm_calls.pop(request_id, None) + del turn.logical_llm_calls[request_id] def _is_cancellation(error: BaseException) -> bool: @@ -939,12 +691,8 @@ def _is_cancellation(error: BaseException) -> bool: def complete_logical_call( - api_request_id: str, - *, - outcome: str, - model_name: str | None = None, - provider_name: str | None = None, - response_model_name: str | None = None, + api_request_id: str, *, outcome: str, model_name: str | None = None, + provider_name: str | None = None, response_model_name: str | None = None, ) -> None: """Complete the active turn's logical LLM call after caller validation.""" turn = relay_runtime.active_turn() @@ -954,30 +702,20 @@ def complete_logical_call( handle = turn.logical_llm_calls.get(api_request_id) if handle is not None: _complete_logical( - (turn, handle, api_request_id), - outcome=outcome, - model_name=model_name, - provider_name=provider_name, - response_model_name=response_model_name, + (turn, handle, api_request_id), outcome=outcome, model_name=model_name, + provider_name=provider_name, response_model_name=response_model_name, ) def _response_model_name(response: Any) -> str | None: """Return a provider-reported model name when one is available.""" - if isinstance(response, dict): - value = response.get("model") - else: - value = getattr(response, "model", None) + value = response.get("model") if isinstance(response, dict) else getattr(response, "model", None) return value if isinstance(value, str) and value.strip() else None def _provider_request( - original: dict[str, Any], - request: Any, - *, - relay_request_body: dict[str, Any], - codec_baseline_body: dict[str, Any] | None, - metadata: dict[str, Any] | None, + original: dict[str, Any], request: Any, *, relay_request_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): @@ -986,146 +724,112 @@ def _provider_request( if codec_baseline_body is not None and not _json_equal(content, relay_request_body): baseline = codec_baseline_body intercepted = _provider_request_body(content, metadata) - # 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. + # 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 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 - ) + _restore_provider_message_extensions(original, final, baseline=baseline, intercepted=intercepted) headers = getattr(request, "headers", None) if isinstance(headers, dict): - headers = { - key: value - for key, value in headers.items() - if str(key).lower() not in _RELAY_INTERNAL_PROVIDER_HEADERS - } + headers = {k: v for k, v in headers.items() if str(k).lower() not in _RELAY_INTERNAL_PROVIDER_HEADERS} if headers: final["extra_headers"] = {**dict(final.get("extra_headers") or {}), **headers} return final +def _rewrite_tools(body: dict[str, Any], match: Callable[[dict], bool], rewrite: Callable[[dict], dict]) -> None: + """Rewrite each dict tool that ``match``es (in place on ``body["tools"]`` when it is a list).""" + tools = body.get("tools") + if isinstance(tools, list): + body["tools"] = [rewrite(t) if isinstance(t, dict) and match(t) else t for t in tools] + + +def _codex_codec_tools(body: dict[str, Any]) -> None: + # The Responses SDK accepts ``tools=None`` as "no tools" while Relay's typed codec + # wants an array or an absent field; only the codec-facing copy is normalized. + if body.get("tools") is None: + body.pop("tools", None) + _rewrite_tools( + body, lambda t: t.get("type") == "function" and "function" not in t, + lambda t: {"type": "function", "function": {k: v for k, v in t.items() if k != "type"}}, + ) + + +def _chat_codec_tools(body: dict[str, Any]) -> None: + _rewrite_tools(body, lambda t: "function" in t and "type" not in t, lambda t: {"type": "function", **t}) + + +# api_mode -> in-place normalizer producing the codec-facing ``tools`` shape. +_CODEC_TOOL_NORMALIZERS = {"codex_responses": _codex_codec_tools, "chat_completions": _chat_codec_tools} + + def _relay_request_body(request: dict[str, Any], metadata: dict[str, Any] | None) -> dict[str, Any]: body = _jsonable_dict(request) - # ``timeout`` configures the provider SDK client, not a wire protocol: - # keep it on the original callback request, never on Relay intercepts. + # ``timeout`` configures the SDK client, not the wire: never expose it to intercepts. body.pop("timeout", None) - api_mode = _api_mode(metadata) - if api_mode == "codex_responses": - # The Responses SDK accepts ``tools=None`` as "no tools" while Relay's - # typed codec expects an array or an absent field; normalize only the - # codec-facing copy (the original request is restored when unchanged). - if body.get("tools") is None: - body.pop("tools", None) - elif isinstance(body.get("tools"), list): - body["tools"] = [ - { - "type": "function", - "function": {key: value for key, value in tool.items() if key != "type"}, - } - if isinstance(tool, dict) - and tool.get("type") == "function" - and "function" not in tool - else tool - for tool in body["tools"] - ] - elif api_mode == "chat_completions": - tools = body.get("tools") - if isinstance(tools, list): - body["tools"] = [ - {"type": "function", **tool} - if isinstance(tool, dict) and "function" in tool and "type" not in tool - else tool - for tool in tools - ] + normalize = _CODEC_TOOL_NORMALIZERS.get(_api_mode(metadata)) + if normalize is not None: + normalize(body) return body def _restore_provider_message_extensions( - original: dict[str, Any], - final: dict[str, Any], - *, - baseline: dict[str, Any], - intercepted: 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.""" message_lists = tuple(body.get("messages") for body in (original, final, baseline, intercepted)) - if not all(isinstance(messages, list) for messages in message_lists): - return - if len({len(messages) for messages in message_lists}) != 1: + if not all(isinstance(m, list) for m in message_lists) or len({len(m) for m in message_lists}) != 1: return for messages in zip(*message_lists, strict=True): if not all(isinstance(message, dict) for message in messages): 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] def _codec_round_trip_request_body( - relay: Any, - relay_request: Any, - *, - relay_request_body: dict[str, Any], - metadata: dict[str, Any] | None, + relay: Any, relay_request: Any, *, relay_request_body: dict[str, Any], metadata: dict[str, Any] | None, ) -> dict[str, Any] | None: """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) + encoded = codec.encode(codec.decode(relay_request), relay_request) content = getattr(encoded, "content", encoded) if isinstance(content, dict): return _provider_request_body(content, metadata) except Exception: - logger.warning( - "NeMo Relay request codec baseline failed; ignoring request rewrites", exc_info=True - ) + 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" - ) + logger.warning("NeMo Relay request codec returned an unsupported baseline; ignoring request rewrites") return None -def _provider_request_body( - content: dict[str, Any], metadata: dict[str, Any] | None -) -> dict[str, Any]: +def _provider_request_body(content: dict[str, Any], metadata: dict[str, Any] | None) -> dict[str, Any]: body = dict(content) - if _api_mode(metadata) != "codex_responses": - return body - tools = body.get("tools") - if not isinstance(tools, list): - return body - body["tools"] = [ - {"type": "function", **dict(tool["function"])} - if isinstance(tool, dict) - and tool.get("type") == "function" - and isinstance(tool.get("function"), dict) - else tool - for tool in tools - ] + if _api_mode(metadata) == "codex_responses": + _rewrite_tools( + body, lambda t: t.get("type") == "function" and isinstance(t.get("function"), dict), + lambda t: {"type": "function", **dict(t["function"])}, + ) return body def _codec(relay: Any, metadata: dict[str, Any] | None) -> Any: - protocol = _relay_protocol(metadata) + protocol = _RELAY_PROTOCOL_BY_API_MODE.get(_api_mode(metadata)) codecs = getattr(relay, "codecs", None) if protocol is None or codecs is None: return None - codec = getattr(codecs, protocol.codec_class, None) + codec = getattr(codecs, protocol[1], None) return codec() if callable(codec) else None @@ -1139,19 +843,16 @@ def _jsonable(value: Any) -> Any: model_dump = getattr(type(value), "model_dump", None) if callable(model_dump): try: - # warnings=False: pydantic warns on generic-union SDK stream events - # and that warning would leak to the user's terminal mid-response. + # warnings=False: pydantic's generic-union warning would leak to the terminal + # mid-response; TypeError = duck-typed model_dump without pydantic's signature. try: return _jsonable(value.model_dump(mode="json", warnings=False)) except TypeError: - # Duck-typed model_dump without pydantic's signature. return _jsonable(value.model_dump()) except Exception: pass try: - attributes = { - str(key): item for key, item in vars(value).items() if not str(key).startswith("_") - } + attributes = {str(key): item for key, item in vars(value).items() if not str(key).startswith("_")} except (TypeError, AttributeError): return str(value) return _jsonable(attributes) if attributes else str(value) @@ -1171,20 +872,22 @@ def _namespace(value: Any) -> Any: return value +def _canonical_json(value: Any, encode: Callable[[Any], Any] = _jsonable) -> str: + return json.dumps(encode(value), sort_keys=True, separators=(",", ":")) + + def _json_equal(left: Any, right: Any) -> bool: try: - return json.dumps( - _jsonable(left), sort_keys=True, separators=(",", ":") - ) == json.dumps(_jsonable(right), sort_keys=True, separators=(",", ":")) + return _canonical_json(left) == _canonical_json(right) except (TypeError, ValueError): return False -def _run_awaitable(value: Any) -> Any: +def _run_awaitable( + value: Any, *, loop_error: str = "Synchronous Relay LLM execution cannot run on an event-loop thread", +) -> Any: if not inspect.isawaitable(value): return value - try: - asyncio.get_running_loop() - except RuntimeError: - return asyncio.run(value) - raise RuntimeError("Synchronous Relay LLM execution cannot run on an event-loop thread") + if _has_running_event_loop(): + raise RuntimeError(loop_error) + return asyncio.run(value) diff --git a/agent/relay_runtime.py b/agent/relay_runtime.py index 06db7a287f..2f28d8c21b 100644 --- a/agent/relay_runtime.py +++ b/agent/relay_runtime.py @@ -4,6 +4,7 @@ from __future__ import annotations import atexit import asyncio +import contextlib import contextvars import importlib import inspect @@ -19,9 +20,7 @@ from pathlib import Path from typing import Any, Callable from hermes_constants import get_hermes_home -from hermes_cli.relay_plugin_cutover import ( - RELAY_PLUGINS_CONFIG_ENV, configured_legacy_relay_env_vars -) +from hermes_cli.relay_plugin_cutover import (RELAY_PLUGINS_CONFIG_ENV, configured_legacy_relay_env_vars) logger = logging.getLogger(__name__) @@ -34,8 +33,8 @@ RUNTIME_INSTANCE_KEY = "hermes.relay.runtime_instance" RELAY_PLUGINS_EXECUTION_CONSUMER = "hermes.nemo_relay.plugins" _PROFILE_KEY_CACHE: dict[str, str] = {} -# Bound for native scope lifecycle ops (push/pop/flush) gating turn/session completion. -# Healthy ops take microseconds; a wedged pipeline costs one lost span, never a blocked agent. +# Bound for native scope ops gating turn/session completion: a wedged pipeline costs one +# lost span, never a blocked agent. _SCOPE_OP_TIMEOUT = 10.0 _SCOPE_OP_EXECUTOR: Any = None @@ -48,21 +47,15 @@ def runtime_metadata(runtime_id: str, **extra: Any) -> dict[str, Any]: def _scope_op_executor(): - """Shared daemon executor for bounded native scope operations. - - Daemon workers (tools.daemon_pool) so a wedged call abandoned at timeout - cannot block interpreter exit. ``Future.result(timeout=...)`` bounds callers - even when every worker is wedged, so exhaustion degrades to fast timeouts. - """ + """Shared daemon executor for bounded native scope ops. + Daemon workers so a wedged call abandoned at timeout cannot block interpreter exit; + ``Future.result(timeout=...)`` still bounds callers when every worker is wedged.""" global _SCOPE_OP_EXECUTOR if _SCOPE_OP_EXECUTOR is None: with _SCOPE_OP_EXECUTOR_LOCK: if _SCOPE_OP_EXECUTOR is None: from tools.daemon_pool import DaemonThreadPoolExecutor - - _SCOPE_OP_EXECUTOR = DaemonThreadPoolExecutor( - max_workers=8, thread_name_prefix="relay-scope-op" - ) + _SCOPE_OP_EXECUTOR = DaemonThreadPoolExecutor(max_workers=8, thread_name_prefix="relay-scope-op") return _SCOPE_OP_EXECUTOR @@ -70,43 +63,32 @@ def _run_on_daemon_thread( fn: Callable[[], Any], *, name: str, timeout: float | None = None, timeout_message: str = "" ) -> Any: """Run ``fn`` on a fresh daemon thread; re-raise its error or return its result. - - With ``timeout`` the join is bounded and a still-running worker is abandoned - with ``TimeoutError`` — a daemon thread cannot block interpreter exit. - """ - result: list[Any] = [] - error: list[BaseException] = [] + With ``timeout`` a still-running worker is abandoned with ``TimeoutError`` (daemon: + cannot block interpreter exit).""" + outcome: dict[str, Any] = {} def _target() -> None: try: - result.append(fn()) + outcome["result"] = fn() except BaseException as exc: # noqa: BLE001 - propagated below - error.append(exc) + outcome["error"] = exc worker = threading.Thread(target=_target, daemon=True, name=name) worker.start() worker.join(timeout) if worker.is_alive(): raise TimeoutError(timeout_message) - if error: - raise error[0] - return result[0] if result else None + if "error" in outcome: + raise outcome["error"] + return outcome.get("result") -def pop_relay_scope( - relay: Any, handle: Any, *, output: Any = None, metadata: Any = None, timestamp: Any = None -) -> Any: +def pop_relay_scope(relay: Any, handle: Any, *, output: Any = None, metadata: Any = None, timestamp: Any = None) -> Any: """Pop a Relay scope, forwarding only the kwargs the live binding accepts. - - ``scope.pop`` gained ``metadata`` in nemo-relay 0.4+; older wheels raise - TypeError on it, which would wedge turn/session close. - """ + ``scope.pop`` gained ``metadata`` in nemo-relay 0.4+; older wheels raise TypeError.""" pop = relay.scope.pop - kwargs = { - key: value - for key, value in (("output", output), ("metadata", metadata), ("timestamp", timestamp)) - if value is not None - } + candidates = (("output", output), ("metadata", metadata), ("timestamp", timestamp)) + kwargs = {key: value for key, value in candidates if value is not None} try: params = inspect.signature(pop).parameters except (TypeError, ValueError): @@ -118,31 +100,23 @@ def pop_relay_scope( def _current_top(relay: Any) -> Any: """Return the current top-of-stack scope handle, or None.""" - # Prefer ``scope.get_handle()``: ``get_scope_stack()`` may return a native - # ScopeStack object that ``scope.pop`` rejects, so never treat it as a handle. + # Prefer scope.get_handle(): get_scope_stack() may return a native ScopeStack + # object that scope.pop rejects, so never treat it as a handle. get_handle = getattr(getattr(relay, "scope", None), "get_handle", None) if callable(get_handle): - try: + with contextlib.suppress(Exception): return get_handle() - except Exception: - pass top = relay.get_scope_stack() - # Some builds return the live stack (list), others the top handle directly - # (including tuple handles from test fakes): only unwrap real lists. - if isinstance(top, list): - return top[-1] if top else None - return top + # Some builds return the live stack (list), others the top handle: only unwrap real lists. + return (top[-1] if top else None) if isinstance(top, list) else top def _same_handle(a: Any, b: Any) -> bool: # Native ScopeHandle has no value __eq__; compare by uuid when both expose one. if a is None or b is None: return a is b - if a is b or a == b: - return True a_uuid = getattr(a, "uuid", None) - b_uuid = getattr(b, "uuid", None) - return a_uuid is not None and a_uuid == b_uuid + return a is b or a == b or (a_uuid is not None and a_uuid == getattr(b, "uuid", None)) class _RelayPluginConfigurationState(Enum): @@ -169,40 +143,33 @@ class RelaySession: closing: bool = False handle: Any = None context: contextvars.Context | None = None - # Session-span segmentation (continuous sessions): rotation closes the - # current session scope and pushes segment N+1 at a turn boundary. + # Session-span segmentation: rotation closes the current session scope and pushes + # segment N+1 at a turn boundary (the only LIFO-safe point). segment: int = 0 # index of the CURRENT session scope (0 = first) segment_turns: int = 0 # turns completed within the current segment rotate_pending: bool = False # set by compaction; consumed at next begin_turn - # A rotating compaction landed while a turn was live on THIS session; closing - # now would pop the session scope under the live turn (LIFO violation), so - # end_turn consumes this and closes the session. + # Rotating compaction landed while a turn was live here; closing now would pop the + # session scope under the live turn, so end_turn consumes this instead. close_pending: bool = False -# Segmentation config (gateway.telemetry.session_segments), cached at first read. -# Both defaults OFF => rotation never fires and the scope lifecycle is identical -# to the pre-segmentation behavior. +# gateway.telemetry.session_segments, cached at first read. Both defaults OFF => +# rotation never fires and the scope lifecycle is unchanged. _SEGMENTS_CONFIG: dict[str, Any] | None = None _SEGMENTS_CONFIG_LOCK = threading.Lock() def _load_segments_config() -> dict[str, Any]: - on_compaction = False - max_turns = 0 - try: + segments: dict[str, Any] = {} + with contextlib.suppress(Exception): # config absence must not crash from gateway.run import _load_gateway_config # late import - telemetry = (_load_gateway_config().get("gateway") or {}).get("telemetry") or {} segments = telemetry.get("session_segments") or {} - on_compaction = bool(segments.get("on_compaction", False)) - try: - max_turns = max(0, int(segments.get("max_turns", 0) or 0)) - except (TypeError, ValueError): - max_turns = 0 - except Exception: # noqa: BLE001 - config absence must not crash - pass - return {"on_compaction": on_compaction, "max_turns": max_turns} + try: + max_turns = max(0, int(segments.get("max_turns", 0) or 0)) + except (TypeError, ValueError): + max_turns = 0 + return {"on_compaction": bool(segments.get("on_compaction", False)), "max_turns": max_turns} def _segments_config() -> dict[str, Any]: @@ -227,21 +194,17 @@ class RelayOperationLease: self._lock = threading.Lock() self._runtime: RelayRuntime | None = runtime - def run_in_session( - self, session: RelaySession, callback: Callable[..., Any], *args: Any, **kwargs: Any - ) -> Any: + def run_in_session(self, session: RelaySession, callback: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: """Run cleanup while this lease still owns the runtime lifetime.""" with self._lock: - runtime = self._runtime - if runtime is None: + if self._runtime is None: raise RuntimeError("Hermes Relay operation lease is released") - return runtime._run_in_session_untracked(session, callback, *args, **kwargs) + return self._runtime._run_in_session_untracked(session, callback, *args, **kwargs) def release(self) -> None: """Release this lease exactly once.""" with self._lock: - runtime = self._runtime - self._runtime = None + runtime, self._runtime = self._runtime, None if runtime is not None: runtime._end_operation() @@ -253,59 +216,57 @@ class _ProcessRelayPluginConfiguration: self._lock = threading.RLock() self._owners: set[int] = set() self._state = _RelayPluginConfigurationState.UNINITIALIZED - self._active = False - self._relay: Any = None + self._relay: Any = None # set while a Hermes-owned configuration is active self._activation: Any = None def acquire(self, owner: Any, relay: Any) -> _RelayPluginConfigurationState: """Join the process configuration, initializing it for the first host.""" - owner_id = id(owner) with self._lock: - if owner_id in self._owners: - return self._state - if self._owners: - self._owners.add(owner_id) - return self._state - if self._active and not self._clear_active(): - logger.warning( - "Hermes Relay plugin cleanup is still pending; refusing to " - "replace the process-global configuration" - ) - return self._remember(owner_id, _RelayPluginConfigurationState.FAILED) + if not self._owners: + # First host decides for the whole process; later hosts just join. + self._state = self._preflight(relay) or self._activate(relay) + if self._state is _RelayPluginConfigurationState.ACTIVE: + logger.info( + "Relay plugins are active process-wide and apply to all profiles hosted by this Hermes process." + ) + self._owners.add(id(owner)) + return self._state - try: - existing_report = relay.plugin.report() - except Exception: - logger.warning( - "Hermes could not determine whether a process-global Relay " - "plugin configuration is already active; refusing to replace it", - exc_info=True, - ) - return self._remember(owner_id, _RelayPluginConfigurationState.FAILED) - if existing_report is not None: - logger.warning( - "A process-global Relay plugin configuration is already active " - "outside Hermes native ownership; leaving it unchanged and " - "disabling Hermes-managed Relay middleware for this process" - ) - return self._remember(owner_id, _RelayPluginConfigurationState.FOREIGN) + def _activate(self, relay: Any) -> _RelayPluginConfigurationState: + try: + if not self._initialize(relay): + return _RelayPluginConfigurationState.DISABLED + except Exception as exc: + self._activation = None + logger.warning("Hermes Relay plugin initialization failed: %s", exc, exc_info=True) + return _RelayPluginConfigurationState.FAILED + self._relay = relay + return _RelayPluginConfigurationState.ACTIVE - try: - if not self._initialize(relay): - return self._remember(owner_id, _RelayPluginConfigurationState.DISABLED) - except Exception as exc: - self._activation = None - logger.warning("Hermes Relay plugin initialization failed: %s", exc, exc_info=True) - return self._remember(owner_id, _RelayPluginConfigurationState.FAILED) - - self._active = True - self._relay = relay - state = self._remember(owner_id, _RelayPluginConfigurationState.ACTIVE) - logger.info( - "Relay plugins are active process-wide and apply to all profiles " - "hosted by this Hermes process." + def _preflight(self, relay: Any) -> _RelayPluginConfigurationState | None: + """Return a terminal state when the process cannot take ownership; None to proceed.""" + if self._relay is not None and not self._clear_active(): + logger.warning( + "Hermes Relay plugin cleanup is still pending; refusing to replace the process-global configuration" ) - return state + return _RelayPluginConfigurationState.FAILED + try: + existing_report = relay.plugin.report() + except Exception: + logger.warning( + "Hermes could not determine whether a process-global Relay " + "plugin configuration is already active; refusing to replace it", + exc_info=True, + ) + return _RelayPluginConfigurationState.FAILED + if existing_report is not None: + logger.warning( + "A process-global Relay plugin configuration is already active " + "outside Hermes native ownership; leaving it unchanged and " + "disabling Hermes-managed Relay middleware for this process" + ) + return _RelayPluginConfigurationState.FOREIGN + return None def _initialize(self, relay: Any) -> bool: """Initialize Relay from the selected plugins.toml; False when none is selected.""" @@ -315,60 +276,41 @@ class _ProcessRelayPluginConfiguration: plugin_config, dynamic_plugins = configured_inputs if dynamic_plugins: try: - activation = _resolve_plugin_awaitable( - relay.plugin.initialize_with_dynamic_plugins(plugin_config, dynamic_plugins) - ) + initialize = relay.plugin.initialize_with_dynamic_plugins + activation = _resolve_plugin_awaitable(initialize(plugin_config, dynamic_plugins)) if activation is None: - raise RuntimeError( - "NeMo Relay dynamic plugin initialization returned no activation handle" - ) + raise RuntimeError("NeMo Relay dynamic plugin initialization returned no activation handle") self._activation = activation except Exception as exc: raise RuntimeError("Hermes Relay dynamic plugin activation failed") from exc if self._activation is None: - # Hermes only enters Relay's initialization path after an - # explicit opt-in; Relay owns any subsequent ambient layering. + # Reached only after explicit opt-in; Relay owns any ambient layering. _resolve_plugin_awaitable(relay.plugin.initialize(plugin_config)) return True - def _remember( - self, owner_id: int, state: _RelayPluginConfigurationState - ) -> _RelayPluginConfigurationState: - """Retain one process decision for all concurrently hosted profiles.""" - self._owners.add(owner_id) - self._state = state - return state - - def _reset_if_cleared(self) -> None: - if self._clear_active(): - self._state = _RelayPluginConfigurationState.UNINITIALIZED - def release(self, owner: Any) -> None: """Release one host and clear Relay after the final host exits.""" - owner_id = id(owner) with self._lock: - if owner_id not in self._owners: + if id(owner) not in self._owners: return - self._owners.remove(owner_id) - if not self._owners: - self._reset_if_cleared() + self._owners.remove(id(owner)) + self.retry_pending_cleanup() def reset_for_tests(self) -> None: """Clear process-global state left by directly constructed test hosts.""" with self._lock: self._owners.clear() - self._reset_if_cleared() + self.retry_pending_cleanup() def retry_pending_cleanup(self) -> None: """Retry a failed final cleanup without disrupting live owners.""" with self._lock: - if not self._owners: - self._reset_if_cleared() + if not self._owners and self._clear_active(): + self._state = _RelayPluginConfigurationState.UNINITIALIZED def _clear_active(self) -> bool: - relay = self._relay - activation = self._activation - if not self._active or relay is None: + relay, activation = self._relay, self._activation + if relay is None: return True try: _resolve_plugin_awaitable(relay.subscribers.flush_async()) @@ -376,19 +318,16 @@ class _ProcessRelayPluginConfiguration: logger.warning("Hermes Relay plugin subscriber flush failed", exc_info=True) return False try: - if activation is not None: - close = getattr(activation, "close", None) - if not callable(close): - raise RuntimeError("NeMo Relay dynamic plugin activation has no close method") + if activation is None: + _resolve_plugin_awaitable(relay.plugin.clear_async()) + elif callable(close := getattr(activation, "close", None)): _resolve_plugin_awaitable(close()) else: - _resolve_plugin_awaitable(relay.plugin.clear_async()) + raise RuntimeError("NeMo Relay dynamic plugin activation has no close method") except Exception: logger.warning("Hermes Relay plugin configuration cleanup failed", exc_info=True) return False - self._active = False - self._relay = None - self._activation = None + self._relay = self._activation = None return True @@ -407,8 +346,7 @@ class RelayRuntime: self._sessions: dict[str, RelaySession] = {} self._subagent_parents: dict[str, str] = {} self._subagent_parent_handles: dict[str, Any] = {} - self._closing = False - self._shutdown_started = False + self._closing = self._shutdown_started = False self._shutdown_complete = threading.Event() self._operations_idle = threading.Event() self._operations_idle.set() @@ -416,10 +354,10 @@ class RelayRuntime: self._execution_consumers_lock = threading.RLock() self._execution_consumers: set[str] = set() self._plugin_configuration_state = _PLUGIN_CONFIGURATION.acquire(self, self.relay) + # Cleared (with the atexit hook) by the first successful _finish_shutdown. self._plugin_configuration_registered = True if self._plugins_active(): self.retain_managed_execution(RELAY_PLUGINS_EXECUTION_CONSUMER) - self._shutdown_registered = True atexit.register(self.shutdown) def _plugins_active(self) -> bool: @@ -442,55 +380,33 @@ class RelayRuntime: with self._execution_consumers_lock: return bool(self._execution_consumers) - def _subagent_parent_handle(self, session: RelaySession) -> Any: - with self._sessions_lock: - return self._subagent_parent_handles.get(session.session_id) - - def _push_session_scope( - self, context: contextvars.Context, *, exit_fallback: bool = False, **push_kwargs: Any - ) -> Any: - """Push a SESSION_SCOPE Agent scope inside ``context``, bounded by ``_SCOPE_OP_TIMEOUT``. - - ``exit_fallback``: at interpreter shutdown the executor refuses new futures - (RuntimeError); push synchronously instead, since no agent turn waits at exit. - """ - args = (self.relay.scope.push, SESSION_SCOPE, self.relay.ScopeType.Agent) - try: - return ( - _scope_op_executor() - .submit(context.run, *args, input={}, **push_kwargs) - .result(timeout=_SCOPE_OP_TIMEOUT) - ) - except RuntimeError: - if not exit_fallback: - raise - return context.run(*args, input={}, **push_kwargs) - def _open_session_scope( - self, - session: RelaySession, - scope_metadata: dict[str, Any], - *, - resolve_parent: bool, - **push_kwargs: Any, + self, session: RelaySession, scope_metadata: dict[str, Any], *, resolve_parent: bool, + exit_fallback: bool = False, **push_kwargs: Any, ) -> None: - """Push a fresh session scope for ``session`` and record its handle + context. - - Subagent sessions parent under their spawning turn/session handle; - ``resolve_parent`` creates the parent session when its handle is unknown. - """ + """Push a fresh SESSION_SCOPE for ``session`` (bounded by ``_SCOPE_OP_TIMEOUT``); record handle + context. + Subagents parent under their spawning turn/session handle; ``resolve_parent`` creates the parent + session when its handle is unknown. ``exit_fallback``: at interpreter shutdown the executor refuses + new futures; push synchronously instead (no agent turn waits at exit).""" parent_handle = None if session.parent_session_id: - parent_handle = self._subagent_parent_handle(session) + with self._sessions_lock: + parent_handle = self._subagent_parent_handles.get(session.session_id) if parent_handle is None and resolve_parent: parent = self.ensure_session({"session_id": session.parent_session_id}) if parent is not None: parent_handle = parent.handle scope_metadata["nemo_relay_scope_role"] = "subagent" context = contextvars.Context() - session.handle = self._push_session_scope( - context, handle=parent_handle, metadata=scope_metadata, **push_kwargs - ) + args = (self.relay.scope.push, SESSION_SCOPE, self.relay.ScopeType.Agent) + push_kwargs.update(handle=parent_handle, metadata=scope_metadata, input={}) + try: + future = _scope_op_executor().submit(context.run, *args, **push_kwargs) + session.handle = future.result(timeout=_SCOPE_OP_TIMEOUT) + except RuntimeError: + if not exit_fallback: + raise + session.handle = context.run(*args, **push_kwargs) session.context = context def ensure_session( @@ -505,10 +421,8 @@ class RelayRuntime: return None session = self._sessions.get(session_id) if session is None: - session = RelaySession( - session_id=session_id, - parent_session_id=self._subagent_parents.get(session_id, ""), - ) + parent_session_id = self._subagent_parents.get(session_id, "") + session = RelaySession(session_id=session_id, parent_session_id=parent_session_id) self._sessions[session_id] = session with session.lock: if session.closing: @@ -516,11 +430,8 @@ class RelayRuntime: if session.handle is None: try: self._open_session_scope( - session, - {**(metadata or {}), **runtime_metadata(self.runtime_id)}, - resolve_parent=True, - data=data, - exit_fallback=True, + session, {**(metadata or {}), **runtime_metadata(self.runtime_id)}, + resolve_parent=True, data=data, exit_fallback=True, ) except Exception: session.context = None @@ -529,55 +440,36 @@ class RelayRuntime: def rotate_session_scope(self, session: RelaySession, *, reason: str) -> None: """Close the current session scope and open the next segment. - - Called ONLY at a turn boundary (before the turn scope pushes): the scope - stack is LIFO and rotating under a live child would close a parent out - of order. Both native calls are bounded by ``_SCOPE_OP_TIMEOUT``, and - segment bookkeeping advances even when a native call fails so a degraded - rotation cannot retry on every turn. - """ + Called ONLY at a turn boundary: the stack is LIFO and rotating under a live child + would close a parent out of order. Bookkeeping advances even when a native call + fails so a degraded rotation cannot retry on every turn.""" with session.lock: if session.closing or session.handle is None: return old_handle = session.handle - # Advance bookkeeping FIRST: a failed native call must not leave - # rotate_pending set (tight rotation loop on every turn). + # Bookkeeping FIRST: a failed native call must not leave rotate_pending set. session.segment += 1 session.segment_turns = 0 session.rotate_pending = False try: self.run_in_session( - session, - self.relay.scope.pop, - old_handle, - output={"hermes.session.segment_reason": reason}, - metadata=runtime_metadata(self.runtime_id), - timeout=_SCOPE_OP_TIMEOUT, + session, self.relay.scope.pop, old_handle, output={"hermes.session.segment_reason": reason}, + metadata=runtime_metadata(self.runtime_id), timeout=_SCOPE_OP_TIMEOUT, ) except Exception: logger.warning( - "Hermes Relay segment close failed (session=%s segment=%d); " - "abandoning the old segment span", - session.session_id, - session.segment - 1, - exc_info=True, + "Hermes Relay segment close failed (session=%s segment=%d); abandoning the old segment span", + session.session_id, session.segment - 1, exc_info=True, ) scope_metadata = runtime_metadata( - self.runtime_id, - **{ - "hermes.session.segment": session.segment, - "hermes.session.segment_reason": reason, - }, + self.runtime_id, **{"hermes.session.segment": session.segment, "hermes.session.segment_reason": reason}, ) try: self._open_session_scope(session, scope_metadata, resolve_parent=False) except Exception: logger.warning( - "Hermes Relay segment open failed (session=%s segment=%d); " - "keeping the prior scope handle", - session.session_id, - session.segment, - exc_info=True, + "Hermes Relay segment open failed (session=%s segment=%d); keeping the prior scope handle", + session.session_id, session.segment, exc_info=True, ) def register_subagent( @@ -591,12 +483,9 @@ class RelayRuntime: parent = self.ensure_session({"session_id": parent_session_id}) parent_handle = None if parent is None else parent.handle turn = active_turn(parent_session_id) + # active_turn() already proved liveness and, for a RelayRuntime host, an open session. if ( - turn is not None - and not turn.closed - and turn.handle is not None - and turn.lease.host is self - and turn.lease.session is not None + turn is not None and turn.handle is not None and turn.lease.host is self and turn.lease.session.session_id == parent_session_id ): parent_handle = turn.handle @@ -611,16 +500,20 @@ class RelayRuntime: def unregister_subagent(self, event: dict[str, Any]) -> None: """Close a delegated session and forget its parent relationship.""" child_session_id = str(event.get("child_session_id") or "") - if not child_session_id: - return - self.close_session({"session_id": child_session_id}) - self._forget_subagent(child_session_id) + if child_session_id: + self.close_session({"session_id": child_session_id}) + self._forget_subagent(child_session_id) def _forget_subagent(self, session_id: str) -> None: with self._sessions_lock: self._subagent_parents.pop(session_id, None) self._subagent_parent_handles.pop(session_id, None) + def _lookup(self, session_id: str) -> RelaySession | None: + """Registry lookup (closing sessions included) without creating one.""" + with self._sessions_lock: + return self._sessions.get(session_id) + def get_session(self, session_id: str) -> RelaySession | None: """Return an active Hermes Relay session without creating one.""" with self._sessions_lock: @@ -630,9 +523,7 @@ class RelayRuntime: with session.lock: return None if session.closing else session - def _session_context( - self, session: RelaySession, *, allow_closing: bool - ) -> contextvars.Context: + def _session_context(self, session: RelaySession, *, allow_closing: bool) -> contextvars.Context: """Copy the current context and overlay the session's saved Relay vars.""" with session.lock: if session.closing and not allow_closing: @@ -640,47 +531,28 @@ class RelayRuntime: if session.context is None or session.handle is None: raise RuntimeError("Hermes Relay session context is unavailable") relay_context = session.context.copy() - # A copy permits a helper called by an existing Relay callback to - # re-enter the same logical session without re-entering Context. + # A copy lets a helper inside a Relay callback re-enter the session's Context. context = contextvars.copy_context() for variable, value in relay_context.items(): context.run(variable.set, value) return context def run_in_session( - self, - session: RelaySession, - callback: Callable[..., Any], - *args: Any, - allow_closing: bool = False, - timeout: float | None = None, - **kwargs: Any, + self, session: RelaySession, callback: Callable[..., Any], *args: Any, + allow_closing: bool = False, timeout: float | None = None, **kwargs: Any, ) -> Any: """Run a Relay operation against a session's isolated scope stack. - - ``timeout`` (seconds) bounds the native call on a shared daemon - executor; ``TimeoutError`` propagates on breach. ``None`` keeps the - synchronous behavior. Scope lifecycle ops that gate turn/session - completion pass ``_SCOPE_OP_TIMEOUT``: the native ``scope.pop`` is - unbounded, and a wedged pipeline must cost at most one span, never the - agent (the abandoned daemon worker cannot block process exit). - """ - self._begin_operation() - try: + ``timeout`` bounds the native call on the daemon executor (``TimeoutError`` on + breach); ``None`` runs synchronously. Lifecycle ops gating turn/session completion + pass ``_SCOPE_OP_TIMEOUT``: a wedged pipeline must cost one span, never the agent.""" + with self._operation(): return self._run_in_session_untracked( session, callback, *args, allow_closing=allow_closing, timeout=timeout, **kwargs ) - finally: - self._end_operation() def _run_in_session_untracked( - self, - session: RelaySession, - callback: Callable[..., Any], - *args: Any, - allow_closing: bool = False, - timeout: float | None = None, - **kwargs: Any, + self, session: RelaySession, callback: Callable[..., Any], *args: Any, + allow_closing: bool = False, timeout: float | None = None, **kwargs: Any, ) -> Any: """Run inside a session whose host-level lifetime is already held.""" context = self._session_context(session, allow_closing=allow_closing) @@ -691,54 +563,39 @@ class RelayRuntime: if timeout is None: return context.run(invoke) + exceeded = f"Relay scope operation exceeded {timeout}s" try: future = _scope_op_executor().submit(context.run, invoke) except RuntimeError: - # Interpreter shutdown: the executor refuses new futures, but the - # atexit close path must still flush — and still bounded, since a - # wedged native call must not block process exit. + # Interpreter shutdown: the executor refuses new futures, but the atexit close + # path must still flush — still bounded so a wedged call cannot block exit. return _run_on_daemon_thread( - lambda: context.run(invoke), - name="relay-scope-op-exit", - timeout=timeout, - timeout_message=( - f"Relay scope operation exceeded {timeout}s during interpreter " - "shutdown; abandoning the native call so process exit can proceed" - ), + lambda: context.run(invoke), name="relay-scope-op-exit", timeout=timeout, + timeout_message=f"{exceeded} during interpreter shutdown; abandoning the native " + "call so process exit can proceed", ) try: return future.result(timeout=timeout) except FuturesTimeoutError as exc: raise TimeoutError( - f"Relay scope operation exceeded {timeout}s " - f"(session={session.session_id}); abandoning the native call " + f"{exceeded} (session={session.session_id}); abandoning the native call " "so the agent can continue — the span for this scope is lost" ) from exc async def run_in_session_async( - self, - session: RelaySession, - callback: Callable[..., Any], - *args: Any, - allow_closing: bool = False, - **kwargs: Any, + self, session: RelaySession, callback: Callable[..., Any], *args: Any, + allow_closing: bool = False, **kwargs: Any, ) -> Any: """Create and await an operation inside the session's saved context.""" - self._begin_operation() - try: + with self._operation(): context = self._session_context(session, allow_closing=allow_closing) async def invoke() -> Any: self.relay.get_scope_stack() result = callback(*args, **kwargs) - if inspect.isawaitable(result): - return await result - return result + return await result if inspect.isawaitable(result) else result - task = context.run(asyncio.create_task, invoke()) - return await task - finally: - self._end_operation() + return await context.run(asyncio.create_task, invoke()) def _begin_operation(self) -> None: """Admit one Relay call while keeping process plugins alive.""" @@ -754,36 +611,24 @@ class RelayRuntime: if self._active_operations == 0: self._operations_idle.set() + @contextlib.contextmanager + def _operation(self): + """``_begin_operation`` / ``_end_operation`` around one tracked Relay call.""" + self._begin_operation() + try: + yield + finally: + self._end_operation() + def acquire_operation_lease(self) -> RelayOperationLease: """Retain plugin lifetime for work that outlives one Relay await.""" self._begin_operation() return RelayOperationLease(self) - def emit_mark( - self, name: str, event: dict[str, Any], *, data: Any = None, metadata: Any = None - ) -> bool: - """Emit a mark parented to the Hermes session identified by ``event``.""" - session = self.ensure_session(event) - if session is None: - return False - self.run_in_session( - session, - self.relay.scope.event, - name, - handle=session.handle, - data=data, - metadata=metadata, - ) - return True - - def apply_tool_request_intercepts( - self, *, session_id: str, tool_name: str, args: dict[str, Any] - ) -> dict[str, Any]: + def apply_tool_request_intercepts(self, *, session_id: str, tool_name: str, args: dict[str, Any]) -> dict[str, Any]: """Apply Relay request rewriting before Hermes authorizes a tool call.""" - if not self.managed_execution_enabled(): - return args request_intercepts = getattr(getattr(self.relay, "tools", None), "request_intercepts", None) - if not callable(request_intercepts): + if not self.managed_execution_enabled() or not callable(request_intercepts): return args session = self.ensure_session({"session_id": session_id}) if session is None: @@ -792,51 +637,31 @@ class RelayRuntime: return result if isinstance(result, dict) else args def _pop_with_drain( - self, - handle: Any, - *, - output: dict[str, Any], - metadata: dict[str, Any], - session_root: Any, - drain_limit: int, + self, handle: Any, *, output: dict[str, Any], metadata: dict[str, Any], session_root: Any, drain_limit: int, ) -> BaseException | None: """Pop ``handle``; if that fails, drain orphans above it and retry once. - - Returns the retry's error (None on success). Must run inside ONE - ``run_in_session`` callback so ContextVar stack views stay consistent. - """ - try: + Returns the retry's error (None on success). Must run inside ONE ``run_in_session`` + callback so ContextVar stack views stay consistent.""" + with contextlib.suppress(Exception): pop_relay_scope(self.relay, handle, output=output, metadata=metadata) return None - except Exception: - pass drained = 0 for _ in range(drain_limit): top = _current_top(self.relay) if top is None or _same_handle(top, handle): break # Never pop the session root while draining for a nested handle. - if ( - session_root is not None - and _same_handle(top, session_root) - and handle is not session_root - ): + if session_root is not None and _same_handle(top, session_root) and handle is not session_root: break try: - pop_relay_scope( - self.relay, - top, - output={"outcome": "cancelled", "hermes.orphan_drain": True}, - metadata=metadata, - ) + orphan_output = {"outcome": "cancelled", "hermes.orphan_drain": True} + pop_relay_scope(self.relay, top, output=orphan_output, metadata=metadata) drained += 1 except Exception: logger.warning("Hermes Relay orphaned scope drain failed", exc_info=True) break if drained: - logger.warning( - "Hermes Relay drained %d orphaned scope(s) before closing %s", drained, handle - ) + logger.warning("Hermes Relay drained %d orphaned scope(s) before closing %s", drained, handle) try: pop_relay_scope(self.relay, handle, output=output, metadata=metadata) return None @@ -844,42 +669,24 @@ class RelayRuntime: return retry_exc def _close_scope_handle( - self, - session: RelaySession, - handle: Any, - *, - output: dict[str, Any] | None = None, - allow_closing: bool = False, - failure_label: str = "scope close failed", - drain_limit: int = 32, - operation_already_held: bool = False, + self, session: RelaySession, handle: Any, *, output: dict[str, Any] | None = None, allow_closing: bool = False, + failure_label: str = "scope close failed", drain_limit: int = 32, operation_already_held: bool = False, ) -> str | None: """Pop ``handle``, draining orphaned children in the same session context. - - Relay scopes are strict LIFO; empty-stream retries + interrupt can - abandon a physical LLM scope above TURN/SESSION. The whole drain+close - is bounded like the direct pops it replaced: a wedged native pipeline - must never block turn/session completion. Returns a failure string. - """ + Relay scopes are strict LIFO; empty-stream retries + interrupt can abandon a + physical LLM scope above TURN/SESSION. Drain+close is bounded so a wedged pipeline + never blocks turn/session completion. Returns a failure string or None.""" if handle is None: return None - run_in_session = ( - self._run_in_session_untracked if operation_already_held else self.run_in_session - ) + run_in_session = (self._run_in_session_untracked if operation_already_held else self.run_in_session) try: failure = run_in_session( - session, - self._pop_with_drain, - handle, - output=output or {}, - metadata=runtime_metadata(self.runtime_id), - session_root=session.handle, - drain_limit=drain_limit, - allow_closing=allow_closing, - timeout=_SCOPE_OP_TIMEOUT, + session, self._pop_with_drain, handle, output=output or {}, + metadata=runtime_metadata(self.runtime_id), session_root=session.handle, + drain_limit=drain_limit, allow_closing=allow_closing, timeout=_SCOPE_OP_TIMEOUT, ) except Exception as exc: - return f"{failure_label}: {exc}" + failure = exc return None if failure is None else f"{failure_label}: {failure}" def close_session(self, event: dict[str, Any]) -> None: @@ -896,8 +703,7 @@ class RelayRuntime: def _close_session(self, event: dict[str, Any]) -> None: """Close one session already admitted by the host lifecycle gate.""" session_id = _session_id(event) - with self._sessions_lock: - session = self._sessions.get(session_id) + session = self._lookup(session_id) if session is None: self._forget_subagent(session_id) return @@ -908,19 +714,14 @@ class RelayRuntime: session.closing = True if session.handle is not None: failure = self._close_scope_handle( - session, - session.handle, - output={}, - allow_closing=True, - failure_label="session scope close failed", - operation_already_held=True, + session, session.handle, output={}, allow_closing=True, + failure_label="session scope close failed", operation_already_held=True, ) - # Subscriber flushing is process-wide and may wait for publications - # owned by other sessions; final plugin teardown flushes once after all - # tracked operations drain. Flushing here can deadlock an asyncio loop. + # No subscriber flush here: it is process-wide, may wait on other sessions' + # publications and can deadlock an asyncio loop; final plugin teardown flushes once. with self._sessions_lock: if self._sessions.get(session_id) is session: - self._sessions.pop(session_id, None) + del self._sessions[session_id] self._forget_subagent(session_id) if failure: logger.warning("Hermes Relay session %s closed with errors: %s", session_id, failure) @@ -930,42 +731,35 @@ class RelayRuntime: with self._sessions_lock: if self._shutdown_started: return - self._shutdown_started = True - self._closing = True + self._shutdown_started = self._closing = True has_active_operations = self._active_operations > 0 - if has_active_operations: - thread = threading.Thread( - target=self._finish_shutdown_after_operations, - name=f"hermes-nemo-relay-shutdown-{self.runtime_id[:8]}", - daemon=True, - ) - try: - thread.start() - except Exception: - with self._sessions_lock: - self._shutdown_started = False - logger.warning("Hermes Relay deferred shutdown could not start", exc_info=True) + if not has_active_operations: + self._finish_shutdown() return - self._finish_shutdown() - - def _finish_shutdown_after_operations(self) -> None: - self._operations_idle.wait() - self._finish_shutdown() + thread = threading.Thread( + target=lambda: (self._operations_idle.wait(), self._finish_shutdown()), + name=f"hermes-nemo-relay-shutdown-{self.runtime_id[:8]}", daemon=True, + ) + try: + thread.start() + except Exception: + with self._sessions_lock: + self._shutdown_started = False + logger.warning("Hermes Relay deferred shutdown could not start", exc_info=True) def _finish_shutdown(self) -> None: try: with self._sessions_lock: session_ids = list(self._sessions) for session_id in session_ids: - self._safe(self._close_session, {"session_id": session_id}) + _warn_on_error("runtime operation", self._close_session, {"session_id": session_id}) if self._plugin_configuration_registered: if self._plugins_active(): self.release_managed_execution(RELAY_PLUGINS_EXECUTION_CONSUMER) _PLUGIN_CONFIGURATION.release(self) self._plugin_configuration_registered = False - if self._shutdown_registered: - self._safe(atexit.unregister, self.shutdown, quiet=True) - self._shutdown_registered = False + with contextlib.suppress(Exception): + atexit.unregister(self.shutdown) except Exception: with self._sessions_lock: self._shutdown_started = False @@ -974,15 +768,6 @@ class RelayRuntime: with self._sessions_lock: self._shutdown_complete.set() - @staticmethod - def _safe(callback: Callable[..., Any], *args: Any, quiet: bool = False, **kwargs: Any) -> Any: - try: - return callback(*args, **kwargs) - except Exception: - if not quiet: - logger.warning("Hermes Relay runtime operation failed", exc_info=True) - return None - @dataclass(frozen=True) class NoopRelayRuntime: @@ -991,24 +776,16 @@ class NoopRelayRuntime: profile_key: str reason: str - def apply_tool_request_intercepts( - self, *, session_id: str, tool_name: str, args: dict[str, Any] - ) -> dict[str, Any]: - del session_id, tool_name + def apply_tool_request_intercepts(self, *, session_id: str, tool_name: str, args: dict[str, Any]) -> dict[str, Any]: return args @staticmethod def retain_managed_execution(consumer: str) -> None: - del consumer + pass release_managed_execution = retain_managed_execution - - @staticmethod - def managed_execution_enabled() -> bool: - return False - - def shutdown(self) -> None: - """No resources are allocated on unsupported platforms.""" + managed_execution_enabled = staticmethod(lambda: False) + shutdown = staticmethod(lambda: None) # no resources are allocated on unsupported platforms RelayHost = RelayRuntime | NoopRelayRuntime @@ -1021,13 +798,8 @@ class RelayHostRegistry: self._lock = threading.RLock() self._hosts: dict[str, RelayHost] = {} - def for_profile( - self, profile_key: str | None = None, *, create: bool = True - ) -> RelayHost | None: + def for_profile(self, profile_key: str | None = None, *, create: bool = True) -> RelayHost | None: key = profile_key or current_profile_key() - host = self._hosts.get(key) - if host is not None or not create: - return host with self._lock: host = self._hosts.get(key) if host is not None or not create: @@ -1065,9 +837,8 @@ class ConversationLease: def live_runtime(self) -> RelayRuntime | None: """Return the real Relay host when this lease owns an open session.""" - if isinstance(self.host, RelayRuntime) and self.session is not None: - return self.host - return None + host = self.host + return host if isinstance(host, RelayRuntime) and self.session is not None else None @dataclass @@ -1091,13 +862,12 @@ _CURRENT_TURN: contextvars.ContextVar[RelayTurnContext | None] = contextvars.Con "hermes_relay_turn", default=None ) -# Depth of managed Relay callbacks on the current logical call path (>0 while the -# native pipeline is mid-dispatch of a Hermes tool/LLM callback). Nested managed -# execution there is structurally broken: the native pipeline binds its Futures to -# the OUTER call's event loop, which is blocked inside the synchronous callback -# ("attached to a different loop" at best, deadlock or "Event loop is closed" at -# worst), so resolve_execution_context() bypasses Relay while set. A ContextVar so -# the marker follows contextvars.copy_context() into worker threads / per-thread loops. +# Depth of managed Relay callbacks on the current call path (>0 while the native pipeline +# is mid-dispatch of a Hermes tool/LLM callback). Nested managed execution there is +# structurally broken: the pipeline binds its Futures to the OUTER call's event loop, which +# is blocked inside the synchronous callback (wrong loop / deadlock / "Event loop is +# closed"), so resolve_execution_context() bypasses Relay while set. A ContextVar so the +# marker follows copy_context() into worker threads / per-thread loops. _MANAGED_CALLBACK_DEPTH: contextvars.ContextVar[int] = contextvars.ContextVar( "hermes_relay_managed_callback_depth", default=0 ) @@ -1105,11 +875,8 @@ _MANAGED_CALLBACK_DEPTH: contextvars.ContextVar[int] = contextvars.ContextVar( class managed_callback_guard: """Mark the current context as inside a managed Relay callback. - - Wrap the ``invoke()`` callbacks handed to the native pipeline; everything - they transitively call (including work forwarded via copy_context()) sees - the marker and runs unmanaged. - """ + Wrap the ``invoke()`` callbacks handed to the native pipeline; everything they + transitively call (incl. work forwarded via copy_context()) runs unmanaged.""" def __enter__(self) -> "managed_callback_guard": self._token = _MANAGED_CALLBACK_DEPTH.set(_MANAGED_CALLBACK_DEPTH.get() + 1) @@ -1119,6 +886,22 @@ class managed_callback_guard: _MANAGED_CALLBACK_DEPTH.reset(self._token) +def _warn_on_error(what: str, callback: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + """Run fail-open telemetry work: log ``Hermes Relay failed`` and return None on error.""" + try: + return callback(*args, **kwargs) + except Exception: + logger.warning("Hermes Relay %s failed", what, exc_info=True) + return None + + +def _flag_open_session(session: RelaySession, flag: str) -> None: + """Set a pending-rotation/close flag unless the session is already closing.""" + with session.lock: + if not session.closing: + setattr(session, flag, True) + + class RelaySessionCoordinator: """Own semantic conversation and turn lifetimes for Hermes core.""" @@ -1129,9 +912,7 @@ class RelaySessionCoordinator: self._active_turns_lock = threading.RLock() self._active_turns: dict[tuple[str, str], set[int]] = {} - def register_session_initializer( - self, name: str, callback: Callable[[RelayRuntime, dict[str, Any]], None] - ) -> None: + def register_session_initializer(self, name: str, callback: Callable[[RelayRuntime, dict[str, Any]], None]) -> None: """Register idempotent profile/session preparation before scope creation.""" with self._initializer_lock: self._session_initializers[name] = callback @@ -1146,91 +927,61 @@ class RelaySessionCoordinator: logger.warning("Hermes Relay session initializer failed: %s", name, exc_info=True) def acquire_conversation( - self, - *, - profile_key: str, - session_id: str, - platform: str, - parent_session_id: str = "", - model: str = "", + self, *, profile_key: str, session_id: str, platform: str, parent_session_id: str = "", model: str = "", ) -> ConversationLease: - host = self.registry.for_profile(profile_key) - if host is None: - host = NoopRelayRuntime(profile_key, "Relay host creation was disabled") + host = self.registry.for_profile(profile_key) or NoopRelayRuntime( + profile_key, "Relay host creation was disabled" + ) session = None if isinstance(host, RelayRuntime): - try: - self._prepare_session(host, { - "profile_key": profile_key, - "session_id": session_id, - "platform": platform, - "parent_session_id": parent_session_id, - "model": model, - }) - metadata = {"hermes.execution_surface": platform or "unknown"} - if parent_session_id and parent_session_id != session_id: - session = host.register_subagent( - {"parent_session_id": parent_session_id, "child_session_id": session_id}, - metadata=metadata, - ) - else: - session = host.ensure_session({"session_id": session_id}, metadata=metadata) - except Exception: - logger.warning("Hermes Relay conversation initialization failed", exc_info=True) + context = { + "profile_key": profile_key, "session_id": session_id, "platform": platform, + "parent_session_id": parent_session_id, "model": model, + } + session = _warn_on_error("conversation initialization", self._open_conversation_session, host, context) return ConversationLease( - profile_key=profile_key, - session_id=session_id, - platform=platform, - host=host, - session=session, - parent_session_id=parent_session_id, + profile_key=profile_key, session_id=session_id, platform=platform, host=host, + session=session, parent_session_id=parent_session_id, ) - def begin_turn( - self, lease: ConversationLease, *, turn_id: str, task_id: str - ) -> RelayTurnContext: + def _open_conversation_session(self, host: RelayRuntime, context: dict[str, Any]) -> RelaySession | None: + self._prepare_session(host, context) + session_id, parent_session_id = context["session_id"], context["parent_session_id"] + metadata = {"hermes.execution_surface": context["platform"] or "unknown"} + if parent_session_id and parent_session_id != session_id: + return host.register_subagent( + {"parent_session_id": parent_session_id, "child_session_id": session_id}, metadata=metadata, + ) + return host.ensure_session({"session_id": session_id}, metadata=metadata) + + def begin_turn(self, lease: ConversationLease, *, turn_id: str, task_id: str) -> RelayTurnContext: if lease.released: raise RuntimeError("Hermes Relay conversation lease is released") turn = RelayTurnContext(lease=lease, turn_id=turn_id, task_id=task_id) key = (lease.profile_key, lease.session_id) with self._active_turns_lock: if self._active_turns.get(key): - # A Relay session owns one physical scope stack; concurrent turns - # would create sibling scopes whose completion order is not LIFO. + # One physical scope stack per session; concurrent turns would create + # sibling scopes whose completion order is not LIFO. turn.relay_enabled = False logger.warning( - "Skipping Relay instrumentation for concurrent Hermes turn " - "%s in session %s", - turn_id, - lease.session_id, + "Skipping Relay instrumentation for concurrent Hermes turn %s in session %s", + turn_id, lease.session_id, ) else: self._active_turns[key] = {id(turn)} turn._active_registered = True host = lease.live_runtime() if turn.relay_enabled else None if host is not None: - # Segment rotation (pending compaction flag or max_turns cap) happens - # HERE — the only point with no live turn scope on the session's - # stack, so the session scope can close/reopen without breaking LIFO. - try: - self._maybe_rotate_segment(host, lease.session) - except Exception: - logger.warning("Hermes Relay segment rotation failed", exc_info=True) - try: - turn.handle = host.run_in_session( - lease.session, - host.relay.scope.push, - TURN_SCOPE, - host.relay.ScopeType.Function, - handle=lease.session.handle, - input={}, - metadata=runtime_metadata( - host.runtime_id, **{"hermes.execution_surface": lease.platform or "unknown"} - ), - timeout=_SCOPE_OP_TIMEOUT, - ) - except Exception: - logger.warning("Hermes Relay turn initialization failed", exc_info=True) + # Segment rotation happens HERE — the only point with no live turn scope on + # the stack, so the session scope can close/reopen without breaking LIFO. + _warn_on_error("segment rotation", self._maybe_rotate_segment, host, lease.session) + turn.handle = _warn_on_error( + "turn initialization", host.run_in_session, lease.session, host.relay.scope.push, + TURN_SCOPE, host.relay.ScopeType.Function, handle=lease.session.handle, input={}, + metadata=runtime_metadata(host.runtime_id, **{"hermes.execution_surface": lease.platform or "unknown"}), + timeout=_SCOPE_OP_TIMEOUT, + ) turn._previous_turn = _CURRENT_TURN.get() _CURRENT_TURN.set(turn) return turn @@ -1256,110 +1007,77 @@ class RelaySessionCoordinator: if host is not None: self._close_turn_scope(host, turn, outcome=outcome) finally: + if turn._active_registered and host is not None: + with contextlib.suppress(Exception), lease.session.lock: # accounting never blocks + lease.session.segment_turns += 1 # max_turns rotation trigger try: - # Segment turn accounting (max_turns rotation trigger). - if turn._active_registered and host is not None: - with lease.session.lock: - lease.session.segment_turns += 1 - except Exception: # noqa: BLE001 - accounting must never block - pass - try: - # Delegated agents own one turn: close their conversation - # while the active-turn guard is still held so a parent - # timeout fallback cannot race this terminal boundary. + # Delegated agents own one turn: close their conversation while the + # active-turn guard is held so a parent timeout fallback cannot race it. if lease.parent_session_id and isinstance(lease.host, RelayRuntime): - lease.host.unregister_subagent({"child_session_id": lease.session_id}) - except Exception: - logger.warning( - "Hermes Relay child conversation finalization failed", exc_info=True - ) + _warn_on_error( + "child conversation finalization", lease.host.unregister_subagent, + {"child_session_id": lease.session_id}, + ) finally: self._unregister_active_turn(turn) self._reset_turn_context(turn) self._consume_deferred_close(lease) - def _close_turn_scope( - self, host: RelayRuntime, turn: RelayTurnContext, *, outcome: str - ) -> None: + def _close_turn_scope(self, host: RelayRuntime, turn: RelayTurnContext, *, outcome: str) -> None: """Pop the turn's logical LLM children, then the turn scope itself (LIFO).""" self._finish_logical_calls(turn, outcome=outcome) - if turn.handle is None: - return failure = host._close_scope_handle( - turn.lease.session, - turn.handle, - output={"outcome": outcome}, - failure_label="turn scope close failed", + turn.lease.session, turn.handle, output={"outcome": outcome}, failure_label="turn scope close failed", ) if failure: logger.warning("Hermes Relay turn finalization failed: %s", failure) def _consume_deferred_close(self, lease: Any) -> None: """Close a session whose rotating-compaction close was deferred. + ``notify_session_compacted`` sets ``close_pending`` when the old session had a live + turn (closing then breaks LIFO). The last live turn consumes it here after its own + scope popped and it left the active-turn table.""" + # Telemetry must never block end_turn. + _warn_on_error("deferred session close", self._consume_deferred_close_unguarded, lease) - ``notify_session_compacted`` sets ``close_pending`` when the old session - still had a live turn (closing then would break LIFO). The turn that was - live consumes it here, after its own scope popped and it left the - active-turn table; if another turn is still live, that turn's end_turn - consumes it instead. - """ - try: - host = lease.live_runtime() - if host is None: - return - session = lease.session - with session.lock: - pending = session.close_pending and not session.closing - if not pending: - return - if self.has_active_turn(profile_key=lease.profile_key, session_id=lease.session_id): - return + def _consume_deferred_close_unguarded(self, lease: ConversationLease) -> None: + host = lease.live_runtime() + if host is None: + return + with lease.session.lock: + pending = lease.session.close_pending and not lease.session.closing + if pending and not self.has_active_turn(profile_key=lease.profile_key, session_id=lease.session_id): host.close_session({"session_id": lease.session_id}) - except Exception: # noqa: BLE001 - telemetry must never block end_turn - logger.warning("Hermes Relay deferred session close failed", exc_info=True) - def notify_session_compacted( - self, *, profile_key: str, session_id: str, old_session_id: str = "" - ) -> None: + def notify_session_compacted(self, *, profile_key: str, session_id: str, old_session_id: str = "") -> None: """React to a completed compaction, per compaction mode. + In-place (``old_session_id`` empty/equal): flag rotation for the next turn boundary + — never rotate immediately, a turn may be live and rotating under it breaks LIFO. + Rotating (ids differ): the next turn gets a fresh session under the new id, so close + the OLD session now or its scope stays an unexported orphan. Unknown sessions and + disabled config are silent no-ops.""" + # Telemetry must never block compaction. + _warn_on_error( + "compaction notification", self._notify_session_compacted_unguarded, profile_key, session_id, old_session_id + ) - In-place compaction (``old_session_id`` empty or equal): flag the session - for rotation at its next turn boundary — never rotate immediately, since - a compaction can finish while a turn is live and rotating under it would - break LIFO; ``begin_turn`` consumes the flag. Rotating compaction (ids - differ): the next turn gets a fresh session under the new id, so close - the OLD session now or its scope stays an unexported orphan. Unknown - sessions and disabled config are silent no-ops. - """ - try: - if not _segments_config()["on_compaction"]: - return - host = self.registry.for_profile(profile_key) - if not isinstance(host, RelayRuntime): - return - if old_session_id and old_session_id != session_id: - # If a turn is still LIVE on the old session, closing now would - # pop the session scope under it (LIFO) — defer to its end_turn. - with host._sessions_lock: - old_session = host._sessions.get(old_session_id) - if old_session is not None and self.has_active_turn( - profile_key=profile_key, session_id=old_session_id - ): - with old_session.lock: - if not old_session.closing: - old_session.close_pending = True - return + def _notify_session_compacted_unguarded(self, profile_key: str, session_id: str, old_session_id: str) -> None: + if not _segments_config()["on_compaction"]: + return + host = self.registry.for_profile(profile_key) + if not isinstance(host, RelayRuntime): + return + if old_session_id and old_session_id != session_id: + # A LIVE turn on the old session: closing now would pop under it (LIFO). + old_session = host._lookup(old_session_id) + if old_session is not None and self.has_active_turn(profile_key=profile_key, session_id=old_session_id): + _flag_open_session(old_session, "close_pending") + else: host.close_session({"session_id": old_session_id}) - return - with host._sessions_lock: - session = host._sessions.get(session_id) - if session is None: - return - with session.lock: - if not session.closing: - session.rotate_pending = True - except Exception: # noqa: BLE001 - telemetry must never block compaction - logger.warning("Hermes Relay compaction notification failed", exc_info=True) + return + session = host._lookup(session_id) + if session is not None: + _flag_open_session(session, "rotate_pending") def has_active_turn(self, *, profile_key: str, session_id: str) -> bool: """Return whether a turn is still running for one profile/session.""" @@ -1375,15 +1093,14 @@ class RelaySessionCoordinator: if active is not None: active.discard(id(turn)) if not active: - self._active_turns.pop(key, None) + del self._active_turns[key] turn._active_registered = False def finish_logical_calls(self, turn: RelayTurnContext, *, outcome: str) -> None: """Close logical LLM children before sibling task aggregation scopes.""" with turn.finalize_lock: - if turn.closed: - return - self._finish_logical_calls(turn, outcome=outcome) + if not turn.closed: + self._finish_logical_calls(turn, outcome=outcome) @staticmethod def _finish_logical_calls(turn: RelayTurnContext, *, outcome: str) -> None: @@ -1394,21 +1111,19 @@ class RelaySessionCoordinator: with turn.logical_llm_lock: logical_calls = list(turn.logical_llm_calls.items()) turn.logical_llm_calls.clear() - for index in range(len(logical_calls) - 1, -1, -1): - request_id, logical_handle = logical_calls[index] + while logical_calls: + _request_id, logical_handle = logical_calls[-1] failure = host._close_scope_handle( - lease.session, - logical_handle, - output={"outcome": outcome}, + lease.session, logical_handle, output={"outcome": outcome}, failure_label="logical LLM scope close failed", ) if failure is None: + logical_calls.pop() continue with turn.logical_llm_lock: - # Relay scopes are stack-owned: if the newest remaining handle - # cannot close even after orphan drain, older ones cannot close - # safely either — retain the unclosed prefix for diagnostics. - for pending_request_id, pending_handle in logical_calls[: index + 1]: + # Stack-owned scopes: if the newest handle cannot close even after orphan + # drain, older ones cannot either — retain the unclosed prefix. + for pending_request_id, pending_handle in logical_calls: turn.logical_llm_calls.setdefault(pending_request_id, pending_handle) logger.warning("Hermes Relay logical LLM finalization failed: %s", failure) break @@ -1418,14 +1133,12 @@ class RelaySessionCoordinator: """Unwind ``turn`` without disturbing a newer context-local turn.""" if _CURRENT_TURN.get() is not turn: return - previous = turn._previous_turn - seen = {id(turn)} - while previous is not None and previous.closed: - if id(previous) in seen: - previous = None - break + previous, seen = turn._previous_turn, {id(turn)} + while previous is not None and previous.closed and id(previous) not in seen: seen.add(id(previous)) previous = previous._previous_turn + if previous is not None and previous.closed: # cycle: no live ancestor + previous = None _CURRENT_TURN.set(previous) @staticmethod @@ -1458,76 +1171,49 @@ def active_turn(session_id: str | None = None) -> RelayTurnContext | None: turn = current_turn() if turn is None or not turn.relay_enabled or turn.closed or turn.lease.released: return None - if turn.lease.profile_key != current_profile_key(): + lease = turn.lease + if lease.profile_key != current_profile_key() or (session_id is not None and lease.session_id != session_id): return None - if session_id is not None and turn.lease.session_id != session_id: + if isinstance(lease.host, RelayRuntime) and ( + lease.session is None or lease.host.get_session(lease.session_id) is not lease.session + ): return None - if isinstance(turn.lease.host, RelayRuntime): - if turn.lease.session is None: - return None - if turn.lease.host.get_session(turn.lease.session_id) is not turn.lease.session: - return None return turn -def resolve_execution_context( - session_id: str, -) -> tuple[RelayRuntime | None, RelaySession | None, Any]: +def resolve_execution_context(session_id: str) -> tuple[RelayRuntime | None, RelaySession | None, Any]: """Resolve one active turn/session parent for managed Relay execution.""" if _MANAGED_CALLBACK_DEPTH.get() > 0: - # Inside a managed Relay callback: nested managed execution is impossible - # (see _MANAGED_CALLBACK_DEPTH). Run unmanaged; the outer scope still - # records the tool-level event for observability. + # Nested managed execution is impossible (see _MANAGED_CALLBACK_DEPTH); the + # outer scope still records the tool-level event. return None, None, None - inherited_turn = current_turn() - if inherited_turn is not None and (not inherited_turn.relay_enabled or inherited_turn.closed): + if not relay_instrumentation_enabled(): return None, None, None turn = active_turn(session_id) host = turn.lease.live_runtime() if turn is not None else None if host is not None: session = turn.lease.session return host, session, turn.handle or session.handle - # Managed-execution consumers create and retain the profile host before - # reaching an out-of-turn adapter; never initialize Relay for the default - # no-consumer path. + # Consumers retain the profile host before reaching an out-of-turn adapter; never + # initialize Relay for the default no-consumer path. runtime = get_runtime(create=False) if runtime is None or not runtime.managed_execution_enabled(): return None, None, None - session = runtime.get_session(session_id) - if session is None: - session = runtime.ensure_session({"session_id": session_id}) + session = runtime.get_session(session_id) or runtime.ensure_session({"session_id": session_id}) return runtime, session, None if session is None else session.handle -def emit_mark(name: str, *, session_id: str, data: Any = None, metadata: Any = None) -> bool: - """Emit a fail-open Relay mark under a Hermes session.""" - runtime = get_runtime(create=False) - if runtime is None: - return False - try: - return runtime.emit_mark(name, {"session_id": session_id}, data=data, metadata=metadata) - except Exception: - logger.warning("Hermes Relay mark failed: %s", name, exc_info=True) - return False - - -def apply_tool_request_intercepts( - *, session_id: str, tool_name: str, args: dict[str, Any] -) -> dict[str, Any]: +def apply_tool_request_intercepts(*, session_id: str, tool_name: str, args: dict[str, Any]) -> dict[str, Any]: """Return Relay-rewritten arguments at Hermes's authorization boundary.""" if not session_id: return args runtime = get_runtime(create=False) if runtime is None: return args - return runtime.apply_tool_request_intercepts( - session_id=session_id, tool_name=tool_name, args=args - ) + return runtime.apply_tool_request_intercepts(session_id=session_id, tool_name=tool_name, args=args) -def _is_relay_wrapped_callback_error( - relay_error: BaseException, callback_error: BaseException -) -> bool: +def _is_relay_wrapped_callback_error(relay_error: BaseException, callback_error: BaseException) -> bool: """Match Relay's native callback wrapper without masking policy errors.""" if relay_error is callback_error: return True @@ -1535,15 +1221,10 @@ def _is_relay_wrapped_callback_error( return False callback_type = callback_error.__class__ type_names = { - callback_type.__name__, - callback_type.__qualname__, - f"{callback_type.__module__}.{callback_type.__qualname__}", + callback_type.__name__, callback_type.__qualname__, f"{callback_type.__module__}.{callback_type.__qualname__}", } message = str(relay_error) - return any( - message.startswith(f"internal error: {type_name}: {callback_error}") - for type_name in type_names - ) + return any(message.startswith(f"internal error: {type_name}: {callback_error}") for type_name in type_names) def get_runtime(*, create: bool = True, profile_key: str | None = None) -> RelayRuntime | None: @@ -1559,10 +1240,7 @@ def current_profile_key() -> str: return str(home.resolve()) raw = str(home) cached = _PROFILE_KEY_CACHE.get(raw) - if cached is not None: - return cached - resolved = str(home.resolve()) - return _PROFILE_KEY_CACHE.setdefault(raw, resolved) + return cached if cached is not None else _PROFILE_KEY_CACHE.setdefault(raw, str(home.resolve())) def _load_nemo_relay() -> Any: @@ -1574,8 +1252,7 @@ def _configured_plugin_inputs(relay: Any) -> tuple[dict[str, Any], list[Any]] | """Load selected plugin inputs, or return ``None`` when none were selected.""" configured = os.environ.get(RELAY_PLUGINS_CONFIG_ENV, "").strip() if not configured: - legacy_vars = configured_legacy_relay_env_vars(os.environ) - if legacy_vars: + if legacy_vars := configured_legacy_relay_env_vars(os.environ): logger.warning( "Legacy NeMo Relay exporter variables are set but no %s was " "provided. %s no longer activate Relay exporters; migrate the " @@ -1584,22 +1261,18 @@ def _configured_plugin_inputs(relay: Any) -> tuple[dict[str, Any], list[Any]] | ", ".join(legacy_vars), ) return None - config_path = Path(configured).expanduser() try: with config_path.open("rb") as config_file: config = tomllib.load(config_file) if "dynamic_plugins" in config: raise ValueError( - "Hermes [[dynamic_plugins]] records are unsupported; use Relay " - "[[plugins.dynamic]] records" + "Hermes [[dynamic_plugins]] records are unsupported; use Relay [[plugins.dynamic]] records" ) dynamic_plugins: list[Any] = [] if "plugins" in config: dynamic_plugins = relay.plugin.load_dynamic_plugin_activation_specs(config_path) - plugin_config = dict(config) - plugin_config.pop("plugins", None) - return plugin_config, dynamic_plugins + return {k: v for k, v in config.items() if k != "plugins"}, dynamic_plugins except Exception as exc: raise _RelayPluginConfigurationLoadError( "Hermes Relay plugin configuration could not be loaded from " @@ -1615,9 +1288,7 @@ def _resolve_plugin_awaitable(value: Any) -> Any: asyncio.get_running_loop() except RuntimeError: return asyncio.run(value) - return _run_on_daemon_thread( - lambda: asyncio.run(value), name="hermes-nemo-relay-plugin-lifecycle" - ) + return _run_on_daemon_thread(lambda: asyncio.run(value), name="hermes-nemo-relay-plugin-lifecycle") def _session_id(event: dict[str, Any]) -> str: diff --git a/agent/relay_tools.py b/agent/relay_tools.py index c38d44068e..42f4e45bc7 100644 --- a/agent/relay_tools.py +++ b/agent/relay_tools.py @@ -2,49 +2,40 @@ from __future__ import annotations -import asyncio import contextvars -import inspect import json import logging from collections.abc import Callable from typing import Any -from agent import relay_runtime +from agent import relay_llm, relay_runtime logger = logging.getLogger(__name__) def execute( - tool_name: str, - args: dict[str, Any], - callback: Callable[[dict[str, Any]], Any], - *, - session_id: str, - metadata: dict[str, Any] | None = None, + tool_name: str, args: dict[str, Any], callback: Callable[[dict[str, Any]], Any], *, + session_id: str, metadata: dict[str, Any] | None = None, ) -> tuple[Any, dict[str, Any]]: """Run one tool call through Relay and return its final arguments.""" runtime, session, parent = relay_runtime.resolve_execution_context(session_id) if runtime is None or session is None or not runtime.managed_execution_enabled(): return callback(args), args - observed_args = args raw_result: dict[str, Any] = {} callback_error: BaseException | None = None callback_context = contextvars.copy_context() + def guarded(final_args: dict[str, Any]) -> Any: + # Everything the tool transitively calls (incl. auxiliary LLM calls on worker + # threads) must bypass managed Relay: the pipeline's Futures bind to THIS loop, + # which is blocked until the tool returns. + with relay_runtime.managed_callback_guard(): + return callback(final_args) + def invoke(next_args: Any) -> Any: nonlocal callback_error, observed_args observed_args = next_args if isinstance(next_args, dict) else args - - def guarded(final_args: dict[str, Any]) -> Any: - # Everything the tool transitively calls (including auxiliary LLM - # calls it forwards to worker threads) must bypass managed Relay - # execution — the native pipeline's Futures bind to THIS loop, - # which is blocked until the tool returns (#77244). - with relay_runtime.managed_callback_guard(): - return callback(final_args) - try: result = callback_context.copy().run(guarded, observed_args) except BaseException as exc: @@ -57,35 +48,23 @@ def execute( try: managed = _run_awaitable( runtime.run_in_session_async( - session, - runtime.relay.tools.execute, - tool_name, - _jsonable(args), - invoke, - handle=parent, - metadata=_jsonable(metadata or {}), + session, runtime.relay.tools.execute, tool_name, _jsonable(args), invoke, + handle=parent, metadata=_jsonable(metadata or {}), ) ) except BaseException as exc: - if ( - callback_error is not None - and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error) - ): + if callback_error is not None and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error): raise callback_error if isinstance(exc, Exception) and callback_error is None and "value" in raw_result: logger.warning( - "NeMo Relay tool post-processing failed after dispatch success; " - "returning the Hermes tool result", + "NeMo Relay tool post-processing failed after dispatch success; returning the Hermes tool result", exc_info=True, ) return raw_result["value"], observed_args raise - if "value" in raw_result and _json_equal(managed, raw_result["json"]): return raw_result["value"], observed_args - if isinstance(managed, str): - return managed, observed_args - return json.dumps(_jsonable(managed), ensure_ascii=False), observed_args + return (managed if isinstance(managed, str) else json.dumps(_jsonable(managed), ensure_ascii=False)), observed_args def _jsonable(value: Any) -> Any: @@ -98,8 +77,7 @@ def _jsonable(value: Any) -> Any: model_dump = getattr(value, "model_dump", None) if callable(model_dump): try: - # warnings=False: suppress pydantic's serializer UserWarnings on - # generic-union SDK models; they would leak to the CLI mid-turn. + # warnings=False: pydantic's generic-union warning would leak to the CLI mid-turn. try: return _jsonable(model_dump(mode="json", warnings=False)) except TypeError: @@ -114,20 +92,12 @@ def _jsonable(value: Any) -> Any: def _json_equal(left: Any, right: Any) -> bool: try: - return json.dumps( - _jsonable(left), sort_keys=True, separators=(",", ":") - ) == json.dumps(_jsonable(right), sort_keys=True, separators=(",", ":")) + return relay_llm._canonical_json(left, _jsonable) == relay_llm._canonical_json(right, _jsonable) except (TypeError, ValueError): return left == right def _run_awaitable(value: Any) -> Any: - if not inspect.isawaitable(value): - return value - try: - asyncio.get_running_loop() - except RuntimeError: - return asyncio.run(value) - raise RuntimeError( - "Synchronous Hermes Relay tool execution cannot run on an active event-loop thread" + return relay_llm._run_awaitable( + value, loop_error="Synchronous Hermes Relay tool execution cannot run on an active event-loop thread", ) diff --git a/agent/transports/__init__.py b/agent/transports/__init__.py index 23557be795..83f31ce9ac 100644 --- a/agent/transports/__init__.py +++ b/agent/transports/__init__.py @@ -1,8 +1,9 @@ """Transport registry for provider response normalization. - transport = get_transport("anthropic_messages") - result = transport.normalize_response(raw_response) -""" + result = transport.normalize_response(raw_response)""" + +import contextlib +import importlib from agent.transports.types import ( # noqa: F401 NormalizedResponse, @@ -24,14 +25,10 @@ def register_transport(api_mode: str, transport_cls: type) -> None: def get_transport(api_mode: str): """Return a transport instance for ``api_mode``, or None so callers can fall back to the legacy path.""" - if not _discovered: + # A directly-imported transport leaves the registry partial; (re)discover on first use and on misses. + if not _discovered or api_mode not in _REGISTRY: _discover_transports() cls = _REGISTRY.get(api_mode) - if cls is None: - # A directly-imported transport module leaves the registry partially - # populated; discover on misses so import order can't hide a valid api_mode. - _discover_transports() - cls = _REGISTRY.get(api_mode) return None if cls is None else cls() @@ -39,10 +36,6 @@ def _discover_transports() -> None: """Import all transport modules to trigger auto-registration.""" global _discovered _discovered = True - import importlib - for name in _TRANSPORT_MODULES: - try: + with contextlib.suppress(ImportError): importlib.import_module(f"agent.transports.{name}") - except ImportError: - pass diff --git a/agent/transports/anthropic.py b/agent/transports/anthropic.py index d6740ac8af..b00469d3af 100644 --- a/agent/transports/anthropic.py +++ b/agent/transports/anthropic.py @@ -1,8 +1,4 @@ -"""Anthropic Messages API transport. - -Delegates format conversion to agent/anthropic_adapter.py; owns normalization, -not client lifecycle. -""" +"""Anthropic Messages API transport: conversion via agent/anthropic_adapter.py, normalization here.""" from typing import Any, Dict, List, Optional @@ -10,20 +6,16 @@ from agent.transports.base import ProviderTransport from agent.transports.types import NormalizedResponse, ToolCall _MCP_PREFIX = "mcp__" +_THINKING_TYPES = ("thinking", "redacted_thinking") def _unprefix_oauth_tool_name(name: str) -> str: """Reverse the OAuth-wire ``mcp__`` prefix back to the registered tool name. - - Two originals map onto one wire name (``mcp__read_file`` <- ``read_file``; - ``mcp__linear_get_issue`` <- ``mcp_linear_get_issue``), so resolve by registry - lookup, never rewriting a name that already resolves natively (GH-25255). - OAuth wire aliases (e.g. chat_history_lookup -> session_search) are checked - LAST so a real tool registered under the wire name still wins. - """ + Two originals map onto one wire name (``read_file`` / ``mcp_linear_get_issue``), so + resolve by registry lookup, never rewriting a name that already resolves natively. + OAuth wire aliases are checked LAST so a real tool under the wire name still wins.""" from agent.anthropic_adapter import _OAUTH_TOOL_NAME_REVERSE_ALIASES from tools.registry import registry as _tool_registry - bare = name[len(_MCP_PREFIX):] for candidate in (name, "mcp_" + bare, bare): if _tool_registry.get_entry(candidate): @@ -31,16 +23,19 @@ def _unprefix_oauth_tool_name(name: str) -> str: return _OAUTH_TOOL_NAME_REVERSE_ALIASES.get(bare, name) +# build_kwargs params forwarded to build_anthropic_kwargs, with the defaults applied when absent. +_BUILD_KWARG_DEFAULTS = { + "max_tokens": 16384, "reasoning_config": None, "tool_choice": None, "is_oauth": False, "preserve_dots": False, + "context_length": None, "base_url": None, "fast_mode": False, "drop_context_1m_beta": False, +} + + class AnthropicTransport(ProviderTransport): """Transport for api_mode='anthropic_messages'.""" _STOP_REASON_MAP = { - "end_turn": "stop", - "tool_use": "tool_calls", - "max_tokens": "length", - "stop_sequence": "stop", - "refusal": "content_filter", - "model_context_window_exceeded": "length", + "end_turn": "stop", "tool_use": "tool_calls", "max_tokens": "length", "stop_sequence": "stop", + "refusal": "content_filter", "model_context_window_exceeded": "length", } @property @@ -50,68 +45,44 @@ class AnthropicTransport(ProviderTransport): def convert_messages(self, messages: List[Dict[str, Any]], **kwargs) -> Any: """Convert OpenAI messages to an Anthropic (system, messages) tuple; ``base_url`` affects thinking-signature handling.""" from agent.anthropic_adapter import convert_messages_to_anthropic - return convert_messages_to_anthropic(messages, base_url=kwargs.get("base_url")) def convert_tools(self, tools: List[Dict[str, Any]]) -> Any: """Convert OpenAI tool schemas to Anthropic input_schema format.""" from agent.anthropic_adapter import convert_tools_to_anthropic - return convert_tools_to_anthropic(tools) def build_kwargs( - self, - model: str, - messages: List[Dict[str, Any]], - tools: Optional[List[Dict[str, Any]]] = None, - **params, + self, model: str, messages: List[Dict[str, Any]], tools: Optional[List[Dict[str, Any]]] = None, **params, ) -> Dict[str, Any]: """Build Anthropic messages.create() kwargs (converts messages and tools internally).""" from agent.anthropic_adapter import build_anthropic_kwargs - return build_anthropic_kwargs( - model=model, - messages=messages, - tools=tools, - max_tokens=params.get("max_tokens", 16384), - reasoning_config=params.get("reasoning_config"), - tool_choice=params.get("tool_choice"), - is_oauth=params.get("is_oauth", False), - preserve_dots=params.get("preserve_dots", False), - context_length=params.get("context_length"), - base_url=params.get("base_url"), - fast_mode=params.get("fast_mode", False), - drop_context_1m_beta=params.get("drop_context_1m_beta", False), + model=model, messages=messages, tools=tools, + **{key: params.get(key, default) for key, default in _BUILD_KWARG_DEFAULTS.items()}, ) def normalize_response(self, response: Any, **kwargs) -> NormalizedResponse: """Parse content blocks (text/thinking/tool_use), map stop_reason, collect reasoning_details.""" import json from agent.anthropic_adapter import _sanitize_replay_block, _to_plain_data - strip_tool_prefix = kwargs.get("strip_tool_prefix", False) text_parts, reasoning_parts, reasoning_details, tool_calls = [], [], [], [] - # Anthropic signs each thinking block against the blocks that PRECEDE it. - # When thinking interleaves with tool_use, the parallel reasoning_details + - # tool_calls lists lose that ordering and replay -> HTTP 400 "thinking ... - # blocks cannot be modified". Keep the exact sequence for the adapter. + # Anthropic signs each thinking block against the blocks PRECEDING it; when thinking + # interleaves with tool_use the parallel lists lose that order and replay -> HTTP 400. ordered_blocks = [] - for block in response.content: block_dict = _to_plain_data(block) - clean_block = None - if isinstance(block_dict, dict): - # Sanitize at capture so output-only SDK fields never persist to - # state.db and leak back as request input on replay (HTTP 400). - clean_block = _sanitize_replay_block(block_dict) - if clean_block is not None: - ordered_blocks.append(clean_block) + # Sanitize at capture so output-only SDK fields never persist and replay (400). + clean_block = _sanitize_replay_block(block_dict) if isinstance(block_dict, dict) else None + if clean_block is not None: + ordered_blocks.append(clean_block) if block.type == "text": text_parts.append(block.text) - elif block.type in ("thinking", "redacted_thinking"): + elif block.type in _THINKING_TYPES: if block.type == "thinking": reasoning_parts.append(block.thinking) - # Prefer the sanitized block (replayed on the non-ordered path); raw only if sanitize dropped it. + # Sanitized block preferred; raw only if sanitize dropped it. if isinstance(clean_block, dict): reasoning_details.append(clean_block) elif isinstance(block_dict, dict): @@ -121,33 +92,28 @@ class AnthropicTransport(ProviderTransport): if strip_tool_prefix and name.startswith(_MCP_PREFIX): name = _unprefix_oauth_tool_name(name) tool_calls.append(ToolCall(id=block.id, name=name, arguments=json.dumps(block.input))) - provider_data = {} if reasoning_details: provider_data["reasoning_details"] = reasoning_details - # Carry the ordered channel only for the one shape the parallel lists - # reconstruct wrongly: signed thinking interleaved with tool_use. - _has_signed_thinking = any( - isinstance(b, dict) and b.get("type") in ("thinking", "redacted_thinking") and (b.get("signature") or b.get("data")) - for b in ordered_blocks + # Ordered channel only for the shape the parallel lists reconstruct wrongly. + kinds = {b.get("type") for b in ordered_blocks if isinstance(b, dict)} + signed = any( + b.get("type") in _THINKING_TYPES and (b.get("signature") or b.get("data")) + for b in ordered_blocks if isinstance(b, dict) ) - if _has_signed_thinking and any(isinstance(b, dict) and b.get("type") == "tool_use" for b in ordered_blocks): + if signed and "tool_use" in kinds: provider_data["anthropic_content_blocks"] = ordered_blocks - return NormalizedResponse( - content="\n".join(text_parts) if text_parts else None, - tool_calls=tool_calls or None, + content="\n".join(text_parts) if text_parts else None, tool_calls=tool_calls or None, finish_reason=self.map_finish_reason(response.stop_reason), - reasoning="\n\n".join(reasoning_parts) if reasoning_parts else None, - usage=None, + reasoning="\n\n".join(reasoning_parts) if reasoning_parts else None, usage=None, provider_data=provider_data or None, ) def validate_response(self, response: Any) -> bool: - """Structural check. An empty content list is legitimate for ``end_turn`` (nothing to add - after a tool turn) and ``refusal`` (Claude 4.5+ declines with empty content); treating - either as invalid would retry a completed/deterministic response forever.""" - content_blocks = getattr(response, "content", None) if response is not None else None + """Structural check; empty content is legitimate for ``end_turn``/``refusal`` (retrying + either would loop forever).""" + content_blocks = getattr(response, "content", None) if not isinstance(content_blocks, list): return False return bool(content_blocks) or getattr(response, "stop_reason", None) in {"end_turn", "refusal"} diff --git a/agent/transports/base.py b/agent/transports/base.py index aae72b5ee0..e53a7265ce 100644 --- a/agent/transports/base.py +++ b/agent/transports/base.py @@ -1,10 +1,7 @@ """Abstract base for provider transports. - -A transport owns the data path for one api_mode: - convert_messages -> convert_tools -> build_kwargs -> normalize_response -It does NOT own client construction, streaming, credential refresh, prompt -caching, interrupt handling, or retry logic — those stay on AIAgent. -""" +A transport owns one api_mode's data path (convert_messages -> convert_tools -> build_kwargs +-> normalize_response), NOT client construction, streaming, credentials, caching, interrupts +or retries — those stay on AIAgent.""" from abc import ABC, abstractmethod from typing import Any, Dict, List, Optional @@ -34,11 +31,8 @@ class ProviderTransport(ABC): @abstractmethod def build_kwargs( - self, - model: str, - messages: List[Dict[str, Any]], - tools: Optional[List[Dict[str, Any]]] = None, - **params, + self, model: str, messages: List[Dict[str, Any]], + tools: Optional[List[Dict[str, Any]]] = None, **params, ) -> Dict[str, Any]: """Primary entry point: convert messages/tools and return kwargs ready for the provider SDK.""" diff --git a/tests/hermes_cli/test_relay_shared_metrics_runtime.py b/tests/hermes_cli/test_relay_shared_metrics_runtime.py index c800862ba0..5520ec8612 100644 --- a/tests/hermes_cli/test_relay_shared_metrics_runtime.py +++ b/tests/hermes_cli/test_relay_shared_metrics_runtime.py @@ -967,7 +967,6 @@ def test_core_runtime_is_fail_open_without_a_published_binding(monkeypatch, capl tool_name="terminal", args={"command": "true"}, ) == {"command": "true"} - assert not relay_runtime.emit_mark("hermes.probe", session_id="s1") assert "Hermes Relay runtime initialization failed" in caplog.text relay_runtime._reset_for_tests()