diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml
index 19c80c7c..c6d5446a 100644
--- a/.pre-commit-config.yaml
+++ b/.pre-commit-config.yaml
@@ -9,7 +9,7 @@ repos:
- id: debug-statements
- repo: https://github.com/astral-sh/ruff-pre-commit
# Ruff version.
- rev: v0.11.4
+ rev: v0.14.3
hooks:
# Run the linter.
- id: ruff
@@ -17,6 +17,6 @@ repos:
- id: ruff-format
- repo: https://github.com/RobertCraigie/pyright-python
- rev: v1.1.399
+ rev: v1.1.407
hooks:
- id: pyright
diff --git a/CHANGELOG.md b/CHANGELOG.md
index 10edf68c..5fa4f7b8 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -1,6 +1,19 @@
# Changelog
## [Unreleased]
+### Added
+
+- Migrated research and deep QA agents to use Pydantic Graph beta API for better graph execution
+- Automatic semaphore-based concurrency control for parallel sub-question processing
+- `max_concurrency` parameter for controlling parallel execution in research and deep QA (default: 1)
+
+### Changed
+
+- **BREAKING**: Research and Deep QA graphs now use `pydantic_graph.beta` instead of the class-based graph implementation
+- Refactored graph common patterns into `graph_common` module
+- Sub-questions now process using `.map()` for true parallel execution
+- Improved graph structure with cleaner node definitions and flow control
+
## [0.14.0] - 2024-11-05
### Added
diff --git a/README.md b/README.md
index 6c285fb6..eb0a429b 100644
--- a/README.md
+++ b/README.md
@@ -88,8 +88,8 @@ To customize settings, create a `haiku.rag.yaml` config file (see [Configuration
```python
from haiku.rag.client import HaikuRAG
+from haiku.rag.config import Config
from haiku.rag.research import (
- PlanNode,
ResearchContext,
ResearchDeps,
ResearchState,
@@ -115,34 +115,22 @@ async with HaikuRAG("database.lancedb") as client:
print(answer)
# Multi‑agent research pipeline (Plan → Search → Evaluate → Synthesize)
- graph = build_research_graph()
+ # Graph settings (provider, model, max_iterations, etc.) come from config
+ graph = build_research_graph(config=Config)
question = (
"What are the main drivers and trends of global temperature "
"anomalies since 1990?"
)
- state = ResearchState(
- context=ResearchContext(original_question=question),
- max_iterations=2,
- confidence_threshold=0.8,
- max_concurrency=2,
- )
+ context = ResearchContext(original_question=question)
+ state = ResearchState.from_config(context=context, config=Config)
deps = ResearchDeps(client=client)
# Blocking run (final result only)
- result = await graph.run(
- PlanNode(provider="openai", model="gpt-4o-mini"),
- state=state,
- deps=deps,
- )
- print(result.output.title)
+ report = await graph.run(state=state, deps=deps)
+ print(report.title)
# Streaming progress (log/report/error events)
- async for event in stream_research_graph(
- graph,
- PlanNode(provider="openai", model="gpt-4o-mini"),
- state,
- deps,
- ):
+ async for event in stream_research_graph(graph, state, deps):
if event.type == "log":
iteration = event.state.iterations if event.state else state.iterations
print(f"[{iteration}] {event.message}")
diff --git a/docs/agents.md b/docs/agents.md
index 935202ec..e9fb688a 100644
--- a/docs/agents.md
+++ b/docs/agents.md
@@ -26,18 +26,17 @@ Python usage:
from haiku.rag.client import HaikuRAG
from haiku.rag.qa.agent import QuestionAnswerAgent
-client = HaikuRAG(path_to_db)
+async with HaikuRAG(path_to_db) as client:
+ # Choose a provider and model (see Configuration for env defaults)
+ agent = QuestionAnswerAgent(
+ client=client,
+ provider="openai", # or "ollama", "vllm", etc.
+ model="gpt-4o-mini",
+ use_citations=False, # set True to bias prompt towards citing sources
+ )
-# Choose a provider and model (see Configuration for env defaults)
-agent = QuestionAnswerAgent(
- client=client,
- provider="openai", # or "ollama", "vllm", etc.
- model="gpt-4o-mini",
- use_citations=False, # set True to bias prompt towards citing sources
-)
-
-answer = await agent.answer("What is climate change?")
-print(answer)
+ answer = await agent.answer("What is climate change?")
+ print(answer)
```
### Deep QA Agent
@@ -49,19 +48,25 @@ Deep QA is a multi-agent system that decomposes complex questions into sub-quest
title: Deep QA graph
---
stateDiagram-v2
- DeepQAPlanNode --> DeepQASearchDispatchNode
- DeepQASearchDispatchNode --> DeepQADecisionNode
- DeepQADecisionNode --> DeepQASearchDispatchNode
- DeepQADecisionNode --> DeepQASynthesizeNode
- DeepQASynthesizeNode --> [*]
+ [*] --> plan
+ plan --> get_batch
+ get_batch --> search_one: Has questions (map)
+ get_batch --> synthesize: No questions
+ search_one --> collect_answers
+ collect_answers --> decide
+ decide --> get_batch: Continue QA
+ decide --> synthesize: Done with QA
+ synthesize --> [*]
```
Key nodes:
-- **Plan**: Decomposes the question into focused sub-questions
-- **Search (batched)**: Answers sub-questions in parallel batches (respects max_concurrency)
-- **Decision**: Evaluates if we have sufficient information or need another iteration
-- **Synthesize**: Generates the final comprehensive answer
+- **plan**: Decomposes the question into focused sub-questions using a presearch tool
+- **get_batch**: Retrieves remaining sub-questions for the current iteration
+- **search_one**: Answers a single sub-question using the knowledge base (mapped in parallel)
+- **collect_answers**: Aggregates search results from parallel executions
+- **decide**: Evaluates if sufficient information has been gathered or if more iterations are needed
+- **synthesize**: Generates the final comprehensive answer from all gathered information
Key differences from Research:
@@ -69,7 +74,12 @@ Key differences from Research:
- **Direct answers**: Returns just the answer (not a full research report)
- **Question-focused**: Optimized for answering specific questions, not open-ended research
- **Supports citations**: Can include inline source citations like `[document.md]`
-- **Configurable iterations**: Control max_iterations (default: 2) and max_concurrency (default: 3)
+- **Configurable iterations**: Control max_iterations (default: 2) and max_concurrency (default: 1)
+
+Note on parallel execution:
+- The `search_one` node is mapped over all questions in a batch
+- Parallelism is controlled via `max_concurrency`
+- All questions in an iteration are processed before evaluation
CLI usage:
@@ -85,33 +95,55 @@ Python usage:
```python
from haiku.rag.client import HaikuRAG
+from haiku.rag.config import Config
from haiku.rag.qa.deep.dependencies import DeepQAContext
from haiku.rag.qa.deep.graph import build_deep_qa_graph
-from haiku.rag.qa.deep.nodes import DeepQAPlanNode
from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
async with HaikuRAG(path_to_db) as client:
- graph = build_deep_qa_graph()
+ # Use global config (recommended)
+ graph = build_deep_qa_graph(config=Config)
context = DeepQAContext(
original_question="What are the main features of haiku.rag?",
use_citations=True
)
- state = DeepQAState(
- context=context,
- max_sub_questions=3,
- max_iterations=2,
- max_concurrency=3
- )
+ state = DeepQAState.from_config(context=context, config=Config)
deps = DeepQADeps(client=client)
result = await graph.run(
- start_node=DeepQAPlanNode(provider="openai", model="gpt-4o-mini"),
state=state,
deps=deps
)
- print(result.output.answer)
- print(result.output.sources)
+ print(result.answer)
+ print(result.sources)
+```
+
+Alternative usage with custom config:
+
+```python
+# Create a custom config with different settings
+from haiku.rag.config.models import AppConfig, QAConfig
+
+custom_config = AppConfig(
+ qa=QAConfig(
+ provider="openai",
+ model="gpt-4o-mini",
+ max_sub_questions=5,
+ max_iterations=3,
+ max_concurrency=2,
+ )
+)
+
+graph = build_deep_qa_graph(config=custom_config)
+context = DeepQAContext(
+ original_question="What are the main features of haiku.rag?",
+ use_citations=True
+)
+state = DeepQAState.from_config(context=context, config=custom_config)
+deps = DeepQADeps(client=client)
+
+result = await graph.run(state=state, deps=deps)
```
### Research Graph
@@ -123,21 +155,27 @@ The research workflow is implemented as a typed pydantic‑graph. It plans, sear
title: Research graph
---
stateDiagram-v2
- PlanNode --> SearchDispatchNode
- SearchDispatchNode --> AnalyzeInsightsNode
- AnalyzeInsightsNode --> DecisionNode
- DecisionNode --> SearchDispatchNode
- DecisionNode --> SynthesizeNode
- SynthesizeNode --> [*]
+ [*] --> plan
+ plan --> get_batch
+ get_batch --> search_one: Has questions (map)
+ get_batch --> synthesize: No questions
+ search_one --> collect_answers
+ collect_answers --> analyze_insights
+ analyze_insights --> decide
+ decide --> get_batch: Continue research
+ decide --> synthesize: Done researching
+ synthesize --> [*]
```
Key nodes:
-- Plan: builds up to 3 standalone sub‑questions (uses an internal presearch tool)
-- Search (batched): answers sub‑questions using the KB with minimal, verbatim context
-- Analyze: aggregates fresh insights, updates gaps, and suggests new sub-questions
-- Decision: checks sufficiency/confidence thresholds and chooses whether to iterate
-- Synthesize: generates a final structured report
+- **plan**: Builds up to 3 standalone sub‑questions (uses an internal presearch tool)
+- **get_batch**: Retrieves remaining sub‑questions for the current iteration
+- **search_one**: Answers a single sub‑question using the KB with minimal, verbatim context (mapped in parallel)
+- **collect_answers**: Aggregates search results from parallel executions
+- **analyze_insights**: Synthesizes fresh insights, updates gaps, and suggests new sub-questions
+- **decide**: Checks sufficiency/confidence thresholds and determines whether to continue research
+- **synthesize**: Generates a final structured research report
Primary models:
@@ -147,77 +185,90 @@ Primary models:
- `EvaluationResult` — insights, new questions, sufficiency, confidence
- `ResearchReport` — final report (title, executive summary, findings, conclusions, …)
+Note on parallel execution:
+- The `search_one` node is mapped over all questions in a batch
+- Parallelism is controlled via `max_concurrency`
+- Analysis and decision nodes process results after each batch completes
+
CLI usage:
```bash
-haiku-rag research "How does haiku.rag organize and query documents?" \
- --max-iterations 2 \
- --confidence-threshold 0.8 \
- --max-concurrency 3 \
- --verbose
+# Basic usage (uses config from file or defaults)
+haiku-rag research "How does haiku.rag organize and query documents?" --verbose
+
+# With custom config file
+haiku-rag --config my-research-config.yaml research "How does haiku.rag organize and query documents?" --verbose
```
Python usage (blocking result):
```python
from haiku.rag.client import HaikuRAG
-from haiku.rag.research import (
- PlanNode,
- ResearchContext,
- ResearchDeps,
- ResearchState,
- build_research_graph,
-)
+from haiku.rag.config import Config
+from haiku.rag.research.dependencies import ResearchContext
+from haiku.rag.research.graph import build_research_graph
+from haiku.rag.research.state import ResearchDeps, ResearchState
async with HaikuRAG(path_to_db) as client:
- graph = build_research_graph()
+ # Use global config (recommended)
+ graph = build_research_graph(config=Config)
question = "What are the main drivers and trends of global temperature anomalies since 1990?"
- state = ResearchState(
- context=ResearchContext(original_question=question),
- max_iterations=2,
- confidence_threshold=0.8,
- max_concurrency=2,
- )
+ context = ResearchContext(original_question=question)
+ state = ResearchState.from_config(context=context, config=Config)
deps = ResearchDeps(client=client)
result = await graph.run(
- PlanNode(provider="openai", model="gpt-4o-mini"),
state=state,
deps=deps,
)
- report = result.output
+ report = result
print(report.title)
print(report.executive_summary)
```
+Alternative usage with custom config:
+
+```python
+from haiku.rag.config.models import AppConfig, ResearchConfig
+
+custom_config = AppConfig(
+ research=ResearchConfig(
+ provider="openai",
+ model="gpt-4o-mini",
+ max_iterations=5,
+ confidence_threshold=0.85,
+ max_concurrency=3,
+ )
+)
+
+graph = build_research_graph(config=custom_config)
+context = ResearchContext(original_question=question)
+state = ResearchState.from_config(context=context, config=custom_config)
+deps = ResearchDeps(client=client)
+
+result = await graph.run(state=state, deps=deps)
+```
+
Python usage (streamed events):
```python
from haiku.rag.client import HaikuRAG
-from haiku.rag.research import (
- PlanNode,
- ResearchContext,
- ResearchDeps,
- ResearchState,
- build_research_graph,
- stream_research_graph,
-)
+from haiku.rag.config import Config
+from haiku.rag.research.dependencies import ResearchContext
+from haiku.rag.research.graph import build_research_graph
+from haiku.rag.research.state import ResearchDeps, ResearchState
+from haiku.rag.research.stream import stream_research_graph
async with HaikuRAG(path_to_db) as client:
- graph = build_research_graph()
+ graph = build_research_graph(config=Config)
question = "What are the main drivers and trends of global temperature anomalies since 1990?"
- state = ResearchState(
- context=ResearchContext(original_question=question),
- max_iterations=2,
- confidence_threshold=0.8,
- max_concurrency=2,
- )
+ context = ResearchContext(original_question=question)
+ state = ResearchState.from_config(context=context, config=Config)
deps = ResearchDeps(client=client)
async for event in stream_research_graph(
graph,
- PlanNode(provider="openai", model="gpt-4o-mini"),
state,
deps,
):
diff --git a/examples/ag-ui-research/backend/agent.py b/examples/ag-ui-research/backend/agent.py
index 049e0dbd..5362696a 100644
--- a/examples/ag-ui-research/backend/agent.py
+++ b/examples/ag-ui-research/backend/agent.py
@@ -8,7 +8,7 @@ from pydantic_ai.ag_ui import StateDeps
from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config
-from haiku.rag.graph.common import get_model
+from haiku.rag.graph_common import get_model
class ResearchState(BaseModel):
diff --git a/haiku_rag_slim/haiku/rag/a2a/__init__.py b/haiku_rag_slim/haiku/rag/a2a/__init__.py
index 892056eb..76e1ee36 100644
--- a/haiku_rag_slim/haiku/rag/a2a/__init__.py
+++ b/haiku_rag_slim/haiku/rag/a2a/__init__.py
@@ -6,7 +6,7 @@ import logfire
from pydantic_ai import Agent, RunContext
from haiku.rag.config import Config
-from haiku.rag.graph.common import get_model
+from haiku.rag.graph_common import get_model
from .context import load_message_history, save_message_history
from .models import AgentDependencies, SearchResult
@@ -138,7 +138,11 @@ def create_a2a_app(
if security_schemes or security:
# Monkey-patch the agent card endpoint to include security
async def _agent_card_endpoint_with_security(request):
- from fasta2a.schema import AgentCapabilities, AgentCard, agent_card_ta
+ from fasta2a.schema import ( # type: ignore
+ AgentCapabilities,
+ AgentCard,
+ agent_card_ta,
+ )
from starlette.responses import Response
if app._agent_card_json_schema is None:
diff --git a/haiku_rag_slim/haiku/rag/app.py b/haiku_rag_slim/haiku/rag/app.py
index ed8b6436..6aa2c178 100644
--- a/haiku_rag_slim/haiku/rag/app.py
+++ b/haiku_rag_slim/haiku/rag/app.py
@@ -13,12 +13,8 @@ from haiku.rag.config import Config
from haiku.rag.mcp import create_mcp_server
from haiku.rag.monitor import FileWatcher
from haiku.rag.research.dependencies import ResearchContext
-from haiku.rag.research.graph import (
- PlanNode,
- ResearchDeps,
- ResearchState,
- build_research_graph,
-)
+from haiku.rag.research.graph import build_research_graph
+from haiku.rag.research.state import ResearchDeps, ResearchState
from haiku.rag.research.stream import stream_research_graph
from haiku.rag.store.models.chunk import Chunk
from haiku.rag.store.models.document import Document
@@ -208,34 +204,35 @@ class HaikuRAGApp:
deep: bool = False,
verbose: bool = False,
):
+ """Ask a question using the RAG system.
+
+ Args:
+ question: The question to ask
+ cite: Include citations in the answer
+ deep: Use deep QA mode (multi-step reasoning)
+ verbose: Show verbose output
+ """
async with HaikuRAG(db_path=self.db_path) as self.client:
try:
if deep:
from rich.console import Console
+ from haiku.rag.config import Config
from haiku.rag.qa.deep.dependencies import DeepQAContext
from haiku.rag.qa.deep.graph import build_deep_qa_graph
- from haiku.rag.qa.deep.nodes import DeepQAPlanNode
from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
- graph = build_deep_qa_graph()
+ graph = build_deep_qa_graph(config=Config)
context = DeepQAContext(
original_question=question, use_citations=cite
)
- state = DeepQAState(context=context)
+ state = DeepQAState.from_config(context=context, config=Config)
deps = DeepQADeps(
client=self.client, console=Console() if verbose else None
)
- start_node = DeepQAPlanNode(
- provider=Config.qa.provider,
- model=Config.qa.model,
- )
-
- result = await graph.run(
- start_node=start_node, state=state, deps=deps
- )
- answer = result.output.answer
+ result = await graph.run(state=state, deps=deps)
+ answer = result.answer
else:
answer = await self.client.ask(question, cite=cite)
@@ -249,37 +246,32 @@ class HaikuRAGApp:
async def research(
self,
question: str,
- max_iterations: int = 3,
- confidence_threshold: float = 0.8,
- max_concurrency: int = 1,
verbose: bool = False,
):
- """Run research via the pydantic-graph pipeline (default)."""
+ """Run research via the pydantic-graph pipeline.
+
+ Args:
+ question: The research question
+ verbose: Show verbose output
+ """
async with HaikuRAG(db_path=self.db_path) as client:
try:
+ from haiku.rag.config import Config
+
if verbose:
self.console.print("[bold cyan]Starting research[/bold cyan]")
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
self.console.print()
- graph = build_research_graph()
+ graph = build_research_graph(config=Config)
context = ResearchContext(original_question=question)
- state = ResearchState(
- context=context,
- max_iterations=max_iterations,
- confidence_threshold=confidence_threshold,
- max_concurrency=max_concurrency,
- )
+ state = ResearchState.from_config(context=context, config=Config)
deps = ResearchDeps(
client=client, console=self.console if verbose else None
)
- start = PlanNode(
- provider=Config.research.provider or Config.qa.provider,
- model=Config.research.model or Config.qa.model,
- )
report = None
- async for event in stream_research_graph(graph, start, state, deps):
+ async for event in stream_research_graph(graph, state, deps):
if event.type == "report":
report = event.report
break
diff --git a/haiku_rag_slim/haiku/rag/cli.py b/haiku_rag_slim/haiku/rag/cli.py
index 5c082986..edefec59 100644
--- a/haiku_rag_slim/haiku/rag/cli.py
+++ b/haiku_rag_slim/haiku/rag/cli.py
@@ -290,22 +290,6 @@ def research(
question: str = typer.Argument(
help="The research question to investigate",
),
- max_iterations: int = typer.Option(
- 3,
- "--max-iterations",
- "-n",
- help="Maximum search/analyze iterations",
- ),
- confidence_threshold: float = typer.Option(
- 0.8,
- "--confidence-threshold",
- help="Minimum confidence (0-1) to stop",
- ),
- max_concurrency: int = typer.Option(
- 1,
- "--max-concurrency",
- help="Max concurrent searches per iteration (planned)",
- ),
db: Path = typer.Option(
Config.storage.data_dir / "haiku.rag.lancedb",
"--db",
@@ -320,15 +304,7 @@ def research(
from haiku.rag.app import HaikuRAGApp
app = HaikuRAGApp(db_path=db)
- asyncio.run(
- app.research(
- question=question,
- max_iterations=max_iterations,
- confidence_threshold=confidence_threshold,
- max_concurrency=max_concurrency,
- verbose=verbose,
- )
- )
+ asyncio.run(app.research(question=question, verbose=verbose))
@cli.command("settings", help="Display current configuration settings")
diff --git a/haiku_rag_slim/haiku/rag/config/models.py b/haiku_rag_slim/haiku/rag/config/models.py
index 4c654e7a..801c6dcb 100644
--- a/haiku_rag_slim/haiku/rag/config/models.py
+++ b/haiku_rag_slim/haiku/rag/config/models.py
@@ -37,11 +37,17 @@ class RerankingConfig(BaseModel):
class QAConfig(BaseModel):
provider: str = "ollama"
model: str = "gpt-oss"
+ max_sub_questions: int = 3
+ max_iterations: int = 2
+ max_concurrency: int = 1
class ResearchConfig(BaseModel):
provider: str = "ollama"
model: str = "gpt-oss"
+ max_iterations: int = 3
+ confidence_threshold: float = 0.8
+ max_concurrency: int = 1
class ProcessingConfig(BaseModel):
diff --git a/haiku_rag_slim/haiku/rag/graph/__init__.py b/haiku_rag_slim/haiku/rag/graph/__init__.py
deleted file mode 100644
index 18013005..00000000
--- a/haiku_rag_slim/haiku/rag/graph/__init__.py
+++ /dev/null
@@ -1 +0,0 @@
-from haiku.rag.graph.models import ResearchPlan, SearchAnswer
diff --git a/haiku_rag_slim/haiku/rag/graph/base.py b/haiku_rag_slim/haiku/rag/graph/base.py
deleted file mode 100644
index 1d78fa7a..00000000
--- a/haiku_rag_slim/haiku/rag/graph/base.py
+++ /dev/null
@@ -1,31 +0,0 @@
-from typing import Protocol, runtime_checkable
-
-from pydantic import BaseModel, Field
-from rich.console import Console
-
-from haiku.rag.client import HaikuRAG
-from haiku.rag.graph.models import SearchAnswer
-
-
-@runtime_checkable
-class GraphContext(Protocol):
- """Protocol for graph context objects."""
-
- original_question: str
- sub_questions: list[str]
- qa_responses: list[SearchAnswer]
-
- def add_qa_response(self, qa: SearchAnswer) -> None: ...
-
-
-class BaseGraphDeps(BaseModel):
- """Base dependencies for graph nodes."""
-
- model_config = {"arbitrary_types_allowed": True}
-
- client: HaikuRAG = Field(description="RAG client for document operations")
- console: Console | None = None
-
- def emit_log(self, message: str) -> None:
- if self.console:
- self.console.print(message)
diff --git a/haiku_rag_slim/haiku/rag/graph/common.py b/haiku_rag_slim/haiku/rag/graph/common.py
deleted file mode 100644
index 31877206..00000000
--- a/haiku_rag_slim/haiku/rag/graph/common.py
+++ /dev/null
@@ -1,33 +0,0 @@
-from typing import Any, Protocol
-
-from pydantic_ai.models.openai import OpenAIChatModel
-from pydantic_ai.providers.ollama import OllamaProvider
-from pydantic_ai.providers.openai import OpenAIProvider
-
-from haiku.rag.config import Config
-
-
-class HasEmitLog(Protocol):
- def emit_log(self, message: str, state: Any = None) -> None: ...
-
-
-def get_model(provider: str, model: str) -> Any:
- if provider == "ollama":
- return OpenAIChatModel(
- model_name=model,
- provider=OllamaProvider(base_url=f"{Config.providers.ollama.base_url}/v1"),
- )
- elif provider == "vllm":
- return OpenAIChatModel(
- model_name=model,
- provider=OpenAIProvider(
- base_url=f"{Config.providers.vllm.research_base_url or Config.providers.vllm.qa_base_url}/v1",
- api_key="none",
- ),
- )
- else:
- return f"{provider}:{model}"
-
-
-def log(deps: HasEmitLog, state: Any, message: str) -> None:
- deps.emit_log(message, state)
diff --git a/haiku_rag_slim/haiku/rag/graph/models.py b/haiku_rag_slim/haiku/rag/graph/models.py
deleted file mode 100644
index 475b6e43..00000000
--- a/haiku_rag_slim/haiku/rag/graph/models.py
+++ /dev/null
@@ -1,24 +0,0 @@
-from pydantic import BaseModel, Field
-
-
-class ResearchPlan(BaseModel):
- main_question: str
- sub_questions: list[str]
-
-
-class SearchAnswer(BaseModel):
- query: str = Field(description="The search query that was performed")
- answer: str = Field(description="The answer generated based on the context")
- context: list[str] = Field(
- description=(
- "Only the minimal set of relevant snippets (verbatim) that directly "
- "support the answer"
- )
- )
- sources: list[str] = Field(
- description=(
- "Document titles (if available) or URIs corresponding to the"
- " snippets actually used in the answer (one per snippet; omit if none)"
- ),
- default_factory=list,
- )
diff --git a/haiku_rag_slim/haiku/rag/graph/nodes/analysis.py b/haiku_rag_slim/haiku/rag/graph/nodes/analysis.py
deleted file mode 100644
index 1d29d935..00000000
--- a/haiku_rag_slim/haiku/rag/graph/nodes/analysis.py
+++ /dev/null
@@ -1,182 +0,0 @@
-from dataclasses import dataclass
-
-from pydantic_ai import Agent
-from pydantic_graph import BaseNode, GraphRunContext
-
-from haiku.rag.graph.common import get_model, log
-from haiku.rag.research.common import (
- format_analysis_for_prompt,
- format_context_for_prompt,
-)
-from haiku.rag.research.dependencies import ResearchDependencies
-from haiku.rag.research.models import EvaluationResult, InsightAnalysis, ResearchReport
-from haiku.rag.research.prompts import DECISION_AGENT_PROMPT, INSIGHT_AGENT_PROMPT
-from haiku.rag.research.state import ResearchDeps, ResearchState
-
-
-@dataclass
-class AnalyzeInsightsNode(BaseNode[ResearchState, ResearchDeps, ResearchReport]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[ResearchState, ResearchDeps]
- ) -> BaseNode[ResearchState, ResearchDeps, ResearchReport]:
- state = ctx.state
- deps = ctx.deps
-
- log(
- deps,
- state,
- "\n[bold cyan]🧭 Synthesizing new insights and gap status...[/bold cyan]",
- )
-
- agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=InsightAnalysis,
- instructions=INSIGHT_AGENT_PROMPT,
- retries=3,
- deps_type=ResearchDependencies,
- )
-
- context_xml = format_context_for_prompt(state.context)
- prompt = (
- "Review the latest research context and update the shared ledger of insights, gaps,"
- " and follow-up questions.\n\n"
- f"{context_xml}"
- )
- agent_deps = ResearchDependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- stream=deps.stream,
- )
- result = await agent.run(prompt, deps=agent_deps)
- analysis: InsightAnalysis = result.output
-
- state.context.integrate_analysis(analysis)
- state.last_analysis = analysis
-
- if analysis.commentary:
- log(deps, state, f" Summary: {analysis.commentary}")
- if analysis.highlights:
- log(deps, state, " [bold]Updated insights:[/bold]")
- for insight in analysis.highlights:
- label = insight.status.value
- log(
- deps,
- state,
- f" • ({label}) {insight.summary}",
- )
- if analysis.gap_assessments:
- log(deps, state, " [bold yellow]Gap updates:[/bold yellow]")
- for gap in analysis.gap_assessments:
- status = "resolved" if gap.resolved else "open"
- severity = gap.severity.value
- log(
- deps,
- state,
- f" • ({severity}/{status}) {gap.description}",
- )
- if analysis.resolved_gaps:
- log(deps, state, " [green]Resolved gaps:[/green]")
- for resolved in analysis.resolved_gaps:
- log(deps, state, f" • {resolved}")
- if analysis.new_questions:
- log(deps, state, " [cyan]Proposed follow-ups:[/cyan]")
- for question in analysis.new_questions:
- log(deps, state, f" • {question}")
-
- from haiku.rag.graph.nodes.analysis import DecisionNode
-
- return DecisionNode(self.provider, self.model)
-
-
-@dataclass
-class DecisionNode(BaseNode[ResearchState, ResearchDeps, ResearchReport]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[ResearchState, ResearchDeps]
- ) -> BaseNode[ResearchState, ResearchDeps, ResearchReport]:
- state = ctx.state
- deps = ctx.deps
-
- log(
- deps,
- state,
- "\n[bold cyan]📊 Evaluating research sufficiency...[/bold cyan]",
- )
-
- agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=EvaluationResult,
- instructions=DECISION_AGENT_PROMPT,
- retries=3,
- deps_type=ResearchDependencies,
- )
-
- context_xml = format_context_for_prompt(state.context)
- analysis_xml = format_analysis_for_prompt(state.last_analysis)
- prompt_parts = [
- "Assess whether the research now answers the original question with adequate confidence.",
- context_xml,
- analysis_xml,
- ]
- if state.last_eval is not None:
- prev = state.last_eval
- prompt_parts.append(
- ""
- f"{prev.confidence_score:.2f}"
- f"{str(prev.is_sufficient).lower()}"
- f"{prev.reasoning}"
- ""
- )
- prompt = "\n\n".join(part for part in prompt_parts if part)
-
- agent_deps = ResearchDependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- stream=deps.stream,
- )
- decision_result = await agent.run(prompt, deps=agent_deps)
- output = decision_result.output
-
- state.last_eval = output
- state.iterations += 1
-
- for new_q in output.new_questions:
- if new_q not in state.context.sub_questions:
- state.context.sub_questions.append(new_q)
-
- if output.key_insights:
- log(deps, state, " [bold]Key insights:[/bold]")
- for insight in output.key_insights:
- log(deps, state, f" • {insight}")
-
- if output.gaps:
- log(deps, state, " [bold yellow]Remaining gaps:[/bold yellow]")
- for gap in output.gaps:
- log(deps, state, f" • {gap}")
-
- log(
- deps,
- state,
- f" Confidence: [yellow]{output.confidence_score:.1%}[/yellow]",
- )
- status = "[green]Yes[/green]" if output.is_sufficient else "[red]No[/red]"
- log(deps, state, f" Sufficient: {status}")
-
- from haiku.rag.graph.nodes.search import SearchDispatchNode
- from haiku.rag.graph.nodes.synthesize import SynthesizeNode
-
- if (
- output.is_sufficient
- and output.confidence_score >= state.confidence_threshold
- ) or state.iterations >= state.max_iterations:
- log(deps, state, "\n[bold green]✅ Stopping research.[/bold green]")
- return SynthesizeNode(self.provider, self.model)
-
- return SearchDispatchNode(self.provider, self.model)
diff --git a/haiku_rag_slim/haiku/rag/graph/nodes/plan.py b/haiku_rag_slim/haiku/rag/graph/nodes/plan.py
deleted file mode 100644
index 5f2cade9..00000000
--- a/haiku_rag_slim/haiku/rag/graph/nodes/plan.py
+++ /dev/null
@@ -1,72 +0,0 @@
-from dataclasses import dataclass
-
-from pydantic_ai import Agent, RunContext
-from pydantic_graph import BaseNode, GraphRunContext
-
-from haiku.rag.graph.common import get_model, log
-from haiku.rag.graph.models import ResearchPlan
-from haiku.rag.graph.prompts import PLAN_PROMPT
-from haiku.rag.research.dependencies import ResearchDependencies
-from haiku.rag.research.models import ResearchReport
-from haiku.rag.research.state import ResearchDeps, ResearchState
-
-
-@dataclass
-class PlanNode(BaseNode[ResearchState, ResearchDeps, ResearchReport]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[ResearchState, ResearchDeps]
- ) -> BaseNode[ResearchState, ResearchDeps, ResearchReport]:
- state = ctx.state
- deps = ctx.deps
-
- log(deps, state, "\n[bold cyan]📋 Creating research plan...[/bold cyan]")
-
- plan_agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=ResearchPlan,
- instructions=(
- PLAN_PROMPT
- + "\n\nUse the gather_context tool once on the main question before planning."
- ),
- retries=3,
- deps_type=ResearchDependencies,
- )
-
- @plan_agent.tool
- async def gather_context(
- ctx2: RunContext[ResearchDependencies], query: str, limit: int = 6
- ) -> str:
- results = await ctx2.deps.client.search(query, limit=limit)
- expanded = await ctx2.deps.client.expand_context(results)
- return "\n\n".join(chunk.content for chunk, _ in expanded)
-
- prompt = (
- "Plan a focused research approach for the main question.\n\n"
- f"Main question: {state.context.original_question}"
- )
-
- agent_deps = ResearchDependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- stream=deps.stream,
- )
- plan_result = await plan_agent.run(prompt, deps=agent_deps)
- state.context.sub_questions = list(plan_result.output.sub_questions)
-
- log(deps, state, "\n[bold green]✅ Research Plan Created:[/bold green]")
- log(
- deps,
- state,
- f" [bold]Main Question:[/bold] {state.context.original_question}",
- )
- log(deps, state, " [bold]Sub-questions:[/bold]")
- for i, sq in enumerate(state.context.sub_questions, 1):
- log(deps, state, f" {i}. {sq}")
-
- from haiku.rag.graph.nodes.search import SearchDispatchNode
-
- return SearchDispatchNode(self.provider, self.model)
diff --git a/haiku_rag_slim/haiku/rag/graph/nodes/search.py b/haiku_rag_slim/haiku/rag/graph/nodes/search.py
deleted file mode 100644
index b7e12bd3..00000000
--- a/haiku_rag_slim/haiku/rag/graph/nodes/search.py
+++ /dev/null
@@ -1,97 +0,0 @@
-import asyncio
-from dataclasses import dataclass
-from typing import Any
-
-from pydantic_ai import Agent, RunContext
-from pydantic_ai.format_prompt import format_as_xml
-from pydantic_ai.output import ToolOutput
-from pydantic_graph import BaseNode, GraphRunContext
-
-from haiku.rag.graph.common import get_model, log
-from haiku.rag.graph.models import SearchAnswer
-from haiku.rag.graph.prompts import SEARCH_AGENT_PROMPT
-from haiku.rag.research.dependencies import ResearchDependencies
-from haiku.rag.research.models import ResearchReport
-from haiku.rag.research.state import ResearchDeps, ResearchState
-
-
-@dataclass
-class SearchDispatchNode(BaseNode[ResearchState, ResearchDeps, ResearchReport]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[ResearchState, ResearchDeps]
- ) -> BaseNode[ResearchState, ResearchDeps, ResearchReport]:
- state = ctx.state
- deps = ctx.deps
- if not state.context.sub_questions:
- from haiku.rag.graph.nodes.analysis import AnalyzeInsightsNode
-
- return AnalyzeInsightsNode(self.provider, self.model)
-
- # Take up to max_concurrency questions and answer them concurrently
- take = max(1, state.max_concurrency)
- batch: list[str] = []
- while state.context.sub_questions and len(batch) < take:
- batch.append(state.context.sub_questions.pop(0))
-
- async def answer_one(sub_q: str) -> SearchAnswer | None:
- log(
- deps,
- state,
- f"\n[bold cyan]🔍 Searching & Answering:[/bold cyan] {sub_q}",
- )
- agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=ToolOutput(SearchAnswer, max_retries=3),
- instructions=SEARCH_AGENT_PROMPT,
- retries=3,
- deps_type=ResearchDependencies,
- )
-
- @agent.tool
- async def search_and_answer(
- ctx2: RunContext[ResearchDependencies], query: str, limit: int = 5
- ) -> str:
- search_results = await ctx2.deps.client.search(query, limit=limit)
- expanded = await ctx2.deps.client.expand_context(search_results)
-
- entries: list[dict[str, Any]] = [
- {
- "text": chunk.content,
- "score": score,
- "document_uri": (
- chunk.document_title or chunk.document_uri or ""
- ),
- }
- for chunk, score in expanded
- ]
- if not entries:
- return f"No relevant information found in the knowledge base for: {query}"
-
- return format_as_xml(entries, root_tag="snippets")
-
- agent_deps = ResearchDependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- stream=deps.stream,
- )
- try:
- result = await agent.run(sub_q, deps=agent_deps)
- except Exception as e:
- log(deps, state, f"[red]Search failed:[/red] {e}")
- return None
-
- return result.output
-
- answers = await asyncio.gather(*(answer_one(q) for q in batch))
- for ans in answers:
- if ans is None:
- continue
- state.context.add_qa_response(ans)
- preview = ans.answer[:150] + ("…" if len(ans.answer) > 150 else "")
- log(deps, state, f" [green]✓[/green] {preview}")
-
- return SearchDispatchNode(self.provider, self.model)
diff --git a/haiku_rag_slim/haiku/rag/graph/nodes/synthesize.py b/haiku_rag_slim/haiku/rag/graph/nodes/synthesize.py
deleted file mode 100644
index 309a933d..00000000
--- a/haiku_rag_slim/haiku/rag/graph/nodes/synthesize.py
+++ /dev/null
@@ -1,54 +0,0 @@
-from dataclasses import dataclass
-
-from pydantic_ai import Agent
-from pydantic_graph import BaseNode, End, GraphRunContext
-
-from haiku.rag.graph.common import get_model, log
-from haiku.rag.research.common import format_context_for_prompt
-from haiku.rag.research.dependencies import ResearchDependencies
-from haiku.rag.research.models import ResearchReport
-from haiku.rag.research.prompts import SYNTHESIS_AGENT_PROMPT
-from haiku.rag.research.state import ResearchDeps, ResearchState
-
-
-@dataclass
-class SynthesizeNode(BaseNode[ResearchState, ResearchDeps, ResearchReport]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[ResearchState, ResearchDeps]
- ) -> End[ResearchReport]:
- state = ctx.state
- deps = ctx.deps
-
- log(
- deps,
- state,
- "\n[bold cyan]📝 Generating final research report...[/bold cyan]",
- )
-
- agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=ResearchReport,
- instructions=SYNTHESIS_AGENT_PROMPT,
- retries=3,
- deps_type=ResearchDependencies,
- )
-
- context_xml = format_context_for_prompt(state.context)
- prompt = (
- "Generate a comprehensive research report based on all gathered information.\n\n"
- f"{context_xml}\n\n"
- "Create a detailed report that synthesizes all findings into a coherent response."
- )
- agent_deps = ResearchDependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- stream=deps.stream,
- )
- result = await agent.run(prompt, deps=agent_deps)
-
- log(deps, state, "[bold green]✅ Research complete![/bold green]")
- return End(result.output)
diff --git a/haiku_rag_slim/haiku/rag/graph_common/__init__.py b/haiku_rag_slim/haiku/rag/graph_common/__init__.py
new file mode 100644
index 00000000..dc47bee0
--- /dev/null
+++ b/haiku_rag_slim/haiku/rag/graph_common/__init__.py
@@ -0,0 +1,5 @@
+"""Common utilities for graph implementations."""
+
+from haiku.rag.graph_common.utils import get_model, log
+
+__all__ = ["get_model", "log"]
diff --git a/haiku_rag_slim/haiku/rag/graph_common/models.py b/haiku_rag_slim/haiku/rag/graph_common/models.py
new file mode 100644
index 00000000..407cb6bb
--- /dev/null
+++ b/haiku_rag_slim/haiku/rag/graph_common/models.py
@@ -0,0 +1,42 @@
+"""Common models used across different graph implementations."""
+
+from pydantic import BaseModel, Field, field_validator
+
+
+class ResearchPlan(BaseModel):
+ """A structured research plan with sub-questions to explore."""
+
+ sub_questions: list[str] = Field(
+ ...,
+ description="Specific questions to research, phrased as complete questions",
+ )
+
+ @field_validator("sub_questions")
+ @classmethod
+ def validate_sub_questions(cls, v: list[str]) -> list[str]:
+ if len(v) < 1:
+ raise ValueError("Must have at least 1 sub-question")
+ if len(v) > 12:
+ raise ValueError("Cannot have more than 12 sub-questions")
+ return v
+
+
+class SearchAnswer(BaseModel):
+ """Answer from a search operation with sources."""
+
+ query: str = Field(..., description="The question that was answered")
+ answer: str = Field(..., description="The comprehensive answer to the question")
+ context: list[str] = Field(
+ default_factory=list,
+ description="Relevant snippets that directly support the answer",
+ )
+ sources: list[str] = Field(
+ default_factory=list,
+ description="Source URIs or titles that contributed to this answer",
+ )
+ confidence: float = Field(
+ default=1.0,
+ description="Confidence score for this answer (0-1)",
+ ge=0.0,
+ le=1.0,
+ )
diff --git a/haiku_rag_slim/haiku/rag/graph/prompts.py b/haiku_rag_slim/haiku/rag/graph_common/prompts.py
similarity index 79%
rename from haiku_rag_slim/haiku/rag/graph/prompts.py
rename to haiku_rag_slim/haiku/rag/graph_common/prompts.py
index 97f42aa8..ed10ec98 100644
--- a/haiku_rag_slim/haiku/rag/graph/prompts.py
+++ b/haiku_rag_slim/haiku/rag/graph_common/prompts.py
@@ -1,23 +1,24 @@
-PLAN_PROMPT = """You are the research orchestrator for a focused, iterative
-workflow.
+"""Common prompts used across different graph implementations."""
+
+PLAN_PROMPT = """You are the research orchestrator for a focused, iterative workflow.
Responsibilities:
1. Understand and decompose the main question
-2. Propose a minimal, high‑leverage plan
+2. Propose a minimal, high-leverage plan
3. Coordinate specialized agents to gather evidence
4. Iterate based on gaps and new findings
Plan requirements:
- Produce at most 3 sub_questions that together cover the main question.
-- Each sub_question must be a standalone, self‑contained query that can run
+- Each sub_question must be a standalone, self-contained query that can run
without extra context. Include concrete entities, scope, timeframe, and any
qualifiers. Avoid ambiguous pronouns (it/they/this/that).
-- Prioritize the highest‑value aspects first; avoid redundancy and overlap.
+- Prioritize the highest-value aspects first; avoid redundancy and overlap.
- Prefer questions that are likely answerable from the current knowledge base;
if coverage is uncertain, make scopes narrower and specific.
- Order sub_questions by execution priority (most valuable first)."""
-SEARCH_AGENT_PROMPT = """You are a search and question‑answering specialist.
+SEARCH_AGENT_PROMPT = """You are a search and question-answering specialist.
Tasks:
1. Search the knowledge base for relevant evidence.
@@ -31,10 +32,10 @@ Tool usage:
- You may call the tool multiple times to refine or broaden context, but do not
exceed 3 total calls. Favor precision over volume.
- Use scores to prioritize evidence, but include only the minimal subset of
- snippet texts (verbatim) in SearchAnswer.context (typically 1‑4).
+ snippet texts (verbatim) in SearchAnswer.context (typically 1-4).
- Set SearchAnswer.sources to the corresponding document identifiers for the
snippets you used (title if available, otherwise URI; one per snippet; same
- order as context). Context must be text‑only.
+ order as context). Context must be text-only.
- If no relevant information is found, clearly say so and return an empty
context list and sources list.
diff --git a/haiku_rag_slim/haiku/rag/graph_common/utils.py b/haiku_rag_slim/haiku/rag/graph_common/utils.py
new file mode 100644
index 00000000..f24fb06c
--- /dev/null
+++ b/haiku_rag_slim/haiku/rag/graph_common/utils.py
@@ -0,0 +1,64 @@
+"""Common utilities for all graph implementations."""
+
+from typing import Any, Protocol
+
+from pydantic_ai.models.openai import OpenAIChatModel
+from pydantic_ai.providers.ollama import OllamaProvider
+from pydantic_ai.providers.openai import OpenAIProvider
+
+from haiku.rag.config import Config
+
+
+class HasEmitLog(Protocol):
+ """Protocol for objects that can emit log messages."""
+
+ def emit_log(self, message: str, state: Any = None) -> None: ...
+
+
+def get_model(provider: str, model: str) -> OpenAIChatModel | str:
+ """
+ Get a model instance for the specified provider and model name.
+
+ Args:
+ provider: The model provider ("ollama", "vllm", or other)
+ model: The model name
+
+ Returns:
+ A configured model instance
+
+ Raises:
+ ValueError: If the provider is unknown
+ """
+ if provider == "ollama":
+ return OpenAIChatModel(
+ model_name=model,
+ provider=OllamaProvider(base_url=f"{Config.providers.ollama.base_url}/v1"),
+ )
+ elif provider == "vllm":
+ return OpenAIChatModel(
+ model_name=model,
+ provider=OpenAIProvider(
+ base_url=f"{Config.providers.vllm.research_base_url or Config.providers.vllm.qa_base_url}/v1",
+ api_key="none",
+ ),
+ )
+ elif provider in ("openai", "anthropic", "gemini", "groq", "bedrock"):
+ # These providers use string format
+ return f"{provider}:{model}"
+ else:
+ raise ValueError(
+ f"Unknown model provider: {provider}. "
+ f"Supported providers: ollama, vllm, openai, anthropic, gemini, groq, bedrock"
+ )
+
+
+def log(deps: HasEmitLog, state: Any, message: str) -> None:
+ """
+ Emit a log message through the dependencies.
+
+ Args:
+ deps: Dependencies object with emit_log method
+ state: Current state (passed to emit_log)
+ message: The message to log
+ """
+ deps.emit_log(message, state)
diff --git a/haiku_rag_slim/haiku/rag/mcp.py b/haiku_rag_slim/haiku/rag/mcp.py
index 61754759..3c18587c 100644
--- a/haiku_rag_slim/haiku/rag/mcp.py
+++ b/haiku_rag_slim/haiku/rag/mcp.py
@@ -194,25 +194,17 @@ def create_mcp_server(db_path: Path) -> FastMCP:
from haiku.rag.config import Config
from haiku.rag.qa.deep.dependencies import DeepQAContext
from haiku.rag.qa.deep.graph import build_deep_qa_graph
- from haiku.rag.qa.deep.nodes import DeepQAPlanNode
from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
- graph = build_deep_qa_graph()
+ graph = build_deep_qa_graph(config=Config)
context = DeepQAContext(
original_question=question, use_citations=cite
)
- state = DeepQAState(context=context)
+ state = DeepQAState.from_config(context=context, config=Config)
deps = DeepQADeps(client=rag)
- start_node = DeepQAPlanNode(
- provider=Config.qa.provider,
- model=Config.qa.model,
- )
-
- result = await graph.run(
- start_node=start_node, state=state, deps=deps
- )
- answer = result.output.answer
+ result = await graph.run(state=state, deps=deps)
+ answer = result.answer
else:
answer = await rag.ask(question, cite=cite)
return answer
@@ -222,9 +214,6 @@ def create_mcp_server(db_path: Path) -> FastMCP:
@mcp.tool()
async def research_question(
question: str,
- max_iterations: int = 3,
- confidence_threshold: float = 0.8,
- max_concurrency: int = 1,
) -> ResearchReport | None:
"""Run multi-agent research to investigate a complex question.
@@ -233,39 +222,24 @@ def create_mcp_server(db_path: Path) -> FastMCP:
Args:
question: The research question to investigate.
- max_iterations: Maximum search/analyze iterations (default: 3).
- confidence_threshold: Minimum confidence score (0-1) to stop early (default: 0.8).
- max_concurrency: Maximum concurrent searches per iteration (default: 1).
Returns:
A research report with findings, or None if an error occurred.
"""
try:
- from haiku.rag.graph.nodes.plan import PlanNode
from haiku.rag.research.dependencies import ResearchContext
from haiku.rag.research.graph import build_research_graph
from haiku.rag.research.state import ResearchDeps, ResearchState
async with HaikuRAG(db_path) as rag:
- graph = build_research_graph()
- state = ResearchState(
- context=ResearchContext(original_question=question),
- max_iterations=max_iterations,
- confidence_threshold=confidence_threshold,
- max_concurrency=max_concurrency,
- )
+ graph = build_research_graph(config=Config)
+ context = ResearchContext(original_question=question)
+ state = ResearchState.from_config(context=context, config=Config)
deps = ResearchDeps(client=rag)
- result = await graph.run(
- PlanNode(
- provider=Config.research.provider or Config.qa.provider,
- model=Config.research.model or Config.qa.model,
- ),
- state=state,
- deps=deps,
- )
+ result = await graph.run(state=state, deps=deps)
- return result.output
+ return result
except Exception:
return None
diff --git a/haiku_rag_slim/haiku/rag/qa/deep/dependencies.py b/haiku_rag_slim/haiku/rag/qa/deep/dependencies.py
index c6017a61..f8bce190 100644
--- a/haiku_rag_slim/haiku/rag/qa/deep/dependencies.py
+++ b/haiku_rag_slim/haiku/rag/qa/deep/dependencies.py
@@ -2,7 +2,7 @@ from pydantic import BaseModel, Field
from rich.console import Console
from haiku.rag.client import HaikuRAG
-from haiku.rag.graph.models import SearchAnswer
+from haiku.rag.graph_common.models import SearchAnswer
class DeepQAContext(BaseModel):
diff --git a/haiku_rag_slim/haiku/rag/qa/deep/graph.py b/haiku_rag_slim/haiku/rag/qa/deep/graph.py
index f5701ebd..9ded283c 100644
--- a/haiku_rag_slim/haiku/rag/qa/deep/graph.py
+++ b/haiku_rag_slim/haiku/rag/qa/deep/graph.py
@@ -1,21 +1,363 @@
-from pydantic_graph import Graph
+from typing import Any
-from haiku.rag.qa.deep.models import DeepQAAnswer
-from haiku.rag.qa.deep.nodes import (
- DeepQADecisionNode,
- DeepQAPlanNode,
- DeepQASearchDispatchNode,
- DeepQASynthesizeNode,
+from pydantic_ai import Agent, RunContext
+from pydantic_ai.format_prompt import format_as_xml
+from pydantic_ai.output import ToolOutput
+from pydantic_graph.beta import Graph, GraphBuilder, StepContext
+from pydantic_graph.beta.join import reduce_list_append
+
+from haiku.rag.config import Config
+from haiku.rag.config.models import AppConfig
+from haiku.rag.graph_common import get_model, log
+from haiku.rag.graph_common.models import ResearchPlan, SearchAnswer
+from haiku.rag.graph_common.prompts import PLAN_PROMPT, SEARCH_AGENT_PROMPT
+from haiku.rag.qa.deep.dependencies import DeepQADependencies
+from haiku.rag.qa.deep.models import DeepQAAnswer, DeepQAEvaluation
+from haiku.rag.qa.deep.prompts import (
+ DECISION_PROMPT,
+ SYNTHESIS_PROMPT,
+ SYNTHESIS_PROMPT_WITH_CITATIONS,
)
from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
-def build_deep_qa_graph() -> Graph[DeepQAState, DeepQADeps, DeepQAAnswer]:
- return Graph(
- nodes=[
- DeepQAPlanNode,
- DeepQASearchDispatchNode,
- DeepQADecisionNode,
- DeepQASynthesizeNode,
- ]
+def build_deep_qa_graph(
+ config: AppConfig = Config,
+) -> Graph[DeepQAState, DeepQADeps, None, DeepQAAnswer]:
+ """Build the Deep QA graph.
+
+ Args:
+ config: AppConfig object (uses config.qa for provider, model, and graph parameters)
+
+ Returns:
+ Configured Deep QA graph
+ """
+ provider = config.qa.provider
+ model = config.qa.model
+ g = GraphBuilder(
+ state_type=DeepQAState,
+ deps_type=DeepQADeps,
+ output_type=DeepQAAnswer,
)
+
+ @g.step
+ async def plan(ctx: StepContext[DeepQAState, DeepQADeps, None]) -> None:
+ state = ctx.state
+ deps = ctx.deps
+
+ log(deps, state, "\n[bold cyan]📋 Planning approach...[/bold cyan]")
+
+ plan_agent = Agent(
+ model=get_model(provider, model),
+ output_type=ResearchPlan,
+ instructions=(
+ PLAN_PROMPT
+ + "\n\nUse the gather_context tool once on the main question before planning."
+ ),
+ retries=3,
+ deps_type=DeepQADependencies,
+ )
+
+ @plan_agent.tool
+ async def gather_context(
+ ctx2: RunContext[DeepQADependencies], query: str, limit: int = 6
+ ) -> str:
+ results = await ctx2.deps.client.search(query, limit=limit)
+ expanded = await ctx2.deps.client.expand_context(results)
+ return "\n\n".join(chunk.content for chunk, _ in expanded)
+
+ prompt = (
+ "Plan a focused approach for the main question.\n\n"
+ f"Main question: {state.context.original_question}"
+ )
+
+ agent_deps = DeepQADependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ )
+ plan_result = await plan_agent.run(prompt, deps=agent_deps)
+ state.context.sub_questions = list(plan_result.output.sub_questions)
+
+ log(deps, state, "\n[bold green]✅ Plan Created:[/bold green]")
+ log(
+ deps,
+ state,
+ f" [bold]Main Question:[/bold] {state.context.original_question}",
+ )
+ log(deps, state, " [bold]Sub-questions:[/bold]")
+ for i, sq in enumerate(state.context.sub_questions, 1):
+ log(deps, state, f" {i}. {sq}")
+
+ @g.step
+ async def search_one(
+ ctx: StepContext[DeepQAState, DeepQADeps, str],
+ ) -> SearchAnswer:
+ state = ctx.state
+ deps = ctx.deps
+ sub_q = ctx.inputs
+
+ # Create semaphore if not already provided
+ if deps.semaphore is None:
+ import asyncio
+
+ deps.semaphore = asyncio.Semaphore(state.max_concurrency)
+
+ # Use semaphore to control concurrency
+ async with deps.semaphore:
+ return await _do_search(state, deps, sub_q)
+
+ async def _do_search(
+ state: DeepQAState,
+ deps: DeepQADeps,
+ sub_q: str,
+ ) -> SearchAnswer:
+ log(
+ deps,
+ state,
+ f"\n[bold cyan]🔍 Searching & Answering:[/bold cyan] {sub_q}",
+ )
+
+ agent = Agent(
+ model=get_model(provider, model),
+ output_type=ToolOutput(SearchAnswer, max_retries=3),
+ instructions=SEARCH_AGENT_PROMPT,
+ retries=3,
+ deps_type=DeepQADependencies,
+ )
+
+ @agent.tool
+ async def search_and_answer(
+ ctx2: RunContext[DeepQADependencies], query: str, limit: int = 5
+ ) -> str:
+ search_results = await ctx2.deps.client.search(query, limit=limit)
+ expanded = await ctx2.deps.client.expand_context(search_results)
+
+ entries: list[dict[str, Any]] = [
+ {
+ "text": chunk.content,
+ "score": score,
+ "document_uri": (chunk.document_title or chunk.document_uri or ""),
+ }
+ for chunk, score in expanded
+ ]
+ if not entries:
+ return (
+ f"No relevant information found in the knowledge base for: {query}"
+ )
+
+ return format_as_xml(entries, root_tag="snippets")
+
+ agent_deps = DeepQADependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ )
+ try:
+ result = await agent.run(sub_q, deps=agent_deps)
+ answer = result.output
+ if answer:
+ state.context.add_qa_response(answer)
+ preview = answer.answer[:150] + (
+ "…" if len(answer.answer) > 150 else ""
+ )
+ log(deps, state, f" [green]✓[/green] {preview}")
+ return answer
+ except Exception as e:
+ log(deps, state, f"[red]Search failed:[/red] {e}")
+ failure_answer = SearchAnswer(
+ query=sub_q,
+ answer=f"Search failed after retries: {str(e)}",
+ confidence=0.0,
+ )
+ return failure_answer
+
+ @g.step
+ async def get_batch(
+ ctx: StepContext[DeepQAState, DeepQADeps, None | bool],
+ ) -> list[str] | None:
+ """Get all remaining questions for this iteration."""
+ state = ctx.state
+
+ if not state.context.sub_questions:
+ return None
+
+ # Take ALL remaining questions - max_concurrency controls parallel execution within .map()
+ batch = list(state.context.sub_questions)
+ state.context.sub_questions.clear()
+ return batch
+
+ @g.step
+ async def decide(
+ ctx: StepContext[DeepQAState, DeepQADeps, list[SearchAnswer]],
+ ) -> bool:
+ state = ctx.state
+ deps = ctx.deps
+
+ log(
+ deps,
+ state,
+ "\n[bold cyan]📊 Evaluating information sufficiency...[/bold cyan]",
+ )
+
+ agent = Agent(
+ model=get_model(provider, model),
+ output_type=DeepQAEvaluation,
+ instructions=DECISION_PROMPT,
+ retries=3,
+ deps_type=DeepQADependencies,
+ )
+
+ context_data = {
+ "original_question": state.context.original_question,
+ "gathered_answers": [
+ {
+ "question": qa.query,
+ "answer": qa.answer,
+ "sources": qa.sources,
+ }
+ for qa in state.context.qa_responses
+ ],
+ }
+ context_xml = format_as_xml(context_data, root_tag="gathered_information")
+
+ prompt = (
+ "Evaluate whether we have sufficient information to answer the question.\n\n"
+ f"{context_xml}"
+ )
+
+ agent_deps = DeepQADependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ )
+ result = await agent.run(prompt, deps=agent_deps)
+ evaluation = result.output
+
+ state.iterations += 1
+
+ log(deps, state, f" [bold]Assessment:[/bold] {evaluation.reasoning}")
+ status = "[green]Yes[/green]" if evaluation.is_sufficient else "[red]No[/red]"
+ log(deps, state, f" Sufficient: {status}")
+
+ for new_q in evaluation.new_questions:
+ if new_q not in state.context.sub_questions:
+ state.context.sub_questions.append(new_q)
+
+ if evaluation.new_questions:
+ log(deps, state, " [cyan]New questions:[/cyan]")
+ for question in evaluation.new_questions:
+ log(deps, state, f" • {question}")
+
+ should_continue = (
+ not evaluation.is_sufficient and state.iterations < state.max_iterations
+ )
+
+ if not should_continue:
+ if state.iterations >= state.max_iterations:
+ log(
+ deps,
+ state,
+ f"\n[bold yellow]⚠️ Reached max iterations ({state.max_iterations})[/bold yellow]",
+ )
+ log(deps, state, "\n[bold green]✅ Moving to synthesis.[/bold green]")
+ else:
+ log(
+ deps,
+ state,
+ f"\n[bold cyan]🔄 Starting iteration {state.iterations + 1}...[/bold cyan]",
+ )
+
+ return should_continue
+
+ @g.step
+ async def synthesize(
+ ctx: StepContext[DeepQAState, DeepQADeps, None | bool],
+ ) -> DeepQAAnswer:
+ state = ctx.state
+ deps = ctx.deps
+
+ log(
+ deps,
+ state,
+ "\n[bold cyan]📝 Synthesizing final answer...[/bold cyan]",
+ )
+
+ prompt_template = (
+ SYNTHESIS_PROMPT_WITH_CITATIONS
+ if state.context.use_citations
+ else SYNTHESIS_PROMPT
+ )
+
+ agent = Agent(
+ model=get_model(provider, model),
+ output_type=DeepQAAnswer,
+ instructions=prompt_template,
+ retries=3,
+ deps_type=DeepQADependencies,
+ )
+
+ context_data = {
+ "original_question": state.context.original_question,
+ "sub_answers": [
+ {
+ "question": qa.query,
+ "answer": qa.answer,
+ "sources": qa.sources,
+ }
+ for qa in state.context.qa_responses
+ ],
+ }
+ context_xml = format_as_xml(context_data, root_tag="gathered_information")
+
+ prompt = f"Synthesize a comprehensive answer to the original question.\n\n{context_xml}"
+
+ agent_deps = DeepQADependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ )
+ result = await agent.run(prompt, deps=agent_deps)
+
+ log(deps, state, "[bold green]✅ Answer complete![/bold green]")
+ return result.output
+
+ # Build the graph structure
+ collect_answers = g.join(
+ reduce_list_append,
+ initial_factory=list[SearchAnswer],
+ )
+
+ g.add(
+ g.edge_from(g.start_node).to(plan),
+ g.edge_from(plan).to(get_batch),
+ )
+
+ # Branch based on whether we have questions
+ g.add(
+ g.edge_from(get_batch).to(
+ g.decision()
+ .branch(g.match(list).label("Has questions").map().to(search_one))
+ .branch(g.match(type(None)).label("No questions").to(synthesize))
+ ),
+ g.edge_from(search_one).to(collect_answers),
+ g.edge_from(collect_answers).to(decide),
+ )
+
+ # Branch based on decision
+ g.add(
+ g.edge_from(decide).to(
+ g.decision()
+ .branch(
+ g.match(bool, matches=lambda x: x).label("Continue QA").to(get_batch)
+ )
+ .branch(
+ g.match(bool, matches=lambda x: not x)
+ .label("Done with QA")
+ .to(synthesize)
+ )
+ ),
+ g.edge_from(synthesize).to(g.end_node),
+ )
+
+ return g.build()
diff --git a/haiku_rag_slim/haiku/rag/qa/deep/nodes.py b/haiku_rag_slim/haiku/rag/qa/deep/nodes.py
deleted file mode 100644
index 9461efe2..00000000
--- a/haiku_rag_slim/haiku/rag/qa/deep/nodes.py
+++ /dev/null
@@ -1,303 +0,0 @@
-import asyncio
-from dataclasses import dataclass
-from typing import Any
-
-from pydantic_ai import Agent, RunContext
-from pydantic_ai.format_prompt import format_as_xml
-from pydantic_ai.output import ToolOutput
-from pydantic_graph import BaseNode, End, GraphRunContext
-
-from haiku.rag.graph.common import get_model, log
-from haiku.rag.graph.models import ResearchPlan, SearchAnswer
-from haiku.rag.graph.prompts import PLAN_PROMPT, SEARCH_AGENT_PROMPT
-from haiku.rag.qa.deep.dependencies import DeepQADependencies
-from haiku.rag.qa.deep.models import DeepQAAnswer, DeepQAEvaluation
-from haiku.rag.qa.deep.prompts import (
- DECISION_PROMPT,
- SYNTHESIS_PROMPT,
- SYNTHESIS_PROMPT_WITH_CITATIONS,
-)
-from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
-
-
-@dataclass
-class DeepQAPlanNode(BaseNode[DeepQAState, DeepQADeps, DeepQAAnswer]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[DeepQAState, DeepQADeps]
- ) -> BaseNode[DeepQAState, DeepQADeps, DeepQAAnswer]:
- state = ctx.state
- deps = ctx.deps
-
- log(deps, state, "\n[bold cyan]📋 Planning approach...[/bold cyan]")
-
- plan_agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=ResearchPlan,
- instructions=(
- PLAN_PROMPT
- + "\n\nUse the gather_context tool once on the main question before planning."
- ),
- retries=3,
- deps_type=DeepQADependencies,
- )
-
- @plan_agent.tool
- async def gather_context(
- ctx2: RunContext[DeepQADependencies], query: str, limit: int = 6
- ) -> str:
- results = await ctx2.deps.client.search(query, limit=limit)
- expanded = await ctx2.deps.client.expand_context(results)
- return "\n\n".join(chunk.content for chunk, _ in expanded)
-
- prompt = (
- "Plan a focused approach for answering the main question.\n\n"
- f"Main question: {state.context.original_question}"
- )
-
- agent_deps = DeepQADependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- )
- plan_result = await plan_agent.run(prompt, deps=agent_deps)
- state.context.sub_questions = list(plan_result.output.sub_questions)[
- : state.max_sub_questions
- ]
-
- log(deps, state, "\n[bold green]✅ Plan Created:[/bold green]")
- log(
- deps,
- state,
- f" [bold]Main Question:[/bold] {state.context.original_question}",
- )
- log(deps, state, " [bold]Sub-questions:[/bold]")
- for i, sq in enumerate(state.context.sub_questions, 1):
- log(deps, state, f" {i}. {sq}")
-
- return DeepQASearchDispatchNode(self.provider, self.model)
-
-
-@dataclass
-class DeepQASearchDispatchNode(BaseNode[DeepQAState, DeepQADeps, DeepQAAnswer]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[DeepQAState, DeepQADeps]
- ) -> BaseNode[DeepQAState, DeepQADeps, DeepQAAnswer]:
- state = ctx.state
- deps = ctx.deps
-
- if not state.context.sub_questions:
- return DeepQADecisionNode(self.provider, self.model)
-
- # Take up to max_concurrency questions and answer them concurrently
- take = max(1, state.max_concurrency)
- batch: list[str] = []
- while state.context.sub_questions and len(batch) < take:
- batch.append(state.context.sub_questions.pop(0))
-
- async def answer_one(sub_q: str) -> SearchAnswer | None:
- log(
- deps,
- state,
- f"\n[bold cyan]🔍 Searching & Answering:[/bold cyan] {sub_q}",
- )
- agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=ToolOutput(SearchAnswer, max_retries=3),
- instructions=SEARCH_AGENT_PROMPT,
- retries=3,
- deps_type=DeepQADependencies,
- )
-
- @agent.tool
- async def search_and_answer(
- ctx2: RunContext[DeepQADependencies], query: str, limit: int = 5
- ) -> str:
- search_results = await ctx2.deps.client.search(query, limit=limit)
- expanded = await ctx2.deps.client.expand_context(search_results)
-
- entries: list[dict[str, Any]] = [
- {
- "text": chunk.content,
- "score": score,
- "document_uri": (
- chunk.document_title or chunk.document_uri or ""
- ),
- }
- for chunk, score in expanded
- ]
- if not entries:
- return f"No relevant information found in the knowledge base for: {query}"
-
- return format_as_xml(entries, root_tag="snippets")
-
- agent_deps = DeepQADependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- )
- try:
- result = await agent.run(sub_q, deps=agent_deps)
- except Exception as e:
- log(deps, state, f"[red]Search failed:[/red] {e}")
- return None
-
- return result.output
-
- answers = await asyncio.gather(*(answer_one(q) for q in batch))
- for ans in answers:
- if ans is None:
- continue
- state.context.add_qa_response(ans)
- preview = ans.answer[:150] + ("…" if len(ans.answer) > 150 else "")
- log(deps, state, f" [green]✓[/green] {preview}")
-
- return DeepQASearchDispatchNode(self.provider, self.model)
-
-
-@dataclass
-class DeepQADecisionNode(BaseNode[DeepQAState, DeepQADeps, DeepQAAnswer]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[DeepQAState, DeepQADeps]
- ) -> BaseNode[DeepQAState, DeepQADeps, DeepQAAnswer]:
- state = ctx.state
- deps = ctx.deps
-
- log(
- deps,
- state,
- "\n[bold cyan]📊 Evaluating information sufficiency...[/bold cyan]",
- )
-
- agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=DeepQAEvaluation,
- instructions=DECISION_PROMPT,
- retries=3,
- deps_type=DeepQADependencies,
- )
-
- context_data = {
- "original_question": state.context.original_question,
- "gathered_answers": [
- {
- "question": qa.query,
- "answer": qa.answer,
- "sources": qa.sources,
- }
- for qa in state.context.qa_responses
- ],
- }
- context_xml = format_as_xml(context_data, root_tag="gathered_information")
-
- prompt = (
- "Evaluate whether we have sufficient information to answer the question.\n\n"
- f"{context_xml}"
- )
-
- agent_deps = DeepQADependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- )
- result = await agent.run(prompt, deps=agent_deps)
- evaluation = result.output
-
- state.iterations += 1
-
- log(deps, state, f" [bold]Assessment:[/bold] {evaluation.reasoning}")
- status = "[green]Yes[/green]" if evaluation.is_sufficient else "[red]No[/red]"
- log(deps, state, f" Sufficient: {status}")
-
- # Add new questions if not sufficient
- for new_q in evaluation.new_questions:
- if new_q not in state.context.sub_questions:
- state.context.sub_questions.append(new_q)
-
- if evaluation.new_questions:
- log(deps, state, " [cyan]New questions:[/cyan]")
- for question in evaluation.new_questions:
- log(deps, state, f" • {question}")
-
- # Decide next step
- if evaluation.is_sufficient or state.iterations >= state.max_iterations:
- if state.iterations >= state.max_iterations:
- log(
- deps,
- state,
- f"\n[bold yellow]⚠️ Reached max iterations ({state.max_iterations})[/bold yellow]",
- )
- log(deps, state, "\n[bold green]✅ Moving to synthesis.[/bold green]")
- return DeepQASynthesizeNode(self.provider, self.model)
-
- log(
- deps,
- state,
- f"\n[bold cyan]🔄 Starting iteration {state.iterations + 1}...[/bold cyan]",
- )
- return DeepQASearchDispatchNode(self.provider, self.model)
-
-
-@dataclass
-class DeepQASynthesizeNode(BaseNode[DeepQAState, DeepQADeps, DeepQAAnswer]):
- provider: str
- model: str
-
- async def run(
- self, ctx: GraphRunContext[DeepQAState, DeepQADeps]
- ) -> End[DeepQAAnswer]:
- state = ctx.state
- deps = ctx.deps
-
- log(
- deps,
- state,
- "\n[bold cyan]📝 Synthesizing final answer...[/bold cyan]",
- )
-
- prompt_template = (
- SYNTHESIS_PROMPT_WITH_CITATIONS
- if state.context.use_citations
- else SYNTHESIS_PROMPT
- )
-
- agent = Agent(
- model=get_model(self.provider, self.model),
- output_type=DeepQAAnswer,
- instructions=prompt_template,
- retries=3,
- deps_type=DeepQADependencies,
- )
-
- context_data = {
- "original_question": state.context.original_question,
- "sub_answers": [
- {
- "question": qa.query,
- "answer": qa.answer,
- "sources": qa.sources,
- }
- for qa in state.context.qa_responses
- ],
- }
- context_xml = format_as_xml(context_data, root_tag="gathered_information")
-
- prompt = f"Synthesize a comprehensive answer to the original question.\n\n{context_xml}"
-
- agent_deps = DeepQADependencies(
- client=deps.client,
- context=state.context,
- console=deps.console,
- )
- result = await agent.run(prompt, deps=agent_deps)
-
- log(deps, state, "[bold green]✅ Answer complete![/bold green]")
- return End(result.output)
diff --git a/haiku_rag_slim/haiku/rag/qa/deep/prompts.py b/haiku_rag_slim/haiku/rag/qa/deep/prompts.py
index 2b0d16b7..75f27fd7 100644
--- a/haiku_rag_slim/haiku/rag/qa/deep/prompts.py
+++ b/haiku_rag_slim/haiku/rag/qa/deep/prompts.py
@@ -1,3 +1,5 @@
+"""Deep QA specific prompts."""
+
SYNTHESIS_PROMPT = """You are an expert at synthesizing information into clear, concise answers.
Task:
diff --git a/haiku_rag_slim/haiku/rag/qa/deep/state.py b/haiku_rag_slim/haiku/rag/qa/deep/state.py
index f76a99a3..0e07098e 100644
--- a/haiku_rag_slim/haiku/rag/qa/deep/state.py
+++ b/haiku_rag_slim/haiku/rag/qa/deep/state.py
@@ -1,15 +1,21 @@
+import asyncio
from dataclasses import dataclass
+from typing import TYPE_CHECKING
from rich.console import Console
from haiku.rag.client import HaikuRAG
from haiku.rag.qa.deep.dependencies import DeepQAContext
+if TYPE_CHECKING:
+ from haiku.rag.config.models import AppConfig
+
@dataclass
class DeepQADeps:
client: HaikuRAG
console: Console | None = None
+ semaphore: asyncio.Semaphore | None = None
def emit_log(self, message: str, state: "DeepQAState | None" = None) -> None:
if self.console:
@@ -21,5 +27,23 @@ class DeepQAState:
context: DeepQAContext
max_sub_questions: int = 3
max_iterations: int = 2
- max_concurrency: int = 3
+ max_concurrency: int = 1
iterations: int = 0
+
+ @classmethod
+ def from_config(cls, context: DeepQAContext, config: "AppConfig") -> "DeepQAState":
+ """Create a DeepQAState from an AppConfig.
+
+ Args:
+ context: The DeepQAContext containing the question and settings
+ config: The AppConfig object (uses config.qa for state parameters)
+
+ Returns:
+ A configured DeepQAState instance
+ """
+ return cls(
+ context=context,
+ max_sub_questions=config.qa.max_sub_questions,
+ max_iterations=config.qa.max_iterations,
+ max_concurrency=config.qa.max_concurrency,
+ )
diff --git a/haiku_rag_slim/haiku/rag/research/__init__.py b/haiku_rag_slim/haiku/rag/research/__init__.py
index e49dd051..9406a89c 100644
--- a/haiku_rag_slim/haiku/rag/research/__init__.py
+++ b/haiku_rag_slim/haiku/rag/research/__init__.py
@@ -1,3 +1,3 @@
-from haiku.rag.graph.models import SearchAnswer
+from haiku.rag.graph_common.models import SearchAnswer
from haiku.rag.research.dependencies import ResearchContext, ResearchDependencies
from haiku.rag.research.models import EvaluationResult, ResearchReport
diff --git a/haiku_rag_slim/haiku/rag/research/dependencies.py b/haiku_rag_slim/haiku/rag/research/dependencies.py
index 9606c658..09376b53 100644
--- a/haiku_rag_slim/haiku/rag/research/dependencies.py
+++ b/haiku_rag_slim/haiku/rag/research/dependencies.py
@@ -4,7 +4,7 @@ from pydantic import BaseModel, Field
from rich.console import Console
from haiku.rag.client import HaikuRAG
-from haiku.rag.graph.models import SearchAnswer
+from haiku.rag.graph_common.models import SearchAnswer
from haiku.rag.research.models import (
GapRecord,
InsightAnalysis,
diff --git a/haiku_rag_slim/haiku/rag/research/graph.py b/haiku_rag_slim/haiku/rag/research/graph.py
index 422ad1fd..a756476e 100644
--- a/haiku_rag_slim/haiku/rag/research/graph.py
+++ b/haiku_rag_slim/haiku/rag/research/graph.py
@@ -1,20 +1,429 @@
-from pydantic_graph import Graph
+from typing import Any
-from haiku.rag.graph.nodes.analysis import AnalyzeInsightsNode, DecisionNode
-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.models import ResearchReport
+from pydantic_ai import Agent, RunContext
+from pydantic_ai.format_prompt import format_as_xml
+from pydantic_ai.output import ToolOutput
+from pydantic_graph.beta import Graph, GraphBuilder, StepContext
+from pydantic_graph.beta.join import reduce_list_append
+
+from haiku.rag.config import Config
+from haiku.rag.config.models import AppConfig
+from haiku.rag.graph_common import get_model, log
+from haiku.rag.graph_common.models import ResearchPlan, SearchAnswer
+from haiku.rag.graph_common.prompts import PLAN_PROMPT, SEARCH_AGENT_PROMPT
+from haiku.rag.research.common import (
+ format_analysis_for_prompt,
+ format_context_for_prompt,
+)
+from haiku.rag.research.dependencies import ResearchDependencies
+from haiku.rag.research.models import (
+ EvaluationResult,
+ InsightAnalysis,
+ ResearchReport,
+)
+from haiku.rag.research.prompts import (
+ DECISION_AGENT_PROMPT,
+ INSIGHT_AGENT_PROMPT,
+ SYNTHESIS_AGENT_PROMPT,
+)
from haiku.rag.research.state import ResearchDeps, ResearchState
-def build_research_graph() -> Graph[ResearchState, ResearchDeps, ResearchReport]:
- return Graph(
- nodes=[
- PlanNode,
- SearchDispatchNode,
- AnalyzeInsightsNode,
- DecisionNode,
- SynthesizeNode,
- ]
+def build_research_graph(
+ config: AppConfig = Config,
+) -> Graph[ResearchState, ResearchDeps, None, ResearchReport]:
+ """Build the Research graph.
+
+ Args:
+ config: AppConfig object (uses config.research for provider, model, and graph parameters)
+
+ Returns:
+ Configured Research graph
+ """
+ provider = config.research.provider
+ model = config.research.model
+ g = GraphBuilder(
+ state_type=ResearchState,
+ deps_type=ResearchDeps,
+ output_type=ResearchReport,
)
+
+ @g.step
+ async def plan(ctx: StepContext[ResearchState, ResearchDeps, None]) -> None:
+ state = ctx.state
+ deps = ctx.deps
+
+ log(deps, state, "\n[bold cyan]📋 Creating research plan...[/bold cyan]")
+
+ plan_agent = Agent(
+ model=get_model(provider, model),
+ output_type=ResearchPlan,
+ instructions=(
+ PLAN_PROMPT
+ + "\n\nUse the gather_context tool once on the main question before planning."
+ ),
+ retries=3,
+ deps_type=ResearchDependencies,
+ )
+
+ @plan_agent.tool
+ async def gather_context(
+ ctx2: RunContext[ResearchDependencies], query: str, limit: int = 6
+ ) -> str:
+ results = await ctx2.deps.client.search(query, limit=limit)
+ expanded = await ctx2.deps.client.expand_context(results)
+ return "\n\n".join(chunk.content for chunk, _ in expanded)
+
+ prompt = (
+ "Plan a focused approach for the main question.\n\n"
+ f"Main question: {state.context.original_question}"
+ )
+
+ agent_deps = ResearchDependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ stream=deps.stream,
+ )
+ plan_result = await plan_agent.run(prompt, deps=agent_deps)
+ state.context.sub_questions = list(plan_result.output.sub_questions)
+
+ log(deps, state, "\n[bold green]✅ Plan Created:[/bold green]")
+ log(
+ deps,
+ state,
+ f" [bold]Main Question:[/bold] {state.context.original_question}",
+ )
+ log(deps, state, " [bold]Sub-questions:[/bold]")
+ for i, sq in enumerate(state.context.sub_questions, 1):
+ log(deps, state, f" {i}. {sq}")
+
+ @g.step
+ async def search_one(
+ ctx: StepContext[ResearchState, ResearchDeps, str],
+ ) -> SearchAnswer:
+ state = ctx.state
+ deps = ctx.deps
+ sub_q = ctx.inputs
+
+ # Create semaphore if not already provided
+ if deps.semaphore is None:
+ import asyncio
+
+ deps.semaphore = asyncio.Semaphore(state.max_concurrency)
+
+ # Use semaphore to control concurrency
+ async with deps.semaphore:
+ return await _do_search(state, deps, sub_q)
+
+ async def _do_search(
+ state: ResearchState,
+ deps: ResearchDeps,
+ sub_q: str,
+ ) -> SearchAnswer:
+ log(
+ deps,
+ state,
+ f"\n[bold cyan]🔍 Searching & Answering:[/bold cyan] {sub_q}",
+ )
+
+ agent = Agent(
+ model=get_model(provider, model),
+ output_type=ToolOutput(SearchAnswer, max_retries=3),
+ instructions=SEARCH_AGENT_PROMPT,
+ retries=3,
+ deps_type=ResearchDependencies,
+ )
+
+ @agent.tool
+ async def search_and_answer(
+ ctx2: RunContext[ResearchDependencies], query: str, limit: int = 5
+ ) -> str:
+ search_results = await ctx2.deps.client.search(query, limit=limit)
+ expanded = await ctx2.deps.client.expand_context(search_results)
+
+ entries: list[dict[str, Any]] = [
+ {
+ "text": chunk.content,
+ "score": score,
+ "document_uri": (chunk.document_title or chunk.document_uri or ""),
+ }
+ for chunk, score in expanded
+ ]
+ if not entries:
+ return (
+ f"No relevant information found in the knowledge base for: {query}"
+ )
+
+ return format_as_xml(entries, root_tag="snippets")
+
+ agent_deps = ResearchDependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ stream=deps.stream,
+ )
+ try:
+ result = await agent.run(sub_q, deps=agent_deps)
+ answer = result.output
+ if answer:
+ state.context.add_qa_response(answer)
+ preview = answer.answer[:150] + (
+ "…" if len(answer.answer) > 150 else ""
+ )
+ log(deps, state, f" [green]✓[/green] {preview}")
+ return answer
+ except Exception as e:
+ log(deps, state, f"[red]Search failed:[/red] {e}")
+ failure_answer = SearchAnswer(
+ query=sub_q,
+ answer=f"Search failed after retries: {str(e)}",
+ confidence=0.0,
+ )
+ return failure_answer
+
+ @g.step
+ async def get_batch(
+ ctx: StepContext[ResearchState, ResearchDeps, None | bool],
+ ) -> list[str] | None:
+ """Get all remaining questions for this iteration."""
+ state = ctx.state
+
+ if not state.context.sub_questions:
+ return None
+
+ # Take ALL remaining questions and process them in parallel
+ batch = list(state.context.sub_questions)
+ state.context.sub_questions.clear()
+ return batch
+
+ @g.step
+ async def analyze_insights(
+ ctx: StepContext[ResearchState, ResearchDeps, list[SearchAnswer]],
+ ) -> None:
+ state = ctx.state
+ deps = ctx.deps
+
+ log(
+ deps,
+ state,
+ "\n[bold cyan]🧭 Synthesizing new insights and gap status...[/bold cyan]",
+ )
+
+ agent = Agent(
+ model=get_model(provider, model),
+ output_type=InsightAnalysis,
+ instructions=INSIGHT_AGENT_PROMPT,
+ retries=3,
+ deps_type=ResearchDependencies,
+ )
+
+ context_xml = format_context_for_prompt(state.context)
+ prompt = (
+ "Review the latest research context and update the shared ledger of insights, gaps,"
+ " and follow-up questions.\n\n"
+ f"{context_xml}"
+ )
+ agent_deps = ResearchDependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ stream=deps.stream,
+ )
+ result = await agent.run(prompt, deps=agent_deps)
+ analysis: InsightAnalysis = result.output
+
+ state.context.integrate_analysis(analysis)
+ state.last_analysis = analysis
+
+ if analysis.commentary:
+ log(deps, state, f" Summary: {analysis.commentary}")
+ if analysis.highlights:
+ log(deps, state, " [bold]Updated insights:[/bold]")
+ for insight in analysis.highlights:
+ label = insight.status.value
+ log(
+ deps,
+ state,
+ f" • ({label}) {insight.summary}",
+ )
+ if analysis.gap_assessments:
+ log(deps, state, " [bold yellow]Gap updates:[/bold yellow]")
+ for gap in analysis.gap_assessments:
+ status = "resolved" if gap.resolved else "open"
+ severity = gap.severity.value
+ log(
+ deps,
+ state,
+ f" • ({severity}/{status}) {gap.description}",
+ )
+ if analysis.resolved_gaps:
+ log(deps, state, " [green]Resolved gaps:[/green]")
+ for resolved in analysis.resolved_gaps:
+ log(deps, state, f" • {resolved}")
+ if analysis.new_questions:
+ log(deps, state, " [cyan]Proposed follow-ups:[/cyan]")
+ for question in analysis.new_questions:
+ log(deps, state, f" • {question}")
+
+ @g.step
+ async def decide(ctx: StepContext[ResearchState, ResearchDeps, None]) -> bool:
+ state = ctx.state
+ deps = ctx.deps
+
+ log(
+ deps,
+ state,
+ "\n[bold cyan]📊 Evaluating research sufficiency...[/bold cyan]",
+ )
+
+ agent = Agent(
+ model=get_model(provider, model),
+ output_type=EvaluationResult,
+ instructions=DECISION_AGENT_PROMPT,
+ retries=3,
+ deps_type=ResearchDependencies,
+ )
+
+ context_xml = format_context_for_prompt(state.context)
+ analysis_xml = format_analysis_for_prompt(state.last_analysis)
+ prompt_parts = [
+ "Assess whether the research now answers the original question with adequate confidence.",
+ context_xml,
+ analysis_xml,
+ ]
+ if state.last_eval is not None:
+ prev = state.last_eval
+ prompt_parts.append(
+ ""
+ f"{prev.confidence_score:.2f}"
+ f"{str(prev.is_sufficient).lower()}"
+ f"{prev.reasoning}"
+ ""
+ )
+ prompt = "\n\n".join(part for part in prompt_parts if part)
+
+ agent_deps = ResearchDependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ stream=deps.stream,
+ )
+ decision_result = await agent.run(prompt, deps=agent_deps)
+ output = decision_result.output
+
+ state.last_eval = output
+ state.iterations += 1
+
+ for new_q in output.new_questions:
+ if new_q not in state.context.sub_questions:
+ state.context.sub_questions.append(new_q)
+
+ if output.key_insights:
+ log(deps, state, " [bold]Key insights:[/bold]")
+ for insight in output.key_insights:
+ log(deps, state, f" • {insight}")
+
+ if output.gaps:
+ log(deps, state, " [bold yellow]Remaining gaps:[/bold yellow]")
+ for gap in output.gaps:
+ log(deps, state, f" • {gap}")
+
+ log(
+ deps,
+ state,
+ f" Confidence: [yellow]{output.confidence_score:.1%}[/yellow]",
+ )
+ status = "[green]Yes[/green]" if output.is_sufficient else "[red]No[/red]"
+ log(deps, state, f" Sufficient: {status}")
+
+ should_continue = (
+ not output.is_sufficient
+ or output.confidence_score < state.confidence_threshold
+ ) and state.iterations < state.max_iterations
+
+ if not should_continue:
+ log(deps, state, "\n[bold green]✅ Stopping research.[/bold green]")
+
+ return should_continue
+
+ @g.step
+ async def synthesize(
+ ctx: StepContext[ResearchState, ResearchDeps, None | bool],
+ ) -> ResearchReport:
+ state = ctx.state
+ deps = ctx.deps
+
+ log(
+ deps,
+ state,
+ "\n[bold cyan]📝 Generating final research report...[/bold cyan]",
+ )
+
+ agent = Agent(
+ model=get_model(provider, model),
+ output_type=ResearchReport,
+ instructions=SYNTHESIS_AGENT_PROMPT,
+ retries=3,
+ deps_type=ResearchDependencies,
+ )
+
+ context_xml = format_context_for_prompt(state.context)
+ prompt = (
+ "Generate a comprehensive research report based on all gathered information.\n\n"
+ f"{context_xml}\n\n"
+ "Create a detailed report that synthesizes all findings into a coherent response."
+ )
+ agent_deps = ResearchDependencies(
+ client=deps.client,
+ context=state.context,
+ console=deps.console,
+ stream=deps.stream,
+ )
+ result = await agent.run(prompt, deps=agent_deps)
+
+ log(deps, state, "[bold green]✅ Research complete![/bold green]")
+ return result.output
+
+ # Build the graph structure
+ collect_answers = g.join(
+ reduce_list_append,
+ initial_factory=list[SearchAnswer],
+ )
+
+ g.add(
+ g.edge_from(g.start_node).to(plan),
+ g.edge_from(plan).to(get_batch),
+ )
+
+ # Branch based on whether we have questions
+ g.add(
+ g.edge_from(get_batch).to(
+ g.decision()
+ .branch(g.match(list).label("Has questions").map().to(search_one))
+ .branch(g.match(type(None)).label("No questions").to(synthesize))
+ ),
+ g.edge_from(search_one).to(collect_answers),
+ g.edge_from(collect_answers).to(analyze_insights),
+ g.edge_from(analyze_insights).to(decide),
+ )
+
+ # Branch based on decision
+ g.add(
+ g.edge_from(decide).to(
+ g.decision()
+ .branch(
+ g.match(bool, matches=lambda x: x)
+ .label("Continue research")
+ .to(get_batch)
+ )
+ .branch(
+ g.match(bool, matches=lambda x: not x)
+ .label("Done researching")
+ .to(synthesize)
+ )
+ ),
+ g.edge_from(synthesize).to(g.end_node),
+ )
+
+ return g.build()
diff --git a/haiku_rag_slim/haiku/rag/research/state.py b/haiku_rag_slim/haiku/rag/research/state.py
index e6df8c41..bfc48cbb 100644
--- a/haiku_rag_slim/haiku/rag/research/state.py
+++ b/haiku_rag_slim/haiku/rag/research/state.py
@@ -1,4 +1,6 @@
+import asyncio
from dataclasses import dataclass
+from typing import TYPE_CHECKING
from rich.console import Console
@@ -7,12 +9,16 @@ from haiku.rag.research.dependencies import ResearchContext
from haiku.rag.research.models import EvaluationResult, InsightAnalysis
from haiku.rag.research.stream import ResearchStream
+if TYPE_CHECKING:
+ from haiku.rag.config.models import AppConfig
+
@dataclass
class ResearchDeps:
client: HaikuRAG
console: Console | None = None
stream: ResearchStream | None = None
+ semaphore: asyncio.Semaphore | None = None
def emit_log(self, message: str, state: "ResearchState | None" = None) -> None:
if self.console:
@@ -26,7 +32,27 @@ class ResearchState:
context: ResearchContext
iterations: int = 0
max_iterations: int = 3
- max_concurrency: int = 1
confidence_threshold: float = 0.8
+ max_concurrency: int = 1
last_eval: EvaluationResult | None = None
last_analysis: InsightAnalysis | None = None
+
+ @classmethod
+ def from_config(
+ cls, context: ResearchContext, config: "AppConfig"
+ ) -> "ResearchState":
+ """Create a ResearchState from an AppConfig.
+
+ Args:
+ context: The ResearchContext containing the question and settings
+ config: The AppConfig object (uses config.research for state parameters)
+
+ Returns:
+ A configured ResearchState instance
+ """
+ return cls(
+ context=context,
+ max_iterations=config.research.max_iterations,
+ confidence_threshold=config.research.confidence_threshold,
+ max_concurrency=config.research.max_concurrency,
+ )
diff --git a/haiku_rag_slim/haiku/rag/research/stream.py b/haiku_rag_slim/haiku/rag/research/stream.py
index 3c1d56e7..5a2b1950 100644
--- a/haiku_rag_slim/haiku/rag/research/stream.py
+++ b/haiku_rag_slim/haiku/rag/research/stream.py
@@ -15,7 +15,6 @@ class ResearchStateSnapshot:
sub_questions: list[str]
iterations: int
max_iterations: int
- max_concurrency: int
confidence_threshold: float
pending_sub_questions: int
answered_questions: int
@@ -38,7 +37,6 @@ class ResearchStateSnapshot:
sub_questions=list(context.sub_questions),
iterations=state.iterations,
max_iterations=state.max_iterations,
- max_concurrency=state.max_concurrency,
confidence_threshold=state.confidence_threshold,
pending_sub_questions=len(context.sub_questions),
answered_questions=len(context.qa_responses),
@@ -124,7 +122,6 @@ class ResearchStream:
async def stream_research_graph(
graph,
- start,
state: "ResearchState",
deps,
) -> AsyncIterator[ResearchStreamEvent]:
@@ -132,7 +129,7 @@ async def stream_research_graph(
from contextlib import suppress
- from haiku.rag.research.state import ResearchDeps # Local import to avoid cycle
+ from haiku.rag.research.state import ResearchDeps
if not isinstance(deps, ResearchDeps):
raise TypeError("deps must be an instance of ResearchDeps")
@@ -142,25 +139,13 @@ async def stream_research_graph(
async def _execute() -> None:
try:
- report = None
- try:
- result = await graph.run(start, state=state, deps=deps)
- report = result.output
- except Exception:
- from pydantic_graph import End
-
- async with graph.iter(start, state=state, deps=deps) as run:
- node = run.next_node
- while not isinstance(node, End):
- node = await run.next(node)
- if run.result:
- report = run.result.output
+ report = await graph.run(state=state, deps=deps)
if report is None:
raise RuntimeError("Graph did not produce a report")
stream.report(report, state)
- except Exception as exc: # pragma: no cover - defensive path
+ except Exception as exc:
stream.error(exc, state)
finally:
await stream.close()
diff --git a/tests/test_app.py b/tests/test_app.py
index 73657202..87c484a5 100644
--- a/tests/test_app.py
+++ b/tests/test_app.py
@@ -383,11 +383,9 @@ async def test_ask_with_deep(app: HaikuRAGApp, monkeypatch):
from haiku.rag.qa.deep.models import DeepQAAnswer
mock_output = DeepQAAnswer(answer="Deep QA answer", sources=["test.md"])
- mock_result = MagicMock()
- mock_result.output = mock_output
mock_graph = AsyncMock()
- mock_graph.run.return_value = mock_result
+ mock_graph.run.return_value = mock_output
mock_client = AsyncMock()
mock_client.__aenter__.return_value = mock_client
@@ -415,11 +413,9 @@ async def test_ask_with_deep_and_cite(app: HaikuRAGApp, monkeypatch):
mock_output = DeepQAAnswer(
answer="Deep QA answer with citations [test.md]", sources=["test.md"]
)
- mock_result = MagicMock()
- mock_result.output = mock_output
mock_graph = AsyncMock()
- mock_graph.run.return_value = mock_result
+ mock_graph.run.return_value = mock_output
mock_client = AsyncMock()
mock_client.__aenter__.return_value = mock_client
@@ -445,11 +441,9 @@ async def test_ask_with_deep_and_verbose(app: HaikuRAGApp, monkeypatch):
from haiku.rag.qa.deep.models import DeepQAAnswer
mock_output = DeepQAAnswer(answer="Deep QA answer", sources=["test.md"])
- mock_result = MagicMock()
- mock_result.output = mock_output
mock_graph = AsyncMock()
- mock_graph.run.return_value = mock_result
+ mock_graph.run.return_value = mock_output
mock_client = AsyncMock()
mock_client.__aenter__.return_value = mock_client
diff --git a/tests/test_deep_qa.py b/tests/test_deep_qa.py
index 2c494b15..c220fcda 100644
--- a/tests/test_deep_qa.py
+++ b/tests/test_deep_qa.py
@@ -2,16 +2,23 @@ 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.graph_common.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.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, temp_db_path):
"""Test deep Q&A graph with mocked LLM using TestModel."""
+
+ # Mock get_model to return TestModel which generates valid schema-compliant data
+ def test_model_factory(provider, model):
+ return TestModel()
+
+ monkeypatch.setattr("haiku.rag.graph_common.utils.get_model", test_model_factory)
+ monkeypatch.setattr("haiku.rag.qa.deep.graph.get_model", test_model_factory)
+
graph = build_deep_qa_graph()
state = DeepQAState(
@@ -25,20 +32,12 @@ async def test_deep_qa_graph_end_to_end(monkeypatch, temp_db_path):
client = HaikuRAG(temp_db_path)
deps = DeepQADeps(client=client, console=None)
- # Mock get_model to return TestModel which generates valid schema-compliant data
- def test_model_factory(provider, model):
- return TestModel()
-
- monkeypatch.setattr("haiku.rag.graph.common.get_model", test_model_factory)
- monkeypatch.setattr("haiku.rag.qa.deep.nodes.get_model", test_model_factory)
-
- start = DeepQAPlanNode(provider="test", model="test")
- result = await graph.run(start_node=start, state=state, deps=deps)
+ result = await graph.run(state=state, deps=deps)
# 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)
+ assert result.answer is not None
+ assert isinstance(result.answer, str)
+ assert isinstance(result.sources, list)
client.close()
@@ -46,6 +45,14 @@ async def test_deep_qa_graph_end_to_end(monkeypatch, temp_db_path):
@pytest.mark.asyncio
async def test_deep_qa_with_citations(monkeypatch, temp_db_path):
"""Test deep Q&A with citations enabled using TestModel."""
+
+ # Mock get_model to return TestModel
+ def test_model_factory(provider, model):
+ return TestModel()
+
+ monkeypatch.setattr("haiku.rag.graph_common.utils.get_model", test_model_factory)
+ monkeypatch.setattr("haiku.rag.qa.deep.graph.get_model", test_model_factory)
+
graph = build_deep_qa_graph()
state = DeepQAState(
@@ -57,20 +64,12 @@ async def test_deep_qa_with_citations(monkeypatch, temp_db_path):
client = HaikuRAG(temp_db_path)
deps = DeepQADeps(client=client, console=None)
- # Mock get_model to return TestModel
- def test_model_factory(provider, model):
- return TestModel()
-
- monkeypatch.setattr("haiku.rag.graph.common.get_model", test_model_factory)
- monkeypatch.setattr("haiku.rag.qa.deep.nodes.get_model", test_model_factory)
-
- start = DeepQAPlanNode(provider="test", model="test")
- result = await graph.run(start_node=start, state=state, deps=deps)
+ result = await graph.run(state=state, deps=deps)
# 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)
+ assert result.answer is not None
+ assert isinstance(result.sources, list)
client.close()
diff --git a/tests/test_mcp.py b/tests/test_mcp.py
index e6a303c6..de2956d3 100644
--- a/tests/test_mcp.py
+++ b/tests/test_mcp.py
@@ -257,7 +257,7 @@ async def test_mcp_ask_question_deep():
mock_graph = AsyncMock()
mock_result = AsyncMock()
- mock_result.output.answer = "Deep answer"
+ mock_result.answer = "Deep answer"
mock_graph.run = AsyncMock(return_value=mock_result)
mock_graph_builder.return_value = mock_graph
@@ -299,9 +299,7 @@ async def test_mcp_research_question():
mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None)
mock_graph = AsyncMock()
- mock_result = AsyncMock()
- mock_result.output = mock_report
- mock_graph.run = AsyncMock(return_value=mock_result)
+ mock_graph.run = AsyncMock(return_value=mock_report)
mock_graph_builder.return_value = mock_graph
tools = await mcp.get_tools()
@@ -311,9 +309,6 @@ async def test_mcp_research_question():
result = await research_tool.fn( # type: ignore[attr-defined]
question="Research question?",
- max_iterations=1,
- confidence_threshold=0.5,
- max_concurrency=1,
)
assert result is not None
diff --git a/tests/test_research_graph.py b/tests/test_research_graph.py
index b986eeba..0a7c0b05 100644
--- a/tests/test_research_graph.py
+++ b/tests/test_research_graph.py
@@ -1,7 +1,8 @@
import asyncio
from haiku.rag.research.dependencies import ResearchContext
-from haiku.rag.research.graph import ResearchState, build_research_graph
+from haiku.rag.research.graph import build_research_graph
+from haiku.rag.research.state import ResearchState
def test_build_graph_and_state():
diff --git a/tests/test_research_graph_integration.py b/tests/test_research_graph_integration.py
index e80ff39e..5bdf40a6 100644
--- a/tests/test_research_graph_integration.py
+++ b/tests/test_research_graph_integration.py
@@ -2,20 +2,24 @@ import pytest
from pydantic_ai.models.test import TestModel
from haiku.rag.client import HaikuRAG
-from haiku.rag.graph.nodes.plan import PlanNode
from haiku.rag.research.dependencies import ResearchContext
-from haiku.rag.research.graph import (
- ResearchDeps,
- ResearchState,
- build_research_graph,
-)
+from haiku.rag.research.graph import build_research_graph
from haiku.rag.research.models import ResearchReport
+from haiku.rag.research.state import ResearchDeps, ResearchState
from haiku.rag.research.stream import stream_research_graph
@pytest.mark.asyncio
async def test_graph_end_to_end_with_test_model(monkeypatch, temp_db_path):
"""Test research graph with mocked LLM using TestModel."""
+
+ # Mock get_model to return TestModel which generates valid schema-compliant data
+ def test_model_factory(provider, model):
+ return TestModel()
+
+ monkeypatch.setattr("haiku.rag.graph_common.utils.get_model", test_model_factory)
+ monkeypatch.setattr("haiku.rag.research.graph.get_model", test_model_factory)
+
graph = build_research_graph()
state = ResearchState(
@@ -29,24 +33,9 @@ async def test_graph_end_to_end_with_test_model(monkeypatch, temp_db_path):
client = HaikuRAG(temp_db_path)
deps = ResearchDeps(client=client, console=None)
- # 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()
-
- 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(
- "haiku.rag.graph.nodes.synthesize.get_model", test_model_factory
- )
-
- start = PlanNode(provider="test", model="test")
-
collected = []
report = None
- async for event in stream_research_graph(graph, start, state, deps):
+ async for event in stream_research_graph(graph, state, deps):
collected.append(event)
if event.type == "report":
report = event.report