diff --git a/docs/agents.md b/docs/agents.md index 9b310476..44b2d1a0 100644 --- a/docs/agents.md +++ b/docs/agents.md @@ -59,7 +59,7 @@ stateDiagram-v2 Key nodes: - **Plan**: Decomposes the question into focused sub-questions -- **Search (parallel)**: Answers all sub-questions in parallel +- **Search (parallel)**: Answers sub-questions in parallel (respects max_concurrency) - **Decision**: Evaluates if we have sufficient information or need another iteration - **Synthesize**: Generates the final comprehensive answer @@ -69,7 +69,7 @@ Key differences from Research: - **Direct answers**: Returns just the answer (not a full research report) - **Question-focused**: Optimized for answering specific questions, not open-ended research - **Supports citations**: Can include inline source citations like `[document.md]` -- **Configurable iterations**: Control max_iterations (default: 2) +- **Configurable iterations**: Control max_iterations (default: 2) and max_concurrency (default: 1) CLI usage: @@ -99,7 +99,8 @@ async with HaikuRAG(path_to_db) as client: state = DeepQAState( context=context, max_sub_questions=3, - max_iterations=2 + max_iterations=2, + max_concurrency=1 ) deps = DeepQADeps(client=client) @@ -175,6 +176,7 @@ async with HaikuRAG(path_to_db) as client: context=ResearchContext(original_question=question), max_iterations=2, confidence_threshold=0.8, + max_concurrency=2, ) deps = ResearchDeps(client=client) @@ -209,6 +211,7 @@ async with HaikuRAG(path_to_db) as client: context=ResearchContext(original_question=question), max_iterations=2, confidence_threshold=0.8, + max_concurrency=2, ) deps = ResearchDeps(client=client) diff --git a/haiku_rag_slim/haiku/rag/mcp.py b/haiku_rag_slim/haiku/rag/mcp.py index 612022ae..8a4573fa 100644 --- a/haiku_rag_slim/haiku/rag/mcp.py +++ b/haiku_rag_slim/haiku/rag/mcp.py @@ -2,11 +2,11 @@ from pathlib import Path from typing import Any from fastmcp import FastMCP +from haiku.rag.client import HaikuRAG +from haiku.rag.research.models import ResearchReport from pydantic import BaseModel -from haiku.rag.client import HaikuRAG from haiku.rag.config import Config -from haiku.rag.research.models import ResearchReport class SearchResult(BaseModel): @@ -191,11 +191,12 @@ def create_mcp_server(db_path: Path) -> FastMCP: try: async with HaikuRAG(db_path) as rag: if deep: - from haiku.rag.config import Config from haiku.rag.qa.deep.dependencies import DeepQAContext from haiku.rag.qa.deep.graph import build_deep_qa_graph from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState + from haiku.rag.config import Config + graph = build_deep_qa_graph( provider=Config.qa.provider, model=Config.qa.model, @@ -219,6 +220,7 @@ def create_mcp_server(db_path: Path) -> FastMCP: question: str, max_iterations: int = 3, confidence_threshold: float = 0.8, + max_concurrency: int = 1, ) -> ResearchReport | None: """Run multi-agent research to investigate a complex question. @@ -229,6 +231,7 @@ def create_mcp_server(db_path: Path) -> FastMCP: question: The research question to investigate. max_iterations: Maximum search/analyze iterations (default: 3). confidence_threshold: Minimum confidence score (0-1) to stop early (default: 0.8). + max_concurrency: Maximum concurrent sub-questions to process (default: 1). Returns: A research report with findings, or None if an error occurred. @@ -247,6 +250,7 @@ def create_mcp_server(db_path: Path) -> FastMCP: context=ResearchContext(original_question=question), max_iterations=max_iterations, confidence_threshold=confidence_threshold, + max_concurrency=max_concurrency, ) deps = ResearchDeps(client=rag) diff --git a/haiku_rag_slim/haiku/rag/qa/deep/graph.py b/haiku_rag_slim/haiku/rag/qa/deep/graph.py index d1e28fe2..4be28cd4 100644 --- a/haiku_rag_slim/haiku/rag/qa/deep/graph.py +++ b/haiku_rag_slim/haiku/rag/qa/deep/graph.py @@ -1,11 +1,5 @@ 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.beta import Graph, GraphBuilder, StepContext -from pydantic_graph.beta.join import reduce_list_append - from haiku.rag.graph_common import get_model, log from haiku.rag.graph_common.models import ResearchPlan, SearchAnswer from haiku.rag.graph_common.prompts import PLAN_PROMPT, SEARCH_AGENT_PROMPT @@ -17,6 +11,11 @@ from haiku.rag.qa.deep.prompts import ( SYNTHESIS_PROMPT_WITH_CITATIONS, ) from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState +from pydantic_ai import Agent, RunContext +from pydantic_ai.format_prompt import format_as_xml +from pydantic_ai.output import ToolOutput +from pydantic_graph.beta import Graph, GraphBuilder, StepContext +from pydantic_graph.beta.join import reduce_list_append def build_deep_qa_graph( @@ -85,6 +84,21 @@ def build_deep_qa_graph( deps = ctx.deps sub_q = ctx.inputs + # Create semaphore if not already provided + if deps.semaphore is None: + import asyncio + + deps.semaphore = asyncio.Semaphore(state.max_concurrency) + + # Use semaphore to control concurrency + async with deps.semaphore: + return await _do_search(state, deps, sub_q) + + async def _do_search( + state: DeepQAState, + deps: DeepQADeps, + sub_q: str, + ) -> SearchAnswer: log( deps, state, diff --git a/haiku_rag_slim/haiku/rag/qa/deep/state.py b/haiku_rag_slim/haiku/rag/qa/deep/state.py index 95c5bd5f..46750242 100644 --- a/haiku_rag_slim/haiku/rag/qa/deep/state.py +++ b/haiku_rag_slim/haiku/rag/qa/deep/state.py @@ -1,15 +1,16 @@ +import asyncio from dataclasses import dataclass -from rich.console import Console - from haiku.rag.client import HaikuRAG from haiku.rag.qa.deep.dependencies import DeepQAContext +from rich.console import Console @dataclass class DeepQADeps: client: HaikuRAG console: Console | None = None + semaphore: asyncio.Semaphore | None = None def emit_log(self, message: str, state: "DeepQAState | None" = None) -> None: if self.console: @@ -21,4 +22,5 @@ class DeepQAState: context: DeepQAContext max_sub_questions: int = 3 max_iterations: int = 2 + max_concurrency: int = 1 iterations: int = 0 diff --git a/haiku_rag_slim/haiku/rag/research/graph.py b/haiku_rag_slim/haiku/rag/research/graph.py index 22914bbf..117bf741 100644 --- a/haiku_rag_slim/haiku/rag/research/graph.py +++ b/haiku_rag_slim/haiku/rag/research/graph.py @@ -1,11 +1,5 @@ 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.beta import Graph, GraphBuilder, StepContext -from pydantic_graph.beta.join import reduce_list_append - from haiku.rag.graph_common import get_model, log from haiku.rag.graph_common.models import ResearchPlan, SearchAnswer from haiku.rag.graph_common.prompts import PLAN_PROMPT, SEARCH_AGENT_PROMPT @@ -25,6 +19,11 @@ from haiku.rag.research.prompts import ( SYNTHESIS_AGENT_PROMPT, ) from haiku.rag.research.state import ResearchDeps, ResearchState +from pydantic_ai import Agent, RunContext +from pydantic_ai.format_prompt import format_as_xml +from pydantic_ai.output import ToolOutput +from pydantic_graph.beta import Graph, GraphBuilder, StepContext +from pydantic_graph.beta.join import reduce_list_append def build_research_graph( @@ -94,6 +93,21 @@ def build_research_graph( deps = ctx.deps sub_q = ctx.inputs + # Create semaphore if not already provided + if deps.semaphore is None: + import asyncio + + deps.semaphore = asyncio.Semaphore(state.max_concurrency) + + # Use semaphore to control concurrency + async with deps.semaphore: + return await _do_search(state, deps, sub_q) + + async def _do_search( + state: ResearchState, + deps: ResearchDeps, + sub_q: str, + ) -> SearchAnswer: log( deps, state, diff --git a/haiku_rag_slim/haiku/rag/research/state.py b/haiku_rag_slim/haiku/rag/research/state.py index 989687f9..2c748103 100644 --- a/haiku_rag_slim/haiku/rag/research/state.py +++ b/haiku_rag_slim/haiku/rag/research/state.py @@ -1,11 +1,11 @@ +import asyncio from dataclasses import dataclass -from rich.console import Console - from haiku.rag.client import HaikuRAG from haiku.rag.research.dependencies import ResearchContext from haiku.rag.research.models import EvaluationResult, InsightAnalysis from haiku.rag.research.stream import ResearchStream +from rich.console import Console @dataclass @@ -13,6 +13,7 @@ class ResearchDeps: client: HaikuRAG console: Console | None = None stream: ResearchStream | None = None + semaphore: asyncio.Semaphore | None = None def emit_log(self, message: str, state: "ResearchState | None" = None) -> None: if self.console: @@ -27,5 +28,6 @@ class ResearchState: iterations: int = 0 max_iterations: int = 3 confidence_threshold: float = 0.8 + max_concurrency: int = 1 last_eval: EvaluationResult | None = None last_analysis: InsightAnalysis | None = None