diff --git a/README.md b/README.md index e374b79c..99637f64 100644 --- a/README.md +++ b/README.md @@ -84,10 +84,10 @@ To customize settings, create a `haiku.rag.yaml` config file (see [Configuration ## Python Usage ```python -from haiku.rag.agui.stream import stream_graph from haiku.rag.client import HaikuRAG from haiku.rag.config import Config -from haiku.rag.research import ( +from haiku.rag.graph.agui import stream_graph +from haiku.rag.graph.research import ( ResearchContext, ResearchDeps, ResearchState, diff --git a/docs/agents.md b/docs/agents.md index 73355100..08b7cf6b 100644 --- a/docs/agents.md +++ b/docs/agents.md @@ -96,9 +96,9 @@ Python usage: ```python from haiku.rag.client import HaikuRAG from haiku.rag.config import Config -from haiku.rag.qa.deep.dependencies import DeepQAContext -from haiku.rag.qa.deep.graph import build_deep_qa_graph -from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState +from haiku.rag.graph.deep_qa.dependencies import DeepQAContext +from haiku.rag.graph.deep_qa.graph import build_deep_qa_graph +from haiku.rag.graph.deep_qa.state import DeepQADeps, DeepQAState async with HaikuRAG(path_to_db) as client: # Use global config (recommended) @@ -205,9 +205,9 @@ Python usage (blocking result): ```python from haiku.rag.client import HaikuRAG from haiku.rag.config import Config -from haiku.rag.research.dependencies import ResearchContext -from haiku.rag.research.graph import build_research_graph -from haiku.rag.research.state import ResearchDeps, ResearchState +from haiku.rag.graph.research.dependencies import ResearchContext +from haiku.rag.graph.research.graph import build_research_graph +from haiku.rag.graph.research.state import ResearchDeps, ResearchState async with HaikuRAG(path_to_db) as client: # Use global config (recommended) @@ -253,12 +253,12 @@ result = await graph.run(state=state, deps=deps) Python usage (streamed AG-UI events): ```python -from haiku.rag.agui.stream import stream_graph +from haiku.rag.graph.agui import stream_graph from haiku.rag.client import HaikuRAG from haiku.rag.config import Config -from haiku.rag.research.dependencies import ResearchContext -from haiku.rag.research.graph import build_research_graph -from haiku.rag.research.state import ResearchDeps, ResearchState +from haiku.rag.graph.research.dependencies import ResearchContext +from haiku.rag.graph.research.graph import build_research_graph +from haiku.rag.graph.research.state import ResearchDeps, ResearchState async with HaikuRAG(path_to_db) as client: graph = build_research_graph(config=Config) diff --git a/docs/cli.md b/docs/cli.md index 9f433cd7..b47c2587 100644 --- a/docs/cli.md +++ b/docs/cli.md @@ -162,7 +162,7 @@ Research parameters like `max_iterations`, `confidence_threshold`, and `max_conc When `--verbose` is set, the CLI consumes the research graph's AG-UI event stream, displaying step events and activity snapshots as agents progress through planning, search, evaluation, and synthesis. Without `--verbose`, only the final research report is displayed. -If you build your own integration, import `stream_graph` from `haiku.rag.agui.stream` to access AG-UI events (`STEP_STARTED`, `ACTIVITY_SNAPSHOT`, `STATE_SNAPSHOT`, `RUN_FINISHED`, etc.) and render them however you like while the graph is running. +If you build your own integration, import `stream_graph` from `haiku.rag.graph.agui` to access AG-UI events (`STEP_STARTED`, `ACTIVITY_SNAPSHOT`, `STATE_SNAPSHOT`, `RUN_FINISHED`, etc.) and render them however you like while the graph is running. ## Server diff --git a/examples/a2a-server/haiku_rag_a2a/a2a/__init__.py b/examples/a2a-server/haiku_rag_a2a/a2a/__init__.py index 03af95ab..d06eaae9 100644 --- a/examples/a2a-server/haiku_rag_a2a/a2a/__init__.py +++ b/examples/a2a-server/haiku_rag_a2a/a2a/__init__.py @@ -6,7 +6,7 @@ import logfire from pydantic_ai import Agent, RunContext from haiku.rag.config import AppConfig, Config -from haiku.rag.graph_common import get_model +from haiku.rag.graph.common import get_model from .context import load_message_history, save_message_history from .models import A2AConfig, AgentDependencies, SearchResult diff --git a/examples/ag-ui-research/backend/agent.py b/examples/ag-ui-research/backend/agent.py index 5362696a..049e0dbd 100644 --- a/examples/ag-ui-research/backend/agent.py +++ b/examples/ag-ui-research/backend/agent.py @@ -8,7 +8,7 @@ from pydantic_ai.ag_ui import StateDeps from haiku.rag.client import HaikuRAG from haiku.rag.config import Config -from haiku.rag.graph_common import get_model +from haiku.rag.graph.common import get_model class ResearchState(BaseModel): diff --git a/haiku_rag_slim/haiku/rag/app.py b/haiku_rag_slim/haiku/rag/app.py index 1ce07208..ba2e77c3 100644 --- a/haiku_rag_slim/haiku/rag/app.py +++ b/haiku_rag_slim/haiku/rag/app.py @@ -8,14 +8,14 @@ from rich.console import Console from rich.markdown import Markdown from rich.progress import Progress -from haiku.rag.agui import AGUIConsoleRenderer, stream_graph from haiku.rag.client import HaikuRAG from haiku.rag.config import AppConfig, Config +from haiku.rag.graph.agui import AGUIConsoleRenderer, stream_graph +from haiku.rag.graph.research.dependencies import ResearchContext +from haiku.rag.graph.research.graph import build_research_graph +from haiku.rag.graph.research.state import ResearchDeps, ResearchState from haiku.rag.mcp import create_mcp_server from haiku.rag.monitor import FileWatcher -from haiku.rag.research.dependencies import ResearchContext -from haiku.rag.research.graph import build_research_graph -from haiku.rag.research.state import ResearchDeps, ResearchState from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.document import Document @@ -216,9 +216,9 @@ class HaikuRAGApp: async with HaikuRAG(db_path=self.db_path, config=self.config) as self.client: try: if deep: - 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.graph.deep_qa.dependencies import DeepQAContext + from haiku.rag.graph.deep_qa.graph import build_deep_qa_graph + from haiku.rag.graph.deep_qa.state import DeepQADeps, DeepQAState graph = build_deep_qa_graph(config=self.config) context = DeepQAContext( @@ -229,7 +229,7 @@ class HaikuRAGApp: if verbose: # Use AG-UI renderer to process and display events - from haiku.rag.agui import AGUIConsoleRenderer + from haiku.rag.graph.agui import AGUIConsoleRenderer renderer = AGUIConsoleRenderer(self.console) result_dict = await renderer.render( @@ -287,7 +287,7 @@ class HaikuRAGApp: return # Convert dict to ResearchReport model - from haiku.rag.research.models import ResearchReport + from haiku.rag.graph.research.models import ResearchReport report = ResearchReport.model_validate(report_dict) @@ -497,7 +497,7 @@ class HaikuRAGApp: async def run_agui(): import uvicorn - from haiku.rag.agui import create_agui_server + from haiku.rag.graph.agui import create_agui_server logger.info( f"Starting AG-UI server on {self.config.agui.host}:{self.config.agui.port}" diff --git a/haiku_rag_slim/haiku/rag/graph/__init__.py b/haiku_rag_slim/haiku/rag/graph/__init__.py new file mode 100644 index 00000000..fff8e10a --- /dev/null +++ b/haiku_rag_slim/haiku/rag/graph/__init__.py @@ -0,0 +1,26 @@ +"""Graph module for haiku.rag. + +This module contains all graph-related functionality including: +- AG-UI protocol for graph streaming +- Common graph utilities and models +- Research graph implementation +- Deep QA graph implementation +""" + +from haiku.rag.graph.agui import ( + AGUIConsoleRenderer, + AGUIEmitter, + create_agui_server, + stream_graph, +) +from haiku.rag.graph.deep_qa.graph import build_deep_qa_graph +from haiku.rag.graph.research.graph import build_research_graph + +__all__ = [ + "AGUIConsoleRenderer", + "AGUIEmitter", + "build_deep_qa_graph", + "build_research_graph", + "create_agui_server", + "stream_graph", +] diff --git a/haiku_rag_slim/haiku/rag/agui/__init__.py b/haiku_rag_slim/haiku/rag/graph/agui/__init__.py similarity index 76% rename from haiku_rag_slim/haiku/rag/agui/__init__.py rename to haiku_rag_slim/haiku/rag/graph/agui/__init__.py index fab7a016..7259a6b0 100644 --- a/haiku_rag_slim/haiku/rag/agui/__init__.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/__init__.py @@ -1,8 +1,8 @@ """Generic AG-UI protocol support for haiku.rag graphs.""" -from haiku.rag.agui.cli_renderer import AGUIConsoleRenderer -from haiku.rag.agui.emitter import AGUIEmitter -from haiku.rag.agui.events import ( +from haiku.rag.graph.agui.cli_renderer import AGUIConsoleRenderer +from haiku.rag.graph.agui.emitter import AGUIEmitter +from haiku.rag.graph.agui.events import ( AGUIEvent, emit_activity, emit_activity_delta, @@ -18,14 +18,14 @@ from haiku.rag.agui.events import ( emit_text_message_end, emit_text_message_start, ) -from haiku.rag.agui.server import ( +from haiku.rag.graph.agui.server import ( RunAgentInput, create_agui_app, create_agui_server, format_sse_event, ) -from haiku.rag.agui.state import compute_state_delta -from haiku.rag.agui.stream import stream_graph +from haiku.rag.graph.agui.state import compute_state_delta +from haiku.rag.graph.agui.stream import stream_graph __all__ = [ "AGUIConsoleRenderer", diff --git a/haiku_rag_slim/haiku/rag/agui/cli_renderer.py b/haiku_rag_slim/haiku/rag/graph/agui/cli_renderer.py similarity index 99% rename from haiku_rag_slim/haiku/rag/agui/cli_renderer.py rename to haiku_rag_slim/haiku/rag/graph/agui/cli_renderer.py index c09e62b6..52e223dc 100644 --- a/haiku_rag_slim/haiku/rag/agui/cli_renderer.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/cli_renderer.py @@ -5,7 +5,7 @@ from typing import Any from rich.console import Console -from haiku.rag.agui.events import AGUIEvent +from haiku.rag.graph.agui.events import AGUIEvent class AGUIConsoleRenderer: diff --git a/haiku_rag_slim/haiku/rag/agui/emitter.py b/haiku_rag_slim/haiku/rag/graph/agui/emitter.py similarity index 99% rename from haiku_rag_slim/haiku/rag/agui/emitter.py rename to haiku_rag_slim/haiku/rag/graph/agui/emitter.py index b634f771..2c284d46 100644 --- a/haiku_rag_slim/haiku/rag/agui/emitter.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/emitter.py @@ -7,7 +7,7 @@ from uuid import uuid4 from pydantic import BaseModel -from haiku.rag.agui.events import ( +from haiku.rag.graph.agui.events import ( AGUIEvent, emit_activity, emit_run_error, diff --git a/haiku_rag_slim/haiku/rag/agui/events.py b/haiku_rag_slim/haiku/rag/graph/agui/events.py similarity index 99% rename from haiku_rag_slim/haiku/rag/agui/events.py rename to haiku_rag_slim/haiku/rag/graph/agui/events.py index d439d005..a7be3d73 100644 --- a/haiku_rag_slim/haiku/rag/agui/events.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/events.py @@ -5,7 +5,7 @@ from uuid import uuid4 from pydantic import BaseModel -from haiku.rag.agui.state import compute_state_delta +from haiku.rag.graph.agui.state import compute_state_delta # Type aliases for AG-UI events (actual types from ag_ui.core will be used at runtime) AGUIEvent = dict[str, Any] diff --git a/haiku_rag_slim/haiku/rag/agui/server.py b/haiku_rag_slim/haiku/rag/graph/agui/server.py similarity index 94% rename from haiku_rag_slim/haiku/rag/agui/server.py rename to haiku_rag_slim/haiku/rag/graph/agui/server.py index ff1be032..e7f7a829 100644 --- a/haiku_rag_slim/haiku/rag/agui/server.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/server.py @@ -13,9 +13,9 @@ from starlette.requests import Request from starlette.responses import JSONResponse, StreamingResponse from starlette.routing import Route -from haiku.rag.agui.events import AGUIEvent -from haiku.rag.agui.stream import stream_graph from haiku.rag.config.models import AGUIConfig +from haiku.rag.graph.agui.events import AGUIEvent +from haiku.rag.graph.agui.stream import stream_graph class GraphDeps(Protocol): @@ -157,12 +157,12 @@ def create_agui_server(config: Any, db_path: Any | None = None) -> Starlette: Starlette app with research and deep ask endpoints """ from haiku.rag.client import HaikuRAG - 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.research.dependencies import ResearchContext - from haiku.rag.research.graph import build_research_graph - from haiku.rag.research.state import ResearchDeps, ResearchState + from haiku.rag.graph.deep_qa.dependencies import DeepQAContext + from haiku.rag.graph.deep_qa.graph import build_deep_qa_graph + from haiku.rag.graph.deep_qa.state import DeepQADeps, DeepQAState + from haiku.rag.graph.research.dependencies import ResearchContext + from haiku.rag.graph.research.graph import build_research_graph + from haiku.rag.graph.research.state import ResearchDeps, ResearchState # Store client reference for proper lifecycle management _client_cache: dict[str, HaikuRAG] = {} diff --git a/haiku_rag_slim/haiku/rag/agui/state.py b/haiku_rag_slim/haiku/rag/graph/agui/state.py similarity index 100% rename from haiku_rag_slim/haiku/rag/agui/state.py rename to haiku_rag_slim/haiku/rag/graph/agui/state.py diff --git a/haiku_rag_slim/haiku/rag/agui/stream.py b/haiku_rag_slim/haiku/rag/graph/agui/stream.py similarity index 95% rename from haiku_rag_slim/haiku/rag/agui/stream.py rename to haiku_rag_slim/haiku/rag/graph/agui/stream.py index e91b9b64..75f6a3e5 100644 --- a/haiku_rag_slim/haiku/rag/agui/stream.py +++ b/haiku_rag_slim/haiku/rag/graph/agui/stream.py @@ -7,8 +7,8 @@ from typing import Any, Protocol from pydantic import BaseModel -from haiku.rag.agui.emitter import AGUIEmitter -from haiku.rag.agui.events import AGUIEvent +from haiku.rag.graph.agui.emitter import AGUIEmitter +from haiku.rag.graph.agui.events import AGUIEvent class GraphDeps(Protocol): diff --git a/haiku_rag_slim/haiku/rag/graph/common/__init__.py b/haiku_rag_slim/haiku/rag/graph/common/__init__.py new file mode 100644 index 00000000..e2a53f16 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/graph/common/__init__.py @@ -0,0 +1,5 @@ +"""Common utilities for graph implementations.""" + +from haiku.rag.graph.common.utils import get_model + +__all__ = ["get_model"] diff --git a/haiku_rag_slim/haiku/rag/graph_common/models.py b/haiku_rag_slim/haiku/rag/graph/common/models.py similarity index 100% rename from haiku_rag_slim/haiku/rag/graph_common/models.py rename to haiku_rag_slim/haiku/rag/graph/common/models.py diff --git a/haiku_rag_slim/haiku/rag/graph_common/prompts.py b/haiku_rag_slim/haiku/rag/graph/common/prompts.py similarity index 100% rename from haiku_rag_slim/haiku/rag/graph_common/prompts.py rename to haiku_rag_slim/haiku/rag/graph/common/prompts.py diff --git a/haiku_rag_slim/haiku/rag/graph_common/utils.py b/haiku_rag_slim/haiku/rag/graph/common/utils.py similarity index 74% rename from haiku_rag_slim/haiku/rag/graph_common/utils.py rename to haiku_rag_slim/haiku/rag/graph/common/utils.py index f24fb06c..19c424c8 100644 --- a/haiku_rag_slim/haiku/rag/graph_common/utils.py +++ b/haiku_rag_slim/haiku/rag/graph/common/utils.py @@ -1,7 +1,5 @@ """Common utilities for all graph implementations.""" -from typing import Any, Protocol - from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.providers.ollama import OllamaProvider from pydantic_ai.providers.openai import OpenAIProvider @@ -9,12 +7,6 @@ from pydantic_ai.providers.openai import OpenAIProvider from haiku.rag.config import Config -class HasEmitLog(Protocol): - """Protocol for objects that can emit log messages.""" - - def emit_log(self, message: str, state: Any = None) -> None: ... - - def get_model(provider: str, model: str) -> OpenAIChatModel | str: """ Get a model instance for the specified provider and model name. @@ -50,15 +42,3 @@ def get_model(provider: str, model: str) -> OpenAIChatModel | str: f"Unknown model provider: {provider}. " f"Supported providers: ollama, vllm, openai, anthropic, gemini, groq, bedrock" ) - - -def log(deps: HasEmitLog, state: Any, message: str) -> None: - """ - Emit a log message through the dependencies. - - Args: - deps: Dependencies object with emit_log method - state: Current state (passed to emit_log) - message: The message to log - """ - deps.emit_log(message, state) diff --git a/haiku_rag_slim/haiku/rag/graph/deep_qa/__init__.py b/haiku_rag_slim/haiku/rag/graph/deep_qa/__init__.py new file mode 100644 index 00000000..aaeb1ae9 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/graph/deep_qa/__init__.py @@ -0,0 +1 @@ +from haiku.rag.graph.deep_qa.models import DeepQAAnswer diff --git a/haiku_rag_slim/haiku/rag/qa/deep/dependencies.py b/haiku_rag_slim/haiku/rag/graph/deep_qa/dependencies.py similarity index 87% rename from haiku_rag_slim/haiku/rag/qa/deep/dependencies.py rename to haiku_rag_slim/haiku/rag/graph/deep_qa/dependencies.py index f8bce190..1c33262b 100644 --- a/haiku_rag_slim/haiku/rag/qa/deep/dependencies.py +++ b/haiku_rag_slim/haiku/rag/graph/deep_qa/dependencies.py @@ -1,8 +1,7 @@ from pydantic import BaseModel, Field -from rich.console import Console from haiku.rag.client import HaikuRAG -from haiku.rag.graph_common.models import SearchAnswer +from haiku.rag.graph.common.models import SearchAnswer class DeepQAContext(BaseModel): @@ -26,4 +25,3 @@ class DeepQADependencies(BaseModel): client: HaikuRAG = Field(description="RAG client for document operations") context: DeepQAContext = Field(description="Shared QA context") - console: Console | None = None diff --git a/haiku_rag_slim/haiku/rag/qa/deep/graph.py b/haiku_rag_slim/haiku/rag/graph/deep_qa/graph.py similarity index 95% rename from haiku_rag_slim/haiku/rag/qa/deep/graph.py rename to haiku_rag_slim/haiku/rag/graph/deep_qa/graph.py index 4ae92bc3..61f2a121 100644 --- a/haiku_rag_slim/haiku/rag/qa/deep/graph.py +++ b/haiku_rag_slim/haiku/rag/graph/deep_qa/graph.py @@ -8,17 +8,17 @@ 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.qa.deep.dependencies import DeepQADependencies -from haiku.rag.qa.deep.models import DeepQAAnswer, DeepQAEvaluation -from haiku.rag.qa.deep.prompts import ( +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.deep_qa.dependencies import DeepQADependencies +from haiku.rag.graph.deep_qa.models import DeepQAAnswer, DeepQAEvaluation +from haiku.rag.graph.deep_qa.prompts import ( DECISION_PROMPT, SYNTHESIS_PROMPT, SYNTHESIS_PROMPT_WITH_CITATIONS, ) -from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState +from haiku.rag.graph.deep_qa.state import DeepQADeps, DeepQAState def build_deep_qa_graph( @@ -77,7 +77,6 @@ def build_deep_qa_graph( agent_deps = DeepQADependencies( client=deps.client, context=state.context, - console=None, ) plan_result = await plan_agent.run(prompt, deps=agent_deps) state.context.sub_questions = list(plan_result.output.sub_questions) @@ -151,7 +150,6 @@ def build_deep_qa_graph( agent_deps = DeepQADependencies( client=deps.client, context=state.context, - console=None, ) try: result = await agent.run(sub_q, deps=agent_deps) @@ -228,7 +226,6 @@ def build_deep_qa_graph( agent_deps = DeepQADependencies( client=deps.client, context=state.context, - console=None, ) result = await agent.run(prompt, deps=agent_deps) evaluation = result.output @@ -302,7 +299,6 @@ def build_deep_qa_graph( agent_deps = DeepQADependencies( client=deps.client, context=state.context, - console=None, ) result = await agent.run(prompt, deps=agent_deps) diff --git a/haiku_rag_slim/haiku/rag/qa/deep/models.py b/haiku_rag_slim/haiku/rag/graph/deep_qa/models.py similarity index 100% rename from haiku_rag_slim/haiku/rag/qa/deep/models.py rename to haiku_rag_slim/haiku/rag/graph/deep_qa/models.py diff --git a/haiku_rag_slim/haiku/rag/qa/deep/prompts.py b/haiku_rag_slim/haiku/rag/graph/deep_qa/prompts.py similarity index 100% rename from haiku_rag_slim/haiku/rag/qa/deep/prompts.py rename to haiku_rag_slim/haiku/rag/graph/deep_qa/prompts.py diff --git a/haiku_rag_slim/haiku/rag/qa/deep/state.py b/haiku_rag_slim/haiku/rag/graph/deep_qa/state.py similarity index 96% rename from haiku_rag_slim/haiku/rag/qa/deep/state.py rename to haiku_rag_slim/haiku/rag/graph/deep_qa/state.py index a340d42d..179542a6 100644 --- a/haiku_rag_slim/haiku/rag/qa/deep/state.py +++ b/haiku_rag_slim/haiku/rag/graph/deep_qa/state.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Any from pydantic import BaseModel, Field from haiku.rag.client import HaikuRAG -from haiku.rag.qa.deep.dependencies import DeepQAContext +from haiku.rag.graph.deep_qa.dependencies import DeepQAContext if TYPE_CHECKING: from haiku.rag.config.models import AppConfig diff --git a/haiku_rag_slim/haiku/rag/graph/research/__init__.py b/haiku_rag_slim/haiku/rag/graph/research/__init__.py new file mode 100644 index 00000000..60f53a90 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/graph/research/__init__.py @@ -0,0 +1,3 @@ +from haiku.rag.graph.common.models import SearchAnswer +from haiku.rag.graph.research.dependencies import ResearchContext, ResearchDependencies +from haiku.rag.graph.research.models import EvaluationResult, ResearchReport diff --git a/haiku_rag_slim/haiku/rag/research/common.py b/haiku_rag_slim/haiku/rag/graph/research/common.py similarity index 95% rename from haiku_rag_slim/haiku/rag/research/common.py rename to haiku_rag_slim/haiku/rag/graph/research/common.py index bd6e349d..28c1aada 100644 --- a/haiku_rag_slim/haiku/rag/research/common.py +++ b/haiku_rag_slim/haiku/rag/graph/research/common.py @@ -1,7 +1,7 @@ from pydantic_ai import format_as_xml -from haiku.rag.research.dependencies import ResearchContext -from haiku.rag.research.models import InsightAnalysis +from haiku.rag.graph.research.dependencies import ResearchContext +from haiku.rag.graph.research.models import InsightAnalysis def format_context_for_prompt(context: ResearchContext) -> str: diff --git a/haiku_rag_slim/haiku/rag/research/dependencies.py b/haiku_rag_slim/haiku/rag/graph/research/dependencies.py similarity index 98% rename from haiku_rag_slim/haiku/rag/research/dependencies.py rename to haiku_rag_slim/haiku/rag/graph/research/dependencies.py index 39bf641c..fe65cc03 100644 --- a/haiku_rag_slim/haiku/rag/research/dependencies.py +++ b/haiku_rag_slim/haiku/rag/graph/research/dependencies.py @@ -3,8 +3,8 @@ from collections.abc import Iterable from pydantic import BaseModel, Field from haiku.rag.client import HaikuRAG -from haiku.rag.graph_common.models import SearchAnswer -from haiku.rag.research.models import ( +from haiku.rag.graph.common.models import SearchAnswer +from haiku.rag.graph.research.models import ( GapRecord, InsightAnalysis, InsightRecord, diff --git a/haiku_rag_slim/haiku/rag/research/graph.py b/haiku_rag_slim/haiku/rag/graph/research/graph.py similarity index 96% rename from haiku_rag_slim/haiku/rag/research/graph.py rename to haiku_rag_slim/haiku/rag/graph/research/graph.py index b2019e00..227a33af 100644 --- a/haiku_rag_slim/haiku/rag/research/graph.py +++ b/haiku_rag_slim/haiku/rag/graph/research/graph.py @@ -8,25 +8,25 @@ 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.research.common import ( +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.research.common import ( format_analysis_for_prompt, format_context_for_prompt, ) -from haiku.rag.research.dependencies import ResearchDependencies -from haiku.rag.research.models import ( +from haiku.rag.graph.research.dependencies import ResearchDependencies +from haiku.rag.graph.research.models import ( EvaluationResult, InsightAnalysis, ResearchReport, ) -from haiku.rag.research.prompts import ( +from haiku.rag.graph.research.prompts import ( DECISION_AGENT_PROMPT, INSIGHT_AGENT_PROMPT, SYNTHESIS_AGENT_PROMPT, ) -from haiku.rag.research.state import ResearchDeps, ResearchState +from haiku.rag.graph.research.state import ResearchDeps, ResearchState def build_research_graph( diff --git a/haiku_rag_slim/haiku/rag/research/models.py b/haiku_rag_slim/haiku/rag/graph/research/models.py similarity index 100% rename from haiku_rag_slim/haiku/rag/research/models.py rename to haiku_rag_slim/haiku/rag/graph/research/models.py diff --git a/haiku_rag_slim/haiku/rag/research/prompts.py b/haiku_rag_slim/haiku/rag/graph/research/prompts.py similarity index 100% rename from haiku_rag_slim/haiku/rag/research/prompts.py rename to haiku_rag_slim/haiku/rag/graph/research/prompts.py diff --git a/haiku_rag_slim/haiku/rag/research/state.py b/haiku_rag_slim/haiku/rag/graph/research/state.py similarity index 91% rename from haiku_rag_slim/haiku/rag/research/state.py rename to haiku_rag_slim/haiku/rag/graph/research/state.py index 72e341f5..8ac49c95 100644 --- a/haiku_rag_slim/haiku/rag/research/state.py +++ b/haiku_rag_slim/haiku/rag/graph/research/state.py @@ -5,12 +5,16 @@ from typing import TYPE_CHECKING from pydantic import BaseModel, Field from haiku.rag.client import HaikuRAG -from haiku.rag.research.dependencies import ResearchContext -from haiku.rag.research.models import EvaluationResult, InsightAnalysis, ResearchReport +from haiku.rag.graph.research.dependencies import ResearchContext +from haiku.rag.graph.research.models import ( + EvaluationResult, + InsightAnalysis, + ResearchReport, +) if TYPE_CHECKING: - from haiku.rag.agui.emitter import AGUIEmitter from haiku.rag.config.models import AppConfig + from haiku.rag.graph.agui.emitter import AGUIEmitter @dataclass diff --git a/haiku_rag_slim/haiku/rag/graph_common/__init__.py b/haiku_rag_slim/haiku/rag/graph_common/__init__.py deleted file mode 100644 index dc47bee0..00000000 --- a/haiku_rag_slim/haiku/rag/graph_common/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -"""Common utilities for graph implementations.""" - -from haiku.rag.graph_common.utils import get_model, log - -__all__ = ["get_model", "log"] diff --git a/haiku_rag_slim/haiku/rag/mcp.py b/haiku_rag_slim/haiku/rag/mcp.py index 25cd2fec..c6c3f97e 100644 --- a/haiku_rag_slim/haiku/rag/mcp.py +++ b/haiku_rag_slim/haiku/rag/mcp.py @@ -6,7 +6,7 @@ from pydantic import BaseModel from haiku.rag.client import HaikuRAG from haiku.rag.config import AppConfig, Config -from haiku.rag.research.models import ResearchReport +from haiku.rag.graph.research.models import ResearchReport class SearchResult(BaseModel): @@ -191,9 +191,9 @@ def create_mcp_server(db_path: Path, config: AppConfig = Config) -> FastMCP: try: async with HaikuRAG(db_path, config=config) as rag: if deep: - 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.graph.deep_qa.dependencies import DeepQAContext + from haiku.rag.graph.deep_qa.graph import build_deep_qa_graph + from haiku.rag.graph.deep_qa.state import DeepQADeps, DeepQAState graph = build_deep_qa_graph(config=config) context = DeepQAContext( @@ -226,9 +226,9 @@ def create_mcp_server(db_path: Path, config: AppConfig = Config) -> FastMCP: A research report with findings, or None if an error occurred. """ try: - from haiku.rag.research.dependencies import ResearchContext - from haiku.rag.research.graph import build_research_graph - from haiku.rag.research.state import ResearchDeps, ResearchState + from haiku.rag.graph.research.dependencies import ResearchContext + from haiku.rag.graph.research.graph import build_research_graph + from haiku.rag.graph.research.state import ResearchDeps, ResearchState async with HaikuRAG(db_path, config=config) as rag: graph = build_research_graph(config=config) diff --git a/haiku_rag_slim/haiku/rag/qa/deep/__init__.py b/haiku_rag_slim/haiku/rag/qa/deep/__init__.py deleted file mode 100644 index 0dc102f6..00000000 --- a/haiku_rag_slim/haiku/rag/qa/deep/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from haiku.rag.qa.deep.models import DeepQAAnswer diff --git a/haiku_rag_slim/haiku/rag/research/__init__.py b/haiku_rag_slim/haiku/rag/research/__init__.py deleted file mode 100644 index 9406a89c..00000000 --- a/haiku_rag_slim/haiku/rag/research/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from haiku.rag.graph_common.models import SearchAnswer -from haiku.rag.research.dependencies import ResearchContext, ResearchDependencies -from haiku.rag.research.models import EvaluationResult, ResearchReport diff --git a/tests/graph/__init__.py b/tests/graph/__init__.py new file mode 100644 index 00000000..8a28b35f --- /dev/null +++ b/tests/graph/__init__.py @@ -0,0 +1 @@ +"""Tests for haiku.rag.graph module.""" diff --git a/tests/agui/__init__.py b/tests/graph/agui/__init__.py similarity index 100% rename from tests/agui/__init__.py rename to tests/graph/agui/__init__.py diff --git a/tests/agui/test_cli_renderer.py b/tests/graph/agui/test_cli_renderer.py similarity index 97% rename from tests/agui/test_cli_renderer.py rename to tests/graph/agui/test_cli_renderer.py index c01d6934..738b0bcc 100644 --- a/tests/agui/test_cli_renderer.py +++ b/tests/graph/agui/test_cli_renderer.py @@ -4,8 +4,8 @@ import pytest from pydantic import BaseModel from rich.console import Console -from haiku.rag.agui.cli_renderer import AGUIConsoleRenderer -from haiku.rag.agui.events import ( +from haiku.rag.graph.agui.cli_renderer import AGUIConsoleRenderer +from haiku.rag.graph.agui.events import ( emit_activity, emit_run_error, emit_run_finished, diff --git a/tests/agui/test_emitter.py b/tests/graph/agui/test_emitter.py similarity index 99% rename from tests/agui/test_emitter.py rename to tests/graph/agui/test_emitter.py index 5a057247..5e6148d1 100644 --- a/tests/agui/test_emitter.py +++ b/tests/graph/agui/test_emitter.py @@ -5,7 +5,7 @@ import asyncio import pytest from pydantic import BaseModel -from haiku.rag.agui.emitter import AGUIEmitter +from haiku.rag.graph.agui.emitter import AGUIEmitter class TestState(BaseModel): diff --git a/tests/agui/test_events.py b/tests/graph/agui/test_events.py similarity index 99% rename from tests/agui/test_events.py rename to tests/graph/agui/test_events.py index 67db2e22..e400f99d 100644 --- a/tests/agui/test_events.py +++ b/tests/graph/agui/test_events.py @@ -2,7 +2,7 @@ from pydantic import BaseModel -from haiku.rag.agui.events import ( +from haiku.rag.graph.agui.events import ( emit_activity, emit_run_error, emit_run_finished, diff --git a/tests/agui/test_server.py b/tests/graph/agui/test_server.py similarity index 98% rename from tests/agui/test_server.py rename to tests/graph/agui/test_server.py index 2752b153..a145f7f6 100644 --- a/tests/agui/test_server.py +++ b/tests/graph/agui/test_server.py @@ -4,8 +4,8 @@ import pytest from pydantic import BaseModel from starlette.testclient import TestClient -from haiku.rag.agui.server import RunAgentInput, create_agui_app, format_sse_event from haiku.rag.config.models import AGUIConfig +from haiku.rag.graph.agui.server import RunAgentInput, create_agui_app, format_sse_event class SimpleState(BaseModel): diff --git a/tests/agui/test_stream.py b/tests/graph/agui/test_stream.py similarity index 98% rename from tests/agui/test_stream.py rename to tests/graph/agui/test_stream.py index 1aca01b4..daa428b0 100644 --- a/tests/agui/test_stream.py +++ b/tests/graph/agui/test_stream.py @@ -5,8 +5,8 @@ from dataclasses import dataclass import pytest from pydantic import BaseModel -from haiku.rag.agui.emitter import AGUIEmitter -from haiku.rag.agui.stream import stream_graph +from haiku.rag.graph.agui.emitter import AGUIEmitter +from haiku.rag.graph.agui.stream import stream_graph class TestState(BaseModel): diff --git a/tests/test_deep_qa.py b/tests/graph/test_deep_qa.py similarity index 83% rename from tests/test_deep_qa.py rename to tests/graph/test_deep_qa.py index b6757498..582b3242 100644 --- a/tests/test_deep_qa.py +++ b/tests/graph/test_deep_qa.py @@ -2,10 +2,10 @@ import pytest from pydantic_ai.models.test import TestModel from haiku.rag.client import HaikuRAG -from haiku.rag.graph_common.models import SearchAnswer -from haiku.rag.qa.deep.dependencies import DeepQAContext -from haiku.rag.qa.deep.graph import build_deep_qa_graph -from haiku.rag.qa.deep.state import DeepQADeps, DeepQAState +from haiku.rag.graph.common.models import SearchAnswer +from haiku.rag.graph.deep_qa.dependencies import DeepQAContext +from haiku.rag.graph.deep_qa.graph import build_deep_qa_graph +from haiku.rag.graph.deep_qa.state import DeepQADeps, DeepQAState @pytest.mark.asyncio @@ -16,8 +16,8 @@ async def test_deep_qa_graph_end_to_end(monkeypatch, temp_db_path): def test_model_factory(provider, model): return TestModel() - monkeypatch.setattr("haiku.rag.graph_common.utils.get_model", test_model_factory) - monkeypatch.setattr("haiku.rag.qa.deep.graph.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.common.utils.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.deep_qa.graph.get_model", test_model_factory) graph = build_deep_qa_graph() @@ -50,8 +50,8 @@ async def test_deep_qa_with_citations(monkeypatch, temp_db_path): def test_model_factory(provider, model): return TestModel() - monkeypatch.setattr("haiku.rag.graph_common.utils.get_model", test_model_factory) - monkeypatch.setattr("haiku.rag.qa.deep.graph.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.common.utils.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.deep_qa.graph.get_model", test_model_factory) graph = build_deep_qa_graph() diff --git a/tests/test_research_graph.py b/tests/graph/test_research_graph.py similarity index 85% rename from tests/test_research_graph.py rename to tests/graph/test_research_graph.py index 94f1f6a3..f452b022 100644 --- a/tests/test_research_graph.py +++ b/tests/graph/test_research_graph.py @@ -3,11 +3,11 @@ import asyncio import pytest from pydantic_ai.models.test import TestModel -from haiku.rag.agui.stream import stream_graph from haiku.rag.client import HaikuRAG -from haiku.rag.research.dependencies import ResearchContext -from haiku.rag.research.graph import build_research_graph -from haiku.rag.research.state import ResearchDeps, ResearchState +from haiku.rag.graph.agui.stream import stream_graph +from haiku.rag.graph.research.dependencies import ResearchContext +from haiku.rag.graph.research.graph import build_research_graph +from haiku.rag.graph.research.state import ResearchDeps, ResearchState def test_build_graph_and_state(): @@ -39,8 +39,8 @@ async def test_graph_end_to_end_with_test_model(monkeypatch, temp_db_path): def test_model_factory(_provider, _model): return TestModel() - monkeypatch.setattr("haiku.rag.graph_common.utils.get_model", test_model_factory) - monkeypatch.setattr("haiku.rag.research.graph.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.common.utils.get_model", test_model_factory) + monkeypatch.setattr("haiku.rag.graph.research.graph.get_model", test_model_factory) graph = build_research_graph() diff --git a/tests/test_app.py b/tests/test_app.py index 51153aa4..e189c022 100644 --- a/tests/test_app.py +++ b/tests/test_app.py @@ -343,7 +343,7 @@ async def test_ask_with_verbose(app: HaikuRAGApp, monkeypatch): @pytest.mark.asyncio async def test_ask_with_deep(app: HaikuRAGApp, monkeypatch): """Test asking a question with deep QA.""" - from haiku.rag.qa.deep.models import DeepQAAnswer + from haiku.rag.graph.deep_qa.models import DeepQAAnswer mock_output = DeepQAAnswer(answer="Deep QA answer", sources=["test.md"]) @@ -358,7 +358,7 @@ async def test_ask_with_deep(app: HaikuRAGApp, monkeypatch): with patch("haiku.rag.app.HaikuRAG", return_value=mock_client): with patch( - "haiku.rag.qa.deep.graph.build_deep_qa_graph", return_value=mock_graph + "haiku.rag.graph.deep_qa.graph.build_deep_qa_graph", return_value=mock_graph ): await app.ask("test question", deep=True) @@ -371,7 +371,7 @@ async def test_ask_with_deep(app: HaikuRAGApp, monkeypatch): @pytest.mark.asyncio async def test_ask_with_deep_and_cite(app: HaikuRAGApp, monkeypatch): """Test asking a question with deep QA and citations.""" - from haiku.rag.qa.deep.models import DeepQAAnswer + from haiku.rag.graph.deep_qa.models import DeepQAAnswer mock_output = DeepQAAnswer( answer="Deep QA answer with citations [test.md]", sources=["test.md"] @@ -388,7 +388,7 @@ async def test_ask_with_deep_and_cite(app: HaikuRAGApp, monkeypatch): with patch("haiku.rag.app.HaikuRAG", return_value=mock_client): with patch( - "haiku.rag.qa.deep.graph.build_deep_qa_graph", return_value=mock_graph + "haiku.rag.graph.deep_qa.graph.build_deep_qa_graph", return_value=mock_graph ): await app.ask("test question", deep=True, cite=True) @@ -417,10 +417,10 @@ async def test_ask_with_deep_and_verbose(app: HaikuRAGApp, monkeypatch): with patch("haiku.rag.app.HaikuRAG", return_value=mock_client): with patch( - "haiku.rag.qa.deep.graph.build_deep_qa_graph", return_value=mock_graph + "haiku.rag.graph.deep_qa.graph.build_deep_qa_graph", return_value=mock_graph ): with patch( - "haiku.rag.agui.AGUIConsoleRenderer", return_value=mock_renderer + "haiku.rag.graph.agui.AGUIConsoleRenderer", return_value=mock_renderer ): await app.ask("test question", deep=True, verbose=True) diff --git a/tests/test_mcp.py b/tests/test_mcp.py index de2956d3..389f11bc 100644 --- a/tests/test_mcp.py +++ b/tests/test_mcp.py @@ -4,8 +4,8 @@ from unittest.mock import AsyncMock, patch import pytest +from haiku.rag.graph.research.models import ResearchReport from haiku.rag.mcp import create_mcp_server -from haiku.rag.research.models import ResearchReport from haiku.rag.store.models.document import Document @@ -249,7 +249,9 @@ async def test_mcp_ask_question_deep(): with ( patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class, - patch("haiku.rag.qa.deep.graph.build_deep_qa_graph") as mock_graph_builder, + patch( + "haiku.rag.graph.deep_qa.graph.build_deep_qa_graph" + ) as mock_graph_builder, ): mock_rag = AsyncMock() mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) @@ -291,7 +293,7 @@ async def test_mcp_research_question(): with ( patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class, patch( - "haiku.rag.research.graph.build_research_graph" + "haiku.rag.graph.research.graph.build_research_graph" ) as mock_graph_builder, ): mock_rag = AsyncMock()