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:
parent
e10762854c
commit
1cb5e296b0
2 changed files with 24 additions and 15 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Reference in a new issue