refactor(agent/relay): fold execute/execute_async invoke into _ManagedAttempt; stream() = ManagedLlmStream; codec tool normalizer table; plugin acquire preflight/activate split; pack call/signature spans
This commit is contained in:
+167
-306
@@ -8,7 +8,6 @@ import inspect
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Callable, Iterator
|
||||
from dataclasses import dataclass
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
@@ -22,16 +21,11 @@ _RELAY_INTERNAL_PROVIDER_HEADERS = frozenset({"x-dynamo-parent-session-id", "x-d
|
||||
_LogicalCall = tuple[relay_runtime.RelayTurnContext, Any, str]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RelayProtocol:
|
||||
operation: str
|
||||
codec_class: str
|
||||
|
||||
|
||||
# api_mode -> (Relay operation name, codec class name on ``relay.codecs``)
|
||||
_RELAY_PROTOCOL_BY_API_MODE = {
|
||||
"chat_completions": _RelayProtocol("openai.chat_completions", "OpenAIChatCodec"),
|
||||
"codex_responses": _RelayProtocol("openai.responses", "OpenAIResponsesCodec"),
|
||||
"anthropic_messages": _RelayProtocol("anthropic.messages", "AnthropicMessagesCodec"),
|
||||
"chat_completions": ("openai.chat_completions", "OpenAIChatCodec"),
|
||||
"codex_responses": ("openai.responses", "OpenAIResponsesCodec"),
|
||||
"anthropic_messages": ("anthropic.messages", "AnthropicMessagesCodec"),
|
||||
}
|
||||
|
||||
|
||||
@@ -39,16 +33,10 @@ def _api_mode(metadata: dict[str, Any] | None) -> str:
|
||||
return str((metadata or {}).get("api_mode") or "")
|
||||
|
||||
|
||||
def _relay_protocol(metadata: dict[str, Any] | None) -> _RelayProtocol | None:
|
||||
"""Return Relay's operation and codec descriptor for an API mode."""
|
||||
api_mode = (metadata or {}).get("api_mode")
|
||||
return _RELAY_PROTOCOL_BY_API_MODE.get(api_mode) if isinstance(api_mode, str) else None
|
||||
|
||||
|
||||
def _relay_operation_name(provider_name: str, metadata: dict[str, Any] | None) -> str:
|
||||
"""Return Relay's canonical operation name when Hermes knows the API mode."""
|
||||
protocol = _relay_protocol(metadata)
|
||||
return protocol.operation if protocol is not None else provider_name
|
||||
protocol = _RELAY_PROTOCOL_BY_API_MODE.get(_api_mode(metadata))
|
||||
return protocol[0] if protocol is not None else provider_name
|
||||
|
||||
|
||||
def _relay_metadata(provider_name: str, metadata: dict[str, Any] | None) -> dict[str, Any]:
|
||||
@@ -63,13 +51,8 @@ class _ManagedAttempt:
|
||||
|
||||
@classmethod
|
||||
def resolve(
|
||||
cls,
|
||||
session_id: str,
|
||||
request: dict[str, Any],
|
||||
metadata: dict[str, Any] | None,
|
||||
*,
|
||||
name: str,
|
||||
model_name: str,
|
||||
cls, session_id: str, request: dict[str, Any], metadata: dict[str, Any] | None, *,
|
||||
name: str, model_name: str,
|
||||
) -> "_ManagedAttempt | None":
|
||||
"""Return the managed attempt for ``session_id``, or None to run unmanaged."""
|
||||
runtime, session, parent = relay_runtime.resolve_execution_context(session_id)
|
||||
@@ -78,15 +61,8 @@ class _ManagedAttempt:
|
||||
return cls(runtime, session, parent, request, metadata, name=name, model_name=model_name)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
runtime: relay_runtime.RelayRuntime,
|
||||
session: Any,
|
||||
parent: Any,
|
||||
request: dict[str, Any],
|
||||
metadata: dict[str, Any] | None,
|
||||
*,
|
||||
name: str,
|
||||
model_name: str,
|
||||
self, runtime: relay_runtime.RelayRuntime, session: Any, parent: Any,
|
||||
request: dict[str, Any], metadata: dict[str, Any] | None, *, name: str, model_name: str,
|
||||
) -> None:
|
||||
self.runtime = runtime
|
||||
self.session = session
|
||||
@@ -101,10 +77,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),
|
||||
"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.
|
||||
@@ -113,11 +87,8 @@ class _ManagedAttempt:
|
||||
|
||||
def provider_request(self, next_request: Any) -> dict[str, Any]:
|
||||
return _provider_request(
|
||||
self.request,
|
||||
next_request,
|
||||
relay_request_body=self.body,
|
||||
codec_baseline_body=self.codec_baseline,
|
||||
metadata=self.metadata,
|
||||
self.request, next_request, relay_request_body=self.body,
|
||||
codec_baseline_body=self.codec_baseline, metadata=self.metadata,
|
||||
)
|
||||
|
||||
def run_callback(self, callback: Callable[..., Any], *args: Any) -> Any:
|
||||
@@ -133,22 +104,40 @@ class _ManagedAttempt:
|
||||
|
||||
return self.context.copy().run(guarded)
|
||||
|
||||
def record(self, raw: Any) -> Any:
|
||||
def _record(self, raw: Any) -> Any:
|
||||
self.raw_response["value"] = raw
|
||||
self.raw_response["json"] = _jsonable(raw)
|
||||
return self.raw_response["json"]
|
||||
|
||||
def fail(self, exc: BaseException) -> None:
|
||||
self.raw_response["error"] = exc
|
||||
def invoke(self, callback: Callable[..., Any], next_request: Any) -> Any:
|
||||
"""Provider callback handed to Relay: run ``callback`` on Relay's (possibly rewritten) request."""
|
||||
try:
|
||||
raw = self.run_callback(callback, self.provider_request(next_request))
|
||||
except BaseException as exc:
|
||||
self.raw_response["error"] = exc
|
||||
raise
|
||||
return self._record(raw)
|
||||
|
||||
async def invoke_async(self, callback: Callable[..., Any], next_request: Any) -> Any:
|
||||
try:
|
||||
final_request = self.provider_request(next_request)
|
||||
|
||||
async def call_provider() -> Any:
|
||||
# Nested relay calls inside a managed provider callback must
|
||||
# run unmanaged — see relay_runtime.managed_callback_guard.
|
||||
with relay_runtime.managed_callback_guard():
|
||||
return await callback(final_request)
|
||||
|
||||
raw = await self.context.copy().run(asyncio.create_task, call_provider())
|
||||
except BaseException as exc:
|
||||
self.raw_response["error"] = exc
|
||||
raise
|
||||
return self._record(raw)
|
||||
|
||||
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.session, relay_call, self.operation, self.relay_request, *callbacks,
|
||||
**self.relay_kwargs,
|
||||
)
|
||||
|
||||
@@ -187,13 +176,8 @@ 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,
|
||||
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."""
|
||||
@@ -204,12 +188,7 @@ def execute(
|
||||
return callback(request)
|
||||
|
||||
def invoke(next_request: Any) -> Any:
|
||||
try:
|
||||
raw = attempt.run_callback(callback, attempt.provider_request(next_request))
|
||||
except BaseException as exc:
|
||||
attempt.fail(exc)
|
||||
raise
|
||||
return attempt.record(raw)
|
||||
return attempt.invoke(callback, next_request)
|
||||
|
||||
try:
|
||||
managed = _run_awaitable(attempt.run_managed(attempt.runtime.relay.llm.execute, invoke))
|
||||
@@ -219,13 +198,8 @@ def execute(
|
||||
|
||||
|
||||
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,
|
||||
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."""
|
||||
@@ -236,20 +210,7 @@ async def execute_async(
|
||||
return await callback(request)
|
||||
|
||||
async def invoke(next_request: Any) -> Any:
|
||||
try:
|
||||
final_request = attempt.provider_request(next_request)
|
||||
|
||||
async def call_provider() -> Any:
|
||||
# Nested relay calls inside a managed provider callback must
|
||||
# run unmanaged — see relay_runtime.managed_callback_guard.
|
||||
with relay_runtime.managed_callback_guard():
|
||||
return await callback(final_request)
|
||||
|
||||
raw = await attempt.context.copy().run(asyncio.create_task, call_provider())
|
||||
except BaseException as exc:
|
||||
attempt.fail(exc)
|
||||
raise
|
||||
return attempt.record(raw)
|
||||
return await attempt.invoke_async(callback, next_request)
|
||||
|
||||
try:
|
||||
managed = await attempt.run_managed(attempt.runtime.relay.llm.execute, invoke)
|
||||
@@ -265,50 +226,30 @@ def _current_session_id() -> str | None:
|
||||
|
||||
|
||||
def execute_current(
|
||||
request: dict[str, Any],
|
||||
callback: Callable[[dict[str, Any]], Any],
|
||||
*,
|
||||
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], *, name: str,
|
||||
model_name: str, metadata: dict[str, Any] | None = None, defer_logical_completion: bool = False,
|
||||
) -> Any:
|
||||
"""Run a provider attempt under the inherited Hermes turn when present."""
|
||||
session_id = _current_session_id()
|
||||
if session_id is None:
|
||||
return callback(request)
|
||||
return execute(
|
||||
request,
|
||||
callback,
|
||||
session_id=session_id,
|
||||
name=name,
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
defer_logical_completion=defer_logical_completion,
|
||||
request, callback, session_id=session_id, name=name, model_name=model_name,
|
||||
metadata=metadata, defer_logical_completion=defer_logical_completion,
|
||||
)
|
||||
|
||||
|
||||
async def execute_current_async(
|
||||
request: dict[str, Any],
|
||||
callback: Callable[[dict[str, Any]], Any],
|
||||
*,
|
||||
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], *, name: str,
|
||||
model_name: str, metadata: dict[str, Any] | None = None, defer_logical_completion: bool = False,
|
||||
) -> Any:
|
||||
"""Run an async provider attempt under the inherited turn when present."""
|
||||
session_id = _current_session_id()
|
||||
if session_id is None:
|
||||
return await callback(request)
|
||||
return await execute_async(
|
||||
request,
|
||||
callback,
|
||||
session_id=session_id,
|
||||
name=name,
|
||||
model_name=model_name,
|
||||
metadata=metadata,
|
||||
defer_logical_completion=defer_logical_completion,
|
||||
request, callback, session_id=session_id, name=name, model_name=model_name,
|
||||
metadata=metadata, defer_logical_completion=defer_logical_completion,
|
||||
)
|
||||
|
||||
|
||||
@@ -321,13 +262,8 @@ 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,
|
||||
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:
|
||||
@@ -351,14 +287,8 @@ def stream_current(
|
||||
# tracks the enclosing attempt and traps a completed response itself.
|
||||
return stream_factory(request)
|
||||
managed = stream(
|
||||
request,
|
||||
stream_factory,
|
||||
session_id=session_id,
|
||||
name=name,
|
||||
model_name=model_name,
|
||||
finalizer=finalizer,
|
||||
metadata=metadata,
|
||||
defer_logical_completion=defer_logical_completion,
|
||||
request, stream_factory, session_id=session_id, name=name, model_name=model_name,
|
||||
finalizer=finalizer, metadata=metadata, defer_logical_completion=defer_logical_completion,
|
||||
completed_response_predicate=completed_response_predicate,
|
||||
)
|
||||
if completed_response_predicate is not None:
|
||||
@@ -371,40 +301,6 @@ def stream_current(
|
||||
return managed
|
||||
|
||||
|
||||
def stream(
|
||||
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,
|
||||
) -> "ManagedLlmStream":
|
||||
"""Return a synchronous view of one Relay-managed provider stream."""
|
||||
return ManagedLlmStream(
|
||||
request,
|
||||
stream_factory,
|
||||
session_id=session_id,
|
||||
name=name,
|
||||
model_name=model_name,
|
||||
finalizer=finalizer,
|
||||
on_stream_created=on_stream_created,
|
||||
on_chunk=on_chunk,
|
||||
chunk_adapter=chunk_adapter,
|
||||
accept_chunk=accept_chunk,
|
||||
completed_response_predicate=completed_response_predicate,
|
||||
metadata=metadata,
|
||||
defer_logical_completion=defer_logical_completion,
|
||||
)
|
||||
|
||||
|
||||
def _aclose_on_loop(loop: asyncio.AbstractEventLoop, stream: Any) -> None:
|
||||
"""Await ``stream.aclose()`` on ``loop`` when the stream exposes one."""
|
||||
close = getattr(stream, "aclose", None)
|
||||
@@ -418,48 +314,42 @@ def _aclose_on_loop(loop: asyncio.AbstractEventLoop, stream: Any) -> None:
|
||||
|
||||
|
||||
class ManagedLlmStream(Iterator[Any]):
|
||||
"""Drive Relay's async stream from Hermes's provider worker thread."""
|
||||
"""Synchronous view of one Relay-managed provider stream, driven from the worker thread."""
|
||||
|
||||
final_response: Any = None
|
||||
output_modified = False
|
||||
_loop: asyncio.AbstractEventLoop | None = None
|
||||
_stream: Any = None
|
||||
_raw_stream_resource: Any = None
|
||||
_closed = False
|
||||
_runtime_lease: relay_runtime.RelayOperationLease | None = None
|
||||
_close_error: BaseException | None = None
|
||||
_callback_error: BaseException | None = None
|
||||
_logical: _LogicalCall | None = None
|
||||
_logical_response_model_name: str | None = None
|
||||
_relay_observes_chunks = False
|
||||
_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,
|
||||
on_chunk: Callable[[Any], None] | None,
|
||||
chunk_adapter: Callable[[Any], Any] | None,
|
||||
accept_chunk: Callable[[Any], bool] | None,
|
||||
completed_response_predicate: Callable[[Any], bool] | None,
|
||||
metadata: dict[str, Any] | None,
|
||||
defer_logical_completion: bool,
|
||||
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:
|
||||
self.final_response: Any = None
|
||||
self._loop: asyncio.AbstractEventLoop | None = None
|
||||
self._stream: Any = None
|
||||
self._raw_stream_resource: Any = None
|
||||
self._closed = False
|
||||
self._runtime_lease: relay_runtime.RelayOperationLease | None = None
|
||||
self._close_error: BaseException | None = None
|
||||
self._callback_error: BaseException | None = None
|
||||
self._logical: _LogicalCall | None = None
|
||||
self._defer_logical_completion = defer_logical_completion
|
||||
# Only auxiliary calls report model/provider on their logical scope.
|
||||
auxiliary = str((metadata or {}).get("call_role") or "").startswith("auxiliary:")
|
||||
self._logical_model_name: str | None = model_name if auxiliary else None
|
||||
self._logical_provider_name: str | None = name if auxiliary else None
|
||||
self._logical_response_model_name: str | None = None
|
||||
self._on_chunk = on_chunk
|
||||
self._chunk_adapter = chunk_adapter or _namespace
|
||||
self._accept_chunk = accept_chunk
|
||||
self._relay_observes_chunks = False
|
||||
self._provider_completed = False
|
||||
self._raw_chunks: list[tuple[Any, Any]] = []
|
||||
self._prefetched_chunks: list[Any] = []
|
||||
self.output_modified = False
|
||||
attempt = _ManagedAttempt.resolve(
|
||||
session_id, request, metadata, name=name, model_name=model_name
|
||||
)
|
||||
@@ -474,9 +364,7 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
)
|
||||
|
||||
def _start_unmanaged(
|
||||
self,
|
||||
request: dict[str, Any],
|
||||
stream_factory: Callable[[dict[str, Any]], Any],
|
||||
self, request: dict[str, Any], stream_factory: Callable[[dict[str, Any]], Any],
|
||||
on_stream_created: Callable[[Any], None] | None,
|
||||
completed_response_predicate: Callable[[Any], bool] | None,
|
||||
) -> None:
|
||||
@@ -491,12 +379,9 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
self._stream = iter(raw_stream)
|
||||
|
||||
def _start_managed(
|
||||
self,
|
||||
attempt: _ManagedAttempt,
|
||||
stream_factory: Callable[[dict[str, Any]], Any],
|
||||
self, attempt: _ManagedAttempt, stream_factory: Callable[[dict[str, Any]], Any],
|
||||
on_stream_created: Callable[[Any], None] | None,
|
||||
completed_response_predicate: Callable[[Any], bool] | None,
|
||||
finalizer: Callable[[], Any],
|
||||
completed_response_predicate: Callable[[Any], bool] | None, finalizer: Callable[[], Any],
|
||||
) -> None:
|
||||
"""Open Relay's stream on a private event loop owned by this iterator."""
|
||||
run_callback = attempt.run_callback
|
||||
@@ -571,9 +456,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,
|
||||
attempt.runtime.relay.llm.stream_execute, provider_stream, observe_chunk,
|
||||
relay_finalizer,
|
||||
)
|
||||
)
|
||||
@@ -603,25 +486,23 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
|
||||
def _recoverable_relay_failure(self, exc: BaseException) -> bool:
|
||||
"""Relay post-processing failed after the provider already succeeded."""
|
||||
if (
|
||||
recoverable = (
|
||||
isinstance(exc, Exception) and self._provider_completed and self._callback_error is None
|
||||
):
|
||||
)
|
||||
if recoverable:
|
||||
logger.warning(
|
||||
"NeMo Relay stream post-processing failed after provider success; "
|
||||
"preserving the provider result",
|
||||
exc_info=True,
|
||||
)
|
||||
return True
|
||||
return False
|
||||
return recoverable
|
||||
|
||||
def _finish_logical(self, outcome: str) -> None:
|
||||
"""Complete the logical LLM scope unless the caller deferred it."""
|
||||
if self._defer_logical_completion:
|
||||
return
|
||||
_complete_logical(
|
||||
self._logical,
|
||||
outcome=outcome,
|
||||
model_name=self._logical_model_name,
|
||||
self._logical, outcome=outcome, model_name=self._logical_model_name,
|
||||
provider_name=self._logical_provider_name,
|
||||
response_model_name=self._logical_response_model_name,
|
||||
operation_lease=self._runtime_lease,
|
||||
@@ -712,14 +593,10 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
|
||||
def _close_provider_resources(self) -> None:
|
||||
"""Close the unmanaged provider stream/resource once each (they may be the same object)."""
|
||||
resources = (self._stream, self._raw_stream_resource)
|
||||
resources = {id(r): r for r in (self._stream, self._raw_stream_resource) if r is not None}
|
||||
self._stream = None
|
||||
self._raw_stream_resource = None
|
||||
closed_ids: set[int] = set()
|
||||
for resource in resources:
|
||||
if resource is None or id(resource) in closed_ids:
|
||||
continue
|
||||
closed_ids.add(id(resource))
|
||||
for resource in resources.values():
|
||||
close = getattr(resource, "close", None)
|
||||
if not callable(close):
|
||||
continue
|
||||
@@ -762,6 +639,9 @@ class ManagedLlmStream(Iterator[Any]):
|
||||
self._close(logical_outcome="cancelled")
|
||||
|
||||
|
||||
stream = ManagedLlmStream
|
||||
|
||||
|
||||
_ANTHROPIC_APPEND_DELTAS = {
|
||||
"text_delta": "text", "thinking_delta": "thinking", "signature_delta": "signature"
|
||||
}
|
||||
@@ -826,23 +706,19 @@ class AnthropicStreamAccumulator:
|
||||
self._message["usage"] = usage
|
||||
|
||||
_EVENT_HANDLERS = {
|
||||
"message_start": _on_message_start,
|
||||
"content_block_start": _on_content_block_start,
|
||||
"content_block_delta": _on_content_block_delta,
|
||||
"message_delta": _on_message_delta,
|
||||
"message_start": _on_message_start, "content_block_start": _on_content_block_start,
|
||||
"content_block_delta": _on_content_block_delta, "message_delta": _on_message_delta,
|
||||
}
|
||||
|
||||
def finalize(self) -> dict[str, Any]:
|
||||
blocks = []
|
||||
for index in sorted(self._blocks):
|
||||
block = dict(self._blocks[index])
|
||||
blocks = [dict(self._blocks[index]) for index in sorted(self._blocks)]
|
||||
for block in blocks:
|
||||
partial = block.pop("_partial_json", None)
|
||||
if partial is not None:
|
||||
try:
|
||||
block["input"] = json.loads(partial)
|
||||
except (TypeError, ValueError):
|
||||
block["input"] = partial
|
||||
blocks.append(block)
|
||||
return {**self._message, "content": blocks}
|
||||
|
||||
def response(self, base: Any = None) -> Any:
|
||||
@@ -887,12 +763,8 @@ def _logical_parent(
|
||||
|
||||
|
||||
def _complete_logical(
|
||||
logical: _LogicalCall | None,
|
||||
*,
|
||||
outcome: str,
|
||||
model_name: str | None = None,
|
||||
provider_name: str | None = None,
|
||||
response_model_name: str | None = None,
|
||||
logical: _LogicalCall | None, *, outcome: str, model_name: str | None = None,
|
||||
provider_name: str | None = None, response_model_name: str | None = None,
|
||||
operation_lease: relay_runtime.RelayOperationLease | None = None,
|
||||
) -> None:
|
||||
if logical is None:
|
||||
@@ -917,12 +789,8 @@ def _complete_logical(
|
||||
if operation_lease is not None:
|
||||
callback = operation_lease.run_in_session
|
||||
callback(
|
||||
lease.session,
|
||||
relay_runtime.pop_relay_scope,
|
||||
lease.host.relay,
|
||||
handle,
|
||||
output=output,
|
||||
metadata=relay_runtime.runtime_metadata(lease.host.runtime_id),
|
||||
lease.session, relay_runtime.pop_relay_scope, lease.host.relay, handle,
|
||||
output=output, metadata=relay_runtime.runtime_metadata(lease.host.runtime_id),
|
||||
)
|
||||
except Exception:
|
||||
# The provider result is authoritative. Retain the handle so turn
|
||||
@@ -939,12 +807,8 @@ def _is_cancellation(error: BaseException) -> bool:
|
||||
|
||||
|
||||
def complete_logical_call(
|
||||
api_request_id: str,
|
||||
*,
|
||||
outcome: str,
|
||||
model_name: str | None = None,
|
||||
provider_name: str | None = None,
|
||||
response_model_name: str | None = None,
|
||||
api_request_id: str, *, outcome: str, model_name: str | None = None,
|
||||
provider_name: str | None = None, response_model_name: str | None = None,
|
||||
) -> None:
|
||||
"""Complete the active turn's logical LLM call after caller validation."""
|
||||
turn = relay_runtime.active_turn()
|
||||
@@ -954,30 +818,20 @@ def complete_logical_call(
|
||||
handle = turn.logical_llm_calls.get(api_request_id)
|
||||
if handle is not None:
|
||||
_complete_logical(
|
||||
(turn, handle, api_request_id),
|
||||
outcome=outcome,
|
||||
model_name=model_name,
|
||||
provider_name=provider_name,
|
||||
response_model_name=response_model_name,
|
||||
(turn, handle, api_request_id), outcome=outcome, model_name=model_name,
|
||||
provider_name=provider_name, response_model_name=response_model_name,
|
||||
)
|
||||
|
||||
|
||||
def _response_model_name(response: Any) -> str | None:
|
||||
"""Return a provider-reported model name when one is available."""
|
||||
if isinstance(response, dict):
|
||||
value = response.get("model")
|
||||
else:
|
||||
value = getattr(response, "model", None)
|
||||
value = response.get("model") if isinstance(response, dict) else getattr(response, "model", None)
|
||||
return value if isinstance(value, str) and value.strip() else None
|
||||
|
||||
|
||||
def _provider_request(
|
||||
original: dict[str, Any],
|
||||
request: Any,
|
||||
*,
|
||||
relay_request_body: dict[str, Any],
|
||||
codec_baseline_body: dict[str, Any] | None,
|
||||
metadata: dict[str, Any] | None,
|
||||
original: dict[str, Any], request: Any, *, relay_request_body: dict[str, Any],
|
||||
codec_baseline_body: dict[str, Any] | None, metadata: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
content = getattr(request, "content", request)
|
||||
if not isinstance(content, dict):
|
||||
@@ -1000,56 +854,62 @@ def _provider_request(
|
||||
headers = getattr(request, "headers", None)
|
||||
if isinstance(headers, dict):
|
||||
headers = {
|
||||
key: value
|
||||
for key, value in headers.items()
|
||||
if str(key).lower() not in _RELAY_INTERNAL_PROVIDER_HEADERS
|
||||
key: value for key,
|
||||
value in headers.items() if str(key).lower() not in _RELAY_INTERNAL_PROVIDER_HEADERS
|
||||
}
|
||||
if headers:
|
||||
final["extra_headers"] = {**dict(final.get("extra_headers") or {}), **headers}
|
||||
return final
|
||||
|
||||
|
||||
def _codex_codec_tools(body: dict[str, Any]) -> None:
|
||||
# The Responses SDK accepts ``tools=None`` as "no tools" while Relay's
|
||||
# typed codec expects an array or an absent field; normalize only the
|
||||
# codec-facing copy (the original request is restored when unchanged).
|
||||
if body.get("tools") is 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"},
|
||||
}
|
||||
if isinstance(tool, dict) and tool.get("type") == "function" and "function" not in tool
|
||||
else tool
|
||||
for tool in body["tools"]
|
||||
]
|
||||
|
||||
|
||||
def _chat_codec_tools(body: dict[str, Any]) -> None:
|
||||
tools = body.get("tools")
|
||||
if isinstance(tools, list):
|
||||
body["tools"] = [
|
||||
{"type": "function", **tool}
|
||||
if isinstance(tool, dict) and "function" in tool and "type" not in tool
|
||||
else tool
|
||||
for tool in tools
|
||||
]
|
||||
|
||||
|
||||
# api_mode -> in-place normalizer producing the codec-facing ``tools`` shape.
|
||||
_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]:
|
||||
body = _jsonable_dict(request)
|
||||
# ``timeout`` configures the provider SDK client, not a wire protocol:
|
||||
# keep it on the original callback request, never on Relay intercepts.
|
||||
body.pop("timeout", None)
|
||||
api_mode = _api_mode(metadata)
|
||||
if api_mode == "codex_responses":
|
||||
# The Responses SDK accepts ``tools=None`` as "no tools" while Relay's
|
||||
# typed codec expects an array or an absent field; normalize only the
|
||||
# codec-facing copy (the original request is restored when unchanged).
|
||||
if body.get("tools") is 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"},
|
||||
}
|
||||
if isinstance(tool, dict)
|
||||
and tool.get("type") == "function"
|
||||
and "function" not in tool
|
||||
else tool
|
||||
for tool in body["tools"]
|
||||
]
|
||||
elif api_mode == "chat_completions":
|
||||
tools = body.get("tools")
|
||||
if isinstance(tools, list):
|
||||
body["tools"] = [
|
||||
{"type": "function", **tool}
|
||||
if isinstance(tool, dict) and "function" in tool and "type" not in tool
|
||||
else tool
|
||||
for tool in tools
|
||||
]
|
||||
normalize = _CODEC_TOOL_NORMALIZERS.get(_api_mode(metadata))
|
||||
if normalize is not None:
|
||||
normalize(body)
|
||||
return body
|
||||
|
||||
|
||||
def _restore_provider_message_extensions(
|
||||
original: dict[str, Any],
|
||||
final: dict[str, Any],
|
||||
*,
|
||||
baseline: 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."""
|
||||
@@ -1073,10 +933,7 @@ def _restore_provider_message_extensions(
|
||||
|
||||
|
||||
def _codec_round_trip_request_body(
|
||||
relay: Any,
|
||||
relay_request: Any,
|
||||
*,
|
||||
relay_request_body: dict[str, Any],
|
||||
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."""
|
||||
@@ -1121,11 +978,11 @@ def _provider_request_body(
|
||||
|
||||
|
||||
def _codec(relay: Any, metadata: dict[str, Any] | None) -> Any:
|
||||
protocol = _relay_protocol(metadata)
|
||||
protocol = _RELAY_PROTOCOL_BY_API_MODE.get(_api_mode(metadata))
|
||||
codecs = getattr(relay, "codecs", None)
|
||||
if protocol is None or codecs is None:
|
||||
return None
|
||||
codec = getattr(codecs, protocol.codec_class, None)
|
||||
codec = getattr(codecs, protocol[1], None)
|
||||
return codec() if callable(codec) else None
|
||||
|
||||
|
||||
@@ -1171,20 +1028,24 @@ def _namespace(value: Any) -> Any:
|
||||
return value
|
||||
|
||||
|
||||
def _canonical_json(value: Any) -> str:
|
||||
return json.dumps(_jsonable(value), sort_keys=True, separators=(",", ":"))
|
||||
|
||||
|
||||
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 _canonical_json(left) == _canonical_json(right)
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
|
||||
|
||||
def _run_awaitable(value: Any) -> Any:
|
||||
def _run_awaitable(
|
||||
value: Any,
|
||||
*,
|
||||
loop_error: str = "Synchronous Relay LLM execution cannot run on an event-loop thread",
|
||||
) -> Any:
|
||||
if not inspect.isawaitable(value):
|
||||
return value
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return asyncio.run(value)
|
||||
raise RuntimeError("Synchronous Relay LLM execution cannot run on an event-loop thread")
|
||||
if _has_running_event_loop():
|
||||
raise RuntimeError(loop_error)
|
||||
return asyncio.run(value)
|
||||
|
||||
+124
-197
@@ -138,11 +138,8 @@ def _same_handle(a: Any, b: Any) -> bool:
|
||||
# Native ScopeHandle has no value __eq__; compare by uuid when both expose one.
|
||||
if a is None or b is None:
|
||||
return a is b
|
||||
if a is b or a == b:
|
||||
return True
|
||||
a_uuid = getattr(a, "uuid", None)
|
||||
b_uuid = getattr(b, "uuid", None)
|
||||
return a_uuid is not None and a_uuid == b_uuid
|
||||
return a is b or a == b or (a_uuid is not None and a_uuid == getattr(b, "uuid", None))
|
||||
|
||||
|
||||
class _RelayPluginConfigurationState(Enum):
|
||||
@@ -188,21 +185,19 @@ _SEGMENTS_CONFIG_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _load_segments_config() -> dict[str, Any]:
|
||||
on_compaction = False
|
||||
max_turns = 0
|
||||
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 {}
|
||||
on_compaction = bool(segments.get("on_compaction", False))
|
||||
try:
|
||||
max_turns = max(0, int(segments.get("max_turns", 0) or 0))
|
||||
except (TypeError, ValueError):
|
||||
max_turns = 0
|
||||
except Exception: # noqa: BLE001 - config absence must not crash
|
||||
pass
|
||||
return {"on_compaction": on_compaction, "max_turns": max_turns}
|
||||
try:
|
||||
max_turns = max(0, int(segments.get("max_turns", 0) or 0))
|
||||
except (TypeError, ValueError):
|
||||
max_turns = 0
|
||||
return {"on_compaction": bool(segments.get("on_compaction", False)), "max_turns": max_turns}
|
||||
|
||||
|
||||
def _segments_config() -> dict[str, Any]:
|
||||
@@ -266,47 +261,55 @@ class _ProcessRelayPluginConfiguration:
|
||||
if self._owners:
|
||||
self._owners.add(owner_id)
|
||||
return self._state
|
||||
if self._active and not self._clear_active():
|
||||
logger.warning(
|
||||
"Hermes Relay plugin cleanup is still pending; refusing to "
|
||||
"replace the process-global configuration"
|
||||
state = self._preflight(relay)
|
||||
if state is None:
|
||||
state = self._activate(relay)
|
||||
state = self._remember(owner_id, state)
|
||||
if state is _RelayPluginConfigurationState.ACTIVE:
|
||||
logger.info(
|
||||
"Relay plugins are active process-wide and apply to all profiles "
|
||||
"hosted by this Hermes process."
|
||||
)
|
||||
return self._remember(owner_id, _RelayPluginConfigurationState.FAILED)
|
||||
|
||||
try:
|
||||
existing_report = relay.plugin.report()
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Hermes could not determine whether a process-global Relay "
|
||||
"plugin configuration is already active; refusing to replace it",
|
||||
exc_info=True,
|
||||
)
|
||||
return self._remember(owner_id, _RelayPluginConfigurationState.FAILED)
|
||||
if existing_report is not None:
|
||||
logger.warning(
|
||||
"A process-global Relay plugin configuration is already active "
|
||||
"outside Hermes native ownership; leaving it unchanged and "
|
||||
"disabling Hermes-managed Relay middleware for this process"
|
||||
)
|
||||
return self._remember(owner_id, _RelayPluginConfigurationState.FOREIGN)
|
||||
|
||||
try:
|
||||
if not self._initialize(relay):
|
||||
return self._remember(owner_id, _RelayPluginConfigurationState.DISABLED)
|
||||
except Exception as exc:
|
||||
self._activation = None
|
||||
logger.warning("Hermes Relay plugin initialization failed: %s", exc, exc_info=True)
|
||||
return self._remember(owner_id, _RelayPluginConfigurationState.FAILED)
|
||||
|
||||
self._active = True
|
||||
self._relay = relay
|
||||
state = self._remember(owner_id, _RelayPluginConfigurationState.ACTIVE)
|
||||
logger.info(
|
||||
"Relay plugins are active process-wide and apply to all profiles "
|
||||
"hosted by this Hermes process."
|
||||
)
|
||||
return state
|
||||
|
||||
def _activate(self, relay: Any) -> _RelayPluginConfigurationState:
|
||||
try:
|
||||
if not self._initialize(relay):
|
||||
return _RelayPluginConfigurationState.DISABLED
|
||||
except Exception as exc:
|
||||
self._activation = None
|
||||
logger.warning("Hermes Relay plugin initialization failed: %s", exc, exc_info=True)
|
||||
return _RelayPluginConfigurationState.FAILED
|
||||
self._active = True
|
||||
self._relay = relay
|
||||
return _RelayPluginConfigurationState.ACTIVE
|
||||
|
||||
def _preflight(self, relay: Any) -> _RelayPluginConfigurationState | None:
|
||||
"""Return a terminal state when the process cannot take ownership; None to proceed."""
|
||||
if self._active and not self._clear_active():
|
||||
logger.warning(
|
||||
"Hermes Relay plugin cleanup is still pending; refusing to "
|
||||
"replace the process-global configuration"
|
||||
)
|
||||
return _RelayPluginConfigurationState.FAILED
|
||||
try:
|
||||
existing_report = relay.plugin.report()
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Hermes could not determine whether a process-global Relay "
|
||||
"plugin configuration is already active; refusing to replace it",
|
||||
exc_info=True,
|
||||
)
|
||||
return _RelayPluginConfigurationState.FAILED
|
||||
if existing_report is not None:
|
||||
logger.warning(
|
||||
"A process-global Relay plugin configuration is already active "
|
||||
"outside Hermes native ownership; leaving it unchanged and "
|
||||
"disabling Hermes-managed Relay middleware for this process"
|
||||
)
|
||||
return _RelayPluginConfigurationState.FOREIGN
|
||||
return None
|
||||
|
||||
def _initialize(self, relay: Any) -> bool:
|
||||
"""Initialize Relay from the selected plugins.toml; False when none is selected."""
|
||||
configured_inputs = _configured_plugin_inputs(relay)
|
||||
@@ -345,13 +348,11 @@ class _ProcessRelayPluginConfiguration:
|
||||
|
||||
def release(self, owner: Any) -> None:
|
||||
"""Release one host and clear Relay after the final host exits."""
|
||||
owner_id = id(owner)
|
||||
with self._lock:
|
||||
if owner_id not in self._owners:
|
||||
if id(owner) not in self._owners:
|
||||
return
|
||||
self._owners.remove(owner_id)
|
||||
if not self._owners:
|
||||
self._reset_if_cleared()
|
||||
self._owners.remove(id(owner))
|
||||
self.retry_pending_cleanup()
|
||||
|
||||
def reset_for_tests(self) -> None:
|
||||
"""Clear process-global state left by directly constructed test hosts."""
|
||||
@@ -416,10 +417,10 @@ class RelayRuntime:
|
||||
self._execution_consumers_lock = threading.RLock()
|
||||
self._execution_consumers: set[str] = set()
|
||||
self._plugin_configuration_state = _PLUGIN_CONFIGURATION.acquire(self, self.relay)
|
||||
# Cleared (with the atexit hook) by the first successful _finish_shutdown.
|
||||
self._plugin_configuration_registered = True
|
||||
if self._plugins_active():
|
||||
self.retain_managed_execution(RELAY_PLUGINS_EXECUTION_CONSUMER)
|
||||
self._shutdown_registered = True
|
||||
atexit.register(self.shutdown)
|
||||
|
||||
def _plugins_active(self) -> bool:
|
||||
@@ -467,11 +468,7 @@ class RelayRuntime:
|
||||
return context.run(*args, input={}, **push_kwargs)
|
||||
|
||||
def _open_session_scope(
|
||||
self,
|
||||
session: RelaySession,
|
||||
scope_metadata: dict[str, Any],
|
||||
*,
|
||||
resolve_parent: bool,
|
||||
self, session: RelaySession, scope_metadata: dict[str, Any], *, resolve_parent: bool,
|
||||
**push_kwargs: Any,
|
||||
) -> None:
|
||||
"""Push a fresh session scope for ``session`` and record its handle + context.
|
||||
@@ -516,11 +513,8 @@ class RelayRuntime:
|
||||
if session.handle is None:
|
||||
try:
|
||||
self._open_session_scope(
|
||||
session,
|
||||
{**(metadata or {}), **runtime_metadata(self.runtime_id)},
|
||||
resolve_parent=True,
|
||||
data=data,
|
||||
exit_fallback=True,
|
||||
session, {**(metadata or {}), **runtime_metadata(self.runtime_id)},
|
||||
resolve_parent=True, data=data, exit_fallback=True,
|
||||
)
|
||||
except Exception:
|
||||
session.context = None
|
||||
@@ -547,12 +541,9 @@ class RelayRuntime:
|
||||
session.rotate_pending = False
|
||||
try:
|
||||
self.run_in_session(
|
||||
session,
|
||||
self.relay.scope.pop,
|
||||
old_handle,
|
||||
session, self.relay.scope.pop, old_handle,
|
||||
output={"hermes.session.segment_reason": reason},
|
||||
metadata=runtime_metadata(self.runtime_id),
|
||||
timeout=_SCOPE_OP_TIMEOUT,
|
||||
metadata=runtime_metadata(self.runtime_id), timeout=_SCOPE_OP_TIMEOUT,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
@@ -625,10 +616,10 @@ class RelayRuntime:
|
||||
"""Return an active Hermes Relay session without creating one."""
|
||||
with self._sessions_lock:
|
||||
session = None if self._closing else self._sessions.get(str(session_id or ""))
|
||||
if session is None:
|
||||
return None
|
||||
with session.lock:
|
||||
return None if session.closing else session
|
||||
if session is not None:
|
||||
with session.lock:
|
||||
return None if session.closing else session
|
||||
return None
|
||||
|
||||
def _session_context(
|
||||
self, session: RelaySession, *, allow_closing: bool
|
||||
@@ -648,13 +639,8 @@ class RelayRuntime:
|
||||
return context
|
||||
|
||||
def run_in_session(
|
||||
self,
|
||||
session: RelaySession,
|
||||
callback: Callable[..., Any],
|
||||
*args: Any,
|
||||
allow_closing: bool = False,
|
||||
timeout: float | None = None,
|
||||
**kwargs: Any,
|
||||
self, session: RelaySession, callback: Callable[..., Any], *args: Any,
|
||||
allow_closing: bool = False, timeout: float | None = None, **kwargs: Any,
|
||||
) -> Any:
|
||||
"""Run a Relay operation against a session's isolated scope stack.
|
||||
|
||||
@@ -674,13 +660,8 @@ class RelayRuntime:
|
||||
self._end_operation()
|
||||
|
||||
def _run_in_session_untracked(
|
||||
self,
|
||||
session: RelaySession,
|
||||
callback: Callable[..., Any],
|
||||
*args: Any,
|
||||
allow_closing: bool = False,
|
||||
timeout: float | None = None,
|
||||
**kwargs: Any,
|
||||
self, session: RelaySession, callback: Callable[..., Any], *args: Any,
|
||||
allow_closing: bool = False, timeout: float | None = None, **kwargs: Any,
|
||||
) -> Any:
|
||||
"""Run inside a session whose host-level lifetime is already held."""
|
||||
context = self._session_context(session, allow_closing=allow_closing)
|
||||
@@ -716,12 +697,8 @@ class RelayRuntime:
|
||||
) from exc
|
||||
|
||||
async def run_in_session_async(
|
||||
self,
|
||||
session: RelaySession,
|
||||
callback: Callable[..., Any],
|
||||
*args: Any,
|
||||
allow_closing: bool = False,
|
||||
**kwargs: Any,
|
||||
self, session: RelaySession, callback: Callable[..., Any], *args: Any,
|
||||
allow_closing: bool = False, **kwargs: Any,
|
||||
) -> Any:
|
||||
"""Create and await an operation inside the session's saved context."""
|
||||
self._begin_operation()
|
||||
@@ -767,11 +744,7 @@ class RelayRuntime:
|
||||
if session is None:
|
||||
return False
|
||||
self.run_in_session(
|
||||
session,
|
||||
self.relay.scope.event,
|
||||
name,
|
||||
handle=session.handle,
|
||||
data=data,
|
||||
session, self.relay.scope.event, name, handle=session.handle, data=data,
|
||||
metadata=metadata,
|
||||
)
|
||||
return True
|
||||
@@ -792,12 +765,7 @@ class RelayRuntime:
|
||||
return result if isinstance(result, dict) else args
|
||||
|
||||
def _pop_with_drain(
|
||||
self,
|
||||
handle: Any,
|
||||
*,
|
||||
output: dict[str, Any],
|
||||
metadata: dict[str, Any],
|
||||
session_root: Any,
|
||||
self, handle: Any, *, output: dict[str, Any], metadata: dict[str, Any], session_root: Any,
|
||||
drain_limit: int,
|
||||
) -> BaseException | None:
|
||||
"""Pop ``handle``; if that fails, drain orphans above it and retry once.
|
||||
@@ -824,9 +792,7 @@ class RelayRuntime:
|
||||
break
|
||||
try:
|
||||
pop_relay_scope(
|
||||
self.relay,
|
||||
top,
|
||||
output={"outcome": "cancelled", "hermes.orphan_drain": True},
|
||||
self.relay, top, output={"outcome": "cancelled", "hermes.orphan_drain": True},
|
||||
metadata=metadata,
|
||||
)
|
||||
drained += 1
|
||||
@@ -844,15 +810,9 @@ class RelayRuntime:
|
||||
return retry_exc
|
||||
|
||||
def _close_scope_handle(
|
||||
self,
|
||||
session: RelaySession,
|
||||
handle: Any,
|
||||
*,
|
||||
output: dict[str, Any] | None = None,
|
||||
allow_closing: bool = False,
|
||||
failure_label: str = "scope close failed",
|
||||
drain_limit: int = 32,
|
||||
operation_already_held: bool = False,
|
||||
self, session: RelaySession, handle: Any, *, output: dict[str, Any] | None = None,
|
||||
allow_closing: bool = False, failure_label: str = "scope close failed",
|
||||
drain_limit: int = 32, operation_already_held: bool = False,
|
||||
) -> str | None:
|
||||
"""Pop ``handle``, draining orphaned children in the same session context.
|
||||
|
||||
@@ -868,18 +828,12 @@ class RelayRuntime:
|
||||
)
|
||||
try:
|
||||
failure = run_in_session(
|
||||
session,
|
||||
self._pop_with_drain,
|
||||
handle,
|
||||
output=output or {},
|
||||
metadata=runtime_metadata(self.runtime_id),
|
||||
session_root=session.handle,
|
||||
drain_limit=drain_limit,
|
||||
allow_closing=allow_closing,
|
||||
timeout=_SCOPE_OP_TIMEOUT,
|
||||
session, self._pop_with_drain, handle, output=output or {},
|
||||
metadata=runtime_metadata(self.runtime_id), session_root=session.handle,
|
||||
drain_limit=drain_limit, allow_closing=allow_closing, timeout=_SCOPE_OP_TIMEOUT,
|
||||
)
|
||||
except Exception as exc:
|
||||
return f"{failure_label}: {exc}"
|
||||
failure = exc
|
||||
return None if failure is None else f"{failure_label}: {failure}"
|
||||
|
||||
def close_session(self, event: dict[str, Any]) -> None:
|
||||
@@ -908,12 +862,8 @@ class RelayRuntime:
|
||||
session.closing = True
|
||||
if session.handle is not None:
|
||||
failure = self._close_scope_handle(
|
||||
session,
|
||||
session.handle,
|
||||
output={},
|
||||
allow_closing=True,
|
||||
failure_label="session scope close failed",
|
||||
operation_already_held=True,
|
||||
session, session.handle, output={}, allow_closing=True,
|
||||
failure_label="session scope close failed", operation_already_held=True,
|
||||
)
|
||||
# Subscriber flushing is process-wide and may wait for publications
|
||||
# owned by other sessions; final plugin teardown flushes once after all
|
||||
@@ -936,8 +886,7 @@ class RelayRuntime:
|
||||
if has_active_operations:
|
||||
thread = threading.Thread(
|
||||
target=self._finish_shutdown_after_operations,
|
||||
name=f"hermes-nemo-relay-shutdown-{self.runtime_id[:8]}",
|
||||
daemon=True,
|
||||
name=f"hermes-nemo-relay-shutdown-{self.runtime_id[:8]}", daemon=True,
|
||||
)
|
||||
try:
|
||||
thread.start()
|
||||
@@ -963,9 +912,7 @@ class RelayRuntime:
|
||||
self.release_managed_execution(RELAY_PLUGINS_EXECUTION_CONSUMER)
|
||||
_PLUGIN_CONFIGURATION.release(self)
|
||||
self._plugin_configuration_registered = False
|
||||
if self._shutdown_registered:
|
||||
self._safe(atexit.unregister, self.shutdown, quiet=True)
|
||||
self._shutdown_registered = False
|
||||
except Exception:
|
||||
with self._sessions_lock:
|
||||
self._shutdown_started = False
|
||||
@@ -1119,6 +1066,13 @@ class managed_callback_guard:
|
||||
_MANAGED_CALLBACK_DEPTH.reset(self._token)
|
||||
|
||||
|
||||
def _flag_open_session(session: RelaySession, flag: str) -> None:
|
||||
"""Set a pending-rotation/close flag unless the session is already closing."""
|
||||
with session.lock:
|
||||
if not session.closing:
|
||||
setattr(session, flag, True)
|
||||
|
||||
|
||||
class RelaySessionCoordinator:
|
||||
"""Own semantic conversation and turn lifetimes for Hermes core."""
|
||||
|
||||
@@ -1146,26 +1100,18 @@ class RelaySessionCoordinator:
|
||||
logger.warning("Hermes Relay session initializer failed: %s", name, exc_info=True)
|
||||
|
||||
def acquire_conversation(
|
||||
self,
|
||||
*,
|
||||
profile_key: str,
|
||||
session_id: str,
|
||||
platform: str,
|
||||
parent_session_id: str = "",
|
||||
self, *, profile_key: str, session_id: str, platform: str, parent_session_id: str = "",
|
||||
model: str = "",
|
||||
) -> ConversationLease:
|
||||
host = self.registry.for_profile(profile_key)
|
||||
if host is None:
|
||||
host = NoopRelayRuntime(profile_key, "Relay host creation was disabled")
|
||||
host = self.registry.for_profile(profile_key) or NoopRelayRuntime(
|
||||
profile_key, "Relay host creation was disabled"
|
||||
)
|
||||
session = None
|
||||
if isinstance(host, RelayRuntime):
|
||||
try:
|
||||
self._prepare_session(host, {
|
||||
"profile_key": profile_key,
|
||||
"session_id": session_id,
|
||||
"platform": platform,
|
||||
"parent_session_id": parent_session_id,
|
||||
"model": model,
|
||||
"profile_key": profile_key, "session_id": session_id, "platform": platform,
|
||||
"parent_session_id": parent_session_id, "model": model,
|
||||
})
|
||||
metadata = {"hermes.execution_surface": platform or "unknown"}
|
||||
if parent_session_id and parent_session_id != session_id:
|
||||
@@ -1178,12 +1124,8 @@ class RelaySessionCoordinator:
|
||||
except Exception:
|
||||
logger.warning("Hermes Relay conversation initialization failed", exc_info=True)
|
||||
return ConversationLease(
|
||||
profile_key=profile_key,
|
||||
session_id=session_id,
|
||||
platform=platform,
|
||||
host=host,
|
||||
session=session,
|
||||
parent_session_id=parent_session_id,
|
||||
profile_key=profile_key, session_id=session_id, platform=platform, host=host,
|
||||
session=session, parent_session_id=parent_session_id,
|
||||
)
|
||||
|
||||
def begin_turn(
|
||||
@@ -1199,10 +1141,8 @@ class RelaySessionCoordinator:
|
||||
# would create sibling scopes whose completion order is not LIFO.
|
||||
turn.relay_enabled = False
|
||||
logger.warning(
|
||||
"Skipping Relay instrumentation for concurrent Hermes turn "
|
||||
"%s in session %s",
|
||||
turn_id,
|
||||
lease.session_id,
|
||||
"Skipping Relay instrumentation for concurrent Hermes turn " "%s in session %s",
|
||||
turn_id, lease.session_id,
|
||||
)
|
||||
else:
|
||||
self._active_turns[key] = {id(turn)}
|
||||
@@ -1286,9 +1226,7 @@ class RelaySessionCoordinator:
|
||||
if turn.handle is None:
|
||||
return
|
||||
failure = host._close_scope_handle(
|
||||
turn.lease.session,
|
||||
turn.handle,
|
||||
output={"outcome": outcome},
|
||||
turn.lease.session, turn.handle, output={"outcome": outcome},
|
||||
failure_label="turn scope close failed",
|
||||
)
|
||||
if failure:
|
||||
@@ -1307,14 +1245,12 @@ class RelaySessionCoordinator:
|
||||
host = lease.live_runtime()
|
||||
if host is None:
|
||||
return
|
||||
session = lease.session
|
||||
with session.lock:
|
||||
pending = session.close_pending and not session.closing
|
||||
if not pending:
|
||||
return
|
||||
if self.has_active_turn(profile_key=lease.profile_key, session_id=lease.session_id):
|
||||
return
|
||||
host.close_session({"session_id": lease.session_id})
|
||||
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
|
||||
):
|
||||
host.close_session({"session_id": lease.session_id})
|
||||
except Exception: # noqa: BLE001 - telemetry must never block end_turn
|
||||
logger.warning("Hermes Relay deferred session close failed", exc_info=True)
|
||||
|
||||
@@ -1345,19 +1281,14 @@ class RelaySessionCoordinator:
|
||||
if old_session is not None and self.has_active_turn(
|
||||
profile_key=profile_key, session_id=old_session_id
|
||||
):
|
||||
with old_session.lock:
|
||||
if not old_session.closing:
|
||||
old_session.close_pending = True
|
||||
return
|
||||
host.close_session({"session_id": old_session_id})
|
||||
_flag_open_session(old_session, "close_pending")
|
||||
else:
|
||||
host.close_session({"session_id": old_session_id})
|
||||
return
|
||||
with host._sessions_lock:
|
||||
session = host._sessions.get(session_id)
|
||||
if session is None:
|
||||
return
|
||||
with session.lock:
|
||||
if not session.closing:
|
||||
session.rotate_pending = True
|
||||
if session is not None:
|
||||
_flag_open_session(session, "rotate_pending")
|
||||
except Exception: # noqa: BLE001 - telemetry must never block compaction
|
||||
logger.warning("Hermes Relay compaction notification failed", exc_info=True)
|
||||
|
||||
@@ -1394,12 +1325,9 @@ class RelaySessionCoordinator:
|
||||
with turn.logical_llm_lock:
|
||||
logical_calls = list(turn.logical_llm_calls.items())
|
||||
turn.logical_llm_calls.clear()
|
||||
for index in range(len(logical_calls) - 1, -1, -1):
|
||||
request_id, logical_handle = logical_calls[index]
|
||||
for index, (request_id, logical_handle) in reversed(list(enumerate(logical_calls))):
|
||||
failure = host._close_scope_handle(
|
||||
lease.session,
|
||||
logical_handle,
|
||||
output={"outcome": outcome},
|
||||
lease.session, logical_handle, output={"outcome": outcome},
|
||||
failure_label="logical LLM scope close failed",
|
||||
)
|
||||
if failure is None:
|
||||
@@ -1458,15 +1386,15 @@ def active_turn(session_id: str | None = None) -> RelayTurnContext | None:
|
||||
turn = current_turn()
|
||||
if turn is None or not turn.relay_enabled or turn.closed or turn.lease.released:
|
||||
return None
|
||||
if turn.lease.profile_key != current_profile_key():
|
||||
lease = turn.lease
|
||||
if lease.profile_key != current_profile_key():
|
||||
return None
|
||||
if session_id is not None and turn.lease.session_id != session_id:
|
||||
if session_id is not None and lease.session_id != session_id:
|
||||
return None
|
||||
if isinstance(lease.host, RelayRuntime) and (
|
||||
lease.session is None or lease.host.get_session(lease.session_id) is not lease.session
|
||||
):
|
||||
return None
|
||||
if isinstance(turn.lease.host, RelayRuntime):
|
||||
if turn.lease.session is None:
|
||||
return None
|
||||
if turn.lease.host.get_session(turn.lease.session_id) is not turn.lease.session:
|
||||
return None
|
||||
return turn
|
||||
|
||||
|
||||
@@ -1535,8 +1463,7 @@ def _is_relay_wrapped_callback_error(
|
||||
return False
|
||||
callback_type = callback_error.__class__
|
||||
type_names = {
|
||||
callback_type.__name__,
|
||||
callback_type.__qualname__,
|
||||
callback_type.__name__, callback_type.__qualname__,
|
||||
f"{callback_type.__module__}.{callback_type.__qualname__}",
|
||||
}
|
||||
message = str(relay_error)
|
||||
|
||||
+8
-24
@@ -2,26 +2,20 @@
|
||||
|
||||
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)
|
||||
@@ -57,13 +51,8 @@ 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:
|
||||
@@ -122,12 +111,7 @@ def _json_equal(left: Any, right: Any) -> bool:
|
||||
|
||||
|
||||
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",
|
||||
)
|
||||
|
||||
@@ -35,11 +35,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",
|
||||
"end_turn": "stop", "tool_use": "tool_calls", "max_tokens": "length",
|
||||
"stop_sequence": "stop", "refusal": "content_filter",
|
||||
"model_context_window_exceeded": "length",
|
||||
}
|
||||
|
||||
@@ -60,26 +57,18 @@ class AnthropicTransport(ProviderTransport):
|
||||
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"),
|
||||
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"),
|
||||
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),
|
||||
)
|
||||
@@ -135,11 +124,9 @@ class AnthropicTransport(ProviderTransport):
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -34,11 +34,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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user