"""Tests for the optional agent completion guard.""" from __future__ import annotations from collections.abc import Sequence import pytest from agents.find_agent.completion_guard import ( configure_find_agent_completion_guard, create_find_completion_guard, ) from supply_agent import Agent from supply_agent.agent.loop import AgentLoop from supply_agent.tools.registry import ToolRegistry from supply_agent.types import Message, Role class _FakeLLM: def __init__(self, responses: list[Message]) -> None: self.responses = list(responses) def chat(self, *_args, **_kwargs) -> Message: return self.responses.pop(0) async def achat(self, *_args, **_kwargs) -> Message: return self.responses.pop(0) def _loop( responses: list[Message], *, max_iterations: int = 3, completion_guard=None, ) -> AgentLoop: return AgentLoop( llm=_FakeLLM(responses), # type: ignore[arg-type] tools=ToolRegistry(), system_message=Message(role=Role.SYSTEM, content="system"), messages=[Message(role=Role.USER, content="task")], max_iterations=max_iterations, completion_guard=completion_guard, ) def test_completion_guard_is_disabled_by_default() -> None: agent = Agent() loop = _loop([Message(role=Role.ASSISTANT, content="final")]) result = loop.run() assert agent.completion_guard is None assert result.content == "final" assert result.iterations == 1 def test_agent_forwards_configured_completion_guard_to_loop() -> None: def guard(_response: Message, _messages: Sequence[Message]) -> None: return None agent = Agent(completion_guard=guard) assert agent.completion_guard is guard assert agent._create_loop([]).completion_guard is guard def test_completion_guard_feedback_continues_sync_loop() -> None: calls = 0 def guard(_response: Message, messages: Sequence[Message]) -> str | None: nonlocal calls calls += 1 assert messages[-1].role == Role.ASSISTANT return "run.status 仍为 running,请先更新为终态" if calls == 1 else None loop = _loop( [ Message(role=Role.ASSISTANT, content="尚未完成"), Message(role=Role.ASSISTANT, content="最终结果"), ], completion_guard=guard, ) result = loop.run() assert result.content == "最终结果" assert result.iterations == 2 assert result.messages[-2].content == "run.status 仍为 running,请先更新为终态" def test_completion_guard_feedback_continues_stream_loop() -> None: calls = 0 def guard(_response: Message, _messages: Sequence[Message]) -> str | None: nonlocal calls calls += 1 return "not ready" if calls == 1 else None loop = _loop( [ Message(role=Role.ASSISTANT, content="premature"), Message(role=Role.ASSISTANT, content="done"), ], completion_guard=guard, ) events = list(loop.stream()) assert events[-1].data["content"] == "done" assert events[-1].data["iterations"] == 2 @pytest.mark.asyncio async def test_completion_guard_feedback_continues_async_loop() -> None: calls = 0 def guard(_response: Message, _messages: Sequence[Message]) -> str | None: nonlocal calls calls += 1 return "not ready" if calls == 1 else None loop = _loop( [ Message(role=Role.ASSISTANT, content="premature"), Message(role=Role.ASSISTANT, content="done"), ], completion_guard=guard, ) result = await loop.arun() assert result.content == "done" assert result.iterations == 2 def test_iteration_limit_takes_priority_over_completion_guard() -> None: def guard(_response: Message, _messages: Sequence[Message]) -> str: return "run.status 仍为 running" loop = _loop( [ Message(role=Role.ASSISTANT, content="premature"), Message(role=Role.ASSISTANT, content="forced final"), ], max_iterations=1, completion_guard=guard, ) result = loop.run() assert result.content == "forced final" assert result.iterations == 1 @pytest.mark.parametrize("status", ["finished", "failed"]) def test_find_agent_completion_guard_accepts_terminal_status( monkeypatch: pytest.MonkeyPatch, status: str, ) -> None: class _Service: @staticmethod def lookup_run(_run_id: str) -> dict[str, str]: return {"status": status} monkeypatch.setattr( "agents.find_agent.completion_guard.get_video_discovery_service", lambda: _Service(), ) guard = create_find_completion_guard("run-1") assert guard( Message(role=Role.ASSISTANT, content="final"), (), ) is None def test_find_agent_completion_guard_rejects_running_status( monkeypatch: pytest.MonkeyPatch, ) -> None: class _Service: @staticmethod def lookup_run(_run_id: str) -> dict[str, str]: return {"status": "running"} monkeypatch.setattr( "agents.find_agent.completion_guard.get_video_discovery_service", lambda: _Service(), ) guard = create_find_completion_guard("run-1") feedback = guard( Message(role=Role.ASSISTANT, content="premature"), (), ) assert feedback is not None assert "run.status 仍为 running" in feedback assert "update_video_discovery_run_status" in feedback def test_configure_find_agent_completion_guard_only_when_missing() -> None: agent = Agent() configure_find_agent_completion_guard(agent, "scheduled-run") assert agent.completion_guard is not None