Browse Source

feat: add configurable tool invocation modes

Problem: The runtime only supported the dual-agent text-event flow, so ChatAgent provider tool calls could not be compared with the existing EventAgent orchestration.

Risk: Direct provider calls can bypass enabled-name controls, exhaust the event budget incorrectly, or create an invalid assistant/tool transcript. Keep dual_agent as the default and enforce provider-resolved execution through the shared EventKernel.
zhenyu.hu 2 weeks ago
parent
commit
aa976f84a8

+ 2 - 1
src/agent_lab/application/contracts.py

@@ -1,5 +1,5 @@
 from copy import deepcopy
-from typing import Any
+from typing import Any, Literal
 
 from pydantic import BaseModel, ConfigDict, Field, model_validator
 
@@ -41,6 +41,7 @@ class DebugRunRequest(BaseModel):
     pre_messages: list[ChatMessage] = Field(default_factory=list)
     chat_agent: AgentParams
     event_agent: EventAgentParams = Field(default_factory=EventAgentParams)
+    tool_invocation_mode: Literal["dual_agent", "chat_agent_tools"] = "dual_agent"
 
     @model_validator(mode="after")
     def reject_tool_pre_messages(self) -> "DebugRunRequest":

+ 243 - 76
src/agent_lab/application/runtime.py

@@ -1,4 +1,5 @@
 import asyncio
+import json
 import logging
 import time
 from collections.abc import AsyncIterator, Callable
@@ -84,34 +85,23 @@ class DebugRuntime:
         self._tasks = [task for task in self._tasks if not task.done()]
 
     async def _produce(self, request: DebugRunRequest, queues: RuntimeQueues) -> None:
-        event_agent = EventAgent(
-            request.event_agent.enabled_tools,
-            registry=self.registry,
-            chat_client=self.chat_client,
-            params=request.event_agent,
-        )
-        event_worker = asyncio.create_task(self._consume_events(queues, event_agent))
+        event_worker = self._start_event_worker(request, queues)
         try:
             await self._run_chat_agent(request, queues)
         except Exception as exc:
             logger.exception("runtime session failed")
             await queues.output.put({"type": "error", "message": str(exc)})
         finally:
-            await queues.events.put(None)
-            await asyncio.gather(event_worker, return_exceptions=True)
+            if event_worker is not None:
+                await queues.events.put(None)
+                await asyncio.gather(event_worker, return_exceptions=True)
 
     async def _produce_session(
         self,
         request: DebugRunRequest,
         queues: RuntimeQueues,
     ) -> None:
-        event_agent = EventAgent(
-            request.event_agent.enabled_tools,
-            registry=self.registry,
-            chat_client=self.chat_client,
-            params=request.event_agent,
-        )
-        event_worker = asyncio.create_task(self._consume_events(queues, event_agent))
+        event_worker = self._start_event_worker(request, queues)
         try:
             await self._run_chat_session(request, queues)
         except asyncio.CancelledError:
@@ -120,8 +110,24 @@ class DebugRuntime:
             logger.exception("runtime session failed")
             await queues.output.put({"type": "error", "message": str(exc)})
         finally:
-            await queues.events.put(None)
-            await asyncio.gather(event_worker, return_exceptions=True)
+            if event_worker is not None:
+                await queues.events.put(None)
+                await asyncio.gather(event_worker, return_exceptions=True)
+
+    def _start_event_worker(
+        self,
+        request: DebugRunRequest,
+        queues: RuntimeQueues,
+    ) -> asyncio.Task[None] | None:
+        if request.tool_invocation_mode != "dual_agent":
+            return None
+        event_agent = EventAgent(
+            request.event_agent.enabled_tools,
+            registry=self.registry,
+            chat_client=self.chat_client,
+            params=request.event_agent,
+        )
+        return asyncio.create_task(self._consume_events(queues, event_agent))
 
     async def _run_chat_agent(
         self,
@@ -133,7 +139,12 @@ class DebugRuntime:
         await queues.input.put(ChatMessage(role="user", content=request.user_message))
 
         await queues.output.put({"type": "session_started"})
-        await self._audit(queues, "session_started", turn_started_at=turn_started_at)
+        await self._audit(
+            queues,
+            "session_started",
+            turn_started_at=turn_started_at,
+            tool_invocation_mode=request.tool_invocation_mode,
+        )
 
         event_loops = 0
         round_index = 0
@@ -158,13 +169,19 @@ class DebugRuntime:
                 if event_loops < request.event_agent.max_event_loops
                 else []
             )
+            chat_tools = self._chat_tools_for_round(request, event_prompt_events)
             raw_chunks: list[dict[str, Any]] = []
-            chat_messages = self._chat_messages_for_round(messages, event_prompt_events)
+            chat_messages = self._chat_messages_for_mode(
+                request,
+                messages,
+                event_prompt_events,
+            )
             await self._audit(
                 queues,
                 "chat_round_started",
                 turn_started_at=turn_started_at,
                 round_index=round_index,
+                tool_invocation_mode=request.tool_invocation_mode,
                 events_enabled=event_prompt_events,
                 configured_events=configured_events,
                 event_generation_enabled=bool(event_prompt_events),
@@ -176,14 +193,15 @@ class DebugRuntime:
                 turn_started_at=turn_started_at,
                 agent="chat_agent",
                 round_index=round_index,
+                tool_invocation_mode=request.tool_invocation_mode,
                 params=self._params_snapshot(request.chat_agent),
                 messages=self._message_snapshots(chat_messages),
-                tools=[],
+                tools=chat_tools,
             )
 
             async for item in self.chat_client.stream_chat(
                 messages=chat_messages,
-                tools=[],
+                tools=chat_tools,
                 params=request.chat_agent,
             ):
                 if item.kind == "raw_chunk" and item.raw_chunk is not None:
@@ -220,9 +238,13 @@ class DebugRuntime:
                     )
                     continue
 
-                if item.kind == "text_event" and item.event is not None:
+                event = self._accepted_chat_event(
+                    request,
+                    item,
+                    budget_available=event_loops < request.event_agent.max_event_loops,
+                )
+                if event is not None:
                     saw_event = True
-                    event = self._event_name_only(item.event)
                     events.append(event)
                     await queues.output.put(
                         {"type": "event", "event": event.model_dump()}
@@ -234,6 +256,12 @@ class DebugRuntime:
                         round_index=round_index,
                         event_id=event.id,
                         event_name=event.name,
+                        event_source=(
+                            "provider_resolved"
+                            if request.tool_invocation_mode == "chat_agent_tools"
+                            else "text_event"
+                        ),
+                        tool_invocation_mode=request.tool_invocation_mode,
                     )
 
             if message_stream_started:
@@ -252,8 +280,10 @@ class DebugRuntime:
                 turn_started_at=turn_started_at,
                 agent="chat_agent",
                 round_index=round_index,
+                tool_invocation_mode=request.tool_invocation_mode,
                 content="".join(assistant_content),
                 event_names=[event.name for event in events],
+                events=self._event_snapshots(events),
                 usage={
                     "prompt_tokens": prompt_tokens,
                     "completion_tokens": completion_tokens,
@@ -267,36 +297,62 @@ class DebugRuntime:
                 assistant_message = ChatMessage(
                     role="assistant",
                     content="".join(assistant_content),
+                    tool_calls=(
+                        events
+                        if request.tool_invocation_mode == "chat_agent_tools"
+                        else []
+                    ),
                 )
                 messages.append(assistant_message)
 
             if events:
-                await queues.events.put(
-                    EventAgentRequest(
-                        events=events,
-                        history=self._event_agent_history(messages),
-                        system_prompt=request.event_agent.system_prompt,
-                        extra_body=request.event_agent.extra_body,
-                        turn_started_at=turn_started_at,
+                if request.tool_invocation_mode == "chat_agent_tools":
+                    tool_replies = await self._execute_provider_tools(
+                        events,
+                        enabled_names=request.event_agent.enabled_tools,
+                    )
+                    messages.extend(tool_replies)
+                else:
+                    await queues.events.put(
+                        EventAgentRequest(
+                            events=events,
+                            history=self._event_agent_history(messages),
+                            system_prompt=request.event_agent.system_prompt,
+                            extra_body=request.event_agent.extra_body,
+                            turn_started_at=turn_started_at,
+                        )
+                    )
+                    tool_replies = await self._wait_for_tool_replies(
+                        queues,
+                        [event.id for event in events],
                     )
-                )
-                tool_replies = await self._wait_for_tool_replies(
-                    queues,
-                    [event.id for event in events],
-                )
                 for reply in tool_replies:
                     await queues.output.put(
                         {"type": "tool_result", "message": reply.model_dump()}
                     )
-                await self._audit(
-                    queues,
-                    "event_agent_completed",
-                    turn_started_at=turn_started_at,
-                    round_index=round_index,
-                    event_names=[event.name for event in events],
-                    result_count=len(tool_replies),
-                    result_summary="\n".join(reply.content for reply in tool_replies),
-                )
+                if request.tool_invocation_mode == "chat_agent_tools":
+                    await self._audit(
+                        queues,
+                        "provider_tools_completed",
+                        turn_started_at=turn_started_at,
+                        round_index=round_index,
+                        tool_invocation_mode=request.tool_invocation_mode,
+                        event_source="provider_resolved",
+                        events=self._event_snapshots(events),
+                        result_count=len(tool_replies),
+                    )
+                else:
+                    await self._audit(
+                        queues,
+                        "event_agent_completed",
+                        turn_started_at=turn_started_at,
+                        round_index=round_index,
+                        event_names=[event.name for event in events],
+                        result_count=len(tool_replies),
+                        result_summary="\n".join(
+                            reply.content for reply in tool_replies
+                        ),
+                    )
 
             elapsed_ms = self._elapsed_ms(started_at)
             await queues.output.put(
@@ -317,6 +373,7 @@ class DebugRuntime:
                 "chat_round_finished",
                 turn_started_at=turn_started_at,
                 round_index=round_index,
+                tool_invocation_mode=request.tool_invocation_mode,
                 had_event=saw_event,
                 elapsed_ms=elapsed_ms,
             )
@@ -351,6 +408,7 @@ class DebugRuntime:
             "session_started",
             turn_started_at=initial_turn_started_at,
             session_id=session_id,
+            tool_invocation_mode=request.tool_invocation_mode,
         )
         await queues.input.put(ChatMessage(role="user", content=request.user_message))
 
@@ -424,8 +482,13 @@ class DebugRuntime:
                 if event_loops < request.event_agent.max_event_loops
                 else []
             )
+            chat_tools = self._chat_tools_for_round(request, event_prompt_events)
             raw_chunks: list[dict[str, Any]] = []
-            chat_messages = self._chat_messages_for_round(messages, event_prompt_events)
+            chat_messages = self._chat_messages_for_mode(
+                request,
+                messages,
+                event_prompt_events,
+            )
             await self._audit(
                 queues,
                 "chat_round_started",
@@ -433,6 +496,7 @@ class DebugRuntime:
                 session_id=session_id,
                 turn_index=turn_index,
                 round_index=round_index,
+                tool_invocation_mode=request.tool_invocation_mode,
                 events_enabled=event_prompt_events,
                 configured_events=configured_events,
                 event_generation_enabled=bool(event_prompt_events),
@@ -446,14 +510,15 @@ class DebugRuntime:
                 agent="chat_agent",
                 turn_index=turn_index,
                 round_index=round_index,
+                tool_invocation_mode=request.tool_invocation_mode,
                 params=self._params_snapshot(request.chat_agent),
                 messages=self._message_snapshots(chat_messages),
-                tools=[],
+                tools=chat_tools,
             )
 
             async for item in self.chat_client.stream_chat(
                 messages=chat_messages,
-                tools=[],
+                tools=chat_tools,
                 params=request.chat_agent,
             ):
                 if item.kind == "raw_chunk" and item.raw_chunk is not None:
@@ -492,9 +557,13 @@ class DebugRuntime:
                     )
                     continue
 
-                if item.kind == "text_event" and item.event is not None:
+                event = self._accepted_chat_event(
+                    request,
+                    item,
+                    budget_available=event_loops < request.event_agent.max_event_loops,
+                )
+                if event is not None:
                     saw_event = True
-                    event = self._event_name_only(item.event)
                     events.append(event)
                     await queues.output.put(
                         {"type": "event", "event": event.model_dump()}
@@ -508,6 +577,12 @@ class DebugRuntime:
                         round_index=round_index,
                         event_id=event.id,
                         event_name=event.name,
+                        event_source=(
+                            "provider_resolved"
+                            if request.tool_invocation_mode == "chat_agent_tools"
+                            else "text_event"
+                        ),
+                        tool_invocation_mode=request.tool_invocation_mode,
                     )
 
             if message_stream_started:
@@ -530,8 +605,10 @@ class DebugRuntime:
                 agent="chat_agent",
                 turn_index=turn_index,
                 round_index=round_index,
+                tool_invocation_mode=request.tool_invocation_mode,
                 content="".join(assistant_content),
                 event_names=[event.name for event in events],
+                events=self._event_snapshots(events),
                 usage={
                     "prompt_tokens": prompt_tokens,
                     "completion_tokens": completion_tokens,
@@ -545,6 +622,11 @@ class DebugRuntime:
                 assistant_message = ChatMessage(
                     role="assistant",
                     content="".join(assistant_content),
+                    tool_calls=(
+                        events
+                        if request.tool_invocation_mode == "chat_agent_tools"
+                        else []
+                    ),
                 )
                 messages.append(assistant_message)
                 self._append_persisted_message(
@@ -554,37 +636,60 @@ class DebugRuntime:
                 )
 
             if events:
-                await queues.events.put(
-                    EventAgentRequest(
-                        events=events,
-                        history=self._event_agent_history(messages),
-                        system_prompt=request.event_agent.system_prompt,
-                        extra_body=request.event_agent.extra_body,
-                        session_id=session_id,
-                        turn_index=turn_index,
-                        round_index=round_index,
-                        turn_started_at=turn_started_at,
+                if request.tool_invocation_mode == "chat_agent_tools":
+                    tool_replies = await self._execute_provider_tools(
+                        events,
+                        enabled_names=request.event_agent.enabled_tools,
+                    )
+                    messages.extend(tool_replies)
+                else:
+                    await queues.events.put(
+                        EventAgentRequest(
+                            events=events,
+                            history=self._event_agent_history(messages),
+                            system_prompt=request.event_agent.system_prompt,
+                            extra_body=request.event_agent.extra_body,
+                            session_id=session_id,
+                            turn_index=turn_index,
+                            round_index=round_index,
+                            turn_started_at=turn_started_at,
+                        )
+                    )
+                    tool_replies = await self._wait_for_tool_replies(
+                        queues,
+                        [event.id for event in events],
                     )
-                )
-                tool_replies = await self._wait_for_tool_replies(
-                    queues,
-                    [event.id for event in events],
-                )
                 for reply in tool_replies:
                     await queues.output.put(
                         {"type": "tool_result", "message": reply.model_dump()}
                     )
-                await self._audit(
-                    queues,
-                    "event_agent_completed",
-                    turn_started_at=turn_started_at,
-                    session_id=session_id,
-                    turn_index=turn_index,
-                    round_index=round_index,
-                    event_names=[event.name for event in events],
-                    result_count=len(tool_replies),
-                    result_summary="\n".join(reply.content for reply in tool_replies),
-                )
+                if request.tool_invocation_mode == "chat_agent_tools":
+                    await self._audit(
+                        queues,
+                        "provider_tools_completed",
+                        turn_started_at=turn_started_at,
+                        session_id=session_id,
+                        turn_index=turn_index,
+                        round_index=round_index,
+                        tool_invocation_mode=request.tool_invocation_mode,
+                        event_source="provider_resolved",
+                        events=self._event_snapshots(events),
+                        result_count=len(tool_replies),
+                    )
+                else:
+                    await self._audit(
+                        queues,
+                        "event_agent_completed",
+                        turn_started_at=turn_started_at,
+                        session_id=session_id,
+                        turn_index=turn_index,
+                        round_index=round_index,
+                        event_names=[event.name for event in events],
+                        result_count=len(tool_replies),
+                        result_summary="\n".join(
+                            reply.content for reply in tool_replies
+                        ),
+                    )
 
             elapsed_ms = self._elapsed_ms(started_at)
             self._append_persisted_usage(
@@ -620,6 +725,7 @@ class DebugRuntime:
                 session_id=session_id,
                 turn_index=turn_index,
                 round_index=round_index,
+                tool_invocation_mode=request.tool_invocation_mode,
                 had_event=saw_event,
                 elapsed_ms=elapsed_ms,
             )
@@ -736,6 +842,66 @@ class DebugRuntime:
         messages.extend(request.pre_messages)
         return messages
 
+    def _chat_messages_for_mode(
+        self,
+        request: DebugRunRequest,
+        messages: list[ChatMessage],
+        enabled_events: list[str],
+    ) -> list[ChatMessage]:
+        if request.tool_invocation_mode == "chat_agent_tools":
+            return list(messages)
+        return self._chat_messages_for_round(messages, enabled_events)
+
+    def _chat_tools_for_round(
+        self,
+        request: DebugRunRequest,
+        enabled_events: list[str],
+    ) -> list[dict[str, Any]]:
+        if request.tool_invocation_mode != "chat_agent_tools":
+            return []
+        return self.registry.provider_tool_schemas(enabled_events)
+
+    def _accepted_chat_event(
+        self,
+        request: DebugRunRequest,
+        item: StreamItem,
+        *,
+        budget_available: bool,
+    ) -> ToolCallEvent | None:
+        if request.tool_invocation_mode == "chat_agent_tools":
+            if (
+                budget_available
+                and item.kind == "provider_tool_call"
+                and item.event is not None
+            ):
+                return item.event
+            return None
+        if item.kind == "text_event" and item.event is not None:
+            return self._event_name_only(item.event)
+        return None
+
+    async def _execute_provider_tools(
+        self,
+        events: list[ToolCallEvent],
+        *,
+        enabled_names: list[str],
+    ) -> list[ChatMessage]:
+        replies: list[ChatMessage] = []
+        for event in events:
+            payload = await self.registry.execute_async(
+                event,
+                enabled_names=enabled_names,
+            )
+            replies.append(
+                ChatMessage(
+                    role="tool",
+                    content=json.dumps(payload, ensure_ascii=False),
+                    name=event.name,
+                    tool_call_id=event.id,
+                )
+            )
+        return replies
+
     def _chat_messages_for_round(
         self,
         messages: list[ChatMessage],
@@ -878,6 +1044,7 @@ class DebugRuntime:
 
     def _session_config_snapshot(self, request: DebugRunRequest) -> dict[str, Any]:
         return {
+            "tool_invocation_mode": request.tool_invocation_mode,
             "system_prompts": list(request.system_prompts),
             "pre_messages": self._message_snapshots(request.pre_messages),
             "chat_agent": request.chat_agent.model_dump(),

+ 21 - 2
src/agent_lab/application/tools.py

@@ -97,6 +97,19 @@ class ToolRegistry:
     def tool_schema(self, name: str) -> dict[str, Any] | None:
         return self.event_registry.tool_schema(name)
 
+    def provider_tool_schemas(
+        self,
+        enabled_names: Iterable[str],
+    ) -> list[dict[str, Any]]:
+        return [
+            schema
+            for schema in (
+                self.tool_schema(name)
+                for name in enabled_names
+            )
+            if schema is not None
+        ]
+
     def handle(
         self,
         event: ToolCallEvent,
@@ -125,9 +138,15 @@ class ToolRegistry:
         )
         return self.tool_payload(result)
 
-    async def execute_async(self, event: ToolCallEvent) -> dict[str, Any]:
+    async def execute_async(
+        self,
+        event: ToolCallEvent,
+        *,
+        enabled_names: Iterable[str] | None = None,
+    ) -> dict[str, Any]:
         result = await self.kernel.execute(
-            self.event_request(event, source=EventSource.PROVIDER_RESOLVED)
+            self.event_request(event, source=EventSource.PROVIDER_RESOLVED),
+            enabled_names=enabled_names,
         )
         return self.tool_payload(result)
 

+ 393 - 0
tests/test_debug_runtime.py

@@ -1579,3 +1579,396 @@ async def test_runtime_continues_persisted_turn_indexes_for_existing_session(tmp
         (2, "new"),
         (2, "hello"),
     ]
+
+
+class ScriptedChatClient:
+    def __init__(self, rounds: list[list[StreamItem]]) -> None:
+        self.rounds = rounds
+        self.calls = 0
+        self.messages_by_call: list[list[ChatMessage]] = []
+        self.tools_by_call: list[list[dict[str, Any]]] = []
+
+    async def stream_chat(
+        self,
+        messages: list[ChatMessage],
+        tools: list[dict],
+        params: AgentParams,
+        tool_choice: dict[str, Any] | None = None,
+    ) -> AsyncIterator[StreamItem]:
+        self.messages_by_call.append(list(messages))
+        self.tools_by_call.append(list(tools))
+        round_items = self.rounds[self.calls]
+        self.calls += 1
+        for item in round_items:
+            yield item
+
+
+def _direct_request(
+    *,
+    enabled_tools: list[str],
+    max_event_loops: int = 1,
+    session_id: str | None = None,
+) -> DebugRunRequest:
+    return DebugRunRequest(
+        session_id=session_id,
+        user_message="debug direct tools",
+        system_prompts=["You are a debugger."],
+        pre_messages=[],
+        chat_agent=AgentParams(model="chat-model"),
+        event_agent=EventAgentParams(
+            enabled_tools=enabled_tools,
+            max_event_loops=max_event_loops,
+        ),
+        tool_invocation_mode="chat_agent_tools",
+    )
+
+
+@pytest.mark.asyncio
+async def test_direct_mode_uses_enabled_schemas_and_valid_ordered_provider_transcript():
+    executed: list[tuple[str, dict[str, Any], str]] = []
+
+    def handle(event: ToolCallEvent) -> dict[str, Any]:
+        executed.append((event.name, event.arguments, event.raw_arguments))
+        return {"tool": event.name, "value": event.arguments["value"]}
+
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="first_tool",
+                description="First direct tool.",
+                parameters={
+                    "type": "object",
+                    "properties": {"value": {"type": "string"}},
+                    "required": ["value"],
+                },
+                handler=handle,
+            ),
+            ToolDefinition(
+                name="second_tool",
+                description="Second direct tool.",
+                parameters={
+                    "type": "object",
+                    "properties": {"value": {"type": "string"}},
+                    "required": ["value"],
+                },
+                handler=handle,
+            ),
+        ]
+    )
+    calls = [
+        ToolCallEvent(
+            id="provider-1",
+            name="first_tool",
+            arguments={"value": "one"},
+            raw_arguments='{"value":"one"}',
+        ),
+        ToolCallEvent(
+            id="provider-2",
+            name="second_tool",
+            arguments={"value": "two"},
+            raw_arguments='{"value":"two"}',
+        ),
+    ]
+    ignored_text_event = ToolCallEvent(
+        id="text-ignored",
+        name="first_tool",
+        arguments={},
+        raw_arguments="{}",
+    )
+    client = ScriptedChatClient(
+        [
+            [
+                StreamItem.message_delta("Visible before tools."),
+                StreamItem.provider_tool_call(calls[0]),
+                StreamItem.text_event(ignored_text_event),
+                StreamItem.provider_tool_call(calls[1]),
+            ],
+            [StreamItem.message_delta("Final answer.")],
+        ]
+    )
+
+    outputs = await _collect_outputs(
+        DebugRuntime(client, registry=registry).run(
+            _direct_request(enabled_tools=["first_tool", "second_tool"])
+        )
+    )
+
+    assert [tool["function"]["name"] for tool in client.tools_by_call[0]] == [
+        "first_tool",
+        "second_tool",
+    ]
+    assert client.tools_by_call[1] == []
+    assert not any(
+        "Available events:" in message.content
+        for message in client.messages_by_call[0]
+        if message.role == "system"
+    )
+    transcript = client.messages_by_call[1]
+    assert [message.role for message in transcript[-4:]] == [
+        "user",
+        "assistant",
+        "tool",
+        "tool",
+    ]
+    assert transcript[-3].content == "Visible before tools."
+    assert transcript[-3].tool_calls == calls
+    assert [message.tool_call_id for message in transcript[-2:]] == [
+        "provider-1",
+        "provider-2",
+    ]
+    assert [message.name for message in transcript[-2:]] == [
+        "first_tool",
+        "second_tool",
+    ]
+    assert executed == [
+        ("first_tool", {"value": "one"}, '{"value":"one"}'),
+        ("second_tool", {"value": "two"}, '{"value":"two"}'),
+    ]
+    emitted_events = [
+        message["event"] for message in outputs if message["type"] == "event"
+    ]
+    assert emitted_events == [call.model_dump() for call in calls]
+    assert "event_agent_request" not in [
+        message.get("event") for message in outputs if message["type"] == "audit"
+    ]
+    assert not any(
+        message.get("event", {}).get("id") == "text-ignored"
+        for message in outputs
+        if message["type"] == "event"
+    )
+    assert all(
+        message["details"]["tool_invocation_mode"] == "chat_agent_tools"
+        for message in outputs
+        if message["type"] == "audit"
+        and message["event"] in {"chat_round_started", "chat_agent_request"}
+    )
+
+
+@pytest.mark.asyncio
+async def test_direct_mode_does_not_execute_unknown_disabled_or_failed_handlers():
+    executed: list[str] = []
+
+    def disabled_handler(event: ToolCallEvent) -> dict[str, Any]:
+        executed.append(event.name)
+        return {"tool": event.name}
+
+    def failing_handler(event: ToolCallEvent) -> dict[str, Any]:
+        executed.append(event.name)
+        raise RuntimeError("direct boom")
+
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="enabled_tool",
+                description="Enabled.",
+                parameters={"type": "object"},
+                handler=lambda event: {"tool": event.name},
+            ),
+            ToolDefinition(
+                name="disabled_tool",
+                description="Disabled.",
+                parameters={"type": "object"},
+                handler=disabled_handler,
+            ),
+            ToolDefinition(
+                name="failing_tool",
+                description="Fails.",
+                parameters={"type": "object"},
+                handler=failing_handler,
+            ),
+        ]
+    )
+    client = ScriptedChatClient(
+        [
+            [
+                StreamItem.provider_tool_call(
+                    ToolCallEvent(
+                        id="unknown-1",
+                        name="unknown_tool",
+                        arguments={"kept": True},
+                        raw_arguments='{"kept":true}',
+                    )
+                ),
+                StreamItem.provider_tool_call(
+                    ToolCallEvent(
+                        id="disabled-1",
+                        name="disabled_tool",
+                        arguments={},
+                        raw_arguments="{}",
+                    )
+                ),
+                StreamItem.provider_tool_call(
+                    ToolCallEvent(
+                        id="failed-1",
+                        name="failing_tool",
+                        arguments={},
+                        raw_arguments="{}",
+                    )
+                ),
+            ],
+            [StreamItem.message_delta("continued")],
+        ]
+    )
+
+    outputs = await _collect_outputs(
+        DebugRuntime(client, registry=registry).run(
+            _direct_request(enabled_tools=["enabled_tool", "failing_tool"])
+        )
+    )
+
+    assert [tool["function"]["name"] for tool in client.tools_by_call[0]] == [
+        "enabled_tool",
+        "failing_tool",
+    ]
+    assert executed == ["failing_tool"]
+    results = [
+        json.loads(message["message"]["content"])
+        for message in outputs
+        if message["type"] == "tool_result"
+    ]
+    assert results == [
+        {"tool": "unknown_tool", "error": "unknown tool"},
+        {"tool": "disabled_tool", "error": "tool disabled"},
+        {"tool": "failing_tool", "error": "tool handler failed: direct boom"},
+    ]
+
+
+@pytest.mark.asyncio
+async def test_direct_mode_ignores_provider_calls_after_event_budget_exhaustion():
+    executed: list[str] = []
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="once_tool",
+                description="Run once.",
+                parameters={"type": "object"},
+                handler=lambda event: executed.append(event.id) or {"tool": event.name},
+            )
+        ]
+    )
+    client = ScriptedChatClient(
+        [
+            [
+                StreamItem.provider_tool_call(
+                    ToolCallEvent(
+                        id="accepted",
+                        name="once_tool",
+                        arguments={},
+                        raw_arguments="{}",
+                    )
+                )
+            ],
+            [
+                StreamItem.provider_tool_call(
+                    ToolCallEvent(
+                        id="ignored",
+                        name="once_tool",
+                        arguments={},
+                        raw_arguments="{}",
+                    )
+                ),
+                StreamItem.message_delta("budget exhausted"),
+            ],
+        ]
+    )
+
+    outputs = await _collect_outputs(
+        DebugRuntime(client, registry=registry).run(
+            _direct_request(enabled_tools=["once_tool"], max_event_loops=1)
+        )
+    )
+
+    assert client.tools_by_call == [
+        [registry.tool_schema("once_tool")],
+        [],
+    ]
+    assert executed == ["accepted"]
+    assert [
+        message["event"]["id"]
+        for message in outputs
+        if message["type"] == "event"
+    ] == ["accepted"]
+    assert outputs[-1] == {"type": "done"}
+
+
+@pytest.mark.asyncio
+async def test_direct_mode_session_turns_reset_budget_and_snapshot_mode(tmp_path):
+    store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="turn_tool",
+                description="Per-turn tool.",
+                parameters={"type": "object"},
+                handler=lambda event: {"tool": event.name, "call_id": event.id},
+            )
+        ]
+    )
+    client = ScriptedChatClient(
+        [
+            [
+                StreamItem.provider_tool_call(
+                    ToolCallEvent(
+                        id="turn-1-call",
+                        name="turn_tool",
+                        arguments={},
+                        raw_arguments="{}",
+                    )
+                )
+            ],
+            [StreamItem.message_delta("turn one done")],
+            [
+                StreamItem.provider_tool_call(
+                    ToolCallEvent(
+                        id="turn-2-call",
+                        name="turn_tool",
+                        arguments={},
+                        raw_arguments="{}",
+                    )
+                )
+            ],
+            [StreamItem.message_delta("turn two done")],
+        ]
+    )
+    runtime = DebugRuntime(client, registry=registry, session_store=store)
+    request = _direct_request(
+        enabled_tools=["turn_tool"],
+        max_event_loops=1,
+        session_id="direct-session",
+    )
+
+    queues = runtime.start_session(request)
+    outputs: list[dict[str, Any]] = []
+    while sum(message["type"] == "turn_completed" for message in outputs) < 1:
+        outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
+    await queues.input.put(ChatMessage(role="user", content="second direct turn"))
+    while sum(message["type"] == "turn_completed" for message in outputs) < 2:
+        outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
+    await runtime.aclose()
+
+    assert [bool(tools) for tools in client.tools_by_call] == [True, False, True, False]
+    assert [message.role for message in client.messages_by_call[1][-3:]] == [
+        "user",
+        "assistant",
+        "tool",
+    ]
+    second_turn_transcript = client.messages_by_call[3]
+    assert not any(message.name == "event_agent" for message in second_turn_transcript)
+    assert [
+        message.tool_call_id
+        for message in second_turn_transcript
+        if message.role == "tool"
+    ] == ["turn-1-call", "turn-2-call"]
+    session = store.get_session("direct-session")
+    assert session is not None
+    assert session["config"]["tool_invocation_mode"] == "chat_agent_tools"
+    request_audits = [
+        audit
+        for audit in store.list_audit_logs("direct-session")
+        if audit["event"] == "chat_agent_request"
+    ]
+    assert request_audits
+    assert all(
+        audit["details"]["tool_invocation_mode"] == "chat_agent_tools"
+        for audit in request_audits
+    )

+ 24 - 0
tests/test_websocket_api.py

@@ -163,6 +163,30 @@ def test_agent_params_defaults_include_extra_body_and_event_agent_defaults_to_on
     assert event_params.system_prompt == ""
 
 
+def test_debug_run_request_defaults_to_dual_agent_tool_invocation_mode():
+    request = DebugRunRequest.model_validate(_request_payload())
+
+    assert request.tool_invocation_mode == "dual_agent"
+
+
+@pytest.mark.parametrize("mode", ["dual_agent", "chat_agent_tools"])
+def test_debug_run_request_accepts_supported_tool_invocation_modes(mode: str):
+    payload = _request_payload()
+    payload["tool_invocation_mode"] = mode
+
+    request = DebugRunRequest.model_validate(payload)
+
+    assert request.tool_invocation_mode == mode
+
+
+def test_debug_run_request_rejects_unknown_tool_invocation_mode():
+    payload = _request_payload()
+    payload["tool_invocation_mode"] = "unknown"
+
+    with pytest.raises(ValidationError, match="tool_invocation_mode"):
+        DebugRunRequest.model_validate(payload)
+
+
 def test_health_returns_ok():
     app = create_app(runtime_factory=FakeRuntime)
     client = TestClient(app)