| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205 |
- """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
|