|
@@ -0,0 +1,482 @@
|
|
|
|
|
+import asyncio
|
|
|
|
|
+import json
|
|
|
|
|
+from collections.abc import AsyncIterator
|
|
|
|
|
+from dataclasses import dataclass
|
|
|
|
|
+from typing import Any
|
|
|
|
|
+
|
|
|
|
|
+import pytest
|
|
|
|
|
+
|
|
|
|
|
+from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
|
|
|
|
|
+from agent_lab.application.runtime import DebugRuntime
|
|
|
|
|
+from agent_lab.application.tools import build_default_tool_registry
|
|
|
|
|
+from agent_lab.domain.events import ToolCallEvent
|
|
|
|
|
+from agent_lab.domain.messages import ChatMessage, StreamItem
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+MODES = ("dual_agent", "chat_agent_tools")
|
|
|
|
|
+SCENARIOS = (
|
|
|
|
|
+ "ordinary_chat",
|
|
|
|
|
+ "device_volume_silent",
|
|
|
|
|
+ "calendar_schedule_template",
|
|
|
|
|
+ "web_search_two_answers",
|
|
|
|
|
+ "session_terminate",
|
|
|
|
|
+ "parallel_volume_schedule",
|
|
|
|
|
+)
|
|
|
|
|
+BUILTIN_EVENT_NAMES = (
|
|
|
|
|
+ "session.terminate",
|
|
|
|
|
+ "device.volume.adjust",
|
|
|
|
|
+ "calendar.schedule.create",
|
|
|
|
|
+ "knowledge.web.search",
|
|
|
|
|
+)
|
|
|
|
|
+SCHEDULE_ARGUMENTS = {
|
|
|
|
|
+ "title": "Design review",
|
|
|
|
|
+ "start_at": "2026-07-14T09:30:00+08:00",
|
|
|
|
|
+ "timezone": "Asia/Shanghai",
|
|
|
|
|
+}
|
|
|
|
|
+SCHEDULE_CONFIRMATION = (
|
|
|
|
|
+ "Scheduled Design review for 2026-07-14T09:30:00+08:00 "
|
|
|
|
|
+ "(Asia/Shanghai)."
|
|
|
|
|
+)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+@dataclass(frozen=True)
|
|
|
|
|
+class ExpectedSemantics:
|
|
|
|
|
+ visible_messages: tuple[str, ...]
|
|
|
|
|
+ tool_names: tuple[str, ...]
|
|
|
|
|
+ model_calls: int
|
|
|
|
|
+ terminal_count: int = 0
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+EXPECTED = {
|
|
|
|
|
+ "ordinary_chat": ExpectedSemantics(("Ordinary answer.",), (), 1),
|
|
|
|
|
+ "device_volume_silent": ExpectedSemantics(
|
|
|
|
|
+ (),
|
|
|
|
|
+ ("device.volume.adjust",),
|
|
|
|
|
+ 1,
|
|
|
|
|
+ ),
|
|
|
|
|
+ "calendar_schedule_template": ExpectedSemantics(
|
|
|
|
|
+ (SCHEDULE_CONFIRMATION,),
|
|
|
|
|
+ ("calendar.schedule.create",),
|
|
|
|
|
+ 1,
|
|
|
|
|
+ ),
|
|
|
|
|
+ "web_search_two_answers": ExpectedSemantics(
|
|
|
|
|
+ ("I will check.", "Grounded search update."),
|
|
|
|
|
+ ("knowledge.web.search",),
|
|
|
|
|
+ 2,
|
|
|
|
|
+ ),
|
|
|
|
|
+ "session_terminate": ExpectedSemantics(
|
|
|
|
|
+ ("Goodbye.",),
|
|
|
|
|
+ ("session.terminate",),
|
|
|
|
|
+ 1,
|
|
|
|
|
+ terminal_count=1,
|
|
|
|
|
+ ),
|
|
|
|
|
+ "parallel_volume_schedule": ExpectedSemantics(
|
|
|
|
|
+ (SCHEDULE_CONFIRMATION,),
|
|
|
|
|
+ ("device.volume.adjust", "calendar.schedule.create"),
|
|
|
|
|
+ 1,
|
|
|
|
|
+ ),
|
|
|
|
|
+}
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+class ScriptedChatClient:
|
|
|
|
|
+ def __init__(self, rounds: list[list[StreamItem]]) -> None:
|
|
|
|
|
+ self.rounds = rounds
|
|
|
|
|
+ self.calls = 0
|
|
|
|
|
+
|
|
|
|
|
+ async def stream_chat(
|
|
|
|
|
+ self,
|
|
|
|
|
+ messages: list[ChatMessage],
|
|
|
|
|
+ tools: list[dict[str, Any]],
|
|
|
|
|
+ params: AgentParams,
|
|
|
|
|
+ tool_choice: dict[str, Any] | None = None,
|
|
|
|
|
+ ) -> AsyncIterator[StreamItem]:
|
|
|
|
|
+ del messages, tools, params, tool_choice
|
|
|
|
|
+ items = self.rounds[self.calls]
|
|
|
|
|
+ self.calls += 1
|
|
|
|
|
+ for item in items:
|
|
|
|
|
+ yield item
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+class RecordingPorts:
|
|
|
|
|
+ def __init__(self, *, parallel_barrier: bool = False) -> None:
|
|
|
|
|
+ self.parallel_barrier = parallel_barrier
|
|
|
|
|
+ self.session_calls: list[tuple[str, str | None]] = []
|
|
|
|
|
+ self.volume_calls: list[tuple[str, str, int | None, int | None]] = []
|
|
|
|
|
+ self.calendar_calls: list[
|
|
|
|
|
+ tuple[str, str, str, str, str | None, int | None]
|
|
|
|
|
+ ] = []
|
|
|
|
|
+ self.search_calls: list[tuple[str, str, int]] = []
|
|
|
|
|
+ self.started: set[str] = set()
|
|
|
|
|
+ self.both_started = asyncio.Event()
|
|
|
|
|
+ self.volume_release = asyncio.Event()
|
|
|
|
|
+ self.schedule_release = asyncio.Event()
|
|
|
|
|
+ self.schedule_finished = asyncio.Event()
|
|
|
|
|
+ if not parallel_barrier:
|
|
|
|
|
+ self.volume_release.set()
|
|
|
|
|
+ self.schedule_release.set()
|
|
|
|
|
+
|
|
|
|
|
+ async def terminate(
|
|
|
|
|
+ self,
|
|
|
|
|
+ event_id: str,
|
|
|
|
|
+ *,
|
|
|
|
|
+ reason: str | None = None,
|
|
|
|
|
+ ) -> dict[str, Any]:
|
|
|
|
|
+ self.session_calls.append((event_id, reason))
|
|
|
|
|
+ return {
|
|
|
|
|
+ "tool": "session.terminate",
|
|
|
|
|
+ "status": "terminated",
|
|
|
|
|
+ "event_id": event_id,
|
|
|
|
|
+ "reason": reason,
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ async def adjust(
|
|
|
|
|
+ self,
|
|
|
|
|
+ event_id: str,
|
|
|
|
|
+ *,
|
|
|
|
|
+ mode: str,
|
|
|
|
|
+ value: int | None = None,
|
|
|
|
|
+ delta: int | None = None,
|
|
|
|
|
+ ) -> dict[str, Any]:
|
|
|
|
|
+ self.volume_calls.append((event_id, mode, value, delta))
|
|
|
|
|
+ await self._wait_at_barrier("volume", self.volume_release)
|
|
|
|
|
+ payload: dict[str, Any] = {
|
|
|
|
|
+ "tool": "device.volume.adjust",
|
|
|
|
|
+ "status": "applied",
|
|
|
|
|
+ "event_id": event_id,
|
|
|
|
|
+ "mode": mode,
|
|
|
|
|
+ }
|
|
|
|
|
+ if value is not None:
|
|
|
|
|
+ payload["value"] = value
|
|
|
|
|
+ if delta is not None:
|
|
|
|
|
+ payload["delta"] = delta
|
|
|
|
|
+ return payload
|
|
|
|
|
+
|
|
|
|
|
+ async def create(
|
|
|
|
|
+ self,
|
|
|
|
|
+ event_id: str,
|
|
|
|
|
+ *,
|
|
|
|
|
+ title: str,
|
|
|
|
|
+ start_at: str,
|
|
|
|
|
+ timezone: str,
|
|
|
|
|
+ recurrence: str | None = None,
|
|
|
|
|
+ reminder_minutes: int | None = None,
|
|
|
|
|
+ ) -> dict[str, Any]:
|
|
|
|
|
+ self.calendar_calls.append(
|
|
|
|
|
+ (
|
|
|
|
|
+ event_id,
|
|
|
|
|
+ title,
|
|
|
|
|
+ start_at,
|
|
|
|
|
+ timezone,
|
|
|
|
|
+ recurrence,
|
|
|
|
|
+ reminder_minutes,
|
|
|
|
|
+ )
|
|
|
|
|
+ )
|
|
|
|
|
+ await self._wait_at_barrier("schedule", self.schedule_release)
|
|
|
|
|
+ self.schedule_finished.set()
|
|
|
|
|
+ return {
|
|
|
|
|
+ "tool": "calendar.schedule.create",
|
|
|
|
|
+ "status": "created",
|
|
|
|
|
+ "event_id": event_id,
|
|
|
|
|
+ "schedule": {
|
|
|
|
|
+ "title": title,
|
|
|
|
|
+ "start_at": start_at,
|
|
|
|
|
+ "timezone": timezone,
|
|
|
|
|
+ },
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ async def search(
|
|
|
|
|
+ self,
|
|
|
|
|
+ event_id: str,
|
|
|
|
|
+ *,
|
|
|
|
|
+ query: str,
|
|
|
|
|
+ max_results: int = 3,
|
|
|
|
|
+ ) -> dict[str, Any]:
|
|
|
|
|
+ self.search_calls.append((event_id, query, max_results))
|
|
|
|
|
+ return {
|
|
|
|
|
+ "tool": "knowledge.web.search",
|
|
|
|
|
+ "query": query,
|
|
|
|
|
+ "sources": [
|
|
|
|
|
+ {
|
|
|
|
|
+ "title": "Matrix source",
|
|
|
|
|
+ "url": "https://example.invalid/matrix",
|
|
|
|
|
+ "snippet": f"Deterministic context for: {query}",
|
|
|
|
|
+ }
|
|
|
|
|
+ ],
|
|
|
|
|
+ "retrieved_at": "2026-07-13T00:00:00Z",
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ async def _wait_at_barrier(
|
|
|
|
|
+ self,
|
|
|
|
|
+ name: str,
|
|
|
|
|
+ release: asyncio.Event,
|
|
|
|
|
+ ) -> None:
|
|
|
|
|
+ if not self.parallel_barrier:
|
|
|
|
|
+ return
|
|
|
|
|
+ self.started.add(name)
|
|
|
|
|
+ if len(self.started) == 2:
|
|
|
|
|
+ self.both_started.set()
|
|
|
|
|
+ await release.wait()
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _event(call_id: str, name: str, arguments: dict[str, Any]) -> ToolCallEvent:
|
|
|
|
|
+ return ToolCallEvent(
|
|
|
|
|
+ id=call_id,
|
|
|
|
|
+ name=name,
|
|
|
|
|
+ arguments=arguments,
|
|
|
|
|
+ raw_arguments=json.dumps(arguments, sort_keys=True, separators=(",", ":")),
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _event_item(mode: str, event: ToolCallEvent) -> StreamItem:
|
|
|
|
|
+ if mode == "dual_agent":
|
|
|
|
|
+ return StreamItem.text_event(event)
|
|
|
|
|
+ return StreamItem.provider_tool_call(event)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _scenario_setup(
|
|
|
|
|
+ mode: str,
|
|
|
|
|
+ scenario: str,
|
|
|
|
|
+) -> tuple[DebugRunRequest, list[list[StreamItem]]]:
|
|
|
|
|
+ if scenario == "ordinary_chat":
|
|
|
|
|
+ user_message = "Hello."
|
|
|
|
|
+ enabled_tools = list(BUILTIN_EVENT_NAMES)
|
|
|
|
|
+ rounds = [[StreamItem.message_delta("Ordinary answer.")]]
|
|
|
|
|
+ elif scenario == "device_volume_silent":
|
|
|
|
|
+ user_message = "set volume to 30"
|
|
|
|
|
+ enabled_tools = ["device.volume.adjust"]
|
|
|
|
|
+ rounds = [
|
|
|
|
|
+ [
|
|
|
|
|
+ _event_item(
|
|
|
|
|
+ mode,
|
|
|
|
|
+ _event(
|
|
|
|
|
+ "volume-1",
|
|
|
|
|
+ "device.volume.adjust",
|
|
|
|
|
+ {"mode": "absolute", "value": 30},
|
|
|
|
|
+ ),
|
|
|
|
|
+ )
|
|
|
|
|
+ ]
|
|
|
|
|
+ ]
|
|
|
|
|
+ elif scenario == "calendar_schedule_template":
|
|
|
|
|
+ user_message = (
|
|
|
|
|
+ 'schedule "Design review" at 2026-07-14T09:30:00+08:00 '
|
|
|
|
|
+ "timezone Asia/Shanghai"
|
|
|
|
|
+ )
|
|
|
|
|
+ enabled_tools = ["calendar.schedule.create"]
|
|
|
|
|
+ rounds = [
|
|
|
|
|
+ [
|
|
|
|
|
+ _event_item(
|
|
|
|
|
+ mode,
|
|
|
|
|
+ _event(
|
|
|
|
|
+ "schedule-1",
|
|
|
|
|
+ "calendar.schedule.create",
|
|
|
|
|
+ SCHEDULE_ARGUMENTS,
|
|
|
|
|
+ ),
|
|
|
|
|
+ )
|
|
|
|
|
+ ]
|
|
|
|
|
+ ]
|
|
|
|
|
+ elif scenario == "web_search_two_answers":
|
|
|
|
|
+ user_message = "event batch safety"
|
|
|
|
|
+ enabled_tools = ["knowledge.web.search"]
|
|
|
|
|
+ rounds = [
|
|
|
|
|
+ [
|
|
|
|
|
+ StreamItem.message_delta("I will check."),
|
|
|
|
|
+ _event_item(
|
|
|
|
|
+ mode,
|
|
|
|
|
+ _event(
|
|
|
|
|
+ "search-1",
|
|
|
|
|
+ "knowledge.web.search",
|
|
|
|
|
+ {"query": user_message},
|
|
|
|
|
+ ),
|
|
|
|
|
+ ),
|
|
|
|
|
+ ],
|
|
|
|
|
+ [StreamItem.message_delta("Grounded search update.")],
|
|
|
|
|
+ ]
|
|
|
|
|
+ elif scenario == "session_terminate":
|
|
|
|
|
+ user_message = "end this session"
|
|
|
|
|
+ enabled_tools = ["session.terminate"]
|
|
|
|
|
+ rounds = [
|
|
|
|
|
+ [
|
|
|
|
|
+ _event_item(
|
|
|
|
|
+ mode,
|
|
|
|
|
+ _event("terminate-1", "session.terminate", {}),
|
|
|
|
|
+ )
|
|
|
|
|
+ ]
|
|
|
|
|
+ ]
|
|
|
|
|
+ elif scenario == "parallel_volume_schedule":
|
|
|
|
|
+ user_message = (
|
|
|
|
|
+ 'set volume to 30 and schedule "Design review" at '
|
|
|
|
|
+ "2026-07-14T09:30:00+08:00 timezone Asia/Shanghai"
|
|
|
|
|
+ )
|
|
|
|
|
+ enabled_tools = ["device.volume.adjust", "calendar.schedule.create"]
|
|
|
|
|
+ rounds = [
|
|
|
|
|
+ [
|
|
|
|
|
+ _event_item(
|
|
|
|
|
+ mode,
|
|
|
|
|
+ _event(
|
|
|
|
|
+ "parallel-volume",
|
|
|
|
|
+ "device.volume.adjust",
|
|
|
|
|
+ {"mode": "absolute", "value": 30},
|
|
|
|
|
+ ),
|
|
|
|
|
+ ),
|
|
|
|
|
+ _event_item(
|
|
|
|
|
+ mode,
|
|
|
|
|
+ _event(
|
|
|
|
|
+ "parallel-schedule",
|
|
|
|
|
+ "calendar.schedule.create",
|
|
|
|
|
+ SCHEDULE_ARGUMENTS,
|
|
|
|
|
+ ),
|
|
|
|
|
+ ),
|
|
|
|
|
+ ]
|
|
|
|
|
+ ]
|
|
|
|
|
+ else:
|
|
|
|
|
+ raise AssertionError(f"unknown comparison scenario: {scenario}")
|
|
|
|
|
+
|
|
|
|
|
+ request = DebugRunRequest(
|
|
|
|
|
+ user_message=user_message,
|
|
|
|
|
+ system_prompts=["Exercise the requested business scenario."],
|
|
|
|
|
+ pre_messages=[],
|
|
|
|
|
+ chat_agent=AgentParams(
|
|
|
|
|
+ model="comparison-model",
|
|
|
|
|
+ temperature=0.0,
|
|
|
|
|
+ max_tokens=128,
|
|
|
|
|
+ ),
|
|
|
|
|
+ event_agent=EventAgentParams(
|
|
|
|
|
+ model="comparison-model",
|
|
|
|
|
+ temperature=0.0,
|
|
|
|
|
+ max_tokens=128,
|
|
|
|
|
+ enabled_tools=enabled_tools,
|
|
|
|
|
+ max_event_loops=1,
|
|
|
|
|
+ max_parallel_events=2,
|
|
|
|
|
+ batch_timeout_seconds=1.0,
|
|
|
|
|
+ ),
|
|
|
|
|
+ tool_invocation_mode=mode,
|
|
|
|
|
+ )
|
|
|
|
|
+ return request, rounds
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+async def _collect_outputs(
|
|
|
|
|
+ stream: AsyncIterator[dict[str, Any]],
|
|
|
|
|
+) -> list[dict[str, Any]]:
|
|
|
|
|
+ return [message async for message in stream]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _tool_payloads(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|
|
|
|
+ return [
|
|
|
|
|
+ json.loads(message["message"]["content"])
|
|
|
|
|
+ for message in outputs
|
|
|
|
|
+ if message["type"] == "tool_result"
|
|
|
|
|
+ ]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def _assert_port_calls(scenario: str, ports: RecordingPorts) -> None:
|
|
|
|
|
+ expected_session: list[tuple[str, str | None]] = []
|
|
|
|
|
+ expected_volume: list[tuple[str, str, int | None, int | None]] = []
|
|
|
|
|
+ expected_calendar: list[
|
|
|
|
|
+ tuple[str, str, str, str, str | None, int | None]
|
|
|
|
|
+ ] = []
|
|
|
|
|
+ expected_search: list[tuple[str, str, int]] = []
|
|
|
|
|
+
|
|
|
|
|
+ if scenario == "device_volume_silent":
|
|
|
|
|
+ expected_volume = [("volume-1", "absolute", 30, None)]
|
|
|
|
|
+ elif scenario == "calendar_schedule_template":
|
|
|
|
|
+ expected_calendar = [
|
|
|
|
|
+ (
|
|
|
|
|
+ "schedule-1",
|
|
|
|
|
+ "Design review",
|
|
|
|
|
+ "2026-07-14T09:30:00+08:00",
|
|
|
|
|
+ "Asia/Shanghai",
|
|
|
|
|
+ None,
|
|
|
|
|
+ None,
|
|
|
|
|
+ )
|
|
|
|
|
+ ]
|
|
|
|
|
+ elif scenario == "web_search_two_answers":
|
|
|
|
|
+ expected_search = [("search-1", "event batch safety", 3)]
|
|
|
|
|
+ elif scenario == "session_terminate":
|
|
|
|
|
+ expected_session = [("terminate-1", None)]
|
|
|
|
|
+ elif scenario == "parallel_volume_schedule":
|
|
|
|
|
+ expected_volume = [("parallel-volume", "absolute", 30, None)]
|
|
|
|
|
+ expected_calendar = [
|
|
|
|
|
+ (
|
|
|
|
|
+ "parallel-schedule",
|
|
|
|
|
+ "Design review",
|
|
|
|
|
+ "2026-07-14T09:30:00+08:00",
|
|
|
|
|
+ "Asia/Shanghai",
|
|
|
|
|
+ None,
|
|
|
|
|
+ None,
|
|
|
|
|
+ )
|
|
|
|
|
+ ]
|
|
|
|
|
+
|
|
|
|
|
+ assert ports.session_calls == expected_session
|
|
|
|
|
+ assert ports.volume_calls == expected_volume
|
|
|
|
|
+ assert ports.calendar_calls == expected_calendar
|
|
|
|
|
+ assert ports.search_calls == expected_search
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+@pytest.mark.asyncio
|
|
|
|
|
+@pytest.mark.parametrize("mode", MODES)
|
|
|
|
|
+@pytest.mark.parametrize("scenario", SCENARIOS)
|
|
|
|
|
+async def test_invocation_modes_share_business_semantics(
|
|
|
|
|
+ mode: str,
|
|
|
|
|
+ scenario: str,
|
|
|
|
|
+) -> None:
|
|
|
|
|
+ request, rounds = _scenario_setup(mode, scenario)
|
|
|
|
|
+ ports = RecordingPorts(parallel_barrier=scenario == "parallel_volume_schedule")
|
|
|
|
|
+ registry = build_default_tool_registry(
|
|
|
|
|
+ session_termination_port=ports,
|
|
|
|
|
+ device_volume_port=ports,
|
|
|
|
|
+ calendar_schedule_port=ports,
|
|
|
|
|
+ web_search_port=ports,
|
|
|
|
|
+ )
|
|
|
|
|
+ client = ScriptedChatClient(rounds)
|
|
|
|
|
+ run_task = asyncio.create_task(
|
|
|
|
|
+ _collect_outputs(DebugRuntime(client, registry=registry).run(request))
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ try:
|
|
|
|
|
+ if scenario == "parallel_volume_schedule":
|
|
|
|
|
+ await asyncio.wait_for(ports.both_started.wait(), timeout=1)
|
|
|
|
|
+ assert ports.started == {"volume", "schedule"}
|
|
|
|
|
+ ports.schedule_release.set()
|
|
|
|
|
+ await asyncio.wait_for(ports.schedule_finished.wait(), timeout=1)
|
|
|
|
|
+ assert not run_task.done()
|
|
|
|
|
+ ports.volume_release.set()
|
|
|
|
|
+ outputs = await asyncio.wait_for(run_task, timeout=2)
|
|
|
|
|
+ finally:
|
|
|
|
|
+ ports.schedule_release.set()
|
|
|
|
|
+ ports.volume_release.set()
|
|
|
|
|
+ if not run_task.done():
|
|
|
|
|
+ run_task.cancel()
|
|
|
|
|
+ await asyncio.gather(run_task, return_exceptions=True)
|
|
|
|
|
+
|
|
|
|
|
+ expected = EXPECTED[scenario]
|
|
|
|
|
+ visible_messages = tuple(
|
|
|
|
|
+ message["content"]
|
|
|
|
|
+ for message in outputs
|
|
|
|
|
+ if message["type"] == "message_delta"
|
|
|
|
|
+ )
|
|
|
|
|
+ tool_payloads = _tool_payloads(outputs)
|
|
|
|
|
+ terminal_count = sum(
|
|
|
|
|
+ message.get("event") == "terminal_completed" for message in outputs
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ assert visible_messages == expected.visible_messages
|
|
|
|
|
+ assert tuple(payload["tool"] for payload in tool_payloads) == expected.tool_names
|
|
|
|
|
+ assert client.calls == expected.model_calls
|
|
|
|
|
+ assert terminal_count == expected.terminal_count
|
|
|
|
|
+ assert outputs[-1] == {"type": "done"}
|
|
|
|
|
+ _assert_port_calls(scenario, ports)
|
|
|
|
|
+
|
|
|
|
|
+ if scenario == "web_search_two_answers":
|
|
|
|
|
+ assert tool_payloads[0]["sources"][0]["title"] == "Matrix source"
|
|
|
|
|
+ business = [message for message in outputs if message["type"] != "audit"]
|
|
|
|
|
+ visible_indexes = [
|
|
|
|
|
+ index
|
|
|
|
|
+ for index, message in enumerate(business)
|
|
|
|
|
+ if message["type"] == "message_delta"
|
|
|
|
|
+ ]
|
|
|
|
|
+ tool_index = next(
|
|
|
|
|
+ index
|
|
|
|
|
+ for index, message in enumerate(business)
|
|
|
|
|
+ if message["type"] == "tool_result"
|
|
|
|
|
+ )
|
|
|
|
|
+ assert visible_indexes[0] < tool_index < visible_indexes[1]
|