Spaces:
Running
Running
File size: 2,076 Bytes
c07ec17 470138f c07ec17 6802438 c07ec17 470138f c07ec17 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 | import pytest
from agent.schemas import Answer
pytest.importorskip("langgraph")
class ScriptedPlanner:
def __init__(self, actions):
self._actions = list(actions)
def __call__(self, prompt, system, provider=None, client=None):
return self._actions.pop(0)
def test_graph_decomposes_and_answers(monkeypatch):
from agent.graph import answer_graph
monkeypatch.setattr(
"agent.llm._raw_completion",
ScriptedPlanner([
'{"action":"search_docs","query":"cnn images"}',
'{"action":"search_docs","query":"dataloader"}',
'{"action":"answer"}',
]),
)
calls = {"n": 0}
def fake_search(query, library=None, kind=None, k=8):
calls["n"] += 1
return {"query": query, "sections": [{"url": f"u{calls['n']}", "anchor": "",
"heading_path": "H", "content": "t"}], "titles": ["H"]}
monkeypatch.setattr("agent.tools.search_docs", fake_search)
captured = {}
def fake_answer(q, s, referrals=None, provider=None, client=None):
captured["s"] = s
return Answer(answer_md="done")
monkeypatch.setattr("agent.grounded.answer_from_sections", fake_answer)
result = answer_graph("how do I build a CNN?")
assert result.answer_md == "done"
# 1 forced seed search + 2 planner-driven searches (parity with the loop)
assert calls["n"] == 3
assert len(captured["s"]) == 3
def test_graph_terminates_on_budget(monkeypatch):
from agent.graph import answer_graph
monkeypatch.setattr(
"agent.llm._raw_completion",
ScriptedPlanner(['{"action":"search_docs","query":"x"}'] * 30),
)
monkeypatch.setattr(
"agent.tools.search_docs",
lambda q, library=None, kind=None, k=8: {"query": q, "sections": [], "titles": []},
)
monkeypatch.setattr(
"agent.grounded.answer_from_sections",
lambda q, s, referrals=None, provider=None, client=None: Answer(answer_md="stopped"),
)
assert answer_graph("q").answer_md == "stopped" # must not recurse forever
|