|
| 1 | +"""测试专用的断言层(TurnResult/wait_reply/turn)。 |
| 2 | +
|
| 3 | +这是从 examples/common/live.py 净移出的断言层——examples 侧不再保留。 |
| 4 | +""" |
| 5 | + |
| 6 | +from __future__ import annotations |
| 7 | + |
| 8 | +from dataclasses import dataclass |
| 9 | +from typing import Any |
| 10 | + |
| 11 | +from tests.support.harness import Run, name |
| 12 | + |
| 13 | + |
| 14 | +@dataclass |
| 15 | +class TurnResult: |
| 16 | + text: str = "" |
| 17 | + last_id: str = "" |
| 18 | + tool_used: bool = False |
| 19 | + complete: bool = False |
| 20 | + |
| 21 | + def observe(self, event: Any) -> None: |
| 22 | + if hasattr(event, "to_dict"): |
| 23 | + event = event.to_dict(mode="json") |
| 24 | + kind = event.get("type") |
| 25 | + if event.get("id"): |
| 26 | + self.last_id = event["id"] |
| 27 | + if kind in ("session.error", "session.status_terminated"): |
| 28 | + raise AssertionError(f"Execution failed: {kind}, event_id={self.last_id}") |
| 29 | + if kind in ("agent.tool_use", "agent.mcp_tool_use"): |
| 30 | + self.tool_used = True |
| 31 | + elif kind == "agent.message": |
| 32 | + # Only the latest completed assistant message can satisfy assertions. |
| 33 | + self.text = "\n".join( |
| 34 | + block.get("text", "") for block in event.get("content", []) if block.get("type") == "text" |
| 35 | + ) |
| 36 | + elif kind == "session.status_idle": |
| 37 | + reason = event.get("stop_reason") |
| 38 | + reason = reason.get("type") if isinstance(reason, dict) else reason |
| 39 | + if reason not in (None, "", "end_turn", "stop_sequence"): |
| 40 | + raise AssertionError(f"Execution stopped early: {reason}, event_id={self.last_id}") |
| 41 | + self.complete = bool(self.text) |
| 42 | + |
| 43 | + def verify(self, expected: list[str], require_tool: bool = False) -> None: |
| 44 | + if not self.complete: |
| 45 | + raise AssertionError(f"No idle state after assistant output; last_event_id={self.last_id}") |
| 46 | + if not all(value in self.text for value in expected): |
| 47 | + raise AssertionError(f"Assistant output is missing expected values; last_event_id={self.last_id}") |
| 48 | + if require_tool and not self.tool_used: |
| 49 | + raise AssertionError(f"No actual tool execution; last_event_id={self.last_id}") |
| 50 | + |
| 51 | + |
| 52 | +def wait_reply(events: Any, run: Run, session_id: str, after: str = "") -> TurnResult: |
| 53 | + result = TurnResult(last_id=after) |
| 54 | + while not result.complete: |
| 55 | + run.remaining() |
| 56 | + page = events.list( |
| 57 | + session_id, |
| 58 | + order="asc", |
| 59 | + limit=100, |
| 60 | + extra_query={"after_id": result.last_id or None, "include_tool_calls": True}, |
| 61 | + timeout=min(run.remaining(), 30), |
| 62 | + ) |
| 63 | + for index, event in enumerate(page): |
| 64 | + if index >= 2000: |
| 65 | + raise AssertionError("Event polling exceeded 2000 events") |
| 66 | + result.observe(event) |
| 67 | + if result.complete: |
| 68 | + break |
| 69 | + if not result.complete: |
| 70 | + run.pause() |
| 71 | + run.output("assistant", result.text) |
| 72 | + return result |
| 73 | + |
| 74 | + |
| 75 | +def turn( |
| 76 | + events: Any, |
| 77 | + run: Run, |
| 78 | + session_id: str, |
| 79 | + prompt: str, |
| 80 | + expected: list[str], |
| 81 | + *, |
| 82 | + require_tool: bool = False, |
| 83 | +) -> TurnResult: |
| 84 | + run.output("user", prompt) |
| 85 | + result = events.send( |
| 86 | + session_id, |
| 87 | + events=[{"type": "user.message", "content": [{"type": "text", "text": prompt}]}], |
| 88 | + extra_headers={"Idempotency-Key": name("event")}, |
| 89 | + ) |
| 90 | + if len(result.data) != 1 or not result.data[0].id: |
| 91 | + raise AssertionError("Send must return exactly one user event ID") |
| 92 | + reply = wait_reply(events, run, session_id, result.data[0].id) |
| 93 | + reply.verify(expected, require_tool) |
| 94 | + return reply |
0 commit comments