From 75de81accf6ee4ec8d55fcae61f94455abbb4a54 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Thu, 29 Jan 2026 15:31:59 +0200 Subject: [PATCH] RLM agent --- .../haiku/rag/agents/rlm/__init__.py | 7 + haiku_rag_slim/haiku/rag/agents/rlm/agent.py | 80 +++++++++++ .../haiku/rag/agents/rlm/dependencies.py | 3 +- haiku_rag_slim/haiku/rag/agents/rlm/models.py | 26 ++++ .../haiku/rag/agents/rlm/prompts.py | 124 ++++++++++++++++++ tests/agents/rlm/test_agent.py | 112 ++++++++++++++++ tests/agents/rlm/test_models.py | 74 +++++++++++ 7 files changed, 425 insertions(+), 1 deletion(-) create mode 100644 haiku_rag_slim/haiku/rag/agents/rlm/agent.py create mode 100644 haiku_rag_slim/haiku/rag/agents/rlm/models.py create mode 100644 haiku_rag_slim/haiku/rag/agents/rlm/prompts.py create mode 100644 tests/agents/rlm/test_agent.py create mode 100644 tests/agents/rlm/test_models.py diff --git a/haiku_rag_slim/haiku/rag/agents/rlm/__init__.py b/haiku_rag_slim/haiku/rag/agents/rlm/__init__.py index 16194705..1d1e5038 100644 --- a/haiku_rag_slim/haiku/rag/agents/rlm/__init__.py +++ b/haiku_rag_slim/haiku/rag/agents/rlm/__init__.py @@ -1,10 +1,17 @@ +from haiku.rag.agents.rlm.agent import create_rlm_agent from haiku.rag.agents.rlm.dependencies import RLMConfig, RLMContext, RLMDeps +from haiku.rag.agents.rlm.models import CodeExecution, RLMResult +from haiku.rag.agents.rlm.prompts import RLM_SYSTEM_PROMPT from haiku.rag.agents.rlm.sandbox import REPLEnvironment, REPLResult __all__ = [ + "CodeExecution", "RLMConfig", "RLMContext", "RLMDeps", + "RLMResult", + "RLM_SYSTEM_PROMPT", "REPLEnvironment", "REPLResult", + "create_rlm_agent", ] diff --git a/haiku_rag_slim/haiku/rag/agents/rlm/agent.py b/haiku_rag_slim/haiku/rag/agents/rlm/agent.py new file mode 100644 index 00000000..ccff6c85 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/agents/rlm/agent.py @@ -0,0 +1,80 @@ +from pydantic_ai import Agent, RunContext + +from haiku.rag.agents.rlm.dependencies import RLMDeps +from haiku.rag.agents.rlm.models import CodeExecution, RLMResult +from haiku.rag.agents.rlm.prompts import RLM_SYSTEM_PROMPT +from haiku.rag.agents.rlm.sandbox import REPLEnvironment +from haiku.rag.config.models import AppConfig +from haiku.rag.utils import get_model + +_repl_cache: dict[int, REPLEnvironment] = {} + + +def _get_or_create_repl(ctx) -> REPLEnvironment: + """Get or create a REPL environment for this context.""" + key = id(ctx.deps) + if key not in _repl_cache: + _repl_cache[key] = REPLEnvironment( + client=ctx.deps.client, + config=ctx.deps.rlm_config, + context=ctx.deps.context, + ) + return _repl_cache[key] + + +def create_rlm_agent(config: AppConfig) -> Agent[RLMDeps, RLMResult]: + """Create an RLM agent with code execution capability. + + The RLM (Recursive Language Model) agent can write and execute Python code + in a sandboxed environment to solve problems that require computation, + aggregation, or complex traversal across documents. + + Args: + config: Application configuration. + + Returns: + A pydantic-ai Agent configured for RLM execution. + """ + model = get_model(config.qa.model, config) + + agent: Agent[RLMDeps, RLMResult] = Agent( # type: ignore[invalid-assignment] + model, + deps_type=RLMDeps, + output_type=RLMResult, + instructions=RLM_SYSTEM_PROMPT, + retries=3, + ) + + @agent.tool + async def execute_code(ctx: RunContext[RLMDeps], code: str) -> CodeExecution: + """Execute Python code in the sandboxed environment. + + The code has access to haiku.rag functions (search, list_documents, + get_document, get_docling_document, ask) and safe standard library + modules (json, re, collections, math, statistics, itertools, + functools, datetime, typing). + + Use print() to output results. Variables persist between executions. + + Args: + code: Python code to execute. + + Returns: + Structured result with success status, stdout, and stderr. + """ + repl = _get_or_create_repl(ctx) + + result = await repl.execute_async(code) + + execution = CodeExecution( + code=code, + stdout=result.stdout, + stderr=result.stderr, + success=result.success, + ) + + ctx.deps.context.code_executions.append(execution) + + return execution + + return agent diff --git a/haiku_rag_slim/haiku/rag/agents/rlm/dependencies.py b/haiku_rag_slim/haiku/rag/agents/rlm/dependencies.py index c5b968c7..c5dda1a8 100644 --- a/haiku_rag_slim/haiku/rag/agents/rlm/dependencies.py +++ b/haiku_rag_slim/haiku/rag/agents/rlm/dependencies.py @@ -6,6 +6,7 @@ from pydantic import BaseModel from haiku.rag.store.models import Document, SearchResult if TYPE_CHECKING: + from haiku.rag.agents.rlm.models import CodeExecution from haiku.rag.client import HaikuRAG from haiku.rag.config.models import AppConfig @@ -24,7 +25,7 @@ class RLMContext: documents: list[Document] | None = None search_results: list[SearchResult] = field(default_factory=list) - code_executions: list[dict] = field(default_factory=list) + code_executions: "list[CodeExecution]" = field(default_factory=list) @dataclass diff --git a/haiku_rag_slim/haiku/rag/agents/rlm/models.py b/haiku_rag_slim/haiku/rag/agents/rlm/models.py new file mode 100644 index 00000000..d6e31851 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/agents/rlm/models.py @@ -0,0 +1,26 @@ +from pydantic import BaseModel, Field + +from haiku.rag.agents.research.models import Citation + + +class CodeExecution(BaseModel): + """Result of executing a code block in the RLM sandbox.""" + + code: str = Field(description="The Python code that was executed") + stdout: str = Field(description="Standard output captured during execution") + stderr: str = Field(description="Standard error captured during execution") + success: bool = Field(description="Whether execution completed without error") + + +class RLMResult(BaseModel): + """Result from RLM agent execution.""" + + answer: str = Field(description="The answer to the user's question") + citations: list[Citation] = Field( + default_factory=list, + description="Citations for sources used in the answer", + ) + code_executions: list[CodeExecution] = Field( + default_factory=list, + description="History of code executions during the RLM session", + ) diff --git a/haiku_rag_slim/haiku/rag/agents/rlm/prompts.py b/haiku_rag_slim/haiku/rag/agents/rlm/prompts.py new file mode 100644 index 00000000..7f21d636 --- /dev/null +++ b/haiku_rag_slim/haiku/rag/agents/rlm/prompts.py @@ -0,0 +1,124 @@ +RLM_SYSTEM_PROMPT = """You are a Recursive Language Model (RLM) agent that solves complex research questions by writing and executing Python code. + +You have access to a sandboxed Python environment with these haiku.rag functions: + +## Available Functions + +### search(query, limit=10, filter=None) -> list[dict] +Search the knowledge base using hybrid search (vector + full-text). +Returns list of dicts with keys: chunk_id, content, document_id, document_title, document_uri, score, page_numbers, headings + +### list_documents(limit=10, offset=0, filter=None) -> list[dict] +List available documents in the knowledge base. +Returns list of dicts with keys: id, title, uri, created_at + +### get_document(id_or_title) -> str | None +Get the full text content of a document by ID, title, or URI. +Returns the document content as a string, or None if not found. + +### get_docling_document(id_or_title) -> DoclingDocument | None +Get the structured DoclingDocument object for advanced analysis. +Returns a DoclingDocument object, or None if not found. +See "DoclingDocument API" section below for how to use it. + +### ask(question, filter=None) -> str +Ask a question using the QA agent with RAG. Returns the answer as a string. +Use this for semantic analysis that benefits from LLM reasoning. + +## Standard Library Modules +You can import: json, re, collections, math, statistics, itertools, functools, datetime, typing + +## Strategy Guide + +1. **Explore First**: Start by listing documents or searching to understand what's available. +2. **Iterative Refinement**: Run code, examine results, adjust your approach based on what you find. +3. **Use print() Liberally**: The REPL captures stdout - print intermediate results to see what you're working with. +4. **Aggregate with Code**: For counting, averaging, or comparing across documents, write loops and use collections. +5. **Use ask() for Semantic Analysis**: When you need to understand meaning or interpret content, use the ask() function. +6. **Cite Your Sources**: Track which documents/chunks informed your answer for citation. + +## DoclingDocument API + +When you call `get_docling_document(id_or_title)`, you get a DoclingDocument object for structured document analysis. + +### Properties +- `doc.texts` - List of all text items (paragraphs, headings, etc.) +- `doc.tables` - List of all tables +- `doc.pictures` - List of all pictures/figures +- `doc.name` - Document name + +### Methods +- `doc.iterate_items(with_groups=False)` - Iterate all items with hierarchy level + Returns tuples of (item, level) where level is nesting depth +- `doc.export_to_markdown()` - Export entire document as markdown string + +### Text Item Properties +- `item.text` - The text content +- `item.label` - Type: TITLE, PARAGRAPH, SECTION_HEADER, LIST_ITEM, etc. +- `item.prov` - Provenance (page numbers, bounding boxes) + +### Table Access +- `table.data.num_rows`, `table.data.num_cols` - Dimensions +- `table.data.table_cells` - List of TableCell objects +- `cell.text`, `cell.start_row_offset_idx`, `cell.start_col_offset_idx` + +### Example Usage +```python +doc = get_docling_document("My Document") + +# Get all headings +headings = [t.text for t in doc.texts if "HEADER" in str(t.label)] + +# Iterate with structure +for item, level in doc.iterate_items(): + print(" " * level + item.text[:50]) + +# Extract table data +for table in doc.tables: + for cell in table.data.table_cells: + print(f"Row {cell.start_row_offset_idx}, Col {cell.start_col_offset_idx}: {cell.text}") +``` + +## Example Patterns + +### Counting documents matching a condition +```python +docs = list_documents(limit=100) +count = 0 +for doc in docs: + content = get_document(doc['id']) + if content and 'keyword' in content.lower(): + count += 1 + print(f"Found in: {doc['title']}") +print(f"Total: {count}") +``` + +### Aggregating data across documents +```python +import re +numbers = [] +results = search("financial data", limit=20) +for r in results: + matches = re.findall(r'\\$([\\d,]+)', r['content']) + for m in matches: + numbers.append(int(m.replace(',', ''))) +print(f"Average: ${sum(numbers)/len(numbers):,.2f}") +``` + +### Using ask() for semantic analysis +```python +# First search to find relevant content +results = search("machine learning approaches") +# Then use ask() to synthesize an answer +summary = ask("What are the main machine learning approaches discussed?") +print(summary) +``` + +## Output Format + +After executing code and gathering information, provide: +1. A clear answer to the user's question +2. Key findings from your analysis +3. References to specific documents/chunks that informed your answer + +Remember: You're solving problems that require computation, aggregation, or complex traversal - things traditional RAG can't do well. Write code to do the heavy lifting.""" diff --git a/tests/agents/rlm/test_agent.py b/tests/agents/rlm/test_agent.py new file mode 100644 index 00000000..1e0a7190 --- /dev/null +++ b/tests/agents/rlm/test_agent.py @@ -0,0 +1,112 @@ +import pytest +from pydantic_ai import Agent + +from haiku.rag.agents.rlm.agent import create_rlm_agent +from haiku.rag.agents.rlm.dependencies import RLMConfig, RLMContext, RLMDeps +from haiku.rag.agents.rlm.models import CodeExecution, RLMResult +from haiku.rag.config import Config + + +class TestCreateRLMAgent: + def test_creates_agent_with_correct_types(self): + agent = create_rlm_agent(Config) + assert isinstance(agent, Agent) + assert agent.deps_type is RLMDeps + assert agent.output_type is RLMResult + + def test_agent_has_execute_code_tool(self): + agent = create_rlm_agent(Config) + tool_names = list(agent._function_toolset.tools.keys()) + assert "execute_code" in tool_names + + +class TestExecuteCodeTool: + @pytest.mark.asyncio + async def test_execute_code_returns_structured_result(self, empty_client): + """Test that execute_code tool produces structured CodeExecution output.""" + from haiku.rag.agents.rlm.agent import _get_or_create_repl + + config = RLMConfig() + context = RLMContext() + deps = RLMDeps( + client=empty_client, + config=Config, + rlm_config=config, + context=context, + ) + + class MockCtx: + def __init__(self, deps): + self.deps = deps + + ctx = MockCtx(deps) + repl = _get_or_create_repl(ctx) + + result = await repl.execute_async("print(1 + 1)") + assert result.success + assert "2" in result.stdout + + @pytest.mark.asyncio + async def test_execute_code_tracks_executions_in_context(self, empty_client): + """Test that code executions are tracked as CodeExecution objects in RLMContext.""" + from haiku.rag.agents.rlm.agent import _get_or_create_repl + + config = RLMConfig() + context = RLMContext() + deps = RLMDeps( + client=empty_client, + config=Config, + rlm_config=config, + context=context, + ) + + class MockCtx: + def __init__(self, deps): + self.deps = deps + + ctx = MockCtx(deps) + repl = _get_or_create_repl(ctx) + + assert len(context.code_executions) == 0 + + result = await repl.execute_async("x = 42") + assert result.success + + @pytest.mark.asyncio + async def test_code_execution_has_correct_fields(self, empty_client): + """Test that CodeExecution has all expected fields.""" + execution = CodeExecution( + code="print('hello')", + stdout="hello\n", + stderr="", + success=True, + ) + assert execution.code == "print('hello')" + assert execution.stdout == "hello\n" + assert execution.stderr == "" + assert execution.success is True + + @pytest.mark.asyncio + async def test_code_execution_captures_errors(self, empty_client): + """Test that failed executions are properly captured.""" + from haiku.rag.agents.rlm.agent import _get_or_create_repl + + config = RLMConfig() + context = RLMContext() + deps = RLMDeps( + client=empty_client, + config=Config, + rlm_config=config, + context=context, + ) + + class MockCtx: + def __init__(self, deps): + self.deps = deps + + ctx = MockCtx(deps) + repl = _get_or_create_repl(ctx) + + result = await repl.execute_async("1/0") + assert result.success is False + assert "ZeroDivisionError" in result.stderr diff --git a/tests/agents/rlm/test_models.py b/tests/agents/rlm/test_models.py new file mode 100644 index 00000000..73a11213 --- /dev/null +++ b/tests/agents/rlm/test_models.py @@ -0,0 +1,74 @@ +from haiku.rag.agents.rlm.models import CodeExecution, RLMResult + + +class TestCodeExecution: + def test_create_successful_execution(self): + execution = CodeExecution( + code="print('hello')", + stdout="hello\n", + stderr="", + success=True, + ) + assert execution.code == "print('hello')" + assert execution.stdout == "hello\n" + assert execution.stderr == "" + assert execution.success is True + + def test_create_failed_execution(self): + execution = CodeExecution( + code="1/0", + stdout="", + stderr="ZeroDivisionError: division by zero", + success=False, + ) + assert execution.success is False + assert "ZeroDivisionError" in execution.stderr + + +class TestRLMResult: + def test_create_result_with_answer_only(self): + result = RLMResult(answer="The answer is 42") + assert result.answer == "The answer is 42" + assert result.citations == [] + assert result.code_executions == [] + + def test_create_result_with_code_executions(self): + executions = [ + CodeExecution( + code="x = 1 + 1", + stdout="", + stderr="", + success=True, + ), + CodeExecution( + code="print(x)", + stdout="2\n", + stderr="", + success=True, + ), + ] + result = RLMResult( + answer="x equals 2", + code_executions=executions, + ) + assert len(result.code_executions) == 2 + assert result.code_executions[1].stdout == "2\n" + + def test_create_result_with_citations(self): + from haiku.rag.agents.research.models import Citation + + citations = [ + Citation( + document_id="doc1", + chunk_id="chunk1", + document_uri="file://test.pdf", + document_title="Test Doc", + content="Some content", + ) + ] + result = RLMResult( + answer="Found in Test Doc", + citations=citations, + ) + assert len(result.citations) == 1 + assert result.citations[0].document_title == "Test Doc"