Преглед на файлове

test: add invocation mode comparison matrix

Problem:
- The two invocation modes lacked one repeatable matrix for the six shared built-in business scenarios.

Risk:
- Protocol-specific differences could hide parity regressions in silent/template/follow-up/terminal policies and parallel result ordering.
zhenyu.hu преди 2 седмици
родител
ревизия
97a98a9e8b
променени са 1 файла, в които са добавени 482 реда и са изтрити 0 реда
  1. 482 0
      tests/test_tool_invocation_comparison.py

+ 482 - 0
tests/test_tool_invocation_comparison.py

@@ -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]