test_completion_guard.py 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205
  1. """Tests for the optional agent completion guard."""
  2. from __future__ import annotations
  3. from collections.abc import Sequence
  4. import pytest
  5. from agents.find_agent.completion_guard import (
  6. configure_find_agent_completion_guard,
  7. create_find_completion_guard,
  8. )
  9. from supply_agent import Agent
  10. from supply_agent.agent.loop import AgentLoop
  11. from supply_agent.tools.registry import ToolRegistry
  12. from supply_agent.types import Message, Role
  13. class _FakeLLM:
  14. def __init__(self, responses: list[Message]) -> None:
  15. self.responses = list(responses)
  16. def chat(self, *_args, **_kwargs) -> Message:
  17. return self.responses.pop(0)
  18. async def achat(self, *_args, **_kwargs) -> Message:
  19. return self.responses.pop(0)
  20. def _loop(
  21. responses: list[Message],
  22. *,
  23. max_iterations: int = 3,
  24. completion_guard=None,
  25. ) -> AgentLoop:
  26. return AgentLoop(
  27. llm=_FakeLLM(responses), # type: ignore[arg-type]
  28. tools=ToolRegistry(),
  29. system_message=Message(role=Role.SYSTEM, content="system"),
  30. messages=[Message(role=Role.USER, content="task")],
  31. max_iterations=max_iterations,
  32. completion_guard=completion_guard,
  33. )
  34. def test_completion_guard_is_disabled_by_default() -> None:
  35. agent = Agent()
  36. loop = _loop([Message(role=Role.ASSISTANT, content="final")])
  37. result = loop.run()
  38. assert agent.completion_guard is None
  39. assert result.content == "final"
  40. assert result.iterations == 1
  41. def test_agent_forwards_configured_completion_guard_to_loop() -> None:
  42. def guard(_response: Message, _messages: Sequence[Message]) -> None:
  43. return None
  44. agent = Agent(completion_guard=guard)
  45. assert agent.completion_guard is guard
  46. assert agent._create_loop([]).completion_guard is guard
  47. def test_completion_guard_feedback_continues_sync_loop() -> None:
  48. calls = 0
  49. def guard(_response: Message, messages: Sequence[Message]) -> str | None:
  50. nonlocal calls
  51. calls += 1
  52. assert messages[-1].role == Role.ASSISTANT
  53. return "run.status 仍为 running,请先更新为终态" if calls == 1 else None
  54. loop = _loop(
  55. [
  56. Message(role=Role.ASSISTANT, content="尚未完成"),
  57. Message(role=Role.ASSISTANT, content="最终结果"),
  58. ],
  59. completion_guard=guard,
  60. )
  61. result = loop.run()
  62. assert result.content == "最终结果"
  63. assert result.iterations == 2
  64. assert result.messages[-2].content == "run.status 仍为 running,请先更新为终态"
  65. def test_completion_guard_feedback_continues_stream_loop() -> None:
  66. calls = 0
  67. def guard(_response: Message, _messages: Sequence[Message]) -> str | None:
  68. nonlocal calls
  69. calls += 1
  70. return "not ready" if calls == 1 else None
  71. loop = _loop(
  72. [
  73. Message(role=Role.ASSISTANT, content="premature"),
  74. Message(role=Role.ASSISTANT, content="done"),
  75. ],
  76. completion_guard=guard,
  77. )
  78. events = list(loop.stream())
  79. assert events[-1].data["content"] == "done"
  80. assert events[-1].data["iterations"] == 2
  81. @pytest.mark.asyncio
  82. async def test_completion_guard_feedback_continues_async_loop() -> None:
  83. calls = 0
  84. def guard(_response: Message, _messages: Sequence[Message]) -> str | None:
  85. nonlocal calls
  86. calls += 1
  87. return "not ready" if calls == 1 else None
  88. loop = _loop(
  89. [
  90. Message(role=Role.ASSISTANT, content="premature"),
  91. Message(role=Role.ASSISTANT, content="done"),
  92. ],
  93. completion_guard=guard,
  94. )
  95. result = await loop.arun()
  96. assert result.content == "done"
  97. assert result.iterations == 2
  98. def test_iteration_limit_takes_priority_over_completion_guard() -> None:
  99. def guard(_response: Message, _messages: Sequence[Message]) -> str:
  100. return "run.status 仍为 running"
  101. loop = _loop(
  102. [
  103. Message(role=Role.ASSISTANT, content="premature"),
  104. Message(role=Role.ASSISTANT, content="forced final"),
  105. ],
  106. max_iterations=1,
  107. completion_guard=guard,
  108. )
  109. result = loop.run()
  110. assert result.content == "forced final"
  111. assert result.iterations == 1
  112. @pytest.mark.parametrize("status", ["finished", "failed"])
  113. def test_find_agent_completion_guard_accepts_terminal_status(
  114. monkeypatch: pytest.MonkeyPatch,
  115. status: str,
  116. ) -> None:
  117. class _Service:
  118. @staticmethod
  119. def lookup_run(_run_id: str) -> dict[str, str]:
  120. return {"status": status}
  121. monkeypatch.setattr(
  122. "agents.find_agent.completion_guard.get_video_discovery_service",
  123. lambda: _Service(),
  124. )
  125. guard = create_find_completion_guard("run-1")
  126. assert guard(
  127. Message(role=Role.ASSISTANT, content="final"),
  128. (),
  129. ) is None
  130. def test_find_agent_completion_guard_rejects_running_status(
  131. monkeypatch: pytest.MonkeyPatch,
  132. ) -> None:
  133. class _Service:
  134. @staticmethod
  135. def lookup_run(_run_id: str) -> dict[str, str]:
  136. return {"status": "running"}
  137. monkeypatch.setattr(
  138. "agents.find_agent.completion_guard.get_video_discovery_service",
  139. lambda: _Service(),
  140. )
  141. guard = create_find_completion_guard("run-1")
  142. feedback = guard(
  143. Message(role=Role.ASSISTANT, content="premature"),
  144. (),
  145. )
  146. assert feedback is not None
  147. assert "run.status 仍为 running" in feedback
  148. assert "update_video_discovery_run_status" in feedback
  149. def test_configure_find_agent_completion_guard_only_when_missing() -> None:
  150. agent = Agent()
  151. configure_find_agent_completion_guard(agent, "scheduled-run")
  152. assert agent.completion_guard is not None