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:
Teknium
2026-09-02 18:35:48 -07:00
parent 113f04616b
commit a560fd840c
5 changed files with 310 additions and 554 deletions
+167 -306
View File
@@ -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
View File
@@ -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
View File
@@ -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",
)
+9 -22
View File
@@ -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,
)
+2 -5
View File
@@ -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."""