668 lines
28 KiB
Python
668 lines
28 KiB
Python
"""Phase E route integration: deterministic runtimes, isolated DBs, no network."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
from dataclasses import dataclass, field
|
|
from types import SimpleNamespace
|
|
|
|
import pytest
|
|
|
|
from app import repository
|
|
from app.config import get_settings
|
|
from app.contracts import IndexRebuildRequest, SearchMode, SearchRequest
|
|
from app.database.db import connect, transaction
|
|
from app.retrieval import routed_vectors
|
|
from app.retrieval.embedding import HashEmbeddingProvider
|
|
from app.retrieval.engine import RetrievalEngine, engine
|
|
from app.retrieval.reranker import LexicalReranker
|
|
from app.retrieval.vectorstore import SqliteVecStore, VectorHit
|
|
from app.services import index_service, note_service
|
|
|
|
|
|
@dataclass
|
|
class FakeRuntime:
|
|
model_id: str = "space-a"
|
|
dimensions: int = 3 # Deliberately differs from sqlite-vec's fixed 128.
|
|
source: str = "api"
|
|
error: BaseException | None = None
|
|
calls: list[list[str]] = field(default_factory=list)
|
|
result_override: object | None = None
|
|
|
|
async def embed(self, texts):
|
|
self.calls.append(list(texts))
|
|
if self.error is not None:
|
|
raise self.error
|
|
if self.result_override is not None:
|
|
return self.result_override
|
|
vectors = []
|
|
for text in texts:
|
|
# The API associates "apple" with banana; hash retrieval picks apple.
|
|
first = text == "apple orchard"
|
|
if self.model_id == "space-b":
|
|
first = not first
|
|
vectors.append(([1.0, 0.0] if first else [0.0, 1.0]) + [0.0] * (self.dimensions - 2))
|
|
return SimpleNamespace(
|
|
vectors=vectors, source=self.source, model_id=self.model_id,
|
|
dimensions=self.dimensions, fallback_reason=None,
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def runtime(monkeypatch):
|
|
runtime = FakeRuntime()
|
|
monkeypatch.setattr(routed_vectors, "get_model_routing", lambda: runtime)
|
|
return runtime
|
|
|
|
|
|
async def seed():
|
|
apple = await note_service.create_note(
|
|
title="Apple", markdown="apple orchard", folder=None, tags=[],
|
|
)
|
|
banana = await note_service.create_note(
|
|
title="Banana", markdown="banana grove", folder=None, tags=[],
|
|
)
|
|
return apple, banana
|
|
|
|
|
|
def test_native_spaces_isolate_dimensions_and_reuse_without_json_scan(runtime, monkeypatch):
|
|
from app.retrieval import space_index
|
|
async def scenario():
|
|
apple, banana = await seed()
|
|
ids = [b.block_id for note in (apple, banana) for b in note.blocks]
|
|
conn = connect()
|
|
try:
|
|
with transaction(conn):
|
|
routed_vectors.store_remote(conn, ids, routed_vectors.RemoteEmbeddings('space-a', 4, [[1., 0., 0., 0.]] * len(ids)))
|
|
assert conn.execute('SELECT COUNT(DISTINCT dimensions) FROM routed_block_vectors').fetchone()[0] == 2
|
|
finally:
|
|
conn.close()
|
|
# A new connection uses the persistent native index, without reading vector JSON.
|
|
def forbidden(*args, **kwargs):
|
|
raise AssertionError('query decoded stored JSON')
|
|
monkeypatch.setattr(space_index.json, 'loads', forbidden)
|
|
hits = await routed_vectors.search_remote('apple orchard', top_k=2, strict=True)
|
|
assert len(hits) == 2
|
|
assert hits[0].id == apple.blocks[0].block_id
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_legacy_vectors_migrate_without_document_embedding(runtime):
|
|
from app.retrieval import space_index
|
|
async def scenario():
|
|
apple, banana = await seed()
|
|
conn = connect()
|
|
table = space_index.table_name('space-a', 3)
|
|
try:
|
|
with transaction(conn):
|
|
conn.execute(f'DROP TRIGGER {table}_delete')
|
|
conn.execute(f'DROP TRIGGER {table}_update')
|
|
conn.execute(f'DROP TABLE {table}')
|
|
conn.execute('ALTER TABLE routed_block_vectors RENAME TO saved_vectors')
|
|
conn.execute('CREATE TABLE routed_block_vectors(space_id TEXT,block_id TEXT REFERENCES blocks(block_id) ON DELETE CASCADE,dimensions INTEGER,vector TEXT,PRIMARY KEY(space_id,block_id))')
|
|
conn.execute('INSERT INTO routed_block_vectors SELECT * FROM saved_vectors')
|
|
conn.execute('DROP TABLE saved_vectors')
|
|
finally:
|
|
conn.close()
|
|
runtime.calls.clear()
|
|
hits = await routed_vectors.search_remote('apple orchard', top_k=2, strict=True)
|
|
assert len(hits) == 2
|
|
assert runtime.calls == [['apple orchard']]
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.mark.parametrize("outcome", ["api", "api_failure", "missing_space"])
|
|
def test_benchmark_reports_actual_embedding_and_fallback(runtime, outcome):
|
|
from app.benchmarks import service
|
|
from app.contracts import RAGRunRequest
|
|
|
|
async def scenario():
|
|
apple, banana = await seed()
|
|
if outcome == "api_failure":
|
|
runtime.result_override = SimpleNamespace(source="local", fallback_reason="PROVIDER_TIMEOUT")
|
|
elif outcome == "missing_space":
|
|
runtime.model_id = "space-without-index"
|
|
directory = get_settings().benchmark_datasets_path
|
|
directory.mkdir(parents=True, exist_ok=True)
|
|
(directory / "routing.json").write_text(json.dumps({
|
|
"dataset_id": "routing", "kind": "rag", "version": "1",
|
|
"cases": [{"case_id": "query", "query": "apple", "expected_note_ids": [banana.note_id]}],
|
|
}), encoding="utf-8")
|
|
run = await service.create_rag_run(RAGRunRequest(
|
|
dataset_id="routing", modes=[SearchMode.fts, SearchMode.vector],
|
|
))
|
|
await service.wait_for_run(run.run_id)
|
|
report = service.get_report(run.run_id)
|
|
assert report.config_snapshot["embedding"]["policy"] == "per_case"
|
|
fts, vector = report.cases
|
|
assert fts.embedding == {"source": "not_used"}
|
|
if outcome == "api":
|
|
assert vector.embedding["source"] == "api"
|
|
assert vector.embedding["model_id"] == "space-a"
|
|
assert vector.embedding["dimensions"] == 3
|
|
assert vector.retrieved_note_ids[0] == banana.note_id
|
|
else:
|
|
assert vector.embedding["source"] == "local"
|
|
assert vector.embedding["model_id"] == "hash-v1"
|
|
assert vector.embedding["dimensions"] == 128
|
|
assert vector.retrieved_note_ids[0] == apple.note_id
|
|
if outcome == "api_failure":
|
|
assert vector.embedding["fallback_reason"] == "PROVIDER_TIMEOUT"
|
|
if outcome == "missing_space":
|
|
assert vector.embedding["fallback_reason"] == "REMOTE_INDEX_UNAVAILABLE"
|
|
assert vector.embedding["attempted_space"]["model_id"] == "space-without-index"
|
|
events = service.get_events(run.run_id)
|
|
case_events = [e for e in events if e.event.value == "CaseCompleted"]
|
|
assert case_events[-1].data["embedding"] == vector.embedding
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_embedding_observations_are_isolated_between_concurrent_searches(runtime, monkeypatch):
|
|
from app.retrieval.provenance import capture_embedding
|
|
|
|
async def scenario():
|
|
await seed()
|
|
original = runtime.embed
|
|
|
|
async def embed(texts):
|
|
await asyncio.sleep(0)
|
|
if texts == ["offline"]:
|
|
raise RuntimeError("private upstream details")
|
|
return await original(texts)
|
|
|
|
monkeypatch.setattr(runtime, "embed", embed)
|
|
|
|
async def query(text):
|
|
with capture_embedding() as observation:
|
|
await engine.search(SearchRequest(query=text, mode=SearchMode.vector))
|
|
return observation
|
|
|
|
remote, local, another = await asyncio.gather(query("apple"), query("offline"), query("apple"))
|
|
assert remote["source"] == another["source"] == "api"
|
|
assert local["source"] == "local"
|
|
assert local["fallback_reason"] == "REMOTE_EMBEDDING_UNAVAILABLE"
|
|
assert "fallback_reason" not in remote or remote["fallback_reason"] is None
|
|
assert "private upstream" not in json.dumps(local)
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.mark.parametrize("failure", ["cancel", "write"])
|
|
def test_rebuild_failure_preserves_concurrent_configuration_and_all_indexes(runtime, monkeypatch, failure):
|
|
from app.container import container
|
|
from app.contracts import ModelRoutingConfig, ProviderConfig, ProviderType
|
|
from app.services import task_service
|
|
|
|
async def scenario():
|
|
apple, _ = await seed()
|
|
task = task_service.create_task(title="before", note_id=apple.note_id)
|
|
before = {table: [tuple(row) for row in rows(f"SELECT * FROM {table}")]
|
|
for table in ("notes", "blocks", "blocks_fts", "vec_blocks", "index_meta", "routed_block_vectors")}
|
|
container.model_routing.update(ModelRoutingConfig())
|
|
entered, release = asyncio.Event(), asyncio.Event()
|
|
original_embed = runtime.embed
|
|
|
|
async def pending_embed(texts):
|
|
entered.set()
|
|
await release.wait()
|
|
return await original_embed(texts)
|
|
|
|
monkeypatch.setattr(runtime, "embed", pending_embed)
|
|
original_index = index_service.index_note
|
|
writes = 0
|
|
|
|
async def fail_write(parsed, **kwargs):
|
|
nonlocal writes
|
|
await original_index(parsed, **kwargs)
|
|
writes += 1
|
|
if writes == 2:
|
|
raise RuntimeError("injected write failure")
|
|
|
|
if failure == "write":
|
|
monkeypatch.setattr(index_service, "index_note", fail_write)
|
|
rebuilding = asyncio.create_task(index_service.rebuild(IndexRebuildRequest()))
|
|
await asyncio.wait_for(entered.wait(), timeout=5)
|
|
saved = container.model_routing.update(container.model_routing.configuration())
|
|
config = ProviderConfig(provider_id="concurrent", provider_type=ProviderType.openai_compatible,
|
|
name="saved during rebuild", base_url="https://unused.invalid/v1")
|
|
container.providers.register(config, container.provider_factory.build(config))
|
|
task_service.update_task(task.task_id, {"title": "saved during rebuild"})
|
|
# Preparation keeps the old searchable index intact while API I/O is pending.
|
|
assert repository.stats()["notes"] == 2
|
|
if failure == "cancel":
|
|
rebuilding.cancel()
|
|
expected = asyncio.CancelledError
|
|
else:
|
|
release.set()
|
|
expected = RuntimeError
|
|
with pytest.raises(expected):
|
|
await rebuilding
|
|
assert container.model_routing.configuration().version == saved.config.version
|
|
assert rows("SELECT provider_id FROM provider_configs")[-1][0] == "concurrent"
|
|
restored = task_service.get_task(task.task_id)
|
|
assert restored.title == "saved during rebuild"
|
|
assert restored.note_id == apple.note_id
|
|
for table, values in before.items():
|
|
assert [tuple(row) for row in rows(f"SELECT * FROM {table}")] == values
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def local_engine():
|
|
return RetrievalEngine(HashEmbeddingProvider(), LexicalReranker(), SqliteVecStore())
|
|
|
|
|
|
def request(mode=SearchMode.vector):
|
|
return SearchRequest(query="apple", mode=mode, limit=10)
|
|
|
|
|
|
def rows(sql, parameters=()):
|
|
conn = connect()
|
|
try:
|
|
return conn.execute(sql, parameters).fetchall()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_api_index_and_query_use_matching_space_and_keep_local_metadata(runtime):
|
|
async def scenario():
|
|
apple, banana = await seed()
|
|
result = await engine.search(request())
|
|
assert result.items[0].note_id == banana.note_id
|
|
baseline = await local_engine().search(request())
|
|
assert baseline.items[0].note_id == apple.note_id
|
|
assert rows("SELECT DISTINCT space_id, dimensions FROM routed_block_vectors")[0][:] == ("space-a", 3)
|
|
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == len(apple.blocks) + len(banana.blocks)
|
|
meta = repository.get_index_meta()
|
|
assert meta["embedding_model"] == "hash-v1"
|
|
assert meta["embedding_dim"] == "128"
|
|
assert len(runtime.calls) == 3
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.mark.parametrize("failure", ["exception", "local", "missing", "dimension", "corrupt"])
|
|
def test_query_falls_back_to_exact_local_results(runtime, failure):
|
|
async def scenario():
|
|
await seed()
|
|
if failure == "exception":
|
|
runtime.error = RuntimeError("offline")
|
|
elif failure == "local":
|
|
runtime.source = "local"
|
|
elif failure == "missing":
|
|
rows("DELETE FROM routed_block_vectors WHERE block_id = (SELECT MIN(block_id) FROM blocks)")
|
|
elif failure == "dimension":
|
|
runtime.dimensions = 4
|
|
else:
|
|
rows("UPDATE routed_block_vectors SET vector = ?", ("[0, 0, 0]",))
|
|
actual = await engine.search(request())
|
|
baseline = await local_engine().search(request())
|
|
assert actual == baseline
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_same_dimension_model_switch_never_combines_partial_spaces(runtime):
|
|
async def scenario():
|
|
apple, banana = await seed()
|
|
baseline = await local_engine().search(request())
|
|
runtime.model_id = "space-b"
|
|
assert await engine.search(request()) == baseline
|
|
await note_service.update_note(apple.note_id, markdown="apple orchard")
|
|
assert {row[0] for row in rows("SELECT DISTINCT space_id FROM routed_block_vectors")} == {"space-a", "space-b"}
|
|
assert await routed_vectors.search_remote("apple", top_k=10) is None
|
|
assert await engine.search(request()) == baseline
|
|
runtime.model_id = "space-a"
|
|
assert await engine.search(request()) == baseline
|
|
runtime.model_id = "space-b"
|
|
await note_service.update_note(banana.note_id, markdown="banana grove")
|
|
hits = await routed_vectors.search_remote("apple", top_k=10)
|
|
assert hits is not None and hits[0].id == banana.blocks[0].block_id
|
|
assert (await engine.search(request())).items[0].note_id == banana.note_id
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_complete_spaces_coexist_but_only_requested_space_is_ranked(runtime):
|
|
async def scenario():
|
|
apple, banana = await seed()
|
|
conn = connect()
|
|
try:
|
|
with transaction(conn):
|
|
routed_vectors.store_remote(
|
|
conn, [apple.blocks[0].block_id, banana.blocks[0].block_id],
|
|
routed_vectors.RemoteEmbeddings("space-b", 3, [[1, 0, 0], [0, 1, 0]]),
|
|
)
|
|
finally:
|
|
conn.close()
|
|
assert (await engine.search(request())).items[0].note_id == banana.note_id
|
|
runtime.model_id = "space-b"
|
|
result = await engine.search(request())
|
|
assert len(result.items) == 2
|
|
assert result.items[0].note_id == apple.note_id
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_failed_note_embedding_preserves_save_and_forces_coverage_fallback(runtime):
|
|
async def scenario():
|
|
apple, banana = await seed()
|
|
runtime.error = RuntimeError("offline")
|
|
await note_service.update_note(banana.note_id, markdown="banana changed")
|
|
assert (await note_service.get_note(banana.note_id)).markdown == "banana changed"
|
|
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == len(apple.blocks)
|
|
runtime.error = None
|
|
assert await engine.search(request()) == await local_engine().search(request())
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.mark.parametrize("vectors, dimensions, space", [
|
|
([], 3, "space-a"),
|
|
([[1, 0]], 3, "space-a"),
|
|
([[0, 0, 0]], 3, "space-a"),
|
|
([[float("nan"), 0, 0]], 3, "space-a"),
|
|
([[float("inf"), 0, 0]], 3, "space-a"),
|
|
([[True, 0, 0]], 3, "space-a"),
|
|
([[1, 0, 0]], 0, "space-a"),
|
|
([[1, 0, 0]], 3, "hash-v1"),
|
|
])
|
|
def test_invalid_remote_batch_does_not_break_note_saving(runtime, vectors, dimensions, space):
|
|
runtime.result_override = SimpleNamespace(
|
|
source="api", vectors=vectors, dimensions=dimensions, model_id=space,
|
|
)
|
|
|
|
async def scenario():
|
|
note = await note_service.create_note(title="Apple", markdown="apple orchard", folder=None, tags=[])
|
|
assert (await local_engine().search(request())).items[0].note_id == note.note_id
|
|
assert await routed_vectors.search_remote("apple", top_k=10) is None
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_remote_storage_failure_rolls_back_batch_but_keeps_local_index(runtime):
|
|
async def scenario():
|
|
await seed()
|
|
rows("""CREATE TRIGGER reject_remote_vector BEFORE INSERT ON routed_block_vectors
|
|
WHEN (SELECT content FROM blocks WHERE block_id = NEW.block_id) = 'second'
|
|
BEGIN SELECT RAISE(ABORT, 'simulated storage failure'); END""")
|
|
note = await note_service.create_note(
|
|
title="Multi", markdown="first\n\nsecond", folder=None, tags=[],
|
|
)
|
|
assert len(note.blocks) == 2
|
|
assert rows(
|
|
"SELECT COUNT(*) FROM routed_block_vectors r JOIN blocks b USING(block_id) WHERE b.note_id = ?",
|
|
(note.note_id,),
|
|
)[0][0] == 0
|
|
assert rows("SELECT COUNT(*) FROM vec_blocks")[0][0] == rows("SELECT COUNT(*) FROM blocks")[0][0]
|
|
assert (get_settings().vault_path / note.file_path).exists()
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_rebuild_and_delete_clear_old_remote_rows_through_foreign_keys(runtime):
|
|
async def scenario():
|
|
apple, _ = await seed()
|
|
await note_service.delete_note(apple.note_id)
|
|
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == 1
|
|
runtime.source = "local"
|
|
job = await index_service.rebuild(IndexRebuildRequest())
|
|
assert job.status == "completed"
|
|
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == 0
|
|
assert rows("SELECT COUNT(*) FROM vec_blocks")[0][0] == 1
|
|
runtime.source = "api"
|
|
runtime.model_id = "space-b"
|
|
await index_service.rebuild(IndexRebuildRequest())
|
|
assert [row[0] for row in rows("SELECT space_id FROM routed_block_vectors")] == ["space-b"]
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.mark.parametrize("operation", ["save", "query", "rebuild"])
|
|
def test_cancellation_propagates_and_mutations_roll_back(runtime, operation):
|
|
async def scenario():
|
|
apple, _ = await seed()
|
|
before = [tuple(row) for row in rows("SELECT * FROM routed_block_vectors ORDER BY block_id")]
|
|
runtime.error = asyncio.CancelledError()
|
|
with pytest.raises(asyncio.CancelledError):
|
|
if operation == "query":
|
|
await engine.search(request())
|
|
elif operation == "rebuild":
|
|
await index_service.rebuild(IndexRebuildRequest())
|
|
else:
|
|
await note_service.update_note(apple.note_id, markdown="changed")
|
|
assert (await note_service.get_note(apple.note_id)).markdown == "apple orchard"
|
|
assert [tuple(row) for row in rows("SELECT * FROM routed_block_vectors ORDER BY block_id")] == before
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.mark.parametrize("injected", ["embedding", "vector_store", "constructor"])
|
|
def test_injected_engine_dependencies_are_respected(runtime, monkeypatch, injected):
|
|
async def scenario():
|
|
apple, _ = await seed()
|
|
target = engine
|
|
if injected == "constructor":
|
|
target = local_engine()
|
|
elif injected == "embedding":
|
|
monkeypatch.setattr(engine, "embedding", HashEmbeddingProvider())
|
|
else:
|
|
class FakeStore:
|
|
async def search(self, vector, *, top_k):
|
|
assert len(vector) == 128
|
|
return [VectorHit(id=apple.blocks[0].block_id, score=1.0)]
|
|
|
|
monkeypatch.setattr(engine, "vector_store", FakeStore())
|
|
runtime.calls.clear()
|
|
assert (await target.search(request())).items[0].note_id == apple.note_id
|
|
assert runtime.calls == []
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_fts_skips_routing_and_hybrid_uses_routed_vector_channel(runtime, monkeypatch):
|
|
async def scenario():
|
|
_, banana = await seed()
|
|
runtime.calls.clear()
|
|
await engine.search(request(SearchMode.fts))
|
|
assert runtime.calls == []
|
|
# Empty lexical channel isolates the vector contribution to hybrid fusion.
|
|
monkeypatch.setattr(repository, "fts_search", lambda *_: [])
|
|
|
|
class PreserveOrder:
|
|
async def rerank(self, query, candidates):
|
|
return sorted(candidates, key=lambda candidate: -candidate.score)
|
|
|
|
monkeypatch.setattr(engine, "reranker", PreserveOrder())
|
|
result = await engine.search(request(SearchMode.hybrid))
|
|
assert result.items[0].note_id == banana.note_id
|
|
assert runtime.calls == [["apple"]]
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_arbitrary_dimensions_and_extreme_finite_values(runtime):
|
|
dimensions = 257
|
|
runtime.result_override = SimpleNamespace(
|
|
source="api", model_id="space-wide", dimensions=dimensions,
|
|
vectors=[[1e308, 1e308] + [0.0] * (dimensions - 2)],
|
|
)
|
|
|
|
async def scenario():
|
|
note = await note_service.create_note(title="Apple", markdown="apple orchard", folder=None, tags=[])
|
|
hits = await routed_vectors.search_remote("apple", top_k=1)
|
|
assert hits is not None and hits[0].id == note.blocks[0].block_id
|
|
assert hits[0].score == pytest.approx(1.0)
|
|
vector = json.loads(rows("SELECT vector FROM routed_block_vectors")[0][0])
|
|
assert len(vector) == dimensions
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_missing_runtime_uses_unchanged_local_retrieval(runtime, monkeypatch):
|
|
monkeypatch.setattr(routed_vectors, "get_model_routing", lambda: None)
|
|
|
|
async def scenario():
|
|
await seed()
|
|
assert await engine.search(request()) == await local_engine().search(request())
|
|
assert runtime.calls == []
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.fixture
|
|
def production_engine(monkeypatch):
|
|
from app.local_models.runtime import LocalEmbedding
|
|
embedding = LocalEmbedding()
|
|
monkeypatch.setattr(note_service, "embedding", embedding)
|
|
return RetrievalEngine(embedding, LexicalReranker(), SqliteVecStore(), route_embeddings=True)
|
|
|
|
|
|
@pytest.mark.parametrize("source", ["api", "local"])
|
|
def test_real_embedding_route_rebuilds_missing_space(runtime, production_engine, source):
|
|
from app.errors import ApiError
|
|
runtime.source = source
|
|
|
|
async def scenario():
|
|
await seed()
|
|
runtime.model_id = "new-configured-space"
|
|
with pytest.raises(ApiError) as error:
|
|
await production_engine.search(request())
|
|
assert error.value.code == "SEMANTIC_INDEX_UNAVAILABLE"
|
|
assert "Embedding 已可用" in error.value.message
|
|
assert error.value.details["source"] == source
|
|
await index_service.rebuild(IndexRebuildRequest())
|
|
assert (await production_engine.search(request())).items
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_real_embedding_failure_is_not_reported_as_missing_configuration(runtime, production_engine):
|
|
from app.errors import ApiError
|
|
|
|
async def scenario():
|
|
await seed()
|
|
runtime.error = ApiError(503, "LOCAL_MODEL_TIMEOUT", "本地模型推理超时。", {"fallback_reason": "PROVIDER_TIMEOUT"})
|
|
with pytest.raises(ApiError) as error:
|
|
await production_engine.search(request())
|
|
assert error.value.code == "LOCAL_MODEL_TIMEOUT"
|
|
assert error.value.details["fallback_reason"] == "PROVIDER_TIMEOUT"
|
|
assert (await production_engine.search(SearchRequest(query="apple", mode=SearchMode.hybrid))).items
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
@pytest.mark.parametrize("failure", ["inference", "storage", "space_change"])
|
|
def test_real_embedding_rebuild_failure_preserves_index(runtime, production_engine, monkeypatch, failure):
|
|
from app.errors import ApiError
|
|
|
|
async def scenario():
|
|
await seed()
|
|
tables = ("notes", "blocks", "blocks_fts", "index_meta", "routed_block_vectors")
|
|
before = {table: [tuple(r) for r in rows(f"SELECT * FROM {table}")] for table in tables}
|
|
if failure == "inference":
|
|
runtime.error = ApiError(503, "LOCAL_MODEL_TIMEOUT", "本地模型推理超时。")
|
|
elif failure == "storage":
|
|
monkeypatch.setattr(routed_vectors, "store_remote", lambda *args: None)
|
|
else:
|
|
original = runtime.embed
|
|
async def changing(texts):
|
|
runtime.model_id += "x"
|
|
return await original(texts)
|
|
monkeypatch.setattr(runtime, "embed", changing)
|
|
with pytest.raises(ApiError):
|
|
await index_service.rebuild(IndexRebuildRequest())
|
|
assert index_service.get_status().status == "failed"
|
|
after = {table: [tuple(r) for r in rows(f"SELECT * FROM {table}")] for table in tables}
|
|
assert before == after
|
|
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_empty_vault_vector_search_returns_empty(runtime, production_engine):
|
|
assert asyncio.run(production_engine.search(request())).items == []
|
|
|
|
|
|
@pytest.fixture
|
|
def policy_runtime(monkeypatch):
|
|
class PolicyRuntime:
|
|
fallback = False
|
|
calls = []
|
|
async def embed(self, texts, *, local_only=False):
|
|
self.calls.append((list(texts), local_only))
|
|
local = local_only or self.fallback
|
|
dim = 3 if local else 2
|
|
return SimpleNamespace(source='local' if local else 'api', model_id='local-space' if local else 'api-space',
|
|
dimensions=dim, vectors=[[1.0] + [0.0] * (dim - 1) for _ in texts],
|
|
fallback_reason='PROVIDER_TIMEOUT' if self.fallback and not local_only else None)
|
|
runtime = PolicyRuntime()
|
|
monkeypatch.setattr(routed_vectors, 'get_model_routing', lambda: runtime)
|
|
return runtime
|
|
|
|
|
|
async def seed_policies():
|
|
normal = await note_service.create_note(title='Normal', markdown='apple public', folder=None, tags=[])
|
|
private = await note_service.create_note(title='Private', markdown='---\nembedding_local_only: true\n---\napple private', folder=None, tags=[])
|
|
return normal, private
|
|
|
|
|
|
@pytest.mark.parametrize('fallback', [False, True])
|
|
def test_mixed_policy_rebuild_and_retrieval(policy_runtime, production_engine, fallback):
|
|
policy_runtime.fallback = fallback
|
|
async def scenario():
|
|
notes = await seed_policies()
|
|
await index_service.rebuild(IndexRebuildRequest())
|
|
for mode in (SearchMode.vector, SearchMode.hybrid):
|
|
result = await production_engine.search(SearchRequest(query='apple', mode=mode))
|
|
assert {item.note_id for item in result.items} == {note.note_id for note in notes}
|
|
for texts, local_only in policy_runtime.calls:
|
|
if any('private' in text for text in texts):
|
|
assert local_only
|
|
if not fallback:
|
|
assert {r[0] for r in rows('SELECT DISTINCT space_id FROM routed_block_vectors')} == {'api-space', 'local-space'}
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_local_only_vault_never_requests_api_for_search(policy_runtime, production_engine):
|
|
async def scenario():
|
|
await note_service.create_note(title='Private', markdown='---\nembedding_local_only: true\n---\napple private', folder=None, tags=[])
|
|
await index_service.rebuild(IndexRebuildRequest())
|
|
assert (await production_engine.search(request())).items
|
|
assert all(local_only for _, local_only in policy_runtime.calls)
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_partition_storage_failure_rolls_back_all_partitions(policy_runtime, production_engine, monkeypatch):
|
|
from app.errors import ApiError
|
|
async def scenario():
|
|
await seed_policies()
|
|
before = [tuple(row) for row in rows('SELECT * FROM routed_block_vectors ORDER BY block_id')]
|
|
original = routed_vectors.store_remote
|
|
def fail_local(conn, ids, batch):
|
|
if batch.source != 'local':
|
|
original(conn, ids, batch)
|
|
monkeypatch.setattr(routed_vectors, 'store_remote', fail_local)
|
|
with pytest.raises(ApiError) as error:
|
|
await index_service.rebuild(IndexRebuildRequest())
|
|
assert error.value.code == 'SEMANTIC_INDEX_WRITE_FAILED'
|
|
assert [tuple(row) for row in rows('SELECT * FROM routed_block_vectors ORDER BY block_id')] == before
|
|
asyncio.run(scenario())
|
|
|
|
|
|
def test_missing_partition_does_not_silently_return_partial_hits(policy_runtime, production_engine):
|
|
from app.errors import ApiError
|
|
async def scenario():
|
|
await seed_policies()
|
|
conn = connect()
|
|
try:
|
|
conn.execute("DELETE FROM routed_block_vectors WHERE space_id='local-space'")
|
|
finally:
|
|
conn.close()
|
|
with pytest.raises(ApiError) as error:
|
|
await production_engine.search(request())
|
|
assert error.value.code == 'SEMANTIC_INDEX_UNAVAILABLE'
|
|
assert (await production_engine.search(request(SearchMode.hybrid))).items
|
|
asyncio.run(scenario())
|