Refactor research, deep ask to work with new beta pyndtic AI graph
This commit is contained in:
parent
8da62e3868
commit
bc71ad6fb0
15 changed files with 785 additions and 882 deletions
|
|
@ -13,12 +13,8 @@ from haiku.rag.config import Config
|
||||||
from haiku.rag.mcp import create_mcp_server
|
from haiku.rag.mcp import create_mcp_server
|
||||||
from haiku.rag.monitor import FileWatcher
|
from haiku.rag.monitor import FileWatcher
|
||||||
from haiku.rag.research.dependencies import ResearchContext
|
from haiku.rag.research.dependencies import ResearchContext
|
||||||
from haiku.rag.research.graph import (
|
from haiku.rag.research.graph import build_research_graph
|
||||||
PlanNode,
|
from haiku.rag.research.state import ResearchDeps, ResearchState
|
||||||
ResearchDeps,
|
|
||||||
ResearchState,
|
|
||||||
build_research_graph,
|
|
||||||
)
|
|
||||||
from haiku.rag.research.stream import stream_research_graph
|
from haiku.rag.research.stream import stream_research_graph
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
from haiku.rag.store.models.document import Document
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
@ -215,10 +211,12 @@ class HaikuRAGApp:
|
||||||
|
|
||||||
from haiku.rag.qa.deep.dependencies import DeepQAContext
|
from haiku.rag.qa.deep.dependencies import DeepQAContext
|
||||||
from haiku.rag.qa.deep.graph import build_deep_qa_graph
|
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
|
from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
|
||||||
|
|
||||||
graph = build_deep_qa_graph()
|
graph = build_deep_qa_graph(
|
||||||
|
provider=Config.qa.provider,
|
||||||
|
model=Config.qa.model,
|
||||||
|
)
|
||||||
context = DeepQAContext(
|
context = DeepQAContext(
|
||||||
original_question=question, use_citations=cite
|
original_question=question, use_citations=cite
|
||||||
)
|
)
|
||||||
|
|
@ -227,15 +225,8 @@ class HaikuRAGApp:
|
||||||
client=self.client, console=Console() if verbose else None
|
client=self.client, console=Console() if verbose else None
|
||||||
)
|
)
|
||||||
|
|
||||||
start_node = DeepQAPlanNode(
|
result = await graph.run(state=state, deps=deps)
|
||||||
provider=Config.qa.provider,
|
answer = result.answer
|
||||||
model=Config.qa.model,
|
|
||||||
)
|
|
||||||
|
|
||||||
result = await graph.run(
|
|
||||||
start_node=start_node, state=state, deps=deps
|
|
||||||
)
|
|
||||||
answer = result.output.answer
|
|
||||||
else:
|
else:
|
||||||
answer = await self.client.ask(question, cite=cite)
|
answer = await self.client.ask(question, cite=cite)
|
||||||
|
|
||||||
|
|
@ -262,7 +253,10 @@ class HaikuRAGApp:
|
||||||
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
|
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
|
||||||
self.console.print()
|
self.console.print()
|
||||||
|
|
||||||
graph = build_research_graph()
|
graph = build_research_graph(
|
||||||
|
provider=Config.research.provider or Config.qa.provider,
|
||||||
|
model=Config.research.model or Config.qa.model,
|
||||||
|
)
|
||||||
context = ResearchContext(original_question=question)
|
context = ResearchContext(original_question=question)
|
||||||
state = ResearchState(
|
state = ResearchState(
|
||||||
context=context,
|
context=context,
|
||||||
|
|
@ -274,12 +268,8 @@ class HaikuRAGApp:
|
||||||
client=client, console=self.console if verbose else None
|
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
|
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":
|
if event.type == "report":
|
||||||
report = event.report
|
report = event.report
|
||||||
break
|
break
|
||||||
|
|
|
||||||
|
|
@ -1 +0,0 @@
|
||||||
from haiku.rag.graph.models import ResearchPlan, SearchAnswer
|
|
||||||
|
|
@ -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)
|
|
||||||
|
|
@ -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(
|
|
||||||
"<previous_evaluation>"
|
|
||||||
f"<confidence>{prev.confidence_score:.2f}</confidence>"
|
|
||||||
f"<is_sufficient>{str(prev.is_sufficient).lower()}</is_sufficient>"
|
|
||||||
f"<reasoning>{prev.reasoning}</reasoning>"
|
|
||||||
"</previous_evaluation>"
|
|
||||||
)
|
|
||||||
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)
|
|
||||||
|
|
@ -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)
|
|
||||||
|
|
@ -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)
|
|
||||||
|
|
@ -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)
|
|
||||||
|
|
@ -194,25 +194,20 @@ def create_mcp_server(db_path: Path) -> FastMCP:
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
from haiku.rag.qa.deep.dependencies import DeepQAContext
|
from haiku.rag.qa.deep.dependencies import DeepQAContext
|
||||||
from haiku.rag.qa.deep.graph import build_deep_qa_graph
|
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
|
from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
|
||||||
|
|
||||||
graph = build_deep_qa_graph()
|
graph = build_deep_qa_graph(
|
||||||
|
provider=Config.qa.provider,
|
||||||
|
model=Config.qa.model,
|
||||||
|
)
|
||||||
context = DeepQAContext(
|
context = DeepQAContext(
|
||||||
original_question=question, use_citations=cite
|
original_question=question, use_citations=cite
|
||||||
)
|
)
|
||||||
state = DeepQAState(context=context)
|
state = DeepQAState(context=context)
|
||||||
deps = DeepQADeps(client=rag)
|
deps = DeepQADeps(client=rag)
|
||||||
|
|
||||||
start_node = DeepQAPlanNode(
|
result = await graph.run(state=state, deps=deps)
|
||||||
provider=Config.qa.provider,
|
answer = result.answer
|
||||||
model=Config.qa.model,
|
|
||||||
)
|
|
||||||
|
|
||||||
result = await graph.run(
|
|
||||||
start_node=start_node, state=state, deps=deps
|
|
||||||
)
|
|
||||||
answer = result.output.answer
|
|
||||||
else:
|
else:
|
||||||
answer = await rag.ask(question, cite=cite)
|
answer = await rag.ask(question, cite=cite)
|
||||||
return answer
|
return answer
|
||||||
|
|
@ -241,13 +236,15 @@ def create_mcp_server(db_path: Path) -> FastMCP:
|
||||||
A research report with findings, or None if an error occurred.
|
A research report with findings, or None if an error occurred.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
from haiku.rag.graph.nodes.plan import PlanNode
|
|
||||||
from haiku.rag.research.dependencies import ResearchContext
|
from haiku.rag.research.dependencies import ResearchContext
|
||||||
from haiku.rag.research.graph import build_research_graph
|
from haiku.rag.research.graph import build_research_graph
|
||||||
from haiku.rag.research.state import ResearchDeps, ResearchState
|
from haiku.rag.research.state import ResearchDeps, ResearchState
|
||||||
|
|
||||||
async with HaikuRAG(db_path) as rag:
|
async with HaikuRAG(db_path) as rag:
|
||||||
graph = build_research_graph()
|
graph = build_research_graph(
|
||||||
|
provider=Config.research.provider or Config.qa.provider,
|
||||||
|
model=Config.research.model or Config.qa.model,
|
||||||
|
)
|
||||||
state = ResearchState(
|
state = ResearchState(
|
||||||
context=ResearchContext(original_question=question),
|
context=ResearchContext(original_question=question),
|
||||||
max_iterations=max_iterations,
|
max_iterations=max_iterations,
|
||||||
|
|
@ -256,16 +253,9 @@ def create_mcp_server(db_path: Path) -> FastMCP:
|
||||||
)
|
)
|
||||||
deps = ResearchDeps(client=rag)
|
deps = ResearchDeps(client=rag)
|
||||||
|
|
||||||
result = await graph.run(
|
result = await graph.run(state=state, deps=deps)
|
||||||
PlanNode(
|
|
||||||
provider=Config.research.provider or Config.qa.provider,
|
|
||||||
model=Config.research.model or Config.qa.model,
|
|
||||||
),
|
|
||||||
state=state,
|
|
||||||
deps=deps,
|
|
||||||
)
|
|
||||||
|
|
||||||
return result.output
|
return result
|
||||||
except Exception:
|
except Exception:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,21 +1,332 @@
|
||||||
from pydantic_graph import Graph
|
from typing import Any
|
||||||
|
|
||||||
from haiku.rag.qa.deep.models import DeepQAAnswer
|
from pydantic_ai import Agent, RunContext
|
||||||
from haiku.rag.qa.deep.nodes import (
|
from pydantic_ai.format_prompt import format_as_xml
|
||||||
DeepQADecisionNode,
|
from pydantic_ai.output import ToolOutput
|
||||||
DeepQAPlanNode,
|
from pydantic_graph.beta import GraphBuilder, StepContext
|
||||||
DeepQASearchDispatchNode,
|
|
||||||
DeepQASynthesizeNode,
|
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
|
from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
|
||||||
|
|
||||||
|
|
||||||
def build_deep_qa_graph() -> Graph[DeepQAState, DeepQADeps, DeepQAAnswer]:
|
def build_deep_qa_graph(provider: str, model: str):
|
||||||
return Graph(
|
g = GraphBuilder(
|
||||||
nodes=[
|
state_type=DeepQAState,
|
||||||
DeepQAPlanNode,
|
deps_type=DeepQADeps,
|
||||||
DeepQASearchDispatchNode,
|
output_type=DeepQAAnswer,
|
||||||
DeepQADecisionNode,
|
|
||||||
DeepQASynthesizeNode,
|
|
||||||
]
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@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 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}")
|
||||||
|
|
||||||
|
@g.step
|
||||||
|
async def search_one(
|
||||||
|
ctx: StepContext[DeepQAState, DeepQADeps, str],
|
||||||
|
) -> SearchAnswer | None:
|
||||||
|
state = ctx.state
|
||||||
|
deps = ctx.deps
|
||||||
|
sub_q = ctx.inputs
|
||||||
|
|
||||||
|
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)
|
||||||
|
except Exception as e:
|
||||||
|
log(deps, state, f"[red]Search failed:[/red] {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
@g.step
|
||||||
|
async def decide(
|
||||||
|
ctx: StepContext[DeepQAState, DeepQADeps, list[SearchAnswer | None]],
|
||||||
|
) -> 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 get_batch(
|
||||||
|
ctx: StepContext[DeepQAState, DeepQADeps, None | bool],
|
||||||
|
) -> list[str] | None:
|
||||||
|
"""Get next batch of questions from state."""
|
||||||
|
state = ctx.state
|
||||||
|
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))
|
||||||
|
return batch if batch else None
|
||||||
|
|
||||||
|
@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
|
||||||
|
def collect_reducer(
|
||||||
|
acc: list[SearchAnswer | None], item: SearchAnswer | None
|
||||||
|
) -> list[SearchAnswer | None]:
|
||||||
|
return acc + [item] if item else acc
|
||||||
|
|
||||||
|
collect_answers = g.join(
|
||||||
|
collect_reducer,
|
||||||
|
initial_factory=lambda: [],
|
||||||
|
)
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
|
||||||
|
|
@ -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)
|
|
||||||
|
|
@ -1,20 +1,396 @@
|
||||||
from pydantic_graph import Graph
|
from typing import Any
|
||||||
|
|
||||||
from haiku.rag.graph.nodes.analysis import AnalyzeInsightsNode, DecisionNode
|
from pydantic_ai import Agent, RunContext
|
||||||
from haiku.rag.graph.nodes.plan import PlanNode
|
from pydantic_ai.format_prompt import format_as_xml
|
||||||
from haiku.rag.graph.nodes.search import SearchDispatchNode
|
from pydantic_ai.output import ToolOutput
|
||||||
from haiku.rag.graph.nodes.synthesize import SynthesizeNode
|
from pydantic_graph.beta import GraphBuilder, StepContext
|
||||||
from haiku.rag.research.models import ResearchReport
|
|
||||||
|
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.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
|
from haiku.rag.research.state import ResearchDeps, ResearchState
|
||||||
|
|
||||||
|
|
||||||
def build_research_graph() -> Graph[ResearchState, ResearchDeps, ResearchReport]:
|
def build_research_graph(provider: str, model: str):
|
||||||
return Graph(
|
g = GraphBuilder(
|
||||||
nodes=[
|
state_type=ResearchState,
|
||||||
PlanNode,
|
deps_type=ResearchDeps,
|
||||||
SearchDispatchNode,
|
output_type=ResearchReport,
|
||||||
AnalyzeInsightsNode,
|
|
||||||
DecisionNode,
|
|
||||||
SynthesizeNode,
|
|
||||||
]
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@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 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}")
|
||||||
|
|
||||||
|
@g.step
|
||||||
|
async def search_one(
|
||||||
|
ctx: StepContext[ResearchState, ResearchDeps, str],
|
||||||
|
) -> SearchAnswer | None:
|
||||||
|
state = ctx.state
|
||||||
|
deps = ctx.deps
|
||||||
|
sub_q = ctx.inputs
|
||||||
|
|
||||||
|
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)
|
||||||
|
except Exception as e:
|
||||||
|
log(deps, state, f"[red]Search failed:[/red] {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
@g.step
|
||||||
|
async def analyze_insights(
|
||||||
|
ctx: StepContext[ResearchState, ResearchDeps, list[SearchAnswer | None]],
|
||||||
|
) -> 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(
|
||||||
|
"<previous_evaluation>"
|
||||||
|
f"<confidence>{prev.confidence_score:.2f}</confidence>"
|
||||||
|
f"<is_sufficient>{str(prev.is_sufficient).lower()}</is_sufficient>"
|
||||||
|
f"<reasoning>{prev.reasoning}</reasoning>"
|
||||||
|
"</previous_evaluation>"
|
||||||
|
)
|
||||||
|
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 get_batch(
|
||||||
|
ctx: StepContext[ResearchState, ResearchDeps, None | bool],
|
||||||
|
) -> list[str] | None:
|
||||||
|
"""Get next batch of questions from state."""
|
||||||
|
state = ctx.state
|
||||||
|
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))
|
||||||
|
return batch if batch else None
|
||||||
|
|
||||||
|
@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
|
||||||
|
def collect_reducer(
|
||||||
|
acc: list[SearchAnswer | None], item: SearchAnswer | None
|
||||||
|
) -> list[SearchAnswer | None]:
|
||||||
|
return acc + [item] if item else acc
|
||||||
|
|
||||||
|
collect_answers = g.join(
|
||||||
|
collect_reducer,
|
||||||
|
initial_factory=lambda: [],
|
||||||
|
)
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
|
||||||
|
|
@ -124,7 +124,6 @@ class ResearchStream:
|
||||||
|
|
||||||
async def stream_research_graph(
|
async def stream_research_graph(
|
||||||
graph,
|
graph,
|
||||||
start,
|
|
||||||
state: "ResearchState",
|
state: "ResearchState",
|
||||||
deps,
|
deps,
|
||||||
) -> AsyncIterator[ResearchStreamEvent]:
|
) -> AsyncIterator[ResearchStreamEvent]:
|
||||||
|
|
@ -132,7 +131,7 @@ async def stream_research_graph(
|
||||||
|
|
||||||
from contextlib import suppress
|
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):
|
if not isinstance(deps, ResearchDeps):
|
||||||
raise TypeError("deps must be an instance of ResearchDeps")
|
raise TypeError("deps must be an instance of ResearchDeps")
|
||||||
|
|
@ -142,25 +141,13 @@ async def stream_research_graph(
|
||||||
|
|
||||||
async def _execute() -> None:
|
async def _execute() -> None:
|
||||||
try:
|
try:
|
||||||
report = None
|
report = await graph.run(state=state, deps=deps)
|
||||||
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
|
|
||||||
|
|
||||||
if report is None:
|
if report is None:
|
||||||
raise RuntimeError("Graph did not produce a report")
|
raise RuntimeError("Graph did not produce a report")
|
||||||
|
|
||||||
stream.report(report, state)
|
stream.report(report, state)
|
||||||
except Exception as exc: # pragma: no cover - defensive path
|
except Exception as exc:
|
||||||
stream.error(exc, state)
|
stream.error(exc, state)
|
||||||
finally:
|
finally:
|
||||||
await stream.close()
|
await stream.close()
|
||||||
|
|
|
||||||
|
|
@ -5,14 +5,21 @@ from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.graph.models import SearchAnswer
|
from haiku.rag.graph.models import SearchAnswer
|
||||||
from haiku.rag.qa.deep.dependencies import DeepQAContext
|
from haiku.rag.qa.deep.dependencies import DeepQAContext
|
||||||
from haiku.rag.qa.deep.graph import build_deep_qa_graph
|
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
|
from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_deep_qa_graph_end_to_end(monkeypatch, temp_db_path):
|
async def test_deep_qa_graph_end_to_end(monkeypatch, temp_db_path):
|
||||||
"""Test deep Q&A graph with mocked LLM using TestModel."""
|
"""Test deep Q&A graph with mocked LLM using TestModel."""
|
||||||
graph = build_deep_qa_graph()
|
|
||||||
|
# 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.graph.get_model", test_model_factory)
|
||||||
|
|
||||||
|
graph = build_deep_qa_graph(provider="test", model="test")
|
||||||
|
|
||||||
state = DeepQAState(
|
state = DeepQAState(
|
||||||
context=DeepQAContext(
|
context=DeepQAContext(
|
||||||
|
|
@ -25,20 +32,12 @@ async def test_deep_qa_graph_end_to_end(monkeypatch, temp_db_path):
|
||||||
client = HaikuRAG(temp_db_path)
|
client = HaikuRAG(temp_db_path)
|
||||||
deps = DeepQADeps(client=client, console=None)
|
deps = DeepQADeps(client=client, console=None)
|
||||||
|
|
||||||
# Mock get_model to return TestModel which generates valid schema-compliant data
|
result = await graph.run(state=state, deps=deps)
|
||||||
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)
|
|
||||||
|
|
||||||
# TestModel will generate valid structured output based on schemas
|
# TestModel will generate valid structured output based on schemas
|
||||||
assert result.output.answer is not None
|
assert result.answer is not None
|
||||||
assert isinstance(result.output.answer, str)
|
assert isinstance(result.answer, str)
|
||||||
assert isinstance(result.output.sources, list)
|
assert isinstance(result.sources, list)
|
||||||
|
|
||||||
client.close()
|
client.close()
|
||||||
|
|
||||||
|
|
@ -46,7 +45,15 @@ async def test_deep_qa_graph_end_to_end(monkeypatch, temp_db_path):
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_deep_qa_with_citations(monkeypatch, temp_db_path):
|
async def test_deep_qa_with_citations(monkeypatch, temp_db_path):
|
||||||
"""Test deep Q&A with citations enabled using TestModel."""
|
"""Test deep Q&A with citations enabled using TestModel."""
|
||||||
graph = build_deep_qa_graph()
|
|
||||||
|
# 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.graph.get_model", test_model_factory)
|
||||||
|
|
||||||
|
graph = build_deep_qa_graph(provider="test", model="test")
|
||||||
|
|
||||||
state = DeepQAState(
|
state = DeepQAState(
|
||||||
context=DeepQAContext(original_question="What is Python?", use_citations=True),
|
context=DeepQAContext(original_question="What is Python?", use_citations=True),
|
||||||
|
|
@ -57,20 +64,12 @@ async def test_deep_qa_with_citations(monkeypatch, temp_db_path):
|
||||||
client = HaikuRAG(temp_db_path)
|
client = HaikuRAG(temp_db_path)
|
||||||
deps = DeepQADeps(client=client, console=None)
|
deps = DeepQADeps(client=client, console=None)
|
||||||
|
|
||||||
# Mock get_model to return TestModel
|
result = await graph.run(state=state, deps=deps)
|
||||||
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)
|
|
||||||
|
|
||||||
# Verify citations flag was used
|
# Verify citations flag was used
|
||||||
assert state.context.use_citations is True
|
assert state.context.use_citations is True
|
||||||
assert result.output.answer is not None
|
assert result.answer is not None
|
||||||
assert isinstance(result.output.sources, list)
|
assert isinstance(result.sources, list)
|
||||||
|
|
||||||
client.close()
|
client.close()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,11 +1,12 @@
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
from haiku.rag.research.dependencies import ResearchContext
|
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():
|
def test_build_graph_and_state():
|
||||||
graph = build_research_graph()
|
graph = build_research_graph(provider="openai", model="gpt-4")
|
||||||
assert graph is not None
|
assert graph is not None
|
||||||
|
|
||||||
state = ResearchState(
|
state = ResearchState(
|
||||||
|
|
|
||||||
|
|
@ -2,21 +2,25 @@ import pytest
|
||||||
from pydantic_ai.models.test import TestModel
|
from pydantic_ai.models.test import TestModel
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
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.dependencies import ResearchContext
|
||||||
from haiku.rag.research.graph import (
|
from haiku.rag.research.graph import build_research_graph
|
||||||
ResearchDeps,
|
|
||||||
ResearchState,
|
|
||||||
build_research_graph,
|
|
||||||
)
|
|
||||||
from haiku.rag.research.models import ResearchReport
|
from haiku.rag.research.models import ResearchReport
|
||||||
|
from haiku.rag.research.state import ResearchDeps, ResearchState
|
||||||
from haiku.rag.research.stream import stream_research_graph
|
from haiku.rag.research.stream import stream_research_graph
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_graph_end_to_end_with_test_model(monkeypatch, temp_db_path):
|
async def test_graph_end_to_end_with_test_model(monkeypatch, temp_db_path):
|
||||||
"""Test research graph with mocked LLM using TestModel."""
|
"""Test research graph with mocked LLM using TestModel."""
|
||||||
graph = build_research_graph()
|
|
||||||
|
# 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.research.graph.get_model", test_model_factory)
|
||||||
|
|
||||||
|
graph = build_research_graph(provider="test", model="test")
|
||||||
|
|
||||||
state = ResearchState(
|
state = ResearchState(
|
||||||
context=ResearchContext(original_question="What is haiku.rag?"),
|
context=ResearchContext(original_question="What is haiku.rag?"),
|
||||||
|
|
@ -29,24 +33,9 @@ async def test_graph_end_to_end_with_test_model(monkeypatch, temp_db_path):
|
||||||
client = HaikuRAG(temp_db_path)
|
client = HaikuRAG(temp_db_path)
|
||||||
deps = ResearchDeps(client=client, console=None)
|
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 = []
|
collected = []
|
||||||
report = None
|
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)
|
collected.append(event)
|
||||||
if event.type == "report":
|
if event.type == "report":
|
||||||
report = event.report
|
report = event.report
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue