From e58175139d046ba53afc75d90a42db0a38afaade Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 19:24:48 -0700 Subject: [PATCH] refactor(agent/relay): drop _safe in favour of _warn_on_error/suppress; pack spans to 110 cols; strip intra-function separator blanks --- agent/relay_llm.py | 121 +++++++++++---------------------- agent/relay_runtime.py | 122 +++++++++------------------------- agent/relay_tools.py | 5 +- agent/transports/anthropic.py | 16 +---- 4 files changed, 76 insertions(+), 188 deletions(-) diff --git a/agent/relay_llm.py b/agent/relay_llm.py index 7b3152e4ef..71db2e0920 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -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 diff --git a/agent/relay_runtime.py b/agent/relay_runtime.py index f135b35f10..e292f931e9 100644 --- a/agent/relay_runtime.py +++ b/agent/relay_runtime.py @@ -4,6 +4,7 @@ from __future__ import annotations import atexit import asyncio +import contextlib import contextvars import importlib import inspect @@ -19,9 +20,7 @@ from pathlib import Path from typing import Any, Callable from hermes_constants import get_hermes_home -from hermes_cli.relay_plugin_cutover import ( - RELAY_PLUGINS_CONFIG_ENV, configured_legacy_relay_env_vars -) +from hermes_cli.relay_plugin_cutover import (RELAY_PLUGINS_CONFIG_ENV, configured_legacy_relay_env_vars) logger = logging.getLogger(__name__) @@ -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: diff --git a/agent/relay_tools.py b/agent/relay_tools.py index d4e7f13a75..ae227a7f3f 100644 --- a/agent/relay_tools.py +++ b/agent/relay_tools.py @@ -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", ) diff --git a/agent/transports/anthropic.py b/agent/transports/anthropic.py index f234115240..fa2350a546 100644 --- a/agent/transports/anthropic.py +++ b/agent/transports/anthropic.py @@ -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),