From a560fd840c958202ae838f37e7c135fffc6fa1e8 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 18:35:48 -0700 Subject: [PATCH] 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 --- agent/relay_llm.py | 473 ++++++++++++---------------------- agent/relay_runtime.py | 321 +++++++++-------------- agent/relay_tools.py | 32 +-- agent/transports/anthropic.py | 31 +-- agent/transports/base.py | 7 +- 5 files changed, 310 insertions(+), 554 deletions(-) diff --git a/agent/relay_llm.py b/agent/relay_llm.py index b49d527ee5..38cb33ec8b 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -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) diff --git a/agent/relay_runtime.py b/agent/relay_runtime.py index 06db7a287f..bf965bd87e 100644 --- a/agent/relay_runtime.py +++ b/agent/relay_runtime.py @@ -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) diff --git a/agent/relay_tools.py b/agent/relay_tools.py index c38d44068e..99f464684d 100644 --- a/agent/relay_tools.py +++ b/agent/relay_tools.py @@ -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", ) diff --git a/agent/transports/anthropic.py b/agent/transports/anthropic.py index d6740ac8af..e3814f8ab3 100644 --- a/agent/transports/anthropic.py +++ b/agent/transports/anthropic.py @@ -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, ) diff --git a/agent/transports/base.py b/agent/transports/base.py index aae72b5ee0..e56b97ed03 100644 --- a/agent/transports/base.py +++ b/agent/transports/base.py @@ -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."""