From 1cb5e296b036a1e9b1699ec2fbe70ce02d6e959b Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Thu, 20 Aug 2026 13:13:34 +0300 Subject: [PATCH] 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. --- haiku_rag_slim/haiku/rag/client/__init__.py | 4 ++- tests/test_hooks.py | 35 ++++++++++++--------- 2 files changed, 24 insertions(+), 15 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/client/__init__.py b/haiku_rag_slim/haiku/rag/client/__init__.py index 209c55f4..93e3318f 100644 --- a/haiku_rag_slim/haiku/rag/client/__init__.py +++ b/haiku_rag_slim/haiku/rag/client/__init__.py @@ -94,7 +94,9 @@ class HaikuRAG: self._vacuum_tasks: set[asyncio.Task] = set() self._last_vacuum_at: float | None = None 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 def is_read_only(self) -> bool: diff --git a/tests/test_hooks.py b/tests/test_hooks.py index d01a3679..624b4cef 100644 --- a/tests/test_hooks.py +++ b/tests/test_hooks.py @@ -3,7 +3,7 @@ import logging import pytest 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.store.models.chunk import Chunk 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): - config = Config.model_copy(deep=True) + config = get_config().model_copy(deep=True) config.hooks = ["missing"] with pytest.raises(ValueError, match="missing"): 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()] 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"] client = HaikuRAG(temp_db_path, config=config, create=True) 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["filter"] == "uri = 'mem://hooked'" # 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): @@ -158,7 +158,7 @@ async def test_before_search_skips_non_text_queries(temp_db_path): client.chunk_repository.search = fake_search 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.supports_images = True @@ -193,7 +193,7 @@ async def test_after_search_transforms_results(temp_db_path): @pytest.mark.asyncio async def test_after_ingest_fires_on_import_batch_update(temp_db_path): 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: client._hooks = [spy] @@ -245,7 +245,7 @@ async def test_after_ingest_fires_on_import_batch_update(temp_db_path): @pytest.mark.asyncio async def test_metadata_only_update_does_not_fire_after_ingest(temp_db_path): 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: client._hooks = [spy] @@ -266,7 +266,7 @@ async def test_metadata_only_update_does_not_fire_after_ingest(temp_db_path): @pytest.mark.asyncio async def test_after_delete_fires_for_cascade(temp_db_path): 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: client._hooks = [spy] @@ -328,7 +328,7 @@ def test_format_for_agent_without_annotations_has_no_notes(): @pytest.mark.asyncio 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.document_item import DocumentItem @@ -362,9 +362,16 @@ async def test_annotations_survive_context_expansion(temp_db_path): annotations=["RCV: receive", "shared note"], ) - expanded = await expand_with_items( - client.document_item_repository, "doc-1", [r1, r2], 5000 - ) + repo = client.document_item_repository + 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 expanded[0].annotations == [ @@ -388,7 +395,7 @@ class ThrowingHook(Hook): @pytest.mark.asyncio async def test_after_ingest_hook_failure_is_logged_not_raised(temp_db_path, caplog): 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: 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 async def test_after_delete_hook_failure_is_logged_not_raised(temp_db_path, caplog): 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: doc = await client.import_document(