Merge branch 'simp/r3-07-C_relay' into simp/r3-07
This commit is contained in:
+325
-622
File diff suppressed because it is too large
Load Diff
+372
-701
File diff suppressed because it is too large
Load Diff
+19
-49
@@ -2,49 +2,40 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
import inspect
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from agent import relay_runtime
|
||||
from agent import relay_llm, relay_runtime
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def execute(
|
||||
tool_name: str,
|
||||
args: dict[str, Any],
|
||||
callback: Callable[[dict[str, Any]], Any],
|
||||
*,
|
||||
session_id: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
tool_name: str, args: dict[str, Any], callback: Callable[[dict[str, Any]], Any], *,
|
||||
session_id: str, metadata: dict[str, Any] | None = None,
|
||||
) -> tuple[Any, dict[str, Any]]:
|
||||
"""Run one tool call through Relay and return its final arguments."""
|
||||
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
|
||||
callback_context = contextvars.copy_context()
|
||||
|
||||
def guarded(final_args: dict[str, Any]) -> Any:
|
||||
# Everything the tool transitively calls (incl. auxiliary LLM calls on worker
|
||||
# threads) must bypass managed Relay: the pipeline's Futures bind to THIS loop,
|
||||
# which is blocked until the tool returns.
|
||||
with relay_runtime.managed_callback_guard():
|
||||
return callback(final_args)
|
||||
|
||||
def invoke(next_args: Any) -> Any:
|
||||
nonlocal callback_error, observed_args
|
||||
observed_args = next_args if isinstance(next_args, dict) else args
|
||||
|
||||
def guarded(final_args: dict[str, Any]) -> Any:
|
||||
# Everything the tool transitively calls (including auxiliary LLM
|
||||
# calls it forwards to worker threads) must bypass managed Relay
|
||||
# execution — the native pipeline's Futures bind to THIS loop,
|
||||
# which is blocked until the tool returns (#77244).
|
||||
with relay_runtime.managed_callback_guard():
|
||||
return callback(final_args)
|
||||
|
||||
try:
|
||||
result = callback_context.copy().run(guarded, observed_args)
|
||||
except BaseException as exc:
|
||||
@@ -57,35 +48,23 @@ def execute(
|
||||
try:
|
||||
managed = _run_awaitable(
|
||||
runtime.run_in_session_async(
|
||||
session,
|
||||
runtime.relay.tools.execute,
|
||||
tool_name,
|
||||
_jsonable(args),
|
||||
invoke,
|
||||
handle=parent,
|
||||
metadata=_jsonable(metadata or {}),
|
||||
session, runtime.relay.tools.execute, tool_name, _jsonable(args), invoke,
|
||||
handle=parent, metadata=_jsonable(metadata or {}),
|
||||
)
|
||||
)
|
||||
except BaseException as exc:
|
||||
if (
|
||||
callback_error is not None
|
||||
and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error)
|
||||
):
|
||||
if callback_error is not None and relay_runtime._is_relay_wrapped_callback_error(exc, callback_error):
|
||||
raise callback_error
|
||||
if isinstance(exc, Exception) and callback_error is None and "value" in raw_result:
|
||||
logger.warning(
|
||||
"NeMo Relay tool post-processing failed after dispatch success; "
|
||||
"returning the Hermes tool result",
|
||||
"NeMo Relay tool post-processing failed after dispatch success; returning the Hermes tool result",
|
||||
exc_info=True,
|
||||
)
|
||||
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):
|
||||
return managed, observed_args
|
||||
return json.dumps(_jsonable(managed), ensure_ascii=False), observed_args
|
||||
return (managed if isinstance(managed, str) else json.dumps(_jsonable(managed), ensure_ascii=False)), observed_args
|
||||
|
||||
|
||||
def _jsonable(value: Any) -> Any:
|
||||
@@ -98,8 +77,7 @@ def _jsonable(value: Any) -> Any:
|
||||
model_dump = getattr(value, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
try:
|
||||
# warnings=False: suppress pydantic's serializer UserWarnings on
|
||||
# generic-union SDK models; they would leak to the CLI mid-turn.
|
||||
# warnings=False: pydantic's generic-union warning would leak to the CLI mid-turn.
|
||||
try:
|
||||
return _jsonable(model_dump(mode="json", warnings=False))
|
||||
except TypeError:
|
||||
@@ -114,20 +92,12 @@ def _jsonable(value: Any) -> Any:
|
||||
|
||||
def _json_equal(left: Any, right: Any) -> bool:
|
||||
try:
|
||||
return json.dumps(
|
||||
_jsonable(left), sort_keys=True, separators=(",", ":")
|
||||
) == json.dumps(_jsonable(right), sort_keys=True, separators=(",", ":"))
|
||||
return relay_llm._canonical_json(left, _jsonable) == relay_llm._canonical_json(right, _jsonable)
|
||||
except (TypeError, ValueError):
|
||||
return left == right
|
||||
|
||||
|
||||
def _run_awaitable(value: Any) -> Any:
|
||||
if not inspect.isawaitable(value):
|
||||
return value
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return asyncio.run(value)
|
||||
raise RuntimeError(
|
||||
"Synchronous Hermes Relay tool execution cannot run on an active event-loop thread"
|
||||
return relay_llm._run_awaitable(
|
||||
value, loop_error="Synchronous Hermes Relay tool execution cannot run on an active event-loop thread",
|
||||
)
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
"""Transport registry for provider response normalization.
|
||||
|
||||
transport = get_transport("anthropic_messages")
|
||||
result = transport.normalize_response(raw_response)
|
||||
"""
|
||||
result = transport.normalize_response(raw_response)"""
|
||||
|
||||
import contextlib
|
||||
import importlib
|
||||
|
||||
from agent.transports.types import ( # noqa: F401
|
||||
NormalizedResponse,
|
||||
@@ -24,14 +25,10 @@ def register_transport(api_mode: str, transport_cls: type) -> None:
|
||||
|
||||
def get_transport(api_mode: str):
|
||||
"""Return a transport instance for ``api_mode``, or None so callers can fall back to the legacy path."""
|
||||
if not _discovered:
|
||||
# A directly-imported transport leaves the registry partial; (re)discover on first use and on misses.
|
||||
if not _discovered or api_mode not in _REGISTRY:
|
||||
_discover_transports()
|
||||
cls = _REGISTRY.get(api_mode)
|
||||
if cls is None:
|
||||
# A directly-imported transport module leaves the registry partially
|
||||
# populated; discover on misses so import order can't hide a valid api_mode.
|
||||
_discover_transports()
|
||||
cls = _REGISTRY.get(api_mode)
|
||||
return None if cls is None else cls()
|
||||
|
||||
|
||||
@@ -39,10 +36,6 @@ def _discover_transports() -> None:
|
||||
"""Import all transport modules to trigger auto-registration."""
|
||||
global _discovered
|
||||
_discovered = True
|
||||
import importlib
|
||||
|
||||
for name in _TRANSPORT_MODULES:
|
||||
try:
|
||||
with contextlib.suppress(ImportError):
|
||||
importlib.import_module(f"agent.transports.{name}")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
@@ -1,8 +1,4 @@
|
||||
"""Anthropic Messages API transport.
|
||||
|
||||
Delegates format conversion to agent/anthropic_adapter.py; owns normalization,
|
||||
not client lifecycle.
|
||||
"""
|
||||
"""Anthropic Messages API transport: conversion via agent/anthropic_adapter.py, normalization here."""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
@@ -10,20 +6,16 @@ from agent.transports.base import ProviderTransport
|
||||
from agent.transports.types import NormalizedResponse, ToolCall
|
||||
|
||||
_MCP_PREFIX = "mcp__"
|
||||
_THINKING_TYPES = ("thinking", "redacted_thinking")
|
||||
|
||||
|
||||
def _unprefix_oauth_tool_name(name: str) -> str:
|
||||
"""Reverse the OAuth-wire ``mcp__`` prefix back to the registered tool name.
|
||||
|
||||
Two originals map onto one wire name (``mcp__read_file`` <- ``read_file``;
|
||||
``mcp__linear_get_issue`` <- ``mcp_linear_get_issue``), so resolve by registry
|
||||
lookup, never rewriting a name that already resolves natively (GH-25255).
|
||||
OAuth wire aliases (e.g. chat_history_lookup -> session_search) are checked
|
||||
LAST so a real tool registered under the wire name still wins.
|
||||
"""
|
||||
Two originals map onto one wire name (``read_file`` / ``mcp_linear_get_issue``), so
|
||||
resolve by registry lookup, never rewriting a name that already resolves natively.
|
||||
OAuth wire aliases are checked LAST so a real tool under the wire name still wins."""
|
||||
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):
|
||||
@@ -31,16 +23,19 @@ def _unprefix_oauth_tool_name(name: str) -> str:
|
||||
return _OAUTH_TOOL_NAME_REVERSE_ALIASES.get(bare, name)
|
||||
|
||||
|
||||
# build_kwargs params forwarded to build_anthropic_kwargs, with the defaults applied when absent.
|
||||
_BUILD_KWARG_DEFAULTS = {
|
||||
"max_tokens": 16384, "reasoning_config": None, "tool_choice": None, "is_oauth": False, "preserve_dots": False,
|
||||
"context_length": None, "base_url": None, "fast_mode": False, "drop_context_1m_beta": False,
|
||||
}
|
||||
|
||||
|
||||
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
|
||||
@@ -50,68 +45,44 @@ 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(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[Dict[str, Any]],
|
||||
tools: Optional[List[Dict[str, Any]]] = None,
|
||||
**params,
|
||||
self, model: str, messages: List[Dict[str, Any]], tools: Optional[List[Dict[str, Any]]] = None, **params,
|
||||
) -> 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),
|
||||
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),
|
||||
model=model, messages=messages, tools=tools,
|
||||
**{key: params.get(key, default) for key, default in _BUILD_KWARG_DEFAULTS.items()},
|
||||
)
|
||||
|
||||
def normalize_response(self, response: Any, **kwargs) -> NormalizedResponse:
|
||||
"""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 that PRECEDE it.
|
||||
# When thinking interleaves with tool_use, the parallel reasoning_details +
|
||||
# tool_calls lists lose that ordering and replay -> HTTP 400 "thinking ...
|
||||
# blocks cannot be modified". Keep the exact sequence for the adapter.
|
||||
# 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
|
||||
if isinstance(block_dict, dict):
|
||||
# Sanitize at capture so output-only SDK fields never persist to
|
||||
# state.db and leak back as request input on replay (HTTP 400).
|
||||
clean_block = _sanitize_replay_block(block_dict)
|
||||
if clean_block is not None:
|
||||
ordered_blocks.append(clean_block)
|
||||
# Sanitize at capture so output-only SDK fields never persist and replay (400).
|
||||
clean_block = _sanitize_replay_block(block_dict) if isinstance(block_dict, dict) else None
|
||||
if clean_block is not None:
|
||||
ordered_blocks.append(clean_block)
|
||||
if block.type == "text":
|
||||
text_parts.append(block.text)
|
||||
elif block.type in ("thinking", "redacted_thinking"):
|
||||
elif block.type in _THINKING_TYPES:
|
||||
if block.type == "thinking":
|
||||
reasoning_parts.append(block.thinking)
|
||||
# Prefer the sanitized block (replayed on the non-ordered path); raw only if sanitize dropped it.
|
||||
# Sanitized block preferred; raw only if sanitize dropped it.
|
||||
if isinstance(clean_block, dict):
|
||||
reasoning_details.append(clean_block)
|
||||
elif isinstance(block_dict, dict):
|
||||
@@ -121,33 +92,28 @@ 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
|
||||
# Carry the ordered channel only for the one shape the parallel lists
|
||||
# reconstruct wrongly: signed thinking interleaved with tool_use.
|
||||
_has_signed_thinking = any(
|
||||
isinstance(b, dict) and b.get("type") in ("thinking", "redacted_thinking") and (b.get("signature") or b.get("data"))
|
||||
for b in ordered_blocks
|
||||
# Ordered channel only for the shape the parallel lists reconstruct wrongly.
|
||||
kinds = {b.get("type") for b in ordered_blocks if isinstance(b, dict)}
|
||||
signed = any(
|
||||
b.get("type") in _THINKING_TYPES and (b.get("signature") or b.get("data"))
|
||||
for b in ordered_blocks if isinstance(b, dict)
|
||||
)
|
||||
if _has_signed_thinking and any(isinstance(b, dict) and b.get("type") == "tool_use" for b in ordered_blocks):
|
||||
if signed and "tool_use" in kinds:
|
||||
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,
|
||||
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),
|
||||
reasoning="\n\n".join(reasoning_parts) if reasoning_parts else None,
|
||||
usage=None,
|
||||
reasoning="\n\n".join(reasoning_parts) if reasoning_parts else None, usage=None,
|
||||
provider_data=provider_data or None,
|
||||
)
|
||||
|
||||
def validate_response(self, response: Any) -> bool:
|
||||
"""Structural check. An empty content list is legitimate for ``end_turn`` (nothing to add
|
||||
after a tool turn) and ``refusal`` (Claude 4.5+ declines with empty content); treating
|
||||
either as invalid would retry a completed/deterministic response forever."""
|
||||
content_blocks = getattr(response, "content", None) if response is not None else None
|
||||
"""Structural check; empty content is legitimate for ``end_turn``/``refusal`` (retrying
|
||||
either would loop forever)."""
|
||||
content_blocks = getattr(response, "content", None)
|
||||
if not isinstance(content_blocks, list):
|
||||
return False
|
||||
return bool(content_blocks) or getattr(response, "stop_reason", None) in {"end_turn", "refusal"}
|
||||
|
||||
@@ -1,10 +1,7 @@
|
||||
"""Abstract base for provider transports.
|
||||
|
||||
A transport owns the data path for one api_mode:
|
||||
convert_messages -> convert_tools -> build_kwargs -> normalize_response
|
||||
It does NOT own client construction, streaming, credential refresh, prompt
|
||||
caching, interrupt handling, or retry logic — those stay on AIAgent.
|
||||
"""
|
||||
A transport owns one api_mode's data path (convert_messages -> convert_tools -> build_kwargs
|
||||
-> normalize_response), NOT client construction, streaming, credentials, caching, interrupts
|
||||
or retries — those stay on AIAgent."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Optional
|
||||
@@ -34,11 +31,8 @@ class ProviderTransport(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def build_kwargs(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[Dict[str, Any]],
|
||||
tools: Optional[List[Dict[str, Any]]] = None,
|
||||
**params,
|
||||
self, model: str, messages: List[Dict[str, Any]],
|
||||
tools: Optional[List[Dict[str, Any]]] = None, **params,
|
||||
) -> Dict[str, Any]:
|
||||
"""Primary entry point: convert messages/tools and return kwargs ready for the provider SDK."""
|
||||
|
||||
|
||||
@@ -967,7 +967,6 @@ def test_core_runtime_is_fail_open_without_a_published_binding(monkeypatch, capl
|
||||
tool_name="terminal",
|
||||
args={"command": "true"},
|
||||
) == {"command": "true"}
|
||||
assert not relay_runtime.emit_mark("hermes.probe", session_id="s1")
|
||||
assert "Hermes Relay runtime initialization failed" in caplog.text
|
||||
relay_runtime._reset_for_tests()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user