From 7cc561d1db17edd436d879fd431f69bc59fa5008 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Mon, 12 Jan 2026 12:30:07 +0200 Subject: [PATCH] Refactor to make agents a top-level module. Bring in the conversational agent from the app --- app/backend/main.py | 7 +- docs/agents.md | 14 +- evaluations/evaluations/benchmark.py | 2 +- haiku_rag_slim/haiku/rag/agents/__init__.py | 48 +++++++ .../haiku/rag/agents/chat/__init__.py | 23 +++ .../haiku/rag/agents/chat}/agent.py | 133 ++---------------- .../haiku/rag/agents/chat/prompts.py | 41 ++++++ .../haiku/rag/agents/chat/search.py | 28 +--- haiku_rag_slim/haiku/rag/agents/chat/state.py | 93 ++++++++++++ .../haiku/rag/{ => agents}/qa/__init__.py | 4 +- .../haiku/rag/{ => agents}/qa/agent.py | 8 +- .../haiku/rag/{ => agents}/qa/prompts.py | 0 .../haiku/rag/agents/research/__init__.py | 6 + .../research/dependencies.py | 2 +- .../rag/{graph => agents}/research/graph.py | 12 +- .../rag/{graph => agents}/research/models.py | 0 .../rag/{graph => agents}/research/prompts.py | 0 .../rag/{graph => agents}/research/state.py | 4 +- haiku_rag_slim/haiku/rag/app.py | 6 +- haiku_rag_slim/haiku/rag/client.py | 4 +- haiku_rag_slim/haiku/rag/graph/__init__.py | 5 - .../haiku/rag/graph/research/__init__.py | 6 - haiku_rag_slim/haiku/rag/mcp.py | 14 +- haiku_rag_slim/haiku/rag/utils.py | 2 +- tests/agents/chat/test_chat_agent.py | 105 ++++++++++++++ tests/{ => agents/qa}/test_qa.py | 4 +- .../test_graph_end_to_end.yaml | 0 ...est_research_graph_uses_search_filter.yaml | 0 .../test_search_filter_none_searches_all.yaml | 0 .../test_search_filter_restricts_results.yaml | 0 .../research}/test_research_graph.py | 8 +- .../research}/test_search_filter.py | 6 +- tests/graph/__init__.py | 1 - tests/test_app.py | 6 +- tests/test_mcp.py | 6 +- 35 files changed, 392 insertions(+), 206 deletions(-) create mode 100644 haiku_rag_slim/haiku/rag/agents/__init__.py create mode 100644 haiku_rag_slim/haiku/rag/agents/chat/__init__.py rename {app/backend => haiku_rag_slim/haiku/rag/agents/chat}/agent.py (66%) create mode 100644 haiku_rag_slim/haiku/rag/agents/chat/prompts.py rename app/backend/search_agent.py => haiku_rag_slim/haiku/rag/agents/chat/search.py (73%) create mode 100644 haiku_rag_slim/haiku/rag/agents/chat/state.py rename haiku_rag_slim/haiku/rag/{ => agents}/qa/__init__.py (89%) rename haiku_rag_slim/haiku/rag/{ => agents}/qa/agent.py (93%) rename haiku_rag_slim/haiku/rag/{ => agents}/qa/prompts.py (100%) create mode 100644 haiku_rag_slim/haiku/rag/agents/research/__init__.py rename haiku_rag_slim/haiku/rag/{graph => agents}/research/dependencies.py (95%) rename haiku_rag_slim/haiku/rag/{graph => agents}/research/graph.py (98%) rename haiku_rag_slim/haiku/rag/{graph => agents}/research/models.py (100%) rename haiku_rag_slim/haiku/rag/{graph => agents}/research/prompts.py (100%) rename haiku_rag_slim/haiku/rag/{graph => agents}/research/state.py (94%) delete mode 100644 haiku_rag_slim/haiku/rag/graph/__init__.py delete mode 100644 haiku_rag_slim/haiku/rag/graph/research/__init__.py create mode 100644 tests/agents/chat/test_chat_agent.py rename tests/{ => agents/qa}/test_qa.py (95%) rename tests/{graph => agents/research}/cassettes/test_research_graph/test_graph_end_to_end.yaml (100%) rename tests/{graph => agents/research}/cassettes/test_search_filter/test_research_graph_uses_search_filter.yaml (100%) rename tests/{graph => agents/research}/cassettes/test_search_filter/test_search_filter_none_searches_all.yaml (100%) rename tests/{graph => agents/research}/cassettes/test_search_filter/test_search_filter_restricts_results.yaml (100%) rename tests/{graph => agents/research}/test_research_graph.py (76%) rename tests/{graph => agents/research}/test_search_filter.py (94%) delete mode 100644 tests/graph/__init__.py diff --git a/app/backend/main.py b/app/backend/main.py index d3c4ebf2..6695be55 100644 --- a/app/backend/main.py +++ b/app/backend/main.py @@ -2,7 +2,6 @@ import logging import os from pathlib import Path -from agent import ChatDeps, ChatSessionState, QAResponse, create_chat_agent from dotenv import load_dotenv from pydantic_ai.ui import SSE_CONTENT_TYPE from pydantic_ai.ui.ag_ui import AGUIAdapter @@ -13,6 +12,12 @@ from starlette.requests import Request from starlette.responses import JSONResponse, Response, StreamingResponse from starlette.routing import Route +from haiku.rag.agents.chat import ( + ChatDeps, + ChatSessionState, + QAResponse, + create_chat_agent, +) from haiku.rag.client import HaikuRAG from haiku.rag.config import load_yaml_config from haiku.rag.config.models import AppConfig diff --git a/docs/agents.md b/docs/agents.md index 25f3aa2c..ca082e90 100644 --- a/docs/agents.md +++ b/docs/agents.md @@ -33,7 +33,7 @@ haiku-rag ask "What are the main features of haiku.rag?" --deep ```python from haiku.rag.client import HaikuRAG -from haiku.rag.qa.agent import QuestionAnswerAgent +from haiku.rag.agents.qa.agent import QuestionAnswerAgent async with HaikuRAG(path_to_db) as client: agent = QuestionAnswerAgent( @@ -105,9 +105,9 @@ haiku-rag research "What are the key findings?" --filter "uri LIKE '%report%'" ```python from haiku.rag.client import HaikuRAG from haiku.rag.config import Config -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.agents.research.dependencies import ResearchContext +from haiku.rag.agents.research.graph import build_research_graph +from haiku.rag.agents.research.state import ResearchDeps, ResearchState async with HaikuRAG(path_to_db) as client: graph = build_research_graph(config=Config) @@ -126,9 +126,9 @@ async with HaikuRAG(path_to_db) as client: ```python from haiku.rag.client import HaikuRAG from haiku.rag.config.models import AppConfig, ResearchConfig -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.agents.research.dependencies import ResearchContext +from haiku.rag.agents.research.graph import build_research_graph +from haiku.rag.agents.research.state import ResearchDeps, ResearchState custom_config = AppConfig( research=ResearchConfig( diff --git a/evaluations/evaluations/benchmark.py b/evaluations/evaluations/benchmark.py index c2e94cd4..064ce40b 100644 --- a/evaluations/evaluations/benchmark.py +++ b/evaluations/evaluations/benchmark.py @@ -20,7 +20,7 @@ from haiku.rag.client import HaikuRAG from haiku.rag.config import AppConfig, find_config_file, load_yaml_config from haiku.rag.config.models import ModelConfig from haiku.rag.logging import configure_cli_logging -from haiku.rag.qa import get_qa_agent +from haiku.rag.agents.qa import get_qa_agent from haiku.rag.utils import get_model load_dotenv() diff --git a/haiku_rag_slim/haiku/rag/agents/__init__.py b/haiku_rag_slim/haiku/rag/agents/__init__.py new file mode 100644 index 00000000..e4ab64c8 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/agents/__init__.py @@ -0,0 +1,48 @@ +from haiku.rag.agents.chat import ( + ChatDeps, + ChatSessionState, + CitationInfo, + QAResponse, + SearchAgent, + SearchDeps, + create_chat_agent, +) +from haiku.rag.agents.qa import QuestionAnswerAgent, get_qa_agent +from haiku.rag.agents.research import ( + EvaluationResult, + ResearchContext, + ResearchDependencies, + ResearchReport, + SearchAnswer, +) +from haiku.rag.agents.research.graph import ( + build_conversational_graph, + build_research_graph, +) +from haiku.rag.agents.research.models import Citation +from haiku.rag.agents.research.state import ResearchDeps, ResearchState + +__all__ = [ + # QA + "get_qa_agent", + "QuestionAnswerAgent", + # Research + "build_research_graph", + "build_conversational_graph", + "ResearchContext", + "ResearchDependencies", + "ResearchDeps", + "ResearchState", + "ResearchReport", + "Citation", + "SearchAnswer", + "EvaluationResult", + # Chat + "create_chat_agent", + "SearchAgent", + "ChatDeps", + "ChatSessionState", + "CitationInfo", + "QAResponse", + "SearchDeps", +] diff --git a/haiku_rag_slim/haiku/rag/agents/chat/__init__.py b/haiku_rag_slim/haiku/rag/agents/chat/__init__.py new file mode 100644 index 00000000..8947cca9 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/agents/chat/__init__.py @@ -0,0 +1,23 @@ +from haiku.rag.agents.chat.agent import create_chat_agent +from haiku.rag.agents.chat.search import SearchAgent +from haiku.rag.agents.chat.state import ( + ChatDeps, + ChatSessionState, + CitationInfo, + QAResponse, + SearchDeps, + build_document_filter, + format_conversation_context, +) + +__all__ = [ + "create_chat_agent", + "SearchAgent", + "ChatDeps", + "ChatSessionState", + "CitationInfo", + "QAResponse", + "SearchDeps", + "build_document_filter", + "format_conversation_context", +] diff --git a/app/backend/agent.py b/haiku_rag_slim/haiku/rag/agents/chat/agent.py similarity index 66% rename from app/backend/agent.py rename to haiku_rag_slim/haiku/rag/agents/chat/agent.py index 9bd00891..78ee1306 100644 --- a/app/backend/agent.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/agent.py @@ -1,120 +1,24 @@ -from dataclasses import dataclass - from ag_ui.core import EventType, StateSnapshotEvent -from pydantic import BaseModel -from pydantic_ai import Agent, RunContext, ToolReturn, format_as_xml +from pydantic_ai import Agent, RunContext, ToolReturn -from haiku.rag.client import HaikuRAG +from haiku.rag.agents.chat.prompts import CHAT_SYSTEM_PROMPT +from haiku.rag.agents.chat.search import SearchAgent +from haiku.rag.agents.chat.state import ( + ChatDeps, + ChatSessionState, + CitationInfo, + QAResponse, + build_document_filter, + format_conversation_context, +) +from haiku.rag.agents.research.dependencies import ResearchContext +from haiku.rag.agents.research.graph import build_conversational_graph +from haiku.rag.agents.research.models import Citation, SearchAnswer +from haiku.rag.agents.research.state import ResearchDeps, ResearchState from haiku.rag.config.models import AppConfig -from haiku.rag.store.models import SearchResult from haiku.rag.utils import get_model -class CitationInfo(BaseModel): - """Citation info for frontend display.""" - - index: int - document_id: str - chunk_id: str - document_uri: str - document_title: str | None = None - page_numbers: list[int] = [] - headings: list[str] | None = None - content: str - - -class QAResponse(BaseModel): - """A Q&A pair from conversation history with citations.""" - - question: str - answer: str - confidence: float = 0.9 - citations: list[CitationInfo] = [] - - @property - def sources(self) -> list[str]: - """Source names for display.""" - return list( - dict.fromkeys(c.document_title or c.document_uri for c in self.citations) - ) - - -class ChatSessionState(BaseModel): - """State shared between frontend and agent via AG-UI.""" - - session_id: str = "" - citations: list[CitationInfo] = [] - qa_history: list[QAResponse] = [] - - -def format_conversation_context(qa_history: list[QAResponse]) -> str: - """Format conversation history as XML for inclusion in prompts.""" - if not qa_history: - return "" - - context_data = { - "previous_qa": [ - { - "question": qa.question, - "answer": qa.answer, - "sources": qa.sources, - } - for qa in qa_history - ], - } - return format_as_xml(context_data, root_tag="conversation_context") - - -@dataclass -class ChatDeps: - """Dependencies for chat agent.""" - - client: HaikuRAG - config: AppConfig - search_results: list[SearchResult] | None = None - session_state: ChatSessionState | None = None - - -def build_document_filter(document_name: str) -> str: - """Build SQL filter for document name matching.""" - escaped = document_name.replace("'", "''") - no_spaces = escaped.replace(" ", "") - return ( - f"LOWER(uri) LIKE LOWER('%{escaped}%') OR LOWER(title) LIKE LOWER('%{escaped}%') " - f"OR LOWER(uri) LIKE LOWER('%{no_spaces}%') OR LOWER(title) LIKE LOWER('%{no_spaces}%')" - ) - - -CHAT_SYSTEM_PROMPT = """You are a helpful research assistant powered by haiku.rag, a knowledge base system. - -You have access to a knowledge base of documents. Use your tools to search and answer questions. - -CRITICAL RULES: -1. For greetings or casual chat: respond directly WITHOUT using any tools -2. For questions: Use the "ask" tool EXACTLY ONCE - it handles query expansion internally -3. For searches: Use the "search" tool EXACTLY ONCE - it handles multi-query expansion internally -4. NEVER call the same tool multiple times for a single user message -5. NEVER make up information - always use tools to get facts from the knowledge base - -How to decide which tool to use: -- "get_document" - Use when the user references a SPECIFIC document by name, title, or URI (e.g., "summarize document X", "get the paper about Y", "fetch 2412.00566"). Retrieves the full document content. -- "ask" - Use for general questions about topics in the knowledge base when no specific document is named. It searches across all documents and returns answers with citations. -- "search" - Use when the user explicitly asks to search/find/explore documents. Call it ONCE. After calling search, copy the ENTIRE tool response to your output INCLUDING the content snippets. Do NOT shorten, summarize, or omit any part of the results. - -IMPORTANT - When user mentions a document in search/ask: -- If user says "search in ", "find in ", "answer from ", or " in ": - - Extract the TOPIC as `query`/`question` - - Extract the DOCUMENT NAME as `document_name` -- Examples for search: - - "search for latrines in TB MED 593" → query="latrines", document_name="TB MED 593" - - "find waste disposal in the army manual" → query="waste disposal", document_name="army manual" -- Examples for ask: - - "what does TB MED 593 say about latrines?" → question="what are the guidelines for latrines?", document_name="TB MED 593" - - "answer from the army manual about sanitation" → question="what are the sanitation guidelines?", document_name="army manual" - -Be friendly and conversational. When you use the "ask" tool, summarize the key findings for the user.""" - - def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]: """Create the chat agent with search and ask tools.""" model = get_model(config.qa.model, config) @@ -141,8 +45,6 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]: query: The search query (what to search for) document_name: Optional document name/title to search within (e.g., "tbmed593", "army manual") """ - from search_agent import SearchAgent - # Build context from conversation history context = None if ctx.deps.session_state and ctx.deps.session_state.qa_history: @@ -228,11 +130,6 @@ def create_chat_agent(config: AppConfig) -> Agent[ChatDeps, str]: question: The question to answer document_name: Optional document name/title to search within (e.g., "tbmed593", "army manual") """ - from haiku.rag.graph.research.dependencies import ResearchContext - from haiku.rag.graph.research.graph import build_conversational_graph - from haiku.rag.graph.research.models import Citation, SearchAnswer - from haiku.rag.graph.research.state import ResearchDeps, ResearchState - # Build filter from document_name doc_filter = build_document_filter(document_name) if document_name else None diff --git a/haiku_rag_slim/haiku/rag/agents/chat/prompts.py b/haiku_rag_slim/haiku/rag/agents/chat/prompts.py new file mode 100644 index 00000000..2d3b44f4 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/agents/chat/prompts.py @@ -0,0 +1,41 @@ +CHAT_SYSTEM_PROMPT = """You are a helpful research assistant powered by haiku.rag, a knowledge base system. + +You have access to a knowledge base of documents. Use your tools to search and answer questions. + +CRITICAL RULES: +1. For greetings or casual chat: respond directly WITHOUT using any tools +2. For questions: Use the "ask" tool EXACTLY ONCE - it handles query expansion internally +3. For searches: Use the "search" tool EXACTLY ONCE - it handles multi-query expansion internally +4. NEVER call the same tool multiple times for a single user message +5. NEVER make up information - always use tools to get facts from the knowledge base + +How to decide which tool to use: +- "get_document" - Use when the user references a SPECIFIC document by name, title, or URI (e.g., "summarize document X", "get the paper about Y", "fetch 2412.00566"). Retrieves the full document content. +- "ask" - Use for general questions about topics in the knowledge base when no specific document is named. It searches across all documents and returns answers with citations. +- "search" - Use when the user explicitly asks to search/find/explore documents. Call it ONCE. After calling search, copy the ENTIRE tool response to your output INCLUDING the content snippets. Do NOT shorten, summarize, or omit any part of the results. + +IMPORTANT - When user mentions a document in search/ask: +- If user says "search in ", "find in ", "answer from ", or " in ": + - Extract the TOPIC as `query`/`question` + - Extract the DOCUMENT NAME as `document_name` +- Examples for search: + - "search for latrines in TB MED 593" → query="latrines", document_name="TB MED 593" + - "find waste disposal in the army manual" → query="waste disposal", document_name="army manual" +- Examples for ask: + - "what does TB MED 593 say about latrines?" → question="what are the guidelines for latrines?", document_name="TB MED 593" + - "answer from the army manual about sanitation" → question="what are the sanitation guidelines?", document_name="army manual" + +Be friendly and conversational. When you use the "ask" tool, summarize the key findings for the user.""" + +SEARCH_SYSTEM_PROMPT = """You are a search query optimizer for a document knowledge base. + +Given a user's search request: +1. ALWAYS run the original query first as-is +2. Then generate 1-2 alternative queries using different keywords or phrasings +3. Keep queries SHORT (2-5 words) - use keywords, not full sentences +4. After all searches, respond with "Search complete" + +Example: User asks "latrines" → queries: "latrines", "latrine sanitation", "field toilet" +Example: User asks "waste disposal" → queries: "waste disposal", "garbage management", "refuse handling" + +Do NOT generate long verbose queries like "environmental impact of waste disposal methods" - keep it simple.""" diff --git a/app/backend/search_agent.py b/haiku_rag_slim/haiku/rag/agents/chat/search.py similarity index 73% rename from app/backend/search_agent.py rename to haiku_rag_slim/haiku/rag/agents/chat/search.py index cd52cc13..689101e2 100644 --- a/app/backend/search_agent.py +++ b/haiku_rag_slim/haiku/rag/agents/chat/search.py @@ -1,37 +1,13 @@ -from dataclasses import dataclass, field - from pydantic_ai import Agent, RunContext +from haiku.rag.agents.chat.prompts import SEARCH_SYSTEM_PROMPT +from haiku.rag.agents.chat.state import SearchDeps from haiku.rag.client import HaikuRAG from haiku.rag.config.models import AppConfig from haiku.rag.store.models import SearchResult from haiku.rag.utils import get_model -@dataclass -class SearchDeps: - """Dependencies for search agent.""" - - client: HaikuRAG - config: AppConfig - filter: str | None = None - search_results: list[SearchResult] = field(default_factory=list) - - -SEARCH_SYSTEM_PROMPT = """You are a search query optimizer for a document knowledge base. - -Given a user's search request: -1. ALWAYS run the original query first as-is -2. Then generate 1-2 alternative queries using different keywords or phrasings -3. Keep queries SHORT (2-5 words) - use keywords, not full sentences -4. After all searches, respond with "Search complete" - -Example: User asks "latrines" → queries: "latrines", "latrine sanitation", "field toilet" -Example: User asks "waste disposal" → queries: "waste disposal", "garbage management", "refuse handling" - -Do NOT generate long verbose queries like "environmental impact of waste disposal methods" - keep it simple.""" - - class SearchAgent: """Agent that generates multiple queries and consolidates results.""" diff --git a/haiku_rag_slim/haiku/rag/agents/chat/state.py b/haiku_rag_slim/haiku/rag/agents/chat/state.py new file mode 100644 index 00000000..d8c41414 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/agents/chat/state.py @@ -0,0 +1,93 @@ +from dataclasses import dataclass, field + +from pydantic import BaseModel +from pydantic_ai import format_as_xml + +from haiku.rag.client import HaikuRAG +from haiku.rag.config.models import AppConfig +from haiku.rag.store.models import SearchResult + + +class CitationInfo(BaseModel): + """Citation info for frontend display.""" + + index: int + document_id: str + chunk_id: str + document_uri: str + document_title: str | None = None + page_numbers: list[int] = [] + headings: list[str] | None = None + content: str + + +class QAResponse(BaseModel): + """A Q&A pair from conversation history with citations.""" + + question: str + answer: str + confidence: float = 0.9 + citations: list[CitationInfo] = [] + + @property + def sources(self) -> list[str]: + """Source names for display.""" + return list( + dict.fromkeys(c.document_title or c.document_uri for c in self.citations) + ) + + +class ChatSessionState(BaseModel): + """State shared between frontend and agent via AG-UI.""" + + session_id: str = "" + citations: list[CitationInfo] = [] + qa_history: list[QAResponse] = [] + + +def format_conversation_context(qa_history: list[QAResponse]) -> str: + """Format conversation history as XML for inclusion in prompts.""" + if not qa_history: + return "" + + context_data = { + "previous_qa": [ + { + "question": qa.question, + "answer": qa.answer, + "sources": qa.sources, + } + for qa in qa_history + ], + } + return format_as_xml(context_data, root_tag="conversation_context") + + +@dataclass +class ChatDeps: + """Dependencies for chat agent.""" + + client: HaikuRAG + config: AppConfig + search_results: list[SearchResult] | None = None + session_state: ChatSessionState | None = None + + +@dataclass +class SearchDeps: + """Dependencies for search agent.""" + + client: HaikuRAG + config: AppConfig + filter: str | None = None + search_results: list[SearchResult] = field(default_factory=list) + + +def build_document_filter(document_name: str) -> str: + """Build SQL filter for document name matching.""" + escaped = document_name.replace("'", "''") + no_spaces = escaped.replace(" ", "") + return ( + f"LOWER(uri) LIKE LOWER('%{escaped}%') OR LOWER(title) LIKE LOWER('%{escaped}%') " + f"OR LOWER(uri) LIKE LOWER('%{no_spaces}%') OR LOWER(title) LIKE LOWER('%{no_spaces}%')" + ) diff --git a/haiku_rag_slim/haiku/rag/qa/__init__.py b/haiku_rag_slim/haiku/rag/agents/qa/__init__.py similarity index 89% rename from haiku_rag_slim/haiku/rag/qa/__init__.py rename to haiku_rag_slim/haiku/rag/agents/qa/__init__.py index 77d7bed7..6af77f02 100644 --- a/haiku_rag_slim/haiku/rag/qa/__init__.py +++ b/haiku_rag_slim/haiku/rag/agents/qa/__init__.py @@ -1,7 +1,7 @@ +from haiku.rag.agents.qa.agent import QuestionAnswerAgent +from haiku.rag.agents.qa.prompts import QA_SYSTEM_PROMPT from haiku.rag.client import HaikuRAG from haiku.rag.config import AppConfig, Config -from haiku.rag.qa.agent import QuestionAnswerAgent -from haiku.rag.qa.prompts import QA_SYSTEM_PROMPT from haiku.rag.utils import build_prompt diff --git a/haiku_rag_slim/haiku/rag/qa/agent.py b/haiku_rag_slim/haiku/rag/agents/qa/agent.py similarity index 93% rename from haiku_rag_slim/haiku/rag/qa/agent.py rename to haiku_rag_slim/haiku/rag/agents/qa/agent.py index 6ce392ea..d2b5194b 100644 --- a/haiku_rag_slim/haiku/rag/qa/agent.py +++ b/haiku_rag_slim/haiku/rag/agents/qa/agent.py @@ -2,10 +2,14 @@ from pydantic import BaseModel from pydantic_ai import Agent, RunContext from pydantic_ai.output import ToolOutput +from haiku.rag.agents.qa.prompts import QA_SYSTEM_PROMPT +from haiku.rag.agents.research.models import ( + Citation, + RawSearchAnswer, + resolve_citations, +) from haiku.rag.client import HaikuRAG from haiku.rag.config.models import AppConfig, ModelConfig -from haiku.rag.graph.research.models import Citation, RawSearchAnswer, resolve_citations -from haiku.rag.qa.prompts import QA_SYSTEM_PROMPT from haiku.rag.store.models import SearchResult from haiku.rag.utils import get_model diff --git a/haiku_rag_slim/haiku/rag/qa/prompts.py b/haiku_rag_slim/haiku/rag/agents/qa/prompts.py similarity index 100% rename from haiku_rag_slim/haiku/rag/qa/prompts.py rename to haiku_rag_slim/haiku/rag/agents/qa/prompts.py diff --git a/haiku_rag_slim/haiku/rag/agents/research/__init__.py b/haiku_rag_slim/haiku/rag/agents/research/__init__.py new file mode 100644 index 00000000..41a039b2 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/agents/research/__init__.py @@ -0,0 +1,6 @@ +from haiku.rag.agents.research.dependencies import ResearchContext, ResearchDependencies +from haiku.rag.agents.research.models import ( + EvaluationResult, + ResearchReport, + SearchAnswer, +) diff --git a/haiku_rag_slim/haiku/rag/graph/research/dependencies.py b/haiku_rag_slim/haiku/rag/agents/research/dependencies.py similarity index 95% rename from haiku_rag_slim/haiku/rag/graph/research/dependencies.py rename to haiku_rag_slim/haiku/rag/agents/research/dependencies.py index a7ecbe22..12f1f154 100644 --- a/haiku_rag_slim/haiku/rag/graph/research/dependencies.py +++ b/haiku_rag_slim/haiku/rag/agents/research/dependencies.py @@ -6,7 +6,7 @@ from haiku.rag.client import HaikuRAG from haiku.rag.store.models import SearchResult if TYPE_CHECKING: - from haiku.rag.graph.research.models import SearchAnswer + from haiku.rag.agents.research.models import SearchAnswer class ResearchContext(BaseModel): diff --git a/haiku_rag_slim/haiku/rag/graph/research/graph.py b/haiku_rag_slim/haiku/rag/agents/research/graph.py similarity index 98% rename from haiku_rag_slim/haiku/rag/graph/research/graph.py rename to haiku_rag_slim/haiku/rag/agents/research/graph.py index 2d62753a..df623b2d 100644 --- a/haiku_rag_slim/haiku/rag/graph/research/graph.py +++ b/haiku_rag_slim/haiku/rag/agents/research/graph.py @@ -5,10 +5,8 @@ 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.research.dependencies import ResearchContext, ResearchDependencies -from haiku.rag.graph.research.models import ( +from haiku.rag.agents.research.dependencies import ResearchContext, ResearchDependencies +from haiku.rag.agents.research.models import ( Citation, ConversationalAnswer, EvaluationResult, @@ -17,7 +15,7 @@ from haiku.rag.graph.research.models import ( ResearchReport, SearchAnswer, ) -from haiku.rag.graph.research.prompts import ( +from haiku.rag.agents.research.prompts import ( CONVERSATIONAL_SYNTHESIS_PROMPT, DECISION_PROMPT, PLAN_PROMPT, @@ -25,7 +23,9 @@ from haiku.rag.graph.research.prompts import ( SEARCH_PROMPT, SYNTHESIS_PROMPT, ) -from haiku.rag.graph.research.state import ResearchDeps, ResearchState +from haiku.rag.agents.research.state import ResearchDeps, ResearchState +from haiku.rag.config import Config +from haiku.rag.config.models import AppConfig from haiku.rag.utils import build_prompt, get_model diff --git a/haiku_rag_slim/haiku/rag/graph/research/models.py b/haiku_rag_slim/haiku/rag/agents/research/models.py similarity index 100% rename from haiku_rag_slim/haiku/rag/graph/research/models.py rename to haiku_rag_slim/haiku/rag/agents/research/models.py diff --git a/haiku_rag_slim/haiku/rag/graph/research/prompts.py b/haiku_rag_slim/haiku/rag/agents/research/prompts.py similarity index 100% rename from haiku_rag_slim/haiku/rag/graph/research/prompts.py rename to haiku_rag_slim/haiku/rag/agents/research/prompts.py diff --git a/haiku_rag_slim/haiku/rag/graph/research/state.py b/haiku_rag_slim/haiku/rag/agents/research/state.py similarity index 94% rename from haiku_rag_slim/haiku/rag/graph/research/state.py rename to haiku_rag_slim/haiku/rag/agents/research/state.py index a28533d6..f305c3a2 100644 --- a/haiku_rag_slim/haiku/rag/graph/research/state.py +++ b/haiku_rag_slim/haiku/rag/agents/research/state.py @@ -4,9 +4,9 @@ from typing import TYPE_CHECKING from pydantic import BaseModel, Field +from haiku.rag.agents.research.dependencies import ResearchContext +from haiku.rag.agents.research.models import EvaluationResult from haiku.rag.client import HaikuRAG -from haiku.rag.graph.research.dependencies import ResearchContext -from haiku.rag.graph.research.models import EvaluationResult if TYPE_CHECKING: from haiku.rag.config.models import AppConfig diff --git a/haiku_rag_slim/haiku/rag/app.py b/haiku_rag_slim/haiku/rag/app.py index 60cfb41e..e570028a 100644 --- a/haiku_rag_slim/haiku/rag/app.py +++ b/haiku_rag_slim/haiku/rag/app.py @@ -18,11 +18,11 @@ from rich.progress import ( TransferSpeedColumn, ) +from haiku.rag.agents.research.dependencies import ResearchContext +from haiku.rag.agents.research.graph import build_research_graph +from haiku.rag.agents.research.state import ResearchDeps, ResearchState from haiku.rag.client import HaikuRAG, RebuildMode from haiku.rag.config import AppConfig, Config -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.store.models.document import Document diff --git a/haiku_rag_slim/haiku/rag/client.py b/haiku_rag_slim/haiku/rag/client.py index fe8fd11f..27aa4482 100644 --- a/haiku_rag_slim/haiku/rag/client.py +++ b/haiku_rag_slim/haiku/rag/client.py @@ -28,7 +28,7 @@ from haiku.rag.store.repositories.settings import SettingsRepository if TYPE_CHECKING: from docling_core.types.doc.document import DoclingDocument - from haiku.rag.graph.research.models import Citation + from haiku.rag.agents.research.models import Citation logger = logging.getLogger(__name__) @@ -1275,7 +1275,7 @@ class HaikuRAG: Returns: Tuple of (answer text, list of resolved citations). """ - from haiku.rag.qa import get_qa_agent + from haiku.rag.agents.qa import get_qa_agent qa_agent = get_qa_agent(self, config=self._config, system_prompt=system_prompt) return await qa_agent.answer(question, filter=filter) diff --git a/haiku_rag_slim/haiku/rag/graph/__init__.py b/haiku_rag_slim/haiku/rag/graph/__init__.py deleted file mode 100644 index 3b9000ef..00000000 --- a/haiku_rag_slim/haiku/rag/graph/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -from haiku.rag.graph.research.graph import build_research_graph - -__all__ = [ - "build_research_graph", -] diff --git a/haiku_rag_slim/haiku/rag/graph/research/__init__.py b/haiku_rag_slim/haiku/rag/graph/research/__init__.py deleted file mode 100644 index 36f46a92..00000000 --- a/haiku_rag_slim/haiku/rag/graph/research/__init__.py +++ /dev/null @@ -1,6 +0,0 @@ -from haiku.rag.graph.research.dependencies import ResearchContext, ResearchDependencies -from haiku.rag.graph.research.models import ( - EvaluationResult, - ResearchReport, - SearchAnswer, -) diff --git a/haiku_rag_slim/haiku/rag/mcp.py b/haiku_rag_slim/haiku/rag/mcp.py index f002c384..4cbc29ab 100644 --- a/haiku_rag_slim/haiku/rag/mcp.py +++ b/haiku_rag_slim/haiku/rag/mcp.py @@ -4,9 +4,9 @@ from typing import Any from fastmcp import FastMCP from pydantic import BaseModel +from haiku.rag.agents.research.models import ResearchReport from haiku.rag.client import HaikuRAG from haiku.rag.config import AppConfig, Config -from haiku.rag.graph.research.models import ResearchReport from haiku.rag.store.models import SearchResult from haiku.rag.utils import format_citations @@ -186,9 +186,9 @@ def create_mcp_server( try: async with HaikuRAG(db_path, config=config, read_only=read_only) as rag: if deep: - 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 ( + from haiku.rag.agents.research.dependencies import ResearchContext + from haiku.rag.agents.research.graph import build_research_graph + from haiku.rag.agents.research.state import ( ResearchDeps, ResearchState, ) @@ -230,9 +230,9 @@ def create_mcp_server( A research report with findings, or None if an error occurred. """ try: - 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.agents.research.dependencies import ResearchContext + from haiku.rag.agents.research.graph import build_research_graph + from haiku.rag.agents.research.state import ResearchDeps, ResearchState async with HaikuRAG(db_path, config=config, read_only=read_only) as rag: graph = build_research_graph(config=config) diff --git a/haiku_rag_slim/haiku/rag/utils.py b/haiku_rag_slim/haiku/rag/utils.py index f6437e43..3d24d50d 100644 --- a/haiku_rag_slim/haiku/rag/utils.py +++ b/haiku_rag_slim/haiku/rag/utils.py @@ -10,8 +10,8 @@ from packaging.version import Version, parse if TYPE_CHECKING: from rich.console import RenderableType + from haiku.rag.agents.research.models import Citation from haiku.rag.config.models import AppConfig, ModelConfig - from haiku.rag.graph.research.models import Citation def parse_datetime(s: str) -> datetime: diff --git a/tests/agents/chat/test_chat_agent.py b/tests/agents/chat/test_chat_agent.py new file mode 100644 index 00000000..33cacae7 --- /dev/null +++ b/tests/agents/chat/test_chat_agent.py @@ -0,0 +1,105 @@ +from haiku.rag.agents.chat import ( + ChatDeps, + ChatSessionState, + CitationInfo, + QAResponse, + SearchAgent, + create_chat_agent, +) +from haiku.rag.client import HaikuRAG +from haiku.rag.config import Config + + +def test_create_chat_agent(): + """Test that create_chat_agent returns a properly configured agent.""" + agent = create_chat_agent(Config) + assert agent is not None + assert agent.name == "chat_agent" or agent.name is None + + +def test_chat_deps_initialization(temp_db_path): + """Test ChatDeps can be initialized with required fields.""" + client = HaikuRAG(temp_db_path, create=True) + deps = ChatDeps(client=client, config=Config) + + assert deps.client is client + assert deps.config is Config + assert deps.search_results is None + assert deps.session_state is None + + client.close() + + +def test_chat_session_state(): + """Test ChatSessionState model.""" + state = ChatSessionState(session_id="test-session") + assert state.session_id == "test-session" + assert state.citations == [] + assert state.qa_history == [] + + +def test_citation_info(): + """Test CitationInfo model.""" + citation = CitationInfo( + index=1, + document_id="doc-123", + chunk_id="chunk-456", + document_uri="test.md", + document_title="Test Document", + page_numbers=[1, 2], + headings=["Section 1"], + content="Test content", + ) + assert citation.index == 1 + assert citation.document_id == "doc-123" + assert citation.chunk_id == "chunk-456" + assert citation.content == "Test content" + + +def test_qa_response(): + """Test QAResponse model.""" + citation = CitationInfo( + index=1, + document_id="doc-123", + chunk_id="chunk-456", + document_uri="test.md", + document_title="Test Document", + content="Test content", + ) + qa = QAResponse( + question="What is this?", + answer="This is a test", + confidence=0.95, + citations=[citation], + ) + assert qa.question == "What is this?" + assert qa.answer == "This is a test" + assert qa.confidence == 0.95 + assert len(qa.citations) == 1 + assert qa.sources == ["Test Document"] + + +def test_qa_response_sources_with_uri_fallback(): + """Test QAResponse.sources falls back to URI when title is None.""" + citation = CitationInfo( + index=1, + document_id="doc-123", + chunk_id="chunk-456", + document_uri="test.md", + document_title=None, + content="Test content", + ) + qa = QAResponse( + question="What is this?", + answer="This is a test", + citations=[citation], + ) + assert qa.sources == ["test.md"] + + +def test_search_agent_initialization(temp_db_path): + """Test SearchAgent can be initialized.""" + client = HaikuRAG(temp_db_path, create=True) + search_agent = SearchAgent(client, Config) + assert search_agent is not None + client.close() diff --git a/tests/test_qa.py b/tests/agents/qa/test_qa.py similarity index 95% rename from tests/test_qa.py rename to tests/agents/qa/test_qa.py index f05204bd..32077d6f 100644 --- a/tests/test_qa.py +++ b/tests/agents/qa/test_qa.py @@ -5,16 +5,16 @@ import pytest from datasets import Dataset from evaluations.evaluators import LLMJudge +from haiku.rag.agents.qa.agent import QuestionAnswerAgent from haiku.rag.client import HaikuRAG from haiku.rag.config.models import ModelConfig -from haiku.rag.qa.agent import QuestionAnswerAgent HAS_ANTHROPIC = importlib.util.find_spec("anthropic") is not None @pytest.fixture(scope="module") def vcr_cassette_dir(): - return str(Path(__file__).parent / "cassettes" / "test_qa") + return str(Path(__file__).parent.parent.parent / "cassettes" / "test_qa") @pytest.mark.vcr() diff --git a/tests/graph/cassettes/test_research_graph/test_graph_end_to_end.yaml b/tests/agents/research/cassettes/test_research_graph/test_graph_end_to_end.yaml similarity index 100% rename from tests/graph/cassettes/test_research_graph/test_graph_end_to_end.yaml rename to tests/agents/research/cassettes/test_research_graph/test_graph_end_to_end.yaml diff --git a/tests/graph/cassettes/test_search_filter/test_research_graph_uses_search_filter.yaml b/tests/agents/research/cassettes/test_search_filter/test_research_graph_uses_search_filter.yaml similarity index 100% rename from tests/graph/cassettes/test_search_filter/test_research_graph_uses_search_filter.yaml rename to tests/agents/research/cassettes/test_search_filter/test_research_graph_uses_search_filter.yaml diff --git a/tests/graph/cassettes/test_search_filter/test_search_filter_none_searches_all.yaml b/tests/agents/research/cassettes/test_search_filter/test_search_filter_none_searches_all.yaml similarity index 100% rename from tests/graph/cassettes/test_search_filter/test_search_filter_none_searches_all.yaml rename to tests/agents/research/cassettes/test_search_filter/test_search_filter_none_searches_all.yaml diff --git a/tests/graph/cassettes/test_search_filter/test_search_filter_restricts_results.yaml b/tests/agents/research/cassettes/test_search_filter/test_search_filter_restricts_results.yaml similarity index 100% rename from tests/graph/cassettes/test_search_filter/test_search_filter_restricts_results.yaml rename to tests/agents/research/cassettes/test_search_filter/test_search_filter_restricts_results.yaml diff --git a/tests/graph/test_research_graph.py b/tests/agents/research/test_research_graph.py similarity index 76% rename from tests/graph/test_research_graph.py rename to tests/agents/research/test_research_graph.py index ca1df9d3..265cff35 100644 --- a/tests/graph/test_research_graph.py +++ b/tests/agents/research/test_research_graph.py @@ -1,10 +1,10 @@ import pytest +from haiku.rag.agents.research.dependencies import ResearchContext +from haiku.rag.agents.research.graph import build_research_graph +from haiku.rag.agents.research.models import ResearchReport +from haiku.rag.agents.research.state import ResearchDeps, ResearchState from haiku.rag.client import HaikuRAG -from haiku.rag.graph.research.dependencies import ResearchContext -from haiku.rag.graph.research.graph import build_research_graph -from haiku.rag.graph.research.models import ResearchReport -from haiku.rag.graph.research.state import ResearchDeps, ResearchState @pytest.mark.vcr() diff --git a/tests/graph/test_search_filter.py b/tests/agents/research/test_search_filter.py similarity index 94% rename from tests/graph/test_search_filter.py rename to tests/agents/research/test_search_filter.py index 046b15bd..6227da94 100644 --- a/tests/graph/test_search_filter.py +++ b/tests/agents/research/test_search_filter.py @@ -1,9 +1,9 @@ import pytest +from haiku.rag.agents.research.dependencies import ResearchContext +from haiku.rag.agents.research.graph import build_research_graph +from haiku.rag.agents.research.state import ResearchDeps, ResearchState from haiku.rag.client import HaikuRAG -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 @pytest.fixture diff --git a/tests/graph/__init__.py b/tests/graph/__init__.py deleted file mode 100644 index 8a28b35f..00000000 --- a/tests/graph/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Tests for haiku.rag.graph module.""" diff --git a/tests/test_app.py b/tests/test_app.py index 998a68b8..259e6f08 100644 --- a/tests/test_app.py +++ b/tests/test_app.py @@ -310,7 +310,7 @@ async def test_ask_without_cite(app: HaikuRAGApp, monkeypatch): @pytest.mark.asyncio async def test_ask_with_cite(app: HaikuRAGApp, monkeypatch): """Test asking a question with citations.""" - from haiku.rag.graph.research.models import Citation + from haiku.rag.agents.research.models import Citation mock_answer = "Test answer with citations" mock_citations = [ @@ -342,7 +342,7 @@ async def test_ask_with_cite(app: HaikuRAGApp, monkeypatch): async def test_ask_with_deep(app: HaikuRAGApp, monkeypatch): """Test asking a question with deep mode uses research graph.""" import haiku.rag.app as app_module - from haiku.rag.graph.research.models import ResearchReport + from haiku.rag.agents.research.models import ResearchReport mock_output = ResearchReport( title="Test", @@ -380,7 +380,7 @@ async def test_ask_with_deep(app: HaikuRAGApp, monkeypatch): async def test_ask_with_deep_and_cite(app: HaikuRAGApp, monkeypatch): """Test asking a question with deep mode (cite is ignored for research graph).""" import haiku.rag.app as app_module - from haiku.rag.graph.research.models import ResearchReport + from haiku.rag.agents.research.models import ResearchReport mock_output = ResearchReport( title="Test", diff --git a/tests/test_mcp.py b/tests/test_mcp.py index 01694757..cae8c528 100644 --- a/tests/test_mcp.py +++ b/tests/test_mcp.py @@ -4,7 +4,7 @@ from unittest.mock import AsyncMock, patch import pytest -from haiku.rag.graph.research.models import ResearchReport +from haiku.rag.agents.research.models import ResearchReport from haiku.rag.mcp import create_mcp_server from haiku.rag.store.models.document import Document @@ -251,7 +251,7 @@ async def test_mcp_ask_question_deep(): with ( patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class, patch( - "haiku.rag.graph.research.graph.build_research_graph" + "haiku.rag.agents.research.graph.build_research_graph" ) as mock_graph_builder, ): mock_rag = AsyncMock() @@ -295,7 +295,7 @@ async def test_mcp_research_question(): with ( patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class, patch( - "haiku.rag.graph.research.graph.build_research_graph" + "haiku.rag.agents.research.graph.build_research_graph" ) as mock_graph_builder, ): mock_rag = AsyncMock()