From 7b371b96ed2cd5d7326f956d3702928e7a49beb1 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Tue, 14 Oct 2025 13:14:19 +0300 Subject: [PATCH] Use TestModel from pydantic AI where appropriate --- tests/test_client.py | 50 ++++++--- tests/test_deep_qa.py | 134 ++++++----------------- tests/test_research_graph_integration.py | 127 +++++---------------- 3 files changed, 98 insertions(+), 213 deletions(-) diff --git a/tests/test_client.py b/tests/test_client.py index 338075b5..6d33ac40 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -538,33 +538,51 @@ async def test_client_create_document_with_custom_chunks(temp_db_path): @pytest.mark.asyncio -async def test_client_ask_without_cite(temp_db_path): +async def test_client_ask_without_cite(monkeypatch, temp_db_path): """Test asking questions without citations.""" + from pydantic_ai.models.test import TestModel + + # Mock OpenAIChatModel to return TestModel + monkeypatch.setattr( + "haiku.rag.qa.agent.OpenAIChatModel", lambda **kwargs: TestModel() + ) + async with HaikuRAG(temp_db_path) as client: - # Mock the QA agent - mock_qa_agent = AsyncMock() - mock_qa_agent.answer.return_value = "Test answer" + # Create a test document for the agent to search + await client.create_document( + content="Python is a high-level programming language.", uri="test.txt" + ) - with patch("haiku.rag.qa.get_qa_agent", return_value=mock_qa_agent): - answer = await client.ask("What is Python?") + # Use real QA agent with TestModel + answer = await client.ask("What is Python?") - assert answer == "Test answer" - mock_qa_agent.answer.assert_called_once_with("What is Python?") + # TestModel will generate a valid string response + assert answer is not None + assert isinstance(answer, str) @pytest.mark.asyncio -async def test_client_ask_with_cite(temp_db_path): +async def test_client_ask_with_cite(monkeypatch, temp_db_path): """Test asking questions with citations.""" + from pydantic_ai.models.test import TestModel + + # Mock OpenAIChatModel to return TestModel + monkeypatch.setattr( + "haiku.rag.qa.agent.OpenAIChatModel", lambda **kwargs: TestModel() + ) + async with HaikuRAG(temp_db_path) as client: - # Mock the QA agent - mock_qa_agent = AsyncMock() - mock_qa_agent.answer.return_value = "Test answer with citations [1]" + # Create a test document + await client.create_document( + content="Python is a high-level programming language.", uri="test.txt" + ) - with patch("haiku.rag.qa.get_qa_agent", return_value=mock_qa_agent): - answer = await client.ask("What is Python?", cite=True) + # Use real QA agent with TestModel + answer = await client.ask("What is Python?", cite=True) - assert answer == "Test answer with citations [1]" - mock_qa_agent.answer.assert_called_once_with("What is Python?") + # TestModel will generate a valid string response + assert answer is not None + assert isinstance(answer, str) @pytest.mark.asyncio diff --git a/tests/test_deep_qa.py b/tests/test_deep_qa.py index ee9b1a1d..2c494b15 100644 --- a/tests/test_deep_qa.py +++ b/tests/test_deep_qa.py @@ -1,22 +1,17 @@ -from typing import Any, cast - import pytest +from pydantic_ai.models.test import TestModel +from haiku.rag.client import HaikuRAG from haiku.rag.graph.models import SearchAnswer from haiku.rag.qa.deep.dependencies import DeepQAContext from haiku.rag.qa.deep.graph import build_deep_qa_graph -from haiku.rag.qa.deep.models import DeepQAAnswer -from haiku.rag.qa.deep.nodes import ( - DeepQADecisionNode, - DeepQAPlanNode, - DeepQASearchDispatchNode, - DeepQASynthesizeNode, -) +from haiku.rag.qa.deep.nodes import DeepQAPlanNode from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState @pytest.mark.asyncio -async def test_deep_qa_graph_end_to_end(monkeypatch): +async def test_deep_qa_graph_end_to_end(monkeypatch, temp_db_path): + """Test deep Q&A graph with mocked LLM using TestModel.""" graph = build_deep_qa_graph() state = DeepQAState( @@ -25,116 +20,59 @@ async def test_deep_qa_graph_end_to_end(monkeypatch): ), max_sub_questions=3, ) - deps = DeepQADeps(client=cast(Any, None), console=None) - async def fake_plan_run(self, ctx) -> Any: - ctx.state.context.sub_questions = [ - "Describe haiku.rag in one sentence", - "List core components of haiku.rag", - ] - return DeepQASearchDispatchNode(self.provider, self.model) + # Use real client but with TestModel for LLM calls + client = HaikuRAG(temp_db_path) + deps = DeepQADeps(client=client, console=None) - async def fake_search_dispatch_run(self, ctx) -> Any: - if not ctx.state.context.sub_questions: - return DeepQADecisionNode(self.provider, self.model) + # Mock get_model to return TestModel which generates valid schema-compliant data + def test_model_factory(provider, model): + return TestModel() - batch = ctx.state.context.sub_questions[:] - ctx.state.context.sub_questions.clear() + monkeypatch.setattr("haiku.rag.graph.common.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.qa.deep.nodes.get_model", test_model_factory) - for question in batch: - ctx.state.context.add_qa_response( - SearchAnswer( - query=question, - answer=f"Answer to: {question}", - context=["Context snippet"], - sources=["test.md"], - ) - ) - return DeepQASearchDispatchNode(self.provider, self.model) - - async def fake_decision_run(self, ctx) -> Any: - ctx.state.iterations += 1 - return DeepQASynthesizeNode(self.provider, self.model) - - async def fake_synthesize_run(self, ctx) -> Any: - from pydantic_graph import End - - return End( - DeepQAAnswer( - answer="haiku.rag is a RAG system with components A, B, C.", - sources=["test.md"], - ) - ) - - monkeypatch.setattr(DeepQAPlanNode, "run", fake_plan_run) - monkeypatch.setattr(DeepQASearchDispatchNode, "run", fake_search_dispatch_run) - monkeypatch.setattr(DeepQADecisionNode, "run", fake_decision_run) - monkeypatch.setattr(DeepQASynthesizeNode, "run", fake_synthesize_run) - - start = DeepQAPlanNode(provider="ollama", model="test") + start = DeepQAPlanNode(provider="test", model="test") result = await graph.run(start_node=start, state=state, deps=deps) - assert result.output.answer == "haiku.rag is a RAG system with components A, B, C." - assert result.output.sources == ["test.md"] - assert len(state.context.qa_responses) == 2 + # TestModel will generate valid structured output based on schemas + assert result.output.answer is not None + assert isinstance(result.output.answer, str) + assert isinstance(result.output.sources, list) + + client.close() @pytest.mark.asyncio -async def test_deep_qa_with_citations(monkeypatch): +async def test_deep_qa_with_citations(monkeypatch, temp_db_path): + """Test deep Q&A with citations enabled using TestModel.""" graph = build_deep_qa_graph() state = DeepQAState( context=DeepQAContext(original_question="What is Python?", use_citations=True), max_sub_questions=2, ) - deps = DeepQADeps(client=cast(Any, None), console=None) - async def fake_plan_run(self, ctx) -> Any: - ctx.state.context.sub_questions = ["What is Python used for?"] - return DeepQASearchDispatchNode(self.provider, self.model) + # Use real client but with TestModel for LLM calls + client = HaikuRAG(temp_db_path) + deps = DeepQADeps(client=client, console=None) - async def fake_search_dispatch_run(self, ctx) -> Any: - if not ctx.state.context.sub_questions: - return DeepQADecisionNode(self.provider, self.model) + # Mock get_model to return TestModel + def test_model_factory(provider, model): + return TestModel() - batch = ctx.state.context.sub_questions[:] - ctx.state.context.sub_questions.clear() + monkeypatch.setattr("haiku.rag.graph.common.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.qa.deep.nodes.get_model", test_model_factory) - for question in batch: - ctx.state.context.add_qa_response( - SearchAnswer( - query=question, - answer="Python is used for web development and data science.", - context=["Python snippet"], - sources=["python.md"], - ) - ) - return DeepQASearchDispatchNode(self.provider, self.model) - - async def fake_decision_run(self, ctx) -> Any: - ctx.state.iterations += 1 - return DeepQASynthesizeNode(self.provider, self.model) - - async def fake_synthesize_run(self, ctx) -> Any: - from pydantic_graph import End - - return End( - DeepQAAnswer( - answer="Python is a programming language [python.md].", - sources=["python.md"], - ) - ) - - monkeypatch.setattr(DeepQAPlanNode, "run", fake_plan_run) - monkeypatch.setattr(DeepQASearchDispatchNode, "run", fake_search_dispatch_run) - monkeypatch.setattr(DeepQADecisionNode, "run", fake_decision_run) - monkeypatch.setattr(DeepQASynthesizeNode, "run", fake_synthesize_run) - - start = DeepQAPlanNode(provider="ollama", model="test") + start = DeepQAPlanNode(provider="test", model="test") result = await graph.run(start_node=start, state=state, deps=deps) - assert "[python.md]" in result.output.answer + # Verify citations flag was used assert state.context.use_citations is True + assert result.output.answer is not None + assert isinstance(result.output.sources, list) + + client.close() @pytest.mark.asyncio diff --git a/tests/test_research_graph_integration.py b/tests/test_research_graph_integration.py index e3fdd153..e80ff39e 100644 --- a/tests/test_research_graph_integration.py +++ b/tests/test_research_graph_integration.py @@ -1,32 +1,21 @@ -from typing import Any, cast - import pytest +from pydantic_ai.models.test import TestModel -from haiku.rag.graph.models import SearchAnswer -from haiku.rag.graph.nodes.analysis import AnalyzeInsightsNode, DecisionNode +from haiku.rag.client import HaikuRAG from haiku.rag.graph.nodes.plan import PlanNode -from haiku.rag.graph.nodes.search import SearchDispatchNode -from haiku.rag.graph.nodes.synthesize import SynthesizeNode from haiku.rag.research.dependencies import ResearchContext from haiku.rag.research.graph import ( ResearchDeps, ResearchState, build_research_graph, ) -from haiku.rag.research.models import ( - EvaluationResult, - GapRecord, - GapSeverity, - InsightAnalysis, - InsightRecord, - InsightStatus, - ResearchReport, -) +from haiku.rag.research.models import ResearchReport from haiku.rag.research.stream import stream_research_graph @pytest.mark.asyncio -async def test_graph_end_to_end_with_patched_nodes(monkeypatch): +async def test_graph_end_to_end_with_test_model(monkeypatch, temp_db_path): + """Test research graph with mocked LLM using TestModel.""" graph = build_research_graph() state = ResearchState( @@ -35,103 +24,43 @@ async def test_graph_end_to_end_with_patched_nodes(monkeypatch): confidence_threshold=0.5, max_concurrency=2, ) - deps = ResearchDeps( - client=cast(Any, None), console=None - ) # client unused in patched nodes - async def fake_plan_run(self, ctx) -> Any: - ctx.state.context.sub_questions = [ - "Describe haiku.rag in one sentence", - "List core components of haiku.rag", - ] - ctx.deps.emit_log("planning", ctx.state) - return SearchDispatchNode(self.provider, self.model) + # Use real client but with TestModel for LLM calls + client = HaikuRAG(temp_db_path) + deps = ResearchDeps(client=client, console=None) - async def fake_search_dispatch_run(self, ctx) -> Any: - # Answer all pending questions deterministically, then move to analysis - while ctx.state.context.sub_questions: - q = ctx.state.context.sub_questions.pop(0) - # pydantic BaseModel kwargs not fully typed for pyright - ctx.state.context.add_qa_response( - SearchAnswer(query=q, answer="A", context=["x"], sources=["s"]) # pyright: ignore[reportCallIssue] - ) - ctx.deps.emit_log(f"answered:{q}", ctx.state) - return AnalyzeInsightsNode(self.provider, self.model) + # Mock get_model to return TestModel which generates valid schema-compliant data + # Need to patch in all modules that import it + def test_model_factory(provider, model): + return TestModel() - async def fake_analyze_run(self, ctx) -> Any: - analysis = InsightAnalysis( - highlights=[ - InsightRecord( - summary="haiku.rag orchestrates research stages", - status=InsightStatus.VALIDATED, - supporting_sources=["s"], - originating_questions=["Describe haiku.rag in one sentence"], - ) - ], - gap_assessments=[ - GapRecord( - description="Need a final summary", - severity=GapSeverity.LOW, - blocking=False, - resolved=False, - ) - ], - resolved_gaps=[], - new_questions=[], - commentary="Insights captured for synthesis", - ) - ctx.state.context.integrate_analysis(analysis) - ctx.state.last_analysis = analysis - ctx.deps.emit_log("analysis", ctx.state) - return DecisionNode(self.provider, self.model) - - async def fake_decision_run(self, ctx) -> Any: - ctx.state.last_eval = EvaluationResult( - key_insights=["haiku.rag coordinates planning, search, and synthesis"], - new_questions=[], - gaps=["Need a final summary"], - confidence_score=1.0, - is_sufficient=True, - reasoning="done", - ) - ctx.state.iterations += 1 - ctx.deps.emit_log("decision", ctx.state) - return SynthesizeNode(self.provider, self.model) - - async def fake_synthesize_run(self, ctx) -> Any: - report = ResearchReport( - title="Haiku RAG", - executive_summary="...", - main_findings=["f1"], - conclusions=["c1"], - limitations=[], - recommendations=[], - sources_summary="s", - ) - from pydantic_graph import End - - return End(report) - - monkeypatch.setattr(PlanNode, "run", fake_plan_run, raising=False) + monkeypatch.setattr("haiku.rag.graph.common.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.nodes.plan.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.nodes.search.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.nodes.analysis.get_model", test_model_factory) monkeypatch.setattr( - SearchDispatchNode, "run", fake_search_dispatch_run, raising=False + "haiku.rag.graph.nodes.synthesize.get_model", test_model_factory ) - monkeypatch.setattr(AnalyzeInsightsNode, "run", fake_analyze_run, raising=False) - monkeypatch.setattr(DecisionNode, "run", fake_decision_run, raising=False) - monkeypatch.setattr(SynthesizeNode, "run", fake_synthesize_run, raising=False) start = PlanNode(provider="test", model="test") collected = [] + report = None async for event in stream_research_graph(graph, start, state, deps): collected.append(event) if event.type == "report": report = event.report break - else: # pragma: no cover - defensive guard - report = None + elif event.type == "error": + pytest.fail(f"Graph execution failed: {event.error}") + # TestModel will generate valid structured output for each node + assert report is not None, ( + f"No report generated. Events collected: {[e.type for e in collected]}" + ) assert isinstance(report, ResearchReport) - assert report.title == "Haiku RAG" - assert len(state.context.qa_responses) == 2 + assert report.title is not None + assert isinstance(report.title, str) assert any(evt.type == "log" for evt in collected) + + client.close()