haiku.rag/haiku_rag_slim/haiku/rag/capabilities/rag.py
2026-07-24 15:26:17 +03:00

97 lines
3.1 KiB
Python

from dataclasses import dataclass
from functools import cache
from pathlib import Path
from typing import Any
from pydantic import BaseModel, Field
from pydantic_ai import RunContext
from pydantic_ai.messages import ToolReturn
from pydantic_ai.toolsets import FunctionToolset
from haiku.rag.capabilities._base import (
RAGCapabilityBase,
resolve_db_path,
)
from haiku.rag.config.models import AppConfig
from haiku.rag.store.models.chunk import SearchResult
from haiku.rag.store.models.citation import Citation
AGENT_PREAMBLE = """You are a helpful research assistant powered by haiku.rag, a knowledge base system.
CRITICAL RULES:
1. For greetings or casual chat: respond directly WITHOUT using any tools
2. NEVER make up information - always use capabilities to get facts from the knowledge base
3. When a capability returns citations, always include them in your response
"""
STATE_NAMESPACE = "rag"
_CAPABILITY_ID = "haiku-rag"
_TOOL_NAMES = frozenset({"rag_search", "rag_cite"})
_instructions_path = Path(__file__).parent / "instructions" / "rag.md"
class RAGState(BaseModel):
citation_index: dict[str, Citation] = Field(default_factory=dict)
citations: list[str] = Field(default_factory=list)
document_filter: str | None = None
searches: dict[str, list[SearchResult]] = Field(default_factory=dict)
@cache
def instructions() -> str:
return _instructions_path.read_text().strip()
@dataclass
class RAGCapability(RAGCapabilityBase[RAGState]):
"""Deferred, native Pydantic AI capability for grounded RAG queries."""
def get_toolset(self) -> FunctionToolset[Any]:
async def rag_search(
ctx: RunContext[Any], query: str, limit: int | None = None
) -> str | ToolReturn:
"""Search the knowledge base using hybrid vector and full-text search."""
return await self._with_state(self._search(query, limit))
async def rag_cite(ctx: RunContext[Any], chunk_ids: list[str]) -> Any:
"""Register exact search-result chunk IDs as citations for the answer."""
return await self._with_state(self._cite(chunk_ids))
return FunctionToolset([rag_search, rag_cite], id=_CAPABILITY_ID, max_retries=3)
def create_capability(
db_path: Path | None = None,
config: AppConfig | None = None,
*,
defer_loading: bool = True,
) -> RAGCapability:
"""Create a native Pydantic AI RAG capability."""
if config is None:
from haiku.rag.config import get_config
config = get_config()
return RAGCapability(
db_path=resolve_db_path(db_path, config),
config=config,
state_type=RAGState,
state_namespace=STATE_NAMESPACE,
instruction_text=instructions(),
model=config.qa.model,
tool_names=_TOOL_NAMES,
id=_CAPABILITY_ID,
description=(
"Search the haiku.rag knowledge base and cite evidence for grounded answers."
),
defer_loading=defer_loading,
)
__all__ = [
"AGENT_PREAMBLE",
"RAGCapability",
"RAGState",
"STATE_NAMESPACE",
"create_capability",
"instructions",
]