diff --git a/agent/relay_llm.py b/agent/relay_llm.py index ea969da43f..2d8c37a18f 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -434,6 +434,7 @@ class ManagedLlmStream(Iterator[Any]): 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: tuple[relay_runtime.RelayTurnContext, Any, str] | None = None @@ -573,7 +574,12 @@ class ManagedLlmStream(Iterator[Any]): self._callback_error = exc raise - loop = asyncio.new_event_loop() + self._runtime_lease = 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: @@ -613,10 +619,14 @@ class ManagedLlmStream(Iterator[Any]): 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 - loop.close() - self._loop = None + try: + loop.close() + finally: + self._loop = None + self._release_runtime_lease() raise def __iter__(self) -> "ManagedLlmStream": @@ -662,6 +672,7 @@ class ManagedLlmStream(Iterator[Any]): 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 self._close(logical_outcome="cancelled") @@ -719,8 +730,75 @@ class ManagedLlmStream(Iterator[Any]): self._stream = iter(pending) self._raw_stream_resource = None self._accept_chunk = None - if loop is not None: - close = getattr(relay_stream, "aclose", None) + try: + if loop is not None: + close = getattr(relay_stream, "aclose", None) + if callable(close): + + async def close_stream() -> None: + await close() + + try: + loop.run_until_complete(close_stream()) + except Exception: + logger.debug( + "Relay stream cleanup failed during provider fallback", + exc_info=True, + ) + loop.close() + if not self._defer_logical_completion: + _complete_logical( + self._logical, + outcome="success", + 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 + finally: + self._release_runtime_lease() + + def _close(self, *, logical_outcome: str) -> None: + if self._closed: + return + self._closed = True + self._prefetched_chunks.clear() + try: + loop = self._loop + self._loop = None + if loop is None: + resources = (self._stream, self._raw_stream_resource) + 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)) + close = getattr(resource, "close", None) + if callable(close): + try: + close() + except Exception as exc: + if self._close_error is None: + self._close_error = exc + logger.debug( + "Provider stream cleanup failed", + exc_info=True, + ) + if not self._defer_logical_completion: + _complete_logical( + self._logical, + outcome=logical_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 + return + close = getattr(self._stream, "aclose", None) if callable(close): async def close_stream() -> None: @@ -728,49 +806,9 @@ class ManagedLlmStream(Iterator[Any]): try: loop.run_until_complete(close_stream()) - except Exception: - logger.debug( - "Relay stream cleanup failed during provider fallback", - exc_info=True, - ) - loop.close() - if not self._defer_logical_completion: - _complete_logical( - self._logical, - outcome="success", - model_name=self._logical_model_name, - provider_name=self._logical_provider_name, - response_model_name=self._logical_response_model_name, - ) - self._logical = None - - def _close(self, *, logical_outcome: str) -> None: - if self._closed: - return - self._closed = True - self._prefetched_chunks.clear() - loop = self._loop - self._loop = None - if loop is None: - resources = (self._stream, self._raw_stream_resource) - 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)) - close = getattr(resource, "close", None) - if callable(close): - try: - close() - except Exception as exc: - if self._close_error is None: - self._close_error = exc - logger.debug( - "Provider stream cleanup failed", - exc_info=True, - ) + except Exception as exc: + if self._close_error is None: + self._close_error = exc if not self._defer_logical_completion: _complete_logical( self._logical, @@ -778,30 +816,18 @@ class ManagedLlmStream(Iterator[Any]): 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 - return - close = getattr(self._stream, "aclose", None) - if callable(close): + loop.close() + finally: + self._release_runtime_lease() - async def close_stream() -> None: - await close() - - try: - loop.run_until_complete(close_stream()) - except Exception as exc: - if self._close_error is None: - self._close_error = exc - if not self._defer_logical_completion: - _complete_logical( - self._logical, - outcome=logical_outcome, - model_name=self._logical_model_name, - provider_name=self._logical_provider_name, - response_model_name=self._logical_response_model_name, - ) - self._logical = None - loop.close() + def _release_runtime_lease(self) -> None: + lease = self._runtime_lease + self._runtime_lease = None + if lease is not None: + lease.release() def __del__(self) -> None: self._close(logical_outcome="cancelled") @@ -941,6 +967,7 @@ def _complete_logical( 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 @@ -960,7 +987,10 @@ def _complete_logical( output.update({"model": model_name, "provider": provider_name}) if response_model_name is not None: output["response_model"] = response_model_name - lease.host.run_in_session( + 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, diff --git a/agent/relay_runtime.py b/agent/relay_runtime.py index 1fecec69be..4f3f0ba5a5 100644 --- a/agent/relay_runtime.py +++ b/agent/relay_runtime.py @@ -8,10 +8,13 @@ import contextvars import importlib import inspect import logging +import os import threading +import tomllib import uuid from concurrent.futures import TimeoutError as FuturesTimeoutError from dataclasses import dataclass, field +from pathlib import Path from typing import Any, Callable from hermes_constants import get_hermes_home @@ -24,6 +27,7 @@ LOGICAL_LLM_SCOPE = "hermes.logical_llm_call" RUNTIME_SCHEMA_KEY = "hermes.relay.schema_version" RUNTIME_SCHEMA_VERSION = "hermes.relay.runtime.v1" RUNTIME_INSTANCE_KEY = "hermes.relay.runtime_instance" +RELAY_PLUGINS_CONFIG_ENV = "HERMES_NEMO_RELAY_PLUGINS_TOML" RELAY_PLUGINS_EXECUTION_CONSUMER = "hermes.nemo_relay.plugins" _PROFILE_KEY_CACHE: dict[str, str] = {} @@ -200,6 +204,41 @@ def _reset_segments_config_for_tests() -> None: _SEGMENTS_CONFIG = None +class RelayOperationLease: + """Keep process-wide Relay plugins alive across a deferred operation.""" + + def __init__(self, runtime: "RelayRuntime") -> None: + self._lock = threading.Lock() + self._runtime: RelayRuntime | None = runtime + + 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: + raise RuntimeError("Hermes Relay operation lease is released") + return 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 + if runtime is not None: + runtime._end_operation() + + class _ProcessRelayPluginConfiguration: """Own one Relay plugin configuration across profile-scoped hosts.""" @@ -208,6 +247,7 @@ class _ProcessRelayPluginConfiguration: self._owners: set[int] = set() self._active = False self._relay: Any = None + self._activation: Any = None def acquire(self, owner: Any, relay: Any) -> bool: """Join the process configuration, initializing it for the first host.""" @@ -218,21 +258,60 @@ class _ProcessRelayPluginConfiguration: if self._owners: self._owners.add(owner_id) return self._active + 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 False - self._owners.add(owner_id) try: plugin_mod = getattr(relay, "plugin", None) - initialize = getattr(plugin_mod, "initialize", None) - if not callable(initialize): - raise RuntimeError( - "installed NeMo Relay binding does not expose " - "plugin.initialize" + plugin_config, dynamic_plugins = _configured_plugin_inputs() + if dynamic_plugins: + initialize_dynamic = getattr( + plugin_mod, + "initialize_with_dynamic_plugins", + None, ) - # An empty override delegates file discovery and layering to - # Relay: /etc, the nearest .nemo-relay/plugins.toml, then the - # user configuration directory. - _resolve_plugin_awaitable(initialize({})) + if callable(initialize_dynamic): + try: + activation = _resolve_plugin_awaitable( + initialize_dynamic(plugin_config, dynamic_plugins) + ) + if activation is None: + raise RuntimeError( + "NeMo Relay dynamic plugin initialization " + "returned no activation handle" + ) + self._activation = activation + except Exception: + logger.warning( + "Hermes Relay dynamic plugin activation failed; " + "continuing with configured and discovered " + "static plugins", + exc_info=True, + ) + else: + logger.warning( + "Hermes Relay dynamic plugins require a binding that " + "exposes plugin.initialize_with_dynamic_plugins; " + "continuing with configured and discovered static " + "plugins" + ) + + if self._activation is None: + initialize = getattr(plugin_mod, "initialize", None) + if not callable(initialize): + raise RuntimeError( + "installed NeMo Relay binding does not expose " + "plugin.initialize" + ) + # Relay owns ambient file discovery and precedence. An explicit + # Hermes file, when configured, is supplied as the final overlay. + _resolve_plugin_awaitable(initialize(plugin_config)) except Exception as exc: + self._activation = None logger.warning( "Hermes Relay plugin initialization failed: %s", exc, @@ -240,6 +319,7 @@ class _ProcessRelayPluginConfiguration: ) return False + self._owners.add(owner_id) self._active = True self._relay = relay return True @@ -261,34 +341,50 @@ class _ProcessRelayPluginConfiguration: self._owners.clear() self._clear_active() - def _clear_active(self) -> None: + def retry_pending_cleanup(self) -> None: + """Retry a failed final cleanup without disrupting live owners.""" + with self._lock: + if not self._owners: + self._clear_active() + + def _clear_active(self) -> bool: relay = self._relay + activation = self._activation active = self._active - self._active = False - self._relay = None if not active or relay is None: - return + return True try: - flush = getattr(getattr(relay, "subscribers", None), "flush", None) - if callable(flush): - flush() + _flush_relay_subscribers(relay) except Exception: logger.warning( "Hermes Relay plugin subscriber flush failed", exc_info=True, ) + return False try: - clear = getattr(getattr(relay, "plugin", None), "clear", None) - if callable(clear): - _resolve_plugin_awaitable(clear()) + 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" + ) + _resolve_plugin_awaitable(close()) + else: + _clear_relay_plugins(relay) except Exception: logger.warning( "Hermes Relay plugin configuration cleanup failed", exc_info=True, ) + return False + self._active = False + self._relay = None + self._activation = None + return True _PLUGIN_CONFIGURATION = _ProcessRelayPluginConfiguration() +atexit.register(_PLUGIN_CONFIGURATION.retry_pending_cleanup) class RelayRuntime: @@ -302,10 +398,19 @@ 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._shutdown_complete = threading.Event() + self._operations_idle = 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_registered = True - if _PLUGIN_CONFIGURATION.acquire(self, self.relay): + self._plugin_configuration_registered = _PLUGIN_CONFIGURATION.acquire( + self, + self.relay, + ) + if self._plugin_configuration_registered: self.retain_managed_execution(RELAY_PLUGINS_EXECUTION_CONSUMER) self._shutdown_registered = True atexit.register(self.shutdown) @@ -339,6 +444,8 @@ class RelayRuntime: if not session_id: return None with self._sessions_lock: + if self._closing: + return None session = self._sessions.get(session_id) if session is None: parent_session_id = self._subagent_parents.get(session_id, "") @@ -499,6 +606,8 @@ class RelayRuntime: ): parent_handle = turn.handle with self._sessions_lock: + if self._closing: + return None self._subagent_parents[child_session_id] = parent_session_id if parent_handle is not None: self._subagent_parent_handles[child_session_id] = parent_handle @@ -520,6 +629,8 @@ class RelayRuntime: def get_session(self, session_id: str) -> RelaySession | None: """Return an active Hermes Relay session without creating one.""" with self._sessions_lock: + if self._closing: + return None session = self._sessions.get(str(session_id or "")) if session is None: return None @@ -553,6 +664,29 @@ class RelayRuntime: span, never the agent. The abandoned daemon worker cannot block process exit (tools.daemon_pool contract). """ + self._begin_operation() + try: + 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, + ) -> Any: + """Run inside a session whose host-level lifetime is already held.""" with session.lock: if session.closing and not allow_closing: raise RuntimeError("Hermes Relay session is closing") @@ -601,26 +735,49 @@ class RelayRuntime: **kwargs: Any, ) -> Any: """Create and await an operation inside the session's saved context.""" - with session.lock: - if session.closing and not allow_closing: - raise RuntimeError("Hermes Relay session is closing") - if session.context is None or session.handle is None: - raise RuntimeError("Hermes Relay session context is unavailable") - relay_context = session.context.copy() + self._begin_operation() + try: + with session.lock: + if session.closing and not allow_closing: + raise RuntimeError("Hermes Relay session is closing") + if session.context is None or session.handle is None: + raise RuntimeError("Hermes Relay session context is unavailable") + relay_context = session.context.copy() - context = contextvars.copy_context() - for variable, value in relay_context.items(): - context.run(variable.set, value) + context = contextvars.copy_context() + for variable, value in relay_context.items(): + context.run(variable.set, value) - async def invoke() -> Any: - self.relay.get_scope_stack() - result = callback(*args, **kwargs) - if inspect.isawaitable(result): - return await result - return result + async def invoke() -> Any: + self.relay.get_scope_stack() + result = callback(*args, **kwargs) + if inspect.isawaitable(result): + return await result + return result - task = context.run(asyncio.create_task, invoke()) - return await task + task = context.run(asyncio.create_task, invoke()) + return await task + finally: + self._end_operation() + + def _begin_operation(self) -> None: + """Admit one Relay call while keeping process plugins alive.""" + with self._sessions_lock: + if self._closing: + raise RuntimeError("Hermes Relay runtime is shutting down") + self._active_operations += 1 + self._operations_idle.clear() + + def _end_operation(self) -> None: + with self._sessions_lock: + self._active_operations -= 1 + if self._active_operations == 0: + self._operations_idle.set() + + 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, @@ -816,6 +973,17 @@ class RelayRuntime: 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: + 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.""" session_id = _session_id(event) with self._sessions_lock: session = self._sessions.get(session_id) @@ -840,17 +1008,7 @@ class RelayRuntime: if failure: failures.append(failure) try: - try: - _scope_op_executor().submit( - self.relay.subscribers.flush - ).result(timeout=_SCOPE_OP_TIMEOUT) - except RuntimeError: - # Interpreter shutdown: executor refuses new futures; flush - # on a bounded exit thread so a wedged pipeline cannot - # block process exit. - _run_bounded_on_exit_thread( - self.relay.subscribers.flush, _SCOPE_OP_TIMEOUT - ) + _flush_relay_subscribers(self.relay) except Exception as exc: failures.append(f"subscriber flush failed: {exc}") with self._sessions_lock: @@ -868,19 +1026,56 @@ class RelayRuntime: def shutdown(self) -> None: """Close core scopes and release process plugin configuration.""" with self._sessions_lock: - session_ids = list(self._sessions) - for session_id in session_ids: - self._safe(self.close_session, {"session_id": session_id}) - if self._plugin_configuration_registered: - self.release_managed_execution(RELAY_PLUGINS_EXECUTION_CONSUMER) - _PLUGIN_CONFIGURATION.release(self) - self._plugin_configuration_registered = False - if self._shutdown_registered: + if self._shutdown_started: + return + self._shutdown_started = True + 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: - atexit.unregister(self.shutdown) + thread.start() except Exception: - pass - self._shutdown_registered = False + with self._sessions_lock: + self._shutdown_started = False + logger.warning( + "Hermes Relay deferred shutdown could not start", + exc_info=True, + ) + return + self._finish_shutdown() + + def _finish_shutdown_after_operations(self) -> None: + self._operations_idle.wait() + self._finish_shutdown() + + 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}) + if self._plugin_configuration_registered: + self.release_managed_execution(RELAY_PLUGINS_EXECUTION_CONSUMER) + _PLUGIN_CONFIGURATION.release(self) + self._plugin_configuration_registered = False + if self._shutdown_registered: + try: + atexit.unregister(self.shutdown) + except Exception: + pass + self._shutdown_registered = False + except Exception: + with self._sessions_lock: + self._shutdown_started = False + logger.warning("Hermes Relay shutdown failed", exc_info=True) + return + with self._sessions_lock: + self._shutdown_complete.set() @staticmethod def _safe(callback: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: @@ -1710,6 +1905,128 @@ def _load_nemo_relay() -> Any: return importlib.import_module("nemo_relay") +def _configured_plugin_inputs() -> tuple[dict[str, Any], list[dict[str, Any]]]: + """Load one explicit config overlay and its Hermes dynamic host specs.""" + configured = os.environ.get(RELAY_PLUGINS_CONFIG_ENV, "").strip() + if not configured: + return {}, [] + + config_path = Path(configured).expanduser() + try: + with config_path.open("rb") as config_file: + config = tomllib.load(config_file) + dynamic_plugins = _dynamic_plugin_specs(config, config_path) + plugin_config = dict(config) + plugin_config.pop("dynamic_plugins", None) + plugin_config.pop("plugins", None) + return plugin_config, dynamic_plugins + except Exception: + logger.warning( + "Hermes Relay plugin configuration could not be loaded from %s; " + "continuing with discovered static plugins", + config_path, + exc_info=True, + ) + return {}, [] + + +def _dynamic_plugin_specs( + config: dict[str, Any], + config_path: Path, +) -> list[dict[str, Any]]: + """Validate Hermes-owned specs for Relay's public dynamic host API.""" + plugins_section = config.get("plugins") + if plugins_section is not None: + if not isinstance(plugins_section, dict): + raise ValueError("[plugins] must be a table") + if plugins_section: + raise ValueError( + "Relay CLI [[plugins.dynamic]] records require lifecycle state; " + "use Hermes [[dynamic_plugins]] activation specs" + ) + + raw_specs = config.get("dynamic_plugins") + if raw_specs is None: + return [] + if not isinstance(raw_specs, list): + raise ValueError("dynamic_plugins must be an array of tables") + + specs: list[dict[str, Any]] = [] + for index, raw_spec in enumerate(raw_specs): + if not isinstance(raw_spec, dict): + raise ValueError(f"dynamic_plugins[{index}] must be a table") + plugin_id = raw_spec.get("plugin_id") + kind = raw_spec.get("kind") + manifest_ref = raw_spec.get("manifest_ref") + environment_ref = raw_spec.get("environment_ref") + plugin_config = raw_spec.get("config", {}) + if not isinstance(plugin_id, str) or not plugin_id.strip(): + raise ValueError(f"dynamic_plugins[{index}].plugin_id is required") + if kind not in {"rust_dynamic", "worker"}: + raise ValueError( + f"dynamic_plugins[{index}].kind must be rust_dynamic or worker" + ) + if not isinstance(manifest_ref, str) or not manifest_ref.strip(): + raise ValueError(f"dynamic_plugins[{index}].manifest_ref is required") + if not isinstance(plugin_config, dict): + raise ValueError(f"dynamic_plugins[{index}].config must be a table") + if environment_ref is not None and ( + not isinstance(environment_ref, str) or not environment_ref.strip() + ): + raise ValueError( + f"dynamic_plugins[{index}].environment_ref must be a non-empty string" + ) + + spec: dict[str, Any] = { + "plugin_id": plugin_id.strip(), + "kind": kind, + "manifest_ref": _config_relative_path( + manifest_ref.strip(), + config_path, + ), + "config": plugin_config, + } + if environment_ref is not None: + spec["environment_ref"] = _config_relative_path( + environment_ref.strip(), + config_path, + ) + specs.append(spec) + return specs + + +def _config_relative_path(value: str, config_path: Path) -> str: + """Resolve one activation path relative to its physical TOML file.""" + path = Path(value).expanduser() + if path.is_absolute(): + return str(path) + return str((config_path.resolve().parent / path).resolve()) + + +def _flush_relay_subscribers(relay: Any) -> None: + """Flush Relay without blocking a 0.7 asyncio event-loop thread.""" + subscribers = getattr(relay, "subscribers", None) + flush_async = getattr(subscribers, "flush_async", None) + if callable(flush_async): + _resolve_plugin_awaitable(flush_async()) + return + flush = getattr(subscribers, "flush", None) + if callable(flush): + _resolve_plugin_awaitable(flush()) + + +def _clear_relay_plugins(relay: Any) -> None: + """Clear Relay plugins through the newest available binding API.""" + plugin_mod = getattr(relay, "plugin", None) + clear_async = getattr(plugin_mod, "clear_async", None) + if callable(clear_async): + _resolve_plugin_awaitable(clear_async()) + return + clear = getattr(plugin_mod, "clear", None) + if callable(clear): + _resolve_plugin_awaitable(clear()) + + def _resolve_plugin_awaitable(value: Any) -> Any: """Resolve Relay's async plugin API from synchronous host construction.""" if not inspect.isawaitable(value): @@ -1730,7 +2047,7 @@ def _resolve_plugin_awaitable(value: Any) -> Any: thread = threading.Thread( target=_runner, - name="hermes-nemo-relay-plugin-init", + name="hermes-nemo-relay-plugin-lifecycle", daemon=True, ) thread.start() diff --git a/tests/agent/test_relay_llm.py b/tests/agent/test_relay_llm.py index 433ac634c6..b1abf03d10 100644 --- a/tests/agent/test_relay_llm.py +++ b/tests/agent/test_relay_llm.py @@ -311,6 +311,39 @@ def test_stream_uses_rewritten_request_and_post_intercept_chunks(relay_turn): assert turn.logical_llm_calls == {} +def test_live_stream_defers_runtime_shutdown_until_exhaustion( + tmp_path, + monkeypatch, +): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "stream-shutdown-profile")) + relay_runtime._reset_for_tests() + host = relay_runtime.get_runtime() + assert host is not None + host.retain_managed_execution("test.live-stream") + assert host.ensure_session({"session_id": "stream-shutdown"}) is not None + chunks = [{"delta": "first"}, {"delta": "second"}] + stream = relay_llm.stream( + {"model": "test-model", "messages": []}, + lambda _request: iter(chunks), + session_id="stream-shutdown", + name="test-provider", + model_name="test-model", + finalizer=lambda: {"content": "complete"}, + metadata={"api_mode": "custom"}, + ) + + try: + host.shutdown() + assert not host._shutdown_complete.is_set() + + assert list(stream) == chunks + assert host._shutdown_complete.wait(5) + finally: + stream.close() + host.release_managed_execution("test.live-stream") + relay_runtime._reset_for_tests() + + diff --git a/tests/agent/test_relay_runtime_plugins.py b/tests/agent/test_relay_runtime_plugins.py index a4ed126082..b486c8efe6 100644 --- a/tests/agent/test_relay_runtime_plugins.py +++ b/tests/agent/test_relay_runtime_plugins.py @@ -4,6 +4,7 @@ from __future__ import annotations import asyncio import json +import threading from types import SimpleNamespace from typing import Any @@ -13,12 +14,21 @@ from agent import relay_runtime class _FakeRelay: - def __init__(self, *, initialize_error: Exception | None = None) -> None: + def __init__( + self, + *, + initialize_error: Exception | None = None, + dynamic_initialize_error: Exception | None = None, + activation_close_error: Exception | None = None, + ) -> None: self.events: list[tuple[Any, ...]] = [] self.initialize_error = initialize_error + self.dynamic_initialize_error = dynamic_initialize_error + self.activation_close_error = activation_close_error self.ScopeType = SimpleNamespace(Agent="agent") self.plugin = SimpleNamespace( initialize=self._initialize_plugins, + initialize_with_dynamic_plugins=self._initialize_dynamic_plugins, clear=self._clear_plugins, ) self.scope = SimpleNamespace( @@ -36,6 +46,25 @@ class _FakeRelay: raise self.initialize_error return {"diagnostics": []} + async def _initialize_dynamic_plugins( + self, + config: dict[str, Any], + dynamic_plugins: list[dict[str, Any]], + ) -> Any: + self.events.append(("plugin.initialize_dynamic", config, dynamic_plugins)) + if self.dynamic_initialize_error is not None: + raise self.dynamic_initialize_error + + relay = self + + class _Activation: + async def close(self) -> None: + relay.events.append(("plugin.activation.close",)) + if relay.activation_close_error is not None: + raise relay.activation_close_error + + return _Activation() + def _clear_plugins(self) -> None: self.events.append(("plugin.clear",)) @@ -51,6 +80,39 @@ class _FakeRelay: self.events.append(("subscribers.flush",)) +class _AsyncCleanupRelay(_FakeRelay): + """Relay 0.7-shaped fake that rejects synchronous loop cleanup.""" + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + self.plugin.clear_async = self._clear_plugins_async + self.subscribers.flush_async = self._flush_async + + def _clear_plugins(self) -> None: + raise AssertionError("synchronous plugin.clear must not be called") + + def _flush(self) -> None: + raise AssertionError("synchronous subscribers.flush must not be called") + + async def _clear_plugins_async(self) -> None: + self.events.append(("plugin.clear_async",)) + + async def _flush_async(self) -> None: + self.events.append(("subscribers.flush_async",)) + + +class _BlockingFlushRelay(_FakeRelay): + def __init__(self) -> None: + super().__init__() + self.flush_started = threading.Event() + self.finish_flush = threading.Event() + + def _flush(self) -> None: + self.events.append(("subscribers.flush",)) + self.flush_started.set() + assert self.finish_flush.wait(5) + + @pytest.fixture(autouse=True) def _reset_runtime(): relay_runtime._reset_for_tests() @@ -84,6 +146,24 @@ def test_initialization_failure_is_fail_open(caplog): host.shutdown() +def test_later_host_retries_after_initialization_failure(): + relay = _FakeRelay(initialize_error=RuntimeError("transient failure")) + failed_host = relay_runtime.RelayRuntime(relay=relay, profile_key="failed") + assert not failed_host.managed_execution_enabled() + + relay.initialize_error = None + recovered_host = relay_runtime.RelayRuntime(relay=relay, profile_key="recovered") + try: + assert recovered_host.managed_execution_enabled() + assert relay.events == [ + ("plugin.initialize", {}), + ("plugin.initialize", {}), + ] + finally: + failed_host.shutdown() + recovered_host.shutdown() + + def test_missing_initialize_api_is_fail_open(caplog): relay = _FakeRelay() del relay.plugin.initialize @@ -138,6 +218,421 @@ def test_plugin_initialization_inside_running_event_loop(): host.shutdown() +def test_static_plugin_cleanup_uses_async_apis_inside_running_event_loop(): + relay = _AsyncCleanupRelay() + + async def run_lifecycle() -> None: + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + host.shutdown() + + asyncio.run(run_lifecycle()) + + assert relay.events == [ + ("plugin.initialize", {}), + ("subscribers.flush_async",), + ("plugin.clear_async",), + ] + + +def test_dynamic_plugins_share_owned_activation_until_final_host_shutdown( + tmp_path, + monkeypatch, +): + config = tmp_path / ".nemo-relay" / "plugins.toml" + config.parent.mkdir() + config.write_text( + """ +version = 1 + +[[components]] +kind = "observability" +enabled = true + +[components.config] +version = 1 + +[[dynamic_plugins]] +plugin_id = "native.policy" +kind = "rust_dynamic" +manifest_ref = "plugins/native/relay-plugin.toml" + +[dynamic_plugins.config] +mode = "strict" + +[[dynamic_plugins]] +plugin_id = "worker.policy" +kind = "worker" +manifest_ref = "plugins/worker/relay-plugin.toml" +environment_ref = "environments/worker" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _FakeRelay() + + host_a = relay_runtime.RelayRuntime(relay=relay, profile_key="profile-a") + host_b = relay_runtime.RelayRuntime(relay=relay, profile_key="profile-b") + + assert host_a.managed_execution_enabled() + assert host_b.managed_execution_enabled() + assert relay.events == [ + ( + "plugin.initialize_dynamic", + { + "version": 1, + "components": [ + { + "kind": "observability", + "enabled": True, + "config": {"version": 1}, + } + ], + }, + [ + { + "plugin_id": "native.policy", + "kind": "rust_dynamic", + "manifest_ref": str( + config.parent / "plugins/native/relay-plugin.toml" + ), + "config": {"mode": "strict"}, + }, + { + "plugin_id": "worker.policy", + "kind": "worker", + "manifest_ref": str( + config.parent / "plugins/worker/relay-plugin.toml" + ), + "environment_ref": str( + config.parent / "environments/worker" + ), + "config": {}, + }, + ], + ) + ] + + host_b.ensure_session({"session_id": "profile-b-session"}) + host_a.shutdown() + assert ("plugin.activation.close",) not in relay.events + + host_b.shutdown() + assert relay.events[-2:] == [ + ("subscribers.flush",), + ("plugin.activation.close",), + ] + assert ("plugin.clear",) not in relay.events + assert relay.events.count(("plugin.activation.close",)) == 1 + pop_index = next( + index for index, event in enumerate(relay.events) if event[0] == "scope.pop" + ) + assert pop_index < relay.events.index(("plugin.activation.close",)) + + +def test_dynamic_activation_failure_falls_back_to_discovered_static_plugins( + tmp_path, + monkeypatch, + caplog, +): + config = tmp_path / "plugins.toml" + config.write_text( + """ +[[dynamic_plugins]] +plugin_id = "worker.policy" +kind = "worker" +manifest_ref = "relay-plugin.toml" +environment_ref = "environment" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _FakeRelay( + dynamic_initialize_error=RuntimeError("worker rejected config") + ) + + with caplog.at_level("WARNING"): + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + try: + assert host.managed_execution_enabled() + assert [event[0] for event in relay.events] == [ + "plugin.initialize_dynamic", + "plugin.initialize", + ] + assert "configured and discovered static plugins" in caplog.text + finally: + host.shutdown() + + assert relay.events[-2:] == [ + ("subscribers.flush",), + ("plugin.clear",), + ] + + +def test_dynamic_activation_lifecycle_inside_running_event_loop( + tmp_path, + monkeypatch, +): + config = tmp_path / "plugins.toml" + config.write_text( + """ +[[dynamic_plugins]] +plugin_id = "native.policy" +kind = "rust_dynamic" +manifest_ref = "relay-plugin.toml" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _AsyncCleanupRelay() + + async def run_lifecycle() -> None: + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + assert host.managed_execution_enabled() + host.shutdown() + + asyncio.run(run_lifecycle()) + + assert [event[0] for event in relay.events] == [ + "plugin.initialize_dynamic", + "subscribers.flush_async", + "plugin.activation.close", + ] + + +def test_shutdown_defers_dynamic_unload_until_async_operation_finishes( + tmp_path, + monkeypatch, +): + config = tmp_path / "plugins.toml" + config.write_text( + """ +[[dynamic_plugins]] +plugin_id = "worker.policy" +kind = "worker" +manifest_ref = "relay-plugin.toml" +environment_ref = "environment" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _AsyncCleanupRelay() + + async def run_lifecycle() -> None: + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + session = host.ensure_session({"session_id": "session"}) + assert session is not None + started = asyncio.Event() + finish = asyncio.Event() + + async def in_flight_call() -> None: + relay.events.append(("operation.start",)) + started.set() + await finish.wait() + relay.events.append(("operation.end",)) + + operation = asyncio.create_task( + host.run_in_session_async(session, in_flight_call) + ) + await started.wait() + host.shutdown() + assert host.ensure_session({"session_id": "late-session"}) is None + assert ("plugin.activation.close",) not in relay.events + + finish.set() + await operation + assert await asyncio.to_thread(host._shutdown_complete.wait, 5) + + asyncio.run(run_lifecycle()) + + assert relay.events.index(("operation.end",)) < relay.events.index( + ("plugin.activation.close",) + ) + + +def test_shutdown_waits_for_concurrent_session_close_before_dynamic_unload( + tmp_path, + monkeypatch, +): + config = tmp_path / "plugins.toml" + config.write_text( + """ +[[dynamic_plugins]] +plugin_id = "native.policy" +kind = "rust_dynamic" +manifest_ref = "relay-plugin.toml" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _BlockingFlushRelay() + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + assert host.ensure_session({"session_id": "session"}) is not None + + close_thread = threading.Thread( + target=host.close_session, + args=({"session_id": "session"},), + ) + close_thread.start() + assert relay.flush_started.wait(5) + + host.shutdown() + assert ("plugin.activation.close",) not in relay.events + + relay.finish_flush.set() + close_thread.join(5) + assert not close_thread.is_alive() + assert host._shutdown_complete.wait(5) + assert relay.events.index(("subscribers.flush",)) < relay.events.index( + ("plugin.activation.close",) + ) + + +def test_failed_dynamic_teardown_retains_activation_and_blocks_replacement( + tmp_path, + monkeypatch, + caplog, +): + config = tmp_path / "plugins.toml" + config.write_text( + """ +[[dynamic_plugins]] +plugin_id = "worker.policy" +kind = "worker" +manifest_ref = "relay-plugin.toml" +environment_ref = "environment" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _FakeRelay(activation_close_error=RuntimeError("worker still busy")) + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + + with caplog.at_level("WARNING"): + host.shutdown() + + assert "plugin configuration cleanup failed" in caplog.text + activation = relay_runtime._PLUGIN_CONFIGURATION._activation + assert activation is not None + + with caplog.at_level("WARNING"): + replacement = relay_runtime.RelayRuntime( + relay=relay, + profile_key="replacement", + ) + try: + assert not replacement.managed_execution_enabled() + assert relay_runtime._PLUGIN_CONFIGURATION._activation is activation + assert relay.events.count(("plugin.initialize_dynamic", {}, [ + { + "plugin_id": "worker.policy", + "kind": "worker", + "manifest_ref": str(tmp_path / "relay-plugin.toml"), + "environment_ref": str(tmp_path / "environment"), + "config": {}, + } + ])) == 1 + assert relay.events.count(("plugin.activation.close",)) == 2 + assert "refusing to replace" in caplog.text + finally: + replacement.shutdown() + # Relay treats a close failure as terminal; only reset the permissive + # fake so this process-global fixture cannot leak into later tests. + relay.activation_close_error = None + relay_runtime._PLUGIN_CONFIGURATION.reset_for_tests() + + +def test_missing_dynamic_initializer_falls_back_to_static_plugins( + tmp_path, + monkeypatch, + caplog, +): + config = tmp_path / "plugins.toml" + config.write_text( + """ +[[dynamic_plugins]] +plugin_id = "native.policy" +kind = "rust_dynamic" +manifest_ref = "relay-plugin.toml" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _FakeRelay() + del relay.plugin.initialize_with_dynamic_plugins + + with caplog.at_level("WARNING"): + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + try: + assert host.managed_execution_enabled() + assert relay.events == [("plugin.initialize", {})] + assert "require a binding" in caplog.text + finally: + host.shutdown() + + +def test_gateway_dynamic_records_are_not_activated_without_cli_lifecycle_state( + tmp_path, + monkeypatch, + caplog, +): + config = tmp_path / "plugins.toml" + config.write_text( + """ +[[plugins.dynamic]] +manifest = "relay-plugin.toml" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _FakeRelay() + + with caplog.at_level("WARNING"): + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + try: + assert host.managed_execution_enabled() + assert relay.events == [("plugin.initialize", {})] + assert "require lifecycle state" in caplog.text + finally: + host.shutdown() + + +def test_invalid_dynamic_spec_rejects_explicit_file_atomically( + tmp_path, + monkeypatch, + caplog, +): + config = tmp_path / "plugins.toml" + config.write_text( + """ +version = 1 + +[[components]] +kind = "observability" + +[[dynamic_plugins]] +plugin_id = "valid.native" +kind = "rust_dynamic" +manifest_ref = "native/relay-plugin.toml" + +[[dynamic_plugins]] +kind = "worker" +manifest_ref = "worker/relay-plugin.toml" +""".strip(), + encoding="utf-8", + ) + monkeypatch.setenv(relay_runtime.RELAY_PLUGINS_CONFIG_ENV, str(config)) + relay = _FakeRelay() + + with caplog.at_level("WARNING"): + host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") + try: + assert host.managed_execution_enabled() + assert relay.events == [("plugin.initialize", {})] + assert "dynamic_plugins[1].plugin_id is required" in caplog.text + finally: + host.shutdown() + + def test_real_binding_discovers_project_config_and_exports_native_activity( tmp_path, monkeypatch, diff --git a/tests/plugins/test_nemo_relay_plugin.py b/tests/plugins/test_nemo_relay_plugin.py index 9da11845d0..78ccd92e48 100644 --- a/tests/plugins/test_nemo_relay_plugin.py +++ b/tests/plugins/test_nemo_relay_plugin.py @@ -422,6 +422,14 @@ def test_real_binding_shares_plugin_configuration_across_two_profiles( monkeypatch.setattr(relay.plugin, "initialize", _initialize) monkeypatch.setattr(relay.plugin, "clear", _clear) + # This test exercises the bundled plugin's legacy configuration owner in + # isolation. Native and bundled-plugin ownership are intentionally not + # combined until their process-global lifetime models are unified. + monkeypatch.setattr( + relay_runtime._PLUGIN_CONFIGURATION, + "acquire", + lambda _owner, _relay: False, + ) monkeypatch.setattr( plugin, "_load_settings", @@ -481,4 +489,3 @@ def test_relay_tool_request_rewrite_precedes_hermes_authorization_boundary( assert result.payload == {"intercepted": True, "value": 1} assert result.trace[0] == {"source": "nemo_relay"} -