| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229 |
- """AgentLoop shared-path behavior tests."""
- from __future__ import annotations
- from typing import Any
- import pytest
- from supply_agent.agent.loop import AgentLoop
- from supply_agent.tools.base import tool
- from supply_agent.tools.registry import ToolRegistry
- from supply_agent.types import (
- AgentEventType,
- Message,
- Role,
- ToolCall,
- )
- class _FakeLLM:
- """Deterministic LLM that returns scripted assistant messages."""
- def __init__(self, responses: list[Message]) -> None:
- self._responses = list(responses)
- self.calls: list[dict[str, Any]] = []
- def chat(
- self,
- messages: list[Message],
- tools: list[Any] | None = None,
- temperature: float | None = None,
- *,
- iteration: int = 0,
- ) -> Message:
- self.calls.append(
- {
- "messages": list(messages),
- "tools": tools,
- "iteration": iteration,
- }
- )
- if not self._responses:
- raise AssertionError("FakeLLM: no more scripted responses")
- return self._responses.pop(0)
- async def achat(
- self,
- messages: list[Message],
- tools: list[Any] | None = None,
- temperature: float | None = None,
- *,
- iteration: int = 0,
- ) -> Message:
- return self.chat(
- messages,
- tools=tools,
- temperature=temperature,
- iteration=iteration,
- )
- def _system() -> Message:
- return Message(role=Role.SYSTEM, content="base system")
- def _assistant_text(content: str) -> Message:
- return Message(role=Role.ASSISTANT, content=content)
- def _assistant_tools(*calls: ToolCall) -> Message:
- return Message(role=Role.ASSISTANT, content=None, tool_calls=list(calls))
- @pytest.fixture
- def skill_tools() -> ToolRegistry:
- registry = ToolRegistry()
- @tool(name="load_skill")
- def load_skill(name: str) -> str:
- return f"skill-body:{name}"
- registry.register(load_skill, name="load_skill")
- return registry
- def test_stream_loads_skill_into_system_message(skill_tools: ToolRegistry) -> None:
- active_skills: list[str] = []
- system_parts = {"value": "base system"}
- def builder() -> Message:
- parts = [system_parts["value"]]
- for name in active_skills:
- parts.append(f"LOADED:{name}")
- return Message(role=Role.SYSTEM, content="\n\n".join(parts))
- llm = _FakeLLM(
- [
- _assistant_tools(
- ToolCall(
- id="c1",
- name="load_skill",
- arguments='{"name": "demo"}',
- )
- ),
- _assistant_text("done"),
- ]
- )
- loop = AgentLoop(
- llm=llm, # type: ignore[arg-type]
- tools=skill_tools,
- system_message=builder(),
- system_message_builder=builder,
- messages=[Message(role=Role.USER, content="hi")],
- max_iterations=5,
- active_skills=active_skills,
- )
- events = list(loop.stream())
- assert any(e.type == AgentEventType.DONE for e in events)
- assert "demo" in active_skills
- # Second LLM call must see refreshed system prompt with skill content.
- assert "LOADED:demo" in (llm.calls[1]["messages"][0].content or "")
- def test_stream_nudges_best_answer_on_max_iterations() -> None:
- registry = ToolRegistry()
- @tool(name="noop")
- def noop() -> str:
- return '{"ok": true}'
- registry.register(noop, name="noop")
- llm = _FakeLLM(
- [
- _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
- _assistant_text("final after nudge"),
- ]
- )
- loop = AgentLoop(
- llm=llm, # type: ignore[arg-type]
- tools=registry,
- system_message=_system(),
- messages=[Message(role=Role.USER, content="hi")],
- max_iterations=1,
- )
- events = list(loop.stream())
- done = next(e for e in events if e.type == AgentEventType.DONE)
- assert done.data["content"] == "final after nudge"
- assert len(llm.calls) == 2
- assert "best answer" in (llm.calls[1]["messages"][-1].content or "").lower()
- def test_run_and_stream_share_max_iteration_nudge() -> None:
- registry = ToolRegistry()
- @tool(name="noop")
- def noop() -> str:
- return '{"ok": true}'
- registry.register(noop, name="noop")
- def make_loop(responses: list[Message]) -> AgentLoop:
- return AgentLoop(
- llm=_FakeLLM(responses), # type: ignore[arg-type]
- tools=registry,
- system_message=_system(),
- messages=[Message(role=Role.USER, content="hi")],
- max_iterations=1,
- )
- run_result = make_loop(
- [
- _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
- _assistant_text("from-run"),
- ]
- ).run()
- stream_events = list(
- make_loop(
- [
- _assistant_tools(ToolCall(id="c1", name="noop", arguments="{}")),
- _assistant_text("from-stream"),
- ]
- ).stream()
- )
- stream_done = next(e for e in stream_events if e.type == AgentEventType.DONE)
- assert run_result.content == "from-run"
- assert stream_done.data["content"] == "from-stream"
- assert run_result.iterations == stream_done.data["iterations"] == 1
- @pytest.mark.asyncio
- async def test_astream_loads_skill(skill_tools: ToolRegistry) -> None:
- active_skills: list[str] = []
- def builder() -> Message:
- parts = ["base"]
- for name in active_skills:
- parts.append(f"LOADED:{name}")
- return Message(role=Role.SYSTEM, content="\n\n".join(parts))
- llm = _FakeLLM(
- [
- _assistant_tools(
- ToolCall(
- id="c1",
- name="load_skill",
- arguments='{"name": "async-demo"}',
- )
- ),
- _assistant_text("ok"),
- ]
- )
- loop = AgentLoop(
- llm=llm, # type: ignore[arg-type]
- tools=skill_tools,
- system_message=builder(),
- system_message_builder=builder,
- messages=[Message(role=Role.USER, content="hi")],
- max_iterations=5,
- active_skills=active_skills,
- )
- events = [event async for event in loop.astream()]
- assert any(e.type == AgentEventType.DONE for e in events)
- assert "async-demo" in active_skills
- assert "LOADED:async-demo" in (llm.calls[1]["messages"][0].content or "")
|