diff --git a/haiku_rag_slim/haiku/rag/graph/common/nodes.py b/haiku_rag_slim/haiku/rag/graph/common/nodes.py new file mode 100644 index 00000000..1097193b --- /dev/null +++ b/haiku_rag_slim/haiku/rag/graph/common/nodes.py @@ -0,0 +1,262 @@ +"""Common node implementations for graph workflows.""" + +import asyncio +from collections.abc import Awaitable, Callable +from typing import Any, Protocol + +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 StepContext + +from haiku.rag.client import HaikuRAG +from haiku.rag.graph.agui.emitter import AGUIEmitter +from haiku.rag.graph.common import get_model +from haiku.rag.graph.common.models import ResearchPlan, SearchAnswer +from haiku.rag.graph.common.prompts import PLAN_PROMPT, SEARCH_AGENT_PROMPT + + +class GraphContext(Protocol): + """Protocol for graph context objects.""" + + original_question: str + sub_questions: list[str] + + def add_qa_response(self, qa: SearchAnswer) -> None: + """Add a QA response to context.""" + ... + + +class GraphState(Protocol): + """Protocol for graph state objects.""" + + context: GraphContext + max_concurrency: int + + +class GraphDeps(Protocol): + """Protocol for graph dependencies.""" + + client: HaikuRAG + agui_emitter: AGUIEmitter[Any, Any] | None + semaphore: asyncio.Semaphore | None + + +class GraphAgentDeps(Protocol): + """Protocol for agent dependencies.""" + + client: HaikuRAG + context: GraphContext + + +def create_plan_node[AgentDepsT: GraphAgentDeps]( + provider: str, + model: str, + deps_type: type[AgentDepsT], + activity_message: str = "Creating plan", + output_retries: int | None = None, +) -> Callable[[StepContext[Any, Any, None]], Awaitable[None]]: + """Create a plan node for any graph. + + Args: + provider: Model provider (e.g., 'openai', 'anthropic') + model: Model name + deps_type: Type of dependencies for the agent (e.g., ResearchDependencies, DeepQADependencies) + activity_message: Message to show during planning activity + output_retries: Number of output retries for the agent (optional) + + Returns: + Async function that can be used as a graph step + """ + + async def plan(ctx: StepContext[Any, Any, None], /) -> None: + state: GraphState = ctx.state # type: ignore[assignment] + deps: GraphDeps = ctx.deps # type: ignore[assignment] + + if deps.agui_emitter: + deps.agui_emitter.start_step("plan") + deps.agui_emitter.update_activity("planning", activity_message) + + try: + # Build agent configuration + agent_config = { + "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": deps_type, + } + if output_retries is not None: + agent_config["output_retries"] = output_retries + + plan_agent = Agent(**agent_config) + + @plan_agent.tool + async def gather_context( + ctx2: RunContext[AgentDepsT], 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) + + # Tool is registered via decorator above + _ = gather_context + + prompt = ( + "Plan a focused approach for the main question.\n\n" + f"Main question: {state.context.original_question}" + ) + + # Create agent dependencies + agent_deps = deps_type(client=deps.client, context=state.context) # type: ignore[call-arg] + plan_result = await plan_agent.run(prompt, deps=agent_deps) + state.context.sub_questions = list(plan_result.output.sub_questions) + + # State now contains the plan - emit state update and narrate + if deps.agui_emitter: + deps.agui_emitter.update_state(state) + count = len(state.context.sub_questions) + deps.agui_emitter.update_activity( + "planning", f"Created plan with {count} sub-questions" + ) + finally: + if deps.agui_emitter: + deps.agui_emitter.finish_step() + + return plan + + +def create_search_node[AgentDepsT: GraphAgentDeps]( + provider: str, + model: str, + deps_type: type[AgentDepsT], + with_step_wrapper: bool = True, + success_message_format: str = "Answered: {sub_q}", + handle_exceptions: bool = False, +) -> Callable[[StepContext[Any, Any, str]], Awaitable[SearchAnswer]]: + """Create a search_one node for any graph. + + Args: + provider: Model provider + model: Model name + deps_type: Type of dependencies for the agent + with_step_wrapper: Whether to wrap with agui_emitter start/finish step + success_message_format: Format string for success activity message + handle_exceptions: Whether to handle exceptions with fallback answer + + Returns: + Async function that can be used as a graph step + """ + + async def search_one(ctx: StepContext[Any, Any, str], /) -> SearchAnswer: + state: GraphState = ctx.state # type: ignore[assignment] + deps: GraphDeps = ctx.deps # type: ignore[assignment] + sub_q = ctx.inputs + + if deps.agui_emitter and with_step_wrapper: + deps.agui_emitter.start_step("search_one") + + try: + # Create semaphore if not already provided + if deps.semaphore is None: + deps.semaphore = asyncio.Semaphore(state.max_concurrency) + + # Use semaphore to control concurrency + async with deps.semaphore: + return await _do_search( + state, + deps, + sub_q, + provider, + model, + deps_type, + success_message_format, + handle_exceptions, + ) + finally: + if deps.agui_emitter and with_step_wrapper: + deps.agui_emitter.finish_step() + + return search_one + + +async def _do_search[AgentDepsT: GraphAgentDeps]( + state: GraphState, + deps: GraphDeps, + sub_q: str, + provider: str, + model: str, + deps_type: type[AgentDepsT], + success_message_format: str, + handle_exceptions: bool, +) -> SearchAnswer: + """Internal search implementation.""" + if deps.agui_emitter: + deps.agui_emitter.update_activity("searching", f"Searching: {sub_q}") + + agent = Agent( + model=get_model(provider, model), + output_type=ToolOutput(SearchAnswer, max_retries=3), + instructions=SEARCH_AGENT_PROMPT, + retries=3, + deps_type=deps_type, + ) + + @agent.tool + async def search_and_answer( + ctx2: RunContext[AgentDepsT], 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") + + # Tool is registered via decorator above + _ = search_and_answer + + agent_deps = deps_type(client=deps.client, context=state.context) # type: ignore[call-arg] + + try: + result = await agent.run(sub_q, deps=agent_deps) + answer = result.output + if answer: + state.context.add_qa_response(answer) + # State updated with new answer - emit state update and narrate + if deps.agui_emitter: + deps.agui_emitter.update_state(state) + # Format the success message + if "{confidence}" in success_message_format: + message = success_message_format.format( + sub_q=sub_q, confidence=answer.confidence + ) + else: + message = success_message_format.format(sub_q=sub_q) + deps.agui_emitter.update_activity("searching", message) + return answer + except Exception as e: + if handle_exceptions: + # Narrate the error + if deps.agui_emitter: + deps.agui_emitter.update_activity("searching", f"Search failed: {e}") + failure_answer = SearchAnswer( + query=sub_q, + answer=f"Search failed after retries: {str(e)}", + confidence=0.0, + ) + return failure_answer + else: + raise diff --git a/haiku_rag_slim/haiku/rag/graph/deep_qa/graph.py b/haiku_rag_slim/haiku/rag/graph/deep_qa/graph.py index 61f2a121..550dc654 100644 --- a/haiku_rag_slim/haiku/rag/graph/deep_qa/graph.py +++ b/haiku_rag_slim/haiku/rag/graph/deep_qa/graph.py @@ -1,16 +1,13 @@ -from typing import Any - -from pydantic_ai import Agent, RunContext +from pydantic_ai import Agent 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 -from haiku.rag.graph.common.models import ResearchPlan, SearchAnswer -from haiku.rag.graph.common.prompts import PLAN_PROMPT, SEARCH_AGENT_PROMPT +from haiku.rag.graph.common.models import SearchAnswer +from haiku.rag.graph.common.nodes import create_plan_node, create_search_node from haiku.rag.graph.deep_qa.dependencies import DeepQADependencies from haiku.rag.graph.deep_qa.models import DeepQAAnswer, DeepQAEvaluation from haiku.rag.graph.deep_qa.prompts import ( @@ -40,133 +37,28 @@ def build_deep_qa_graph( output_type=DeepQAAnswer, ) - @g.step - async def plan(ctx: StepContext[DeepQAState, DeepQADeps, None]) -> None: - state = ctx.state - deps = ctx.deps - - if deps.agui_emitter: - deps.agui_emitter.start_step("plan") - deps.agui_emitter.update_activity("planning", "Planning approach") - - try: - 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, - ) - plan_result = await plan_agent.run(prompt, deps=agent_deps) - state.context.sub_questions = list(plan_result.output.sub_questions) - - if deps.agui_emitter: - deps.agui_emitter.update_state(state) - count = len(state.context.sub_questions) - deps.agui_emitter.update_activity( - "planning", f"Created plan with {count} sub-questions" - ) - finally: - if deps.agui_emitter: - deps.agui_emitter.finish_step() - - @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: - if deps.agui_emitter: - deps.agui_emitter.update_activity("searching", f"Searching: {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, + # Create and register the plan node using the factory + plan = g.step( + create_plan_node( + provider=provider, + model=model, + deps_type=DeepQADependencies, # type: ignore[arg-type] + activity_message="Planning approach", + output_retries=None, # Deep QA doesn't use output_retries ) + ) # type: ignore[arg-type] - @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, + # Create and register the search_one node using the factory + search_one = g.step( + create_search_node( + provider=provider, + model=model, + deps_type=DeepQADependencies, # type: ignore[arg-type] + with_step_wrapper=False, # Deep QA doesn't wrap with agui_emitter step + success_message_format="Answered: {sub_q}", + handle_exceptions=True, ) - try: - result = await agent.run(sub_q, deps=agent_deps) - answer = result.output - if answer: - state.context.add_qa_response(answer) - if deps.agui_emitter: - deps.agui_emitter.update_state(state) - deps.agui_emitter.update_activity("searching", f"Answered: {sub_q}") - return answer - except Exception as e: - failure_answer = SearchAnswer( - query=sub_q, - answer=f"Search failed after retries: {str(e)}", - confidence=0.0, - ) - return failure_answer + ) # type: ignore[arg-type] @g.step async def get_batch( diff --git a/haiku_rag_slim/haiku/rag/graph/research/graph.py b/haiku_rag_slim/haiku/rag/graph/research/graph.py index 227a33af..843bab6b 100644 --- a/haiku_rag_slim/haiku/rag/graph/research/graph.py +++ b/haiku_rag_slim/haiku/rag/graph/research/graph.py @@ -1,16 +1,12 @@ -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_ai import Agent 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 -from haiku.rag.graph.common.models import ResearchPlan, SearchAnswer -from haiku.rag.graph.common.prompts import PLAN_PROMPT, SEARCH_AGENT_PROMPT +from haiku.rag.graph.common.models import SearchAnswer +from haiku.rag.graph.common.nodes import create_plan_node, create_search_node from haiku.rag.graph.research.common import ( format_analysis_for_prompt, format_context_for_prompt, @@ -48,149 +44,28 @@ def build_research_graph( output_type=ResearchReport, ) - @g.step - async def plan(ctx: StepContext[ResearchState, ResearchDeps, None]) -> None: - state = ctx.state - deps = ctx.deps - - if deps.agui_emitter: - deps.agui_emitter.start_step("plan") - deps.agui_emitter.update_activity("planning", "Creating research plan") - - try: - 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, - output_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, - ) - plan_result = await plan_agent.run(prompt, deps=agent_deps) - state.context.sub_questions = list(plan_result.output.sub_questions) - - # State now contains the plan - emit state update and narrate - if deps.agui_emitter: - deps.agui_emitter.update_state(state) - count = len(state.context.sub_questions) - deps.agui_emitter.update_activity( - "planning", f"Created plan with {count} sub-questions" - ) - finally: - if deps.agui_emitter: - deps.agui_emitter.finish_step() - - @g.step - async def search_one( - ctx: StepContext[ResearchState, ResearchDeps, str], - ) -> SearchAnswer: - state = ctx.state - deps = ctx.deps - sub_q = ctx.inputs - - if deps.agui_emitter: - deps.agui_emitter.start_step("search_one") - - try: - # 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) - finally: - if deps.agui_emitter: - deps.agui_emitter.finish_step() - - async def _do_search( - state: ResearchState, - deps: ResearchDeps, - sub_q: str, - ) -> SearchAnswer: - if deps.agui_emitter: - deps.agui_emitter.update_activity("searching", f"Searching: {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, + # Create and register the plan node using the factory + plan = g.step( + create_plan_node( + provider=provider, + model=model, + deps_type=ResearchDependencies, # type: ignore[arg-type] + activity_message="Creating research plan", + output_retries=3, ) + ) # type: ignore[arg-type] - @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, + # Create and register the search_one node using the factory + search_one = g.step( + create_search_node( + provider=provider, + model=model, + deps_type=ResearchDependencies, # type: ignore[arg-type] + with_step_wrapper=True, + success_message_format="Found answer with {confidence:.0%} confidence", + handle_exceptions=True, ) - try: - result = await agent.run(sub_q, deps=agent_deps) - answer = result.output - if answer: - state.context.add_qa_response(answer) - # State updated with new answer - emit state update and narrate - if deps.agui_emitter: - deps.agui_emitter.update_state(state) - deps.agui_emitter.update_activity( - "searching", - f"Found answer with {answer.confidence:.0%} confidence", - ) - return answer - except Exception as e: - # Narrate the error - if deps.agui_emitter: - deps.agui_emitter.update_activity("searching", f"Search failed: {e}") - failure_answer = SearchAnswer( - query=sub_q, - answer=f"Search failed after retries: {str(e)}", - confidence=0.0, - ) - return failure_answer + ) # type: ignore[arg-type] @g.step async def get_batch(