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:
Teknium
2026-09-02 19:24:48 -07:00
parent 2ed4135398
commit e58175139d
4 changed files with 76 additions and 188 deletions

View File

@@ -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

View File

@@ -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:

View File

@@ -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",
)

View File

@@ -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),