Hooks build from the resolved client config

HaikuRAG resolves `config=None` to the global config for everything
else; the hook list read the raw argument, so a client constructed
without an explicit config never loaded its configured hooks.

Tests follow the config and context-expansion APIs: `get_config()`
replaces the removed `Config` instance, and the expansion test does
its own `resolve_refs_grouped` / `get_items_in_ranges` fetch now that
`expand_with_items` takes the items rather than fetching them.
This commit is contained in:
Yiorgis Gozadinos 2026-08-20 13:13:34 +03:00
parent e10762854c
commit 1cb5e296b0
No known key found for this signature in database
2 changed files with 24 additions and 15 deletions

View file

@ -94,7 +94,9 @@ class HaikuRAG:
self._vacuum_tasks: set[asyncio.Task] = set() self._vacuum_tasks: set[asyncio.Task] = set()
self._last_vacuum_at: float | None = None self._last_vacuum_at: float | None = None
self._vacuum_dirty = False self._vacuum_dirty = False
self._hooks = build_hooks(config.hooks, load_hooks()) if config.hooks else [] self._hooks = (
build_hooks(self._config.hooks, load_hooks()) if self._config.hooks else []
)
@property @property
def is_read_only(self) -> bool: def is_read_only(self) -> bool:

View file

@ -3,7 +3,7 @@ import logging
import pytest import pytest
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config from haiku.rag.config import get_config
from haiku.rag.hooks import ENTRY_POINT_GROUP, Hook, build_hooks from haiku.rag.hooks import ENTRY_POINT_GROUP, Hook, build_hooks
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
from tests.test_client import _docling_doc, _import from tests.test_client import _docling_doc, _import
@ -72,7 +72,7 @@ def test_build_hooks_loads_lazily_in_configured_order():
def test_client_init_unknown_hook_raises(temp_db_path): def test_client_init_unknown_hook_raises(temp_db_path):
config = Config.model_copy(deep=True) config = get_config().model_copy(deep=True)
config.hooks = ["missing"] config.hooks = ["missing"]
with pytest.raises(ValueError, match="missing"): with pytest.raises(ValueError, match="missing"):
HaikuRAG(temp_db_path, config=config, create=True) HaikuRAG(temp_db_path, config=config, create=True)
@ -90,7 +90,7 @@ def test_client_builds_hooks_from_entry_points(temp_db_path, monkeypatch):
return [_NamedEntryPoint()] return [_NamedEntryPoint()]
monkeypatch.setattr("haiku.rag.hooks.entry_points", fake_entry_points) monkeypatch.setattr("haiku.rag.hooks.entry_points", fake_entry_points)
config = Config.model_copy(deep=True) config = get_config().model_copy(deep=True)
config.hooks = ["recording"] config.hooks = ["recording"]
client = HaikuRAG(temp_db_path, config=config, create=True) client = HaikuRAG(temp_db_path, config=config, create=True)
assert len(client._hooks) == 1 assert len(client._hooks) == 1
@ -131,7 +131,7 @@ async def test_before_search_hooks_chain_in_order(temp_db_path):
assert captured["query"] == "alpha one two" assert captured["query"] == "alpha one two"
assert captured["filter"] == "uri = 'mem://hooked'" assert captured["filter"] == "uri = 'mem://hooked'"
# The request carries the resolved search parameters. # The request carries the resolved search parameters.
assert spy.requests == [("alpha one two", "hybrid", Config.search.limit)] assert spy.requests == [("alpha one two", "hybrid", get_config().search.limit)]
class SpyBeforeSearchHook(Hook): class SpyBeforeSearchHook(Hook):
@ -158,7 +158,7 @@ async def test_before_search_skips_non_text_queries(temp_db_path):
client.chunk_repository.search = fake_search client.chunk_repository.search = fake_search
async def fake_embed_image(image): async def fake_embed_image(image):
return [0.1] * Config.embeddings.model.vector_dim return [0.1] * get_config().embeddings.model.vector_dim
client.store.embedder.embed_image = fake_embed_image client.store.embedder.embed_image = fake_embed_image
client.store.embedder.supports_images = True client.store.embedder.supports_images = True
@ -193,7 +193,7 @@ async def test_after_search_transforms_results(temp_db_path):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_after_ingest_fires_on_import_batch_update(temp_db_path): async def test_after_ingest_fires_on_import_batch_update(temp_db_path):
spy = RecordingHook() spy = RecordingHook()
dim = Config.embeddings.model.vector_dim dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(temp_db_path, create=True) as client: async with HaikuRAG(temp_db_path, create=True) as client:
client._hooks = [spy] client._hooks = [spy]
@ -245,7 +245,7 @@ async def test_after_ingest_fires_on_import_batch_update(temp_db_path):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_metadata_only_update_does_not_fire_after_ingest(temp_db_path): async def test_metadata_only_update_does_not_fire_after_ingest(temp_db_path):
spy = RecordingHook() spy = RecordingHook()
dim = Config.embeddings.model.vector_dim dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(temp_db_path, create=True) as client: async with HaikuRAG(temp_db_path, create=True) as client:
client._hooks = [spy] client._hooks = [spy]
@ -266,7 +266,7 @@ async def test_metadata_only_update_does_not_fire_after_ingest(temp_db_path):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_after_delete_fires_for_cascade(temp_db_path): async def test_after_delete_fires_for_cascade(temp_db_path):
spy = RecordingHook() spy = RecordingHook()
dim = Config.embeddings.model.vector_dim dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(temp_db_path, create=True) as client: async with HaikuRAG(temp_db_path, create=True) as client:
client._hooks = [spy] client._hooks = [spy]
@ -328,7 +328,7 @@ def test_format_for_agent_without_annotations_has_no_notes():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_annotations_survive_context_expansion(temp_db_path): async def test_annotations_survive_context_expansion(temp_db_path):
from haiku.rag.context import expand_with_items from haiku.rag.context import expand_with_items, window_for
from haiku.rag.store.models.chunk import SearchResult from haiku.rag.store.models.chunk import SearchResult
from haiku.rag.store.models.document_item import DocumentItem from haiku.rag.store.models.document_item import DocumentItem
@ -362,9 +362,16 @@ async def test_annotations_survive_context_expansion(temp_db_path):
annotations=["RCV: receive", "shared note"], annotations=["RCV: receive", "shared note"],
) )
expanded = await expand_with_items( repo = client.document_item_repository
client.document_item_repository, "doc-1", [r1, r2], 5000 positions = (
) await repo.resolve_refs_grouped(
{"doc-1": [ref for r in (r1, r2) for ref in r.doc_item_refs]}
)
)["doc-1"]
window_items = (
await repo.get_items_in_ranges({"doc-1": window_for(positions)})
)["doc-1"]
expanded = expand_with_items([r1, r2], 5000, positions, window_items)
assert len(expanded) == 1 assert len(expanded) == 1
assert expanded[0].annotations == [ assert expanded[0].annotations == [
@ -388,7 +395,7 @@ class ThrowingHook(Hook):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_after_ingest_hook_failure_is_logged_not_raised(temp_db_path, caplog): async def test_after_ingest_hook_failure_is_logged_not_raised(temp_db_path, caplog):
spy = RecordingHook() spy = RecordingHook()
dim = Config.embeddings.model.vector_dim dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(temp_db_path, create=True) as client: async with HaikuRAG(temp_db_path, create=True) as client:
client._hooks = [ThrowingHook(), spy] client._hooks = [ThrowingHook(), spy]
@ -416,7 +423,7 @@ async def test_after_ingest_hook_failure_is_logged_not_raised(temp_db_path, capl
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_after_delete_hook_failure_is_logged_not_raised(temp_db_path, caplog): async def test_after_delete_hook_failure_is_logged_not_raised(temp_db_path, caplog):
spy = RecordingHook() spy = RecordingHook()
dim = Config.embeddings.model.vector_dim dim = get_config().embeddings.model.vector_dim
async with HaikuRAG(temp_db_path, create=True) as client: async with HaikuRAG(temp_db_path, create=True) as client:
doc = await client.import_document( doc = await client.import_document(