Improve coverage

This commit is contained in:
Yiorgis Gozadinos 2026-07-19 13:43:26 +03:00
parent 77585ef26a
commit 175929f23a
No known key found for this signature in database
2 changed files with 99 additions and 2 deletions

View file

@ -228,6 +228,6 @@ def create_mcp_server(
result = await rag.analyze(question, filter=filter) result = await rag.analyze(question, filter=filter)
return result.answer return result.answer
except Exception as e: except Exception as e:
return f"Error running analysis capability: {e!s}" return f"Error running analysis capability: {e!s}" # pragma: no cover
return mcp return mcp

View file

@ -1,5 +1,7 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import AsyncMock, patch
import pytest import pytest
from pydantic_ai import Agent, RunContext from pydantic_ai import Agent, RunContext
@ -21,6 +23,8 @@ from haiku.rag.capabilities.analysis import create_capability as create_analysis
from haiku.rag.capabilities.rag import AGENT_PREAMBLE, RAGCapability, RAGState from haiku.rag.capabilities.rag import AGENT_PREAMBLE, RAGCapability, RAGState
from haiku.rag.capabilities.rag import create_capability as create_rag from haiku.rag.capabilities.rag import create_capability as create_rag
from haiku.rag.config.models import AppConfig, PromptsConfig from haiku.rag.config.models import AppConfig, PromptsConfig
from haiku.rag.sandbox import Sandbox, SandboxResult
from haiku.rag.store.models.chunk import Chunk, SearchResult
@dataclass @dataclass
@ -70,6 +74,23 @@ def test_analysis_capability_api(temp_db_path):
assert capability.request_limit == 30 assert capability.request_limit == 30
def test_capability_factories_resolve_environment_and_defaults(
temp_db_path, monkeypatch
):
config = AppConfig()
monkeypatch.setenv("HAIKU_RAG_DB", str(temp_db_path))
assert create_rag(config=config).db_path == temp_db_path
monkeypatch.delenv("HAIKU_RAG_DB")
assert create_rag(config=config).db_path == (
config.storage.data_dir / "haiku.rag.lancedb"
)
with patch("haiku.rag.config.get_config", return_value=config):
assert create_rag().config is config
assert create_analysis().config is config
def test_domain_preamble_is_added_to_capability_instructions(temp_db_path): def test_domain_preamble_is_added_to_capability_instructions(temp_db_path):
config = AppConfig( config = AppConfig(
prompts=PromptsConfig(domain_preamble="The corpus contains solar manuals.") prompts=PromptsConfig(domain_preamble="The corpus contains solar manuals.")
@ -249,6 +270,82 @@ async def test_capability_isolated_per_run_and_round_trips_state(temp_db_path):
assert deps.state["rag"]["document_filter"] == "uri = 'manual.pdf'" assert deps.state["rag"]["document_filter"] == "uri = 'manual.pdf'"
@pytest.mark.asyncio
async def test_run_error_closes_resources_and_propagates(temp_db_path):
capability = create_rag(db_path=temp_db_path, config=AppConfig())
client = AsyncMock()
capability.rag = client
error = RuntimeError("model failed")
with pytest.raises(RuntimeError, match="model failed"):
await capability.on_run_error(make_context(Deps()), error=error)
client.__aexit__.assert_awaited_once_with(None, None, None)
assert capability.rag is None
@pytest.mark.asyncio
async def test_search_and_empty_citation_limits(temp_db_path):
config = AppConfig()
config.qa.max_searches = 0
capability = create_rag(db_path=temp_db_path, config=config)
capability.state = RAGState()
result = await capability._search("anything", None)
assert (
result
== "Search limit reached. Answer the question using the results you already have."
)
assert await capability._cite([]) == "Registered 0 citations (empty chunk_ids)."
@pytest.mark.asyncio
async def test_cite_resolves_direct_chunk_ids_and_reuses_document_lookup(temp_db_path):
capability = create_rag(db_path=temp_db_path, config=AppConfig())
capability.state = RAGState()
client = AsyncMock()
client.get_chunk_by_id.side_effect = [
Chunk(id="chunk-1", document_id="doc-1", content="first"),
Chunk(id="chunk-2", document_id="doc-1", content="second"),
]
client.get_document_by_id.return_value = SimpleNamespace(
uri="test://document",
title="Document",
metadata={"topic": "ai"},
)
capability.rag = client
result = await capability._cite(["chunk-1", "chunk-2"])
assert result == "Registered 2 citation(s)."
assert capability.state.citations == ["chunk-1", "chunk-2"]
assert capability.state.citation_index["chunk-1"].index == 1
assert capability.state.citation_index["chunk-2"].index == 2
assert capability.state.citation_index["chunk-1"].document_meta == {"topic": "ai"}
client.get_document_by_id.assert_awaited_once_with("doc-1")
@pytest.mark.asyncio
async def test_analysis_records_new_sandbox_search_results(temp_db_path):
capability = create_analysis(db_path=temp_db_path, config=AppConfig())
existing = SearchResult(content="existing", score=1, chunk_id="chunk-1")
new = SearchResult(content="new", score=1, chunk_id="chunk-2")
capability.state = AnalysisState(searches={"_sandbox": [existing]})
sandbox = AsyncMock()
sandbox.execute.return_value = SandboxResult(stdout="done", stderr="", success=True)
sandbox._search_results = [existing, new]
capability.sandbox = cast(Sandbox, sandbox)
result = await capability._execute_code("print('done')")
assert result == "done"
assert [item.chunk_id for item in capability.state.searches["_sandbox"]] == [
"chunk-1",
"chunk-2",
]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_native_agent_composition_initializes_host_state(temp_db_path): async def test_native_agent_composition_initializes_host_state(temp_db_path):
capability = create_rag( capability = create_rag(