refactor(agent/relay): drop _safe in favour of _warn_on_error/suppress; pack spans to 110 cols; strip intra-function separator blanks
This commit is contained in:
@@ -80,9 +80,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] = {}
|
||||
@@ -138,8 +137,7 @@ class _ManagedAttempt:
|
||||
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:
|
||||
@@ -153,10 +151,7 @@ class _ManagedAttempt:
|
||||
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; "
|
||||
@@ -176,40 +171,34 @@ class _ManagedAttempt:
|
||||
|
||||
|
||||
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,
|
||||
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
|
||||
)
|
||||
attempt = _ManagedAttempt.resolve(session_id, request, metadata, name=name, model_name=model_name)
|
||||
if attempt is None:
|
||||
return callback(request)
|
||||
|
||||
invoke = partial(attempt.invoke, callback)
|
||||
try:
|
||||
managed = _run_awaitable(attempt.run_managed(attempt.runtime.relay.llm.execute, invoke))
|
||||
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_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,
|
||||
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
|
||||
)
|
||||
attempt = _ManagedAttempt.resolve(session_id, request, metadata, name=name, model_name=model_name)
|
||||
if attempt is None:
|
||||
return await callback(request)
|
||||
|
||||
invoke = partial(attempt.invoke_async, callback)
|
||||
try:
|
||||
managed = await attempt.run_managed(attempt.runtime.relay.llm.execute, invoke)
|
||||
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)
|
||||
@@ -252,10 +241,9 @@ 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.
|
||||
|
||||
@@ -317,12 +305,10 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
_provider_completed = False
|
||||
|
||||
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 = None,
|
||||
on_chunk: Callable[[Any], None] | None = None,
|
||||
chunk_adapter: Callable[[Any], Any] | None = None,
|
||||
accept_chunk: Callable[[Any], bool] | None = None,
|
||||
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:
|
||||
@@ -336,13 +322,9 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
self._accept_chunk = accept_chunk
|
||||
self._raw_chunks: list[tuple[Any, Any]] = []
|
||||
self._prefetched_chunks: list[Any] = []
|
||||
attempt = _ManagedAttempt.resolve(
|
||||
session_id, request, metadata, name=name, model_name=model_name
|
||||
)
|
||||
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, stream_factory, on_stream_created, completed_response_predicate)
|
||||
return
|
||||
self._logical = attempt.logical
|
||||
self._start_managed(
|
||||
@@ -390,9 +372,7 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
chunk = run_callback(next, raw_iterator)
|
||||
except StopIteration:
|
||||
break
|
||||
if self._accept_chunk is not None and not run_callback(
|
||||
self._accept_chunk, chunk
|
||||
):
|
||||
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))
|
||||
@@ -441,8 +421,7 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
try:
|
||||
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, provider_stream, observe_chunk, relay_finalizer,
|
||||
)
|
||||
)
|
||||
except BaseException as exc:
|
||||
@@ -488,8 +467,7 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
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,
|
||||
provider_name=self._logical_provider_name, response_model_name=self._logical_response_model_name,
|
||||
operation_lease=self._runtime_lease,
|
||||
)
|
||||
self._logical = None
|
||||
@@ -568,9 +546,7 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
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:
|
||||
@@ -833,9 +809,7 @@ def _provider_request(
|
||||
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 = {
|
||||
@@ -854,10 +828,7 @@ def _codex_codec_tools(body: dict[str, Any]) -> 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"},
|
||||
}
|
||||
{"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"]
|
||||
@@ -875,9 +846,7 @@ def _chat_codec_tools(body: dict[str, Any]) -> None:
|
||||
|
||||
|
||||
# api_mode -> in-place normalizer producing the codec-facing ``tools`` shape.
|
||||
_CODEC_TOOL_NORMALIZERS = {
|
||||
"codex_responses": _codex_codec_tools, "chat_completions": _chat_codec_tools
|
||||
}
|
||||
_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]:
|
||||
@@ -891,8 +860,7 @@ def _relay_request_body(request: dict[str, Any], metadata: dict[str, Any] | None
|
||||
|
||||
|
||||
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))
|
||||
@@ -913,8 +881,7 @@ def _restore_provider_message_extensions(
|
||||
|
||||
|
||||
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)
|
||||
@@ -927,19 +894,13 @@ def _codec_round_trip_request_body(
|
||||
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
|
||||
@@ -984,9 +945,7 @@ def _jsonable(value: Any) -> Any:
|
||||
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)
|
||||
@@ -1018,9 +977,7 @@ def _json_equal(left: Any, right: Any) -> bool:
|
||||
|
||||
|
||||
def _run_awaitable(
|
||||
value: Any,
|
||||
*,
|
||||
loop_error: str = "Synchronous Relay LLM execution cannot run on an event-loop thread",
|
||||
value: Any, *, loop_error: str = "Synchronous Relay LLM execution cannot run on an event-loop thread",
|
||||
) -> Any:
|
||||
if not inspect.isawaitable(value):
|
||||
return value
|
||||
|
||||
@@ -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__)
|
||||
|
||||
@@ -58,7 +57,6 @@ def _scope_op_executor():
|
||||
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"
|
||||
)
|
||||
@@ -101,9 +99,8 @@ def pop_relay_scope(
|
||||
"""
|
||||
pop = relay.scope.pop
|
||||
kwargs = {
|
||||
key: value
|
||||
for key, value in (("output", output), ("metadata", metadata), ("timestamp", timestamp))
|
||||
if value is not None
|
||||
key: value for key,
|
||||
value in (("output", output), ("metadata", metadata), ("timestamp", timestamp)) if value is not None
|
||||
}
|
||||
try:
|
||||
params = inspect.signature(pop).parameters
|
||||
@@ -181,7 +178,6 @@ def _load_segments_config() -> dict[str, Any]:
|
||||
segments: dict[str, Any] = {}
|
||||
try:
|
||||
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 {}
|
||||
except Exception: # noqa: BLE001 - config absence must not crash
|
||||
@@ -484,10 +480,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:
|
||||
@@ -528,16 +522,11 @@ class RelayRuntime:
|
||||
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,
|
||||
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,
|
||||
},
|
||||
**{"hermes.session.segment": session.segment, "hermes.session.segment_reason": reason},
|
||||
)
|
||||
try:
|
||||
self._open_session_scope(session, scope_metadata, resolve_parent=False)
|
||||
@@ -545,9 +534,7 @@ class RelayRuntime:
|
||||
logger.warning(
|
||||
"Hermes Relay segment open failed (session=%s segment=%d); "
|
||||
"keeping the prior scope handle",
|
||||
session.session_id,
|
||||
session.segment,
|
||||
exc_info=True,
|
||||
session.session_id, session.segment, exc_info=True,
|
||||
)
|
||||
|
||||
def register_subagent(
|
||||
@@ -597,9 +584,7 @@ class RelayRuntime:
|
||||
return None if session.closing else session
|
||||
return None
|
||||
|
||||
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:
|
||||
@@ -707,16 +692,13 @@ class RelayRuntime:
|
||||
self._begin_operation()
|
||||
return RelayOperationLease(self)
|
||||
|
||||
def emit_mark(
|
||||
self, name: str, event: dict[str, Any], *, data: Any = None, metadata: Any = None
|
||||
) -> bool:
|
||||
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,
|
||||
session, self.relay.scope.event, name, handle=session.handle, data=data, metadata=metadata,
|
||||
)
|
||||
return True
|
||||
|
||||
@@ -755,24 +737,17 @@ class RelayRuntime:
|
||||
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
|
||||
@@ -792,9 +767,7 @@ class RelayRuntime:
|
||||
"""
|
||||
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 {},
|
||||
@@ -874,13 +847,14 @@ class RelayRuntime:
|
||||
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
|
||||
self._safe(atexit.unregister, self.shutdown, quiet=True)
|
||||
with contextlib.suppress(Exception):
|
||||
atexit.unregister(self.shutdown)
|
||||
except Exception:
|
||||
with self._sessions_lock:
|
||||
self._shutdown_started = False
|
||||
@@ -889,15 +863,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:
|
||||
@@ -935,16 +900,14 @@ 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:
|
||||
if host is not None:
|
||||
return host
|
||||
try:
|
||||
host = RelayRuntime(profile_key=key)
|
||||
@@ -1082,9 +1045,8 @@ class RelaySessionCoordinator:
|
||||
session = None
|
||||
if isinstance(host, RelayRuntime):
|
||||
session = _warn_on_error(
|
||||
"conversation initialization", self._open_conversation_session, host,
|
||||
profile_key=profile_key, session_id=session_id, platform=platform,
|
||||
parent_session_id=parent_session_id, model=model,
|
||||
"conversation initialization", self._open_conversation_session, host, profile_key=profile_key,
|
||||
session_id=session_id, platform=platform, parent_session_id=parent_session_id, model=model,
|
||||
)
|
||||
return ConversationLease(
|
||||
profile_key=profile_key, session_id=session_id, platform=platform, host=host,
|
||||
@@ -1102,14 +1064,11 @@ class RelaySessionCoordinator:
|
||||
metadata = {"hermes.execution_surface": 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,
|
||||
{"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:
|
||||
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)
|
||||
@@ -1184,9 +1143,7 @@ class RelaySessionCoordinator:
|
||||
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:
|
||||
@@ -1214,9 +1171,7 @@ class RelaySessionCoordinator:
|
||||
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
|
||||
):
|
||||
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})
|
||||
|
||||
def notify_session_compacted(
|
||||
@@ -1365,9 +1320,7 @@ def active_turn(session_id: str | None = None) -> RelayTurnContext | 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:
|
||||
# Nested managed execution is impossible (see _MANAGED_CALLBACK_DEPTH); the
|
||||
@@ -1404,23 +1357,17 @@ def emit_mark(name: str, *, session_id: str, data: Any = None, metadata: Any = N
|
||||
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
|
||||
@@ -1476,7 +1423,6 @@ 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:
|
||||
@@ -1507,9 +1453,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:
|
||||
|
||||
@@ -21,7 +21,6 @@ def execute(
|
||||
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
|
||||
@@ -68,7 +67,6 @@ def execute(
|
||||
)
|
||||
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):
|
||||
@@ -110,6 +108,5 @@ def _json_equal(left: Any, right: Any) -> bool:
|
||||
|
||||
def _run_awaitable(value: Any) -> Any:
|
||||
return relay_llm._run_awaitable(
|
||||
value,
|
||||
loop_error="Synchronous Hermes Relay tool execution cannot run on an active event-loop thread",
|
||||
value, loop_error="Synchronous Hermes Relay tool execution cannot run on an active event-loop thread",
|
||||
)
|
||||
|
||||
@@ -17,7 +17,6 @@ def _unprefix_oauth_tool_name(name: str) -> str:
|
||||
"""
|
||||
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):
|
||||
@@ -29,9 +28,8 @@ 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
|
||||
@@ -41,13 +39,11 @@ 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(
|
||||
@@ -56,12 +52,10 @@ class AnthropicTransport(ProviderTransport):
|
||||
) -> 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),
|
||||
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),
|
||||
@@ -71,13 +65,11 @@ class AnthropicTransport(ProviderTransport):
|
||||
"""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 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
|
||||
@@ -101,7 +93,6 @@ 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
|
||||
@@ -112,7 +103,6 @@ class AnthropicTransport(ProviderTransport):
|
||||
)
|
||||
if _has_signed_thinking and any(isinstance(b, dict) and 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,
|
||||
finish_reason=self.map_finish_reason(response.stop_reason),
|
||||
|
||||
Reference in New Issue
Block a user