diff --git a/agent/relay_llm.py b/agent/relay_llm.py index 6f6a3358f4..048735d159 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -56,7 +56,9 @@ class _ManagedAttempt: 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.""" + """Return the managed attempt for ``session_id`` (None: the inherited turn's), or None to run unmanaged.""" + if session_id is None: + session_id = _current_session_id() if not session_id: return None runtime, session, parent = relay_runtime.resolve_execution_context(session_id) @@ -68,12 +70,9 @@ class _ManagedAttempt: 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 + self.runtime, self.session, self.request, self.metadata = runtime, session, request, metadata self.logical = _logical_parent(runtime, session, parent, metadata) self.parent = self.logical[1] if self.logical is not None else parent - self.request = request - self.metadata = metadata self.body = _relay_request_body(request, metadata) self.relay_request = runtime.relay.LLMRequest({}, self.body) self.codec_baseline = _codec_round_trip_request_body( @@ -109,28 +108,28 @@ class _ManagedAttempt: self.raw_response["json"] = _jsonable(raw) return self.raw_response["json"] - def invoke(self, callback: Callable[..., Any], next_request: Any) -> Any: - """Provider callback handed to Relay: run ``callback`` on Relay's (possibly rewritten) request.""" + @contextlib.contextmanager + def _recording_errors(self) -> Iterator[None]: try: - raw = self.run_callback(callback, self.provider_request(next_request)) + yield except BaseException as exc: self.raw_response["error"] = exc raise + + def invoke(self, callback: Callable[..., Any], next_request: Any) -> Any: + """Provider callback handed to Relay: run ``callback`` on Relay's (possibly rewritten) request.""" + with self._recording_errors(): + raw = self.run_callback(callback, self.provider_request(next_request)) return self._record(raw) async def invoke_async(self, callback: Callable[..., Any], next_request: Any) -> Any: - try: + async def call_provider() -> Any: + with relay_runtime.managed_callback_guard(): # nested relay calls run unmanaged + return await callback(final_request) + + with self._recording_errors(): 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: @@ -177,8 +176,6 @@ def execute( ) -> Any: """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) @@ -196,8 +193,6 @@ async def execute_async( session_id: str | None = None, metadata: dict[str, Any] | None = None, defer_logical_completion: bool = False, ) -> Any: """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) @@ -214,11 +209,9 @@ execute_current_async = execute_async def _has_running_event_loop() -> bool: - try: - asyncio.get_running_loop() - except RuntimeError: - return False - return True + with contextlib.suppress(RuntimeError): + return asyncio.get_running_loop() is not None + return False def stream_current( @@ -227,15 +220,13 @@ def stream_current( 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 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 + With ``completed_response_predicate`` set, a factory that ignores ``stream=True`` and returns a + complete response is unwrapped and returned directly (pre-Relay behavior). 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() - # 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. + # Inside a managed callback (on the Relay session's loop) a nested ManagedLlmStream would be + # iterated synchronously on that loop, which asyncio forbids; the outer stream tracks this attempt. if session_id is None or _has_running_event_loop(): return stream_factory(request) managed = stream( @@ -258,7 +249,7 @@ def _aclose_on_loop(loop: asyncio.AbstractEventLoop, stream: Any) -> None: if not callable(close): return - async def close_stream() -> None: + async def close_stream() -> None: # create the coroutine on ``loop``, not the caller's thread await close() loop.run_until_complete(close_stream()) @@ -287,17 +278,12 @@ class ManagedLlmStream(Iterator[Any]): 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._on_chunk = on_chunk - self._chunk_adapter = chunk_adapter or _namespace - self._accept_chunk = accept_chunk + self._logical_model_name, self._logical_provider_name = (model_name, name) if auxiliary else (None, None) + self._on_chunk, self._chunk_adapter, self._accept_chunk = on_chunk, chunk_adapter or _namespace, accept_chunk + self._stream_factory, self._on_stream_created, self._finalizer = stream_factory, on_stream_created, finalizer + self._completed_response_predicate = completed_response_predicate self._raw_chunks: list[tuple[Any, Any]] = [] self._prefetched_chunks: list[Any] = [] - 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) @@ -355,8 +341,8 @@ class ManagedLlmStream(Iterator[Any]): 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". + # 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: @@ -410,7 +396,7 @@ class ManagedLlmStream(Iterator[Any]): def _recoverable_relay_failure(self, exc: BaseException) -> bool: """Relay post-processing failed after the provider already succeeded.""" - recoverable = (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", @@ -510,10 +496,9 @@ class ManagedLlmStream(Iterator[Any]): self._raw_stream_resource = None for resource in resources.values(): close = getattr(resource, "close", None) - if not callable(close): - continue try: - close() + if callable(close): + close() except Exception as exc: self._keep_first_close_error(exc) logger.debug("Provider stream cleanup failed", exc_info=True) @@ -524,18 +509,17 @@ class ManagedLlmStream(Iterator[Any]): self._closed = True self._prefetched_chunks.clear() try: - loop = self._loop - self._loop = None + loop, self._loop = self._loop, None if loop is None: self._close_provider_resources() - self._finish_logical(logical_outcome) - return - try: - _aclose_on_loop(loop, self._stream) - except Exception as exc: - self._keep_first_close_error(exc) + else: + try: + _aclose_on_loop(loop, self._stream) + except Exception as exc: + self._keep_first_close_error(exc) self._finish_logical(logical_outcome) - loop.close() + if loop is not None: + loop.close() finally: self._release_runtime_lease() @@ -724,9 +708,8 @@ 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/normalize unknown arguments. for key in baseline.keys() | intercepted.keys(): if key not in intercepted: final.pop(key, None) diff --git a/agent/relay_runtime.py b/agent/relay_runtime.py index add842fc8f..1bfd01a2e3 100644 --- a/agent/relay_runtime.py +++ b/agent/relay_runtime.py @@ -16,7 +16,7 @@ import tomllib import uuid from concurrent.futures import TimeoutError as FuturesTimeoutError from dataclasses import dataclass, field -from enum import Enum, auto +from enum import Enum from pathlib import Path from typing import Any, Callable @@ -113,8 +113,7 @@ def pop_relay_scope(relay: Any, handle: Any, *, output: Any = None, metadata: An 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 that scope.pop rejects. get_handle = getattr(getattr(relay, "scope", None), "get_handle", None) if callable(get_handle): with contextlib.suppress(Exception): @@ -132,14 +131,8 @@ def _same_handle(a: Any, b: Any) -> bool: return a is b or a == b or (a_uuid is not None and a_uuid == getattr(b, "uuid", None)) -class _RelayPluginConfigurationState(Enum): - """Process-wide result shared by every currently hosted profile.""" - - UNINITIALIZED = auto() - DISABLED = auto() - ACTIVE = auto() - FOREIGN = auto() - FAILED = auto() +# Process-wide plugin-configuration result shared by every currently hosted profile. +_RelayPluginConfigurationState = Enum("_RelayPluginConfigurationState", "UNINITIALIZED DISABLED ACTIVE FOREIGN FAILED") class _RelayPluginConfigurationLoadError(RuntimeError): @@ -250,16 +243,14 @@ class _ProcessRelayPluginConfiguration: 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, + "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" + "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 @@ -307,21 +298,23 @@ class _ProcessRelayPluginConfiguration: relay, activation = self._relay, self._activation if relay is None: return True - try: - _resolve_plugin_awaitable(relay.subscribers.flush_async()) - except Exception: - logger.warning("Hermes Relay plugin subscriber flush failed", exc_info=True) - return False - try: + + def close_configuration() -> Any: if activation is None: - _resolve_plugin_awaitable(relay.plugin.clear_async()) - elif callable(close := getattr(activation, "close", None)): - _resolve_plugin_awaitable(close()) - else: - 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 + return _resolve_plugin_awaitable(relay.plugin.clear_async()) + if callable(close := getattr(activation, "close", None)): + return _resolve_plugin_awaitable(close()) + raise RuntimeError("NeMo Relay dynamic plugin activation has no close method") + + for what, step in ( + ("subscriber flush", lambda: _resolve_plugin_awaitable(relay.subscribers.flush_async())), + ("configuration cleanup", close_configuration), + ): + try: + step() + except Exception: + logger.warning("Hermes Relay plugin %s failed", what, exc_info=True) + return False self._relay = self._activation = None return True @@ -337,17 +330,15 @@ class RelayRuntime: self.relay = relay or _load_nemo_relay() self.profile_key = profile_key or current_profile_key() self.runtime_id = uuid.uuid4().hex - self._sessions_lock = threading.RLock() + self._sessions_lock, self._execution_consumers_lock = threading.RLock(), threading.RLock() self._sessions: dict[str, RelaySession] = {} self._subagent_parents: dict[str, str] = {} self._subagent_parent_handles: dict[str, Any] = {} + self._execution_consumers: set[str] = set() self._closing = self._shutdown_started = False - self._shutdown_complete = threading.Event() - self._operations_idle = threading.Event() + self._shutdown_complete, self._operations_idle = threading.Event(), threading.Event() self._operations_idle.set() self._active_operations = 0 - 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 @@ -554,8 +545,8 @@ class RelayRuntime: 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 — still bounded so a wedged call cannot block exit. + # Interpreter shutdown: the executor refuses 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"{exceeded} during interpreter shutdown; abandoning the native " @@ -675,15 +666,9 @@ class RelayRuntime: return None if failure is None else f"{failure_label}: {failure}" def close_session(self, event: dict[str, Any]) -> None: - """Close one session scope and remove it from the core registry.""" - try: - self._begin_operation() - except RuntimeError: - return - try: + """Close one session scope and remove it from the core registry (no-op once shutting down).""" + with contextlib.suppress(RuntimeError), self._operation(): # _close_session itself never raises self._close_session(event) - finally: - self._end_operation() def _close_session(self, event: dict[str, Any]) -> None: """Close one session already admitted by the host lifecycle gate.""" @@ -702,8 +687,8 @@ class RelayRuntime: session, session.handle, output={}, allow_closing=True, failure_label="session scope close failed", operation_already_held=True, ) - # 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. + # No subscriber flush here: process-wide, may wait on other sessions and deadlock an asyncio + # loop; final plugin teardown flushes once. with self._sessions_lock: if self._sessions.get(session_id) is session: del self._sessions[session_id] @@ -855,16 +840,15 @@ _MANAGED_CALLBACK_DEPTH: contextvars.ContextVar[int] = contextvars.ContextVar( ) -class managed_callback_guard: +@contextlib.contextmanager +def managed_callback_guard(): """Mark the current context as inside a managed Relay callback: everything the wrapped ``invoke()`` transitively calls (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) - return self - - def __exit__(self, *exc_info: Any) -> None: - _MANAGED_CALLBACK_DEPTH.reset(self._token) + token = _MANAGED_CALLBACK_DEPTH.set(_MANAGED_CALLBACK_DEPTH.get() + 1) + try: + yield + finally: + _MANAGED_CALLBACK_DEPTH.reset(token) def _warn_on_error(what: str, callback: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: @@ -952,8 +936,7 @@ class RelaySessionCoordinator: key = (lease.profile_key, lease.session_id) with self._active_turns_lock: if self._active_turns.get(key): - # One physical scope stack per session; concurrent turns would create - # sibling scopes whose completion order is not LIFO. + # One physical scope stack per session; concurrent turns' sibling scopes would not close LIFO. turn.relay_enabled = False logger.warning( "Skipping Relay instrumentation for concurrent Hermes turn %s in session %s", @@ -964,8 +947,7 @@ class RelaySessionCoordinator: turn._active_registered = True host = lease.live_runtime() if turn.relay_enabled else None if host is not None: - # 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. + # Rotation happens HERE: no live turn scope on the stack, so the session scope can close/reopen 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, @@ -1002,8 +984,8 @@ class RelaySessionCoordinator: with contextlib.suppress(Exception), lease.session.lock: # accounting never blocks lease.session.segment_turns += 1 # max_turns rotation trigger try: - # Delegated agents own one turn: close their conversation while the - # active-turn guard is held so a parent timeout fallback cannot race it. + # 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): _warn_on_error( "child conversation finalization", lease.host.unregister_subagent, @@ -1100,8 +1082,7 @@ class RelaySessionCoordinator: logical_calls.pop() continue with turn.logical_llm_lock: - # Stack-owned scopes: if the newest handle cannot close even after orphan - # drain, older ones cannot either — retain the unclosed prefix. + # Stack-owned: if the newest handle cannot close even after drain, older ones cannot either. 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) @@ -1162,11 +1143,9 @@ def active_turn(session_id: str | None = None) -> RelayTurnContext | None: 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: - # Nested managed execution is impossible (see _MANAGED_CALLBACK_DEPTH); the - # outer scope still records the tool-level event. - return None, None, None - if not relay_instrumentation_enabled(): + # Nested managed execution is impossible (see _MANAGED_CALLBACK_DEPTH); the outer scope + # still records the tool-level event. + if _MANAGED_CALLBACK_DEPTH.get() > 0 or 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 @@ -1228,11 +1207,9 @@ def _configured_plugin_inputs(relay: Any) -> tuple[dict[str, Any], list[Any]] | if not configured: 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 " - "exporter configuration to a Relay plugins.toml file.", - RELAY_PLUGINS_CONFIG_ENV, - ", ".join(legacy_vars), + "Legacy NeMo Relay exporter variables are set but no %s was provided. %s no longer activate " + "Relay exporters; migrate the exporter configuration to a Relay plugins.toml file.", + RELAY_PLUGINS_CONFIG_ENV, ", ".join(legacy_vars), ) return None config_path = Path(configured).expanduser() @@ -1240,17 +1217,12 @@ def _configured_plugin_inputs(relay: Any) -> tuple[dict[str, Any], list[Any]] | 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" - ) - dynamic_plugins: list[Any] = [] - if "plugins" in config: - dynamic_plugins = relay.plugin.load_dynamic_plugin_activation_specs(config_path) + raise ValueError("Hermes [[dynamic_plugins]] records are unsupported; use Relay [[plugins.dynamic]] records") + dynamic_plugins = relay.plugin.load_dynamic_plugin_activation_specs(config_path) if "plugins" in config else [] 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 " - f"{config_path}; continuing without Relay plugins" + f"Hermes Relay plugin configuration could not be loaded from {config_path}; continuing without Relay plugins" ) from exc diff --git a/agent/relay_tools.py b/agent/relay_tools.py index 42f4e45bc7..54ccb74bd4 100644 --- a/agent/relay_tools.py +++ b/agent/relay_tools.py @@ -41,8 +41,7 @@ def execute( except BaseException as exc: callback_error = exc raise - raw_result["value"] = result - raw_result["json"] = _jsonable(result) + raw_result.update(value=result, json=_jsonable(result)) return raw_result["json"] try: diff --git a/agent/transports/anthropic.py b/agent/transports/anthropic.py index b00469d3af..c80ad07099 100644 --- a/agent/transports/anthropic.py +++ b/agent/transports/anthropic.py @@ -82,26 +82,18 @@ class AnthropicTransport(ProviderTransport): elif block.type in _THINKING_TYPES: if block.type == "thinking": reasoning_parts.append(block.thinking) - # Sanitized block preferred; raw only if sanitize dropped it. - if isinstance(clean_block, dict): - reasoning_details.append(clean_block) - elif isinstance(block_dict, dict): - reasoning_details.append(block_dict) + detail = clean_block if clean_block is not None else block_dict # raw only if sanitize dropped it + if isinstance(detail, dict): + reasoning_details.append(detail) elif block.type == "tool_use": name = block.name 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 + provider_data = {"reasoning_details": reasoning_details} if reasoning_details else {} # 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 signed and "tool_use" in kinds: + signed = any(b.get("type") in _THINKING_TYPES and (b.get("signature") or b.get("data")) for b in ordered_blocks) + if signed and any(b.get("type") == "tool_use" for b in ordered_blocks): 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, @@ -114,9 +106,9 @@ class AnthropicTransport(ProviderTransport): """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"} + return isinstance(content_blocks, list) and ( + bool(content_blocks) or getattr(response, "stop_reason", None) in {"end_turn", "refusal"} + ) def extract_cache_stats(self, response: Any) -> Optional[Dict[str, int]]: """Anthropic cache_read / cache_creation token counts."""