fix: coordinate vector migration with concurrent searches and note saves
This commit is contained in:
@@ -190,9 +190,49 @@ async def search_remote(query: str, *, top_k: int, accept_local=False, strict=Fa
|
|||||||
if batch is None:
|
if batch is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
if not await _prepare_for_search([batch], strict):
|
||||||
|
return None
|
||||||
return await asyncio.to_thread(_search_space, batch, top_k, strict)
|
return await asyncio.to_thread(_search_space, batch, top_k, strict)
|
||||||
|
|
||||||
|
|
||||||
|
async def _prepare_indexes(batches):
|
||||||
|
from app.services.coordination import vault_mutation_lock
|
||||||
|
def prepare(check_only=False):
|
||||||
|
conn = connect()
|
||||||
|
try:
|
||||||
|
if check_only:
|
||||||
|
return space_index.is_ready(conn, batches)
|
||||||
|
space_index.prepare(conn, batches)
|
||||||
|
finally:
|
||||||
|
conn.close()
|
||||||
|
if await asyncio.to_thread(prepare, True):
|
||||||
|
return
|
||||||
|
# Share the cooperative gate with saves: never block the event loop on a
|
||||||
|
# SQLite write lock while a migration owns it in another thread.
|
||||||
|
async with vault_mutation_lock():
|
||||||
|
work = asyncio.create_task(asyncio.to_thread(prepare))
|
||||||
|
cancelled = False
|
||||||
|
while not work.done():
|
||||||
|
try:
|
||||||
|
await asyncio.shield(work)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
cancelled = True
|
||||||
|
work.result()
|
||||||
|
if cancelled:
|
||||||
|
raise asyncio.CancelledError
|
||||||
|
|
||||||
|
|
||||||
|
async def _prepare_for_search(batches, strict):
|
||||||
|
try:
|
||||||
|
await _prepare_indexes(batches)
|
||||||
|
return True
|
||||||
|
except Exception as exc:
|
||||||
|
record_embedding(fallback_reason='REMOTE_INDEX_UNAVAILABLE')
|
||||||
|
if strict:
|
||||||
|
raise ApiError(409, 'SEMANTIC_INDEX_UNAVAILABLE', '向量索引准备失败,请检查索引状态。') from exc
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def _search_space(batch, top_k, strict):
|
def _search_space(batch, top_k, strict):
|
||||||
record_embedding(attempted_space={"model_id": batch.space_id, "dimensions": batch.dimensions})
|
record_embedding(attempted_space={"model_id": batch.space_id, "dimensions": batch.dimensions})
|
||||||
try:
|
try:
|
||||||
@@ -209,7 +249,6 @@ def _search_space(batch, top_k, strict):
|
|||||||
if strict:
|
if strict:
|
||||||
raise ValueError("semantic index missing")
|
raise ValueError("semantic index missing")
|
||||||
return None
|
return None
|
||||||
_ensure_table(conn)
|
|
||||||
result = space_index.search(conn, batch, top_k)
|
result = space_index.search(conn, batch, top_k)
|
||||||
record_embedding(source=batch.source, model_id=batch.space_id,
|
record_embedding(source=batch.source, model_id=batch.space_id,
|
||||||
dimensions=batch.dimensions, fallback_reason=None)
|
dimensions=batch.dimensions, fallback_reason=None)
|
||||||
@@ -234,6 +273,8 @@ async def _search_partitioned(query: str, policies: set[bool], *, top_k: int, st
|
|||||||
if batch is None:
|
if batch is None:
|
||||||
return None
|
return None
|
||||||
batches[policy] = batch
|
batches[policy] = batch
|
||||||
|
if not await _prepare_for_search(list(batches.values()), strict):
|
||||||
|
return None
|
||||||
return await asyncio.to_thread(_search_partitions, batches, policies, top_k, strict)
|
return await asyncio.to_thread(_search_partitions, batches, policies, top_k, strict)
|
||||||
|
|
||||||
|
|
||||||
@@ -247,7 +288,6 @@ def _search_partitions(batches, policies, top_k, strict):
|
|||||||
raise ValueError("embedding policies changed while querying")
|
raise ValueError("embedding policies changed while querying")
|
||||||
ranked = []
|
ranked = []
|
||||||
for policy, batch in batches.items():
|
for policy, batch in batches.items():
|
||||||
_ensure_table(conn)
|
|
||||||
ranked.append(space_index.search(conn, batch, top_k, policy))
|
ranked.append(space_index.search(conn, batch, top_k, policy))
|
||||||
spaces = [{"source": b.source, "model_id": b.space_id, "dimensions": b.dimensions,
|
spaces = [{"source": b.source, "model_id": b.space_id, "dimensions": b.dimensions,
|
||||||
"local_only": policy} for policy, b in batches.items()]
|
"local_only": policy} for policy, b in batches.items()]
|
||||||
|
|||||||
@@ -1,12 +1,42 @@
|
|||||||
"""Persistent vec0 indexes derived from durable routed vectors, one per space/dimension."""
|
"""Persistent vec0 indexes derived from durable routed vectors, one per space/dimension."""
|
||||||
import hashlib
|
import hashlib
|
||||||
import json
|
import json
|
||||||
|
import threading
|
||||||
|
|
||||||
import sqlite_vec
|
import sqlite_vec
|
||||||
|
|
||||||
from app.retrieval.vectorstore import VectorHit
|
from app.retrieval.vectorstore import VectorHit
|
||||||
|
|
||||||
|
|
||||||
|
_migration_lock = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
|
def is_ready(conn, batches):
|
||||||
|
return all(conn.execute('SELECT 1 FROM sqlite_master WHERE name=?',
|
||||||
|
(table_name(batch.space_id, batch.dimensions),)).fetchone() for batch in batches)
|
||||||
|
|
||||||
|
|
||||||
|
def prepare(conn, batches):
|
||||||
|
"""Finish lazy writes before opening a search snapshot. Warm searches do not write."""
|
||||||
|
from app.retrieval.routed_vectors import _ensure_table
|
||||||
|
batches = list(batches)
|
||||||
|
if is_ready(conn, batches):
|
||||||
|
return
|
||||||
|
# Waiting holds no read transaction, so a concurrent migration can commit.
|
||||||
|
with _migration_lock:
|
||||||
|
if is_ready(conn, batches):
|
||||||
|
return
|
||||||
|
conn.execute('BEGIN IMMEDIATE')
|
||||||
|
try:
|
||||||
|
_ensure_table(conn)
|
||||||
|
for batch in batches:
|
||||||
|
ensure(conn, batch.space_id, batch.dimensions)
|
||||||
|
conn.execute('COMMIT')
|
||||||
|
except BaseException:
|
||||||
|
conn.execute('ROLLBACK')
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
def table_name(space, dimensions):
|
def table_name(space, dimensions):
|
||||||
return 'routed_vec_' + hashlib.sha256(json.dumps([space, dimensions]).encode()).hexdigest()
|
return 'routed_vec_' + hashlib.sha256(json.dumps([space, dimensions]).encode()).hexdigest()
|
||||||
|
|
||||||
@@ -40,7 +70,7 @@ def upsert(conn, block_ids, batch):
|
|||||||
|
|
||||||
|
|
||||||
def search(conn, batch, top_k, policy=None):
|
def search(conn, batch, top_k, policy=None):
|
||||||
table = ensure(conn, batch.space_id, batch.dimensions)
|
table = table_name(batch.space_id, batch.dimensions)
|
||||||
# Coverage checks stay relational; no JSON decoding or Python dot products on the hot path.
|
# Coverage checks stay relational; no JSON decoding or Python dot products on the hot path.
|
||||||
where = '' if policy is None else ' AND b.embedding_local_only=?'
|
where = '' if policy is None else ' AND b.embedding_local_only=?'
|
||||||
params = () if policy is None else (int(policy),)
|
params = () if policy is None else (int(policy),)
|
||||||
|
|||||||
@@ -1,7 +1,14 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
|
from weakref import WeakKeyDictionary
|
||||||
|
|
||||||
_vault_mutation_lock = asyncio.Lock()
|
_vault_locks = WeakKeyDictionary()
|
||||||
|
|
||||||
|
|
||||||
|
def vault_mutation_lock():
|
||||||
|
# Service/test lifecycle restarts must not reuse a lock bound to a closed loop.
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
return _vault_locks.setdefault(loop, asyncio.Lock())
|
||||||
|
|
||||||
|
|
||||||
def serialized_vault_mutation(operation):
|
def serialized_vault_mutation(operation):
|
||||||
@@ -9,7 +16,7 @@ def serialized_vault_mutation(operation):
|
|||||||
|
|
||||||
@wraps(operation)
|
@wraps(operation)
|
||||||
async def wrapped(*args, **kwargs):
|
async def wrapped(*args, **kwargs):
|
||||||
async with _vault_mutation_lock:
|
async with vault_mutation_lock():
|
||||||
return await operation(*args, **kwargs)
|
return await operation(*args, **kwargs)
|
||||||
|
|
||||||
return wrapped
|
return wrapped
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from app.errors import ApiError
|
|||||||
from app.knowledge.parser import parse_note
|
from app.knowledge.parser import parse_note
|
||||||
from app.services.note_service import index_note, prepare_note_index
|
from app.services.note_service import index_note, prepare_note_index
|
||||||
from app.database.db import connect, transaction
|
from app.database.db import connect, transaction
|
||||||
from app.services.coordination import _vault_mutation_lock
|
from app.services.coordination import vault_mutation_lock
|
||||||
from app.retrieval.vectorstore import SqliteVecStore
|
from app.retrieval.vectorstore import SqliteVecStore
|
||||||
from app.local_models.runtime import LocalEmbedding
|
from app.local_models.runtime import LocalEmbedding
|
||||||
from app.services import note_service
|
from app.services import note_service
|
||||||
@@ -116,7 +116,7 @@ async def rebuild(request: IndexRebuildRequest) -> IndexJob:
|
|||||||
prepared_notes.append((parsed, prepared))
|
prepared_notes.append((parsed, prepared))
|
||||||
# All network/model awaits precede the transaction. The concrete SQLite
|
# All network/model awaits precede the transaction. The concrete SQLite
|
||||||
# methods below complete synchronously despite their async interfaces.
|
# methods below complete synchronously despite their async interfaces.
|
||||||
async with _vault_mutation_lock:
|
async with vault_mutation_lock():
|
||||||
if _scan_vault() != docs or saved_records != {key: repository.get_note_record(key) for key in _pending_notes()}:
|
if _scan_vault() != docs or saved_records != {key: repository.get_note_record(key) for key in _pending_notes()}:
|
||||||
raise ApiError(409, "INDEX_SNAPSHOT_CHANGED", "笔记在计算期间发生变化,稍后重新计算。")
|
raise ApiError(409, "INDEX_SNAPSHOT_CHANGED", "笔记在计算期间发生变化,稍后重新计算。")
|
||||||
conn = connect()
|
conn = connect()
|
||||||
@@ -263,7 +263,7 @@ async def _refresh_saved_note(note_id: str) -> None:
|
|||||||
prepared = await prepare_note_index(parsed, strict=True)
|
prepared = await prepare_note_index(parsed, strict=True)
|
||||||
if isinstance(note_service.embedding, LocalEmbedding) and parsed.blocks and prepared[1] is None:
|
if isinstance(note_service.embedding, LocalEmbedding) and parsed.blocks and prepared[1] is None:
|
||||||
raise ApiError(503, "EMBEDDING_UNAVAILABLE", "笔记已保存,后台向量计算未完成。")
|
raise ApiError(503, "EMBEDDING_UNAVAILABLE", "笔记已保存,后台向量计算未完成。")
|
||||||
async with _vault_mutation_lock:
|
async with vault_mutation_lock():
|
||||||
current = repository.get_note_record(note_id)
|
current = repository.get_note_record(note_id)
|
||||||
if current != record or note_service._read_markdown(record.file_path) != markdown:
|
if current != record or note_service._read_markdown(record.file_path) != markdown:
|
||||||
# Another save or rename won the race; leave the durable queue entry intact.
|
# Another save or rename won the race; leave the durable queue entry intact.
|
||||||
|
|||||||
@@ -112,6 +112,115 @@ def test_legacy_vectors_migrate_without_document_embedding(runtime):
|
|||||||
asyncio.run(scenario())
|
asyncio.run(scenario())
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize('partitioned', [False, True])
|
||||||
|
def test_concurrent_first_search_serializes_migration_and_warm_search_is_read_only(runtime, monkeypatch, partitioned):
|
||||||
|
import threading
|
||||||
|
from app.retrieval import space_index
|
||||||
|
async def scenario():
|
||||||
|
await seed()
|
||||||
|
table = space_index.table_name('space-a', 3)
|
||||||
|
conn = connect()
|
||||||
|
try:
|
||||||
|
with transaction(conn):
|
||||||
|
conn.execute(f'DROP TRIGGER {table}_delete')
|
||||||
|
conn.execute(f'DROP TRIGGER {table}_update')
|
||||||
|
conn.execute(f'DROP TABLE {table}')
|
||||||
|
finally:
|
||||||
|
conn.close()
|
||||||
|
entered, release, second = threading.Event(), threading.Event(), threading.Event()
|
||||||
|
original_ensure, original_prepare = space_index.ensure, space_index.prepare
|
||||||
|
calls = []
|
||||||
|
def ensure(*args):
|
||||||
|
calls.append(1)
|
||||||
|
entered.set()
|
||||||
|
assert release.wait(5)
|
||||||
|
return original_ensure(*args)
|
||||||
|
def prepare(*args):
|
||||||
|
if entered.is_set():
|
||||||
|
second.set()
|
||||||
|
return original_prepare(*args)
|
||||||
|
monkeypatch.setattr(space_index, 'ensure', ensure)
|
||||||
|
monkeypatch.setattr(space_index, 'prepare', prepare)
|
||||||
|
batch = routed_vectors.RemoteEmbeddings('space-a', 3, [[1., 0., 0.]])
|
||||||
|
async def search():
|
||||||
|
if entered.is_set():
|
||||||
|
second.set()
|
||||||
|
await routed_vectors._prepare_indexes([batch])
|
||||||
|
if partitioned:
|
||||||
|
return await asyncio.to_thread(routed_vectors._search_partitions, {False: batch}, {False}, 2, True)
|
||||||
|
return await asyncio.to_thread(routed_vectors._search_space, batch, 2, True)
|
||||||
|
tasks = []
|
||||||
|
try:
|
||||||
|
tasks.append(asyncio.create_task(search()))
|
||||||
|
assert await asyncio.to_thread(entered.wait, 5)
|
||||||
|
tasks.append(asyncio.create_task(search()))
|
||||||
|
assert await asyncio.to_thread(second.wait, 5)
|
||||||
|
release.set()
|
||||||
|
first, other = await asyncio.gather(*tasks)
|
||||||
|
assert first == other and len(first) == 2
|
||||||
|
assert len(calls) == 1
|
||||||
|
# Prepared indexes are reusable even with SQLite query_only enforced.
|
||||||
|
original_connect = routed_vectors.connect
|
||||||
|
def read_only():
|
||||||
|
connection = original_connect()
|
||||||
|
connection.execute('PRAGMA query_only=ON')
|
||||||
|
return connection
|
||||||
|
monkeypatch.setattr(routed_vectors, 'connect', read_only)
|
||||||
|
assert await search() == first
|
||||||
|
finally:
|
||||||
|
release.set()
|
||||||
|
await asyncio.gather(*tasks, return_exceptions=True)
|
||||||
|
asyncio.run(scenario())
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize('cancel_search', [False, True])
|
||||||
|
def test_save_waits_for_migration_even_when_search_is_cancelled(runtime, monkeypatch, cancel_search):
|
||||||
|
import threading
|
||||||
|
from app.retrieval import space_index
|
||||||
|
async def scenario():
|
||||||
|
apple, _ = await seed()
|
||||||
|
table = space_index.table_name('space-a', 3)
|
||||||
|
conn = connect()
|
||||||
|
try:
|
||||||
|
with transaction(conn):
|
||||||
|
conn.execute(f'DROP TRIGGER {table}_delete')
|
||||||
|
conn.execute(f'DROP TRIGGER {table}_update')
|
||||||
|
conn.execute(f'DROP TABLE {table}')
|
||||||
|
finally:
|
||||||
|
conn.close()
|
||||||
|
entered, release = threading.Event(), threading.Event()
|
||||||
|
original = space_index.ensure
|
||||||
|
def slow(*args):
|
||||||
|
entered.set()
|
||||||
|
assert release.wait(5)
|
||||||
|
return original(*args)
|
||||||
|
monkeypatch.setattr(space_index, 'ensure', slow)
|
||||||
|
# Keep the subsequent vector job queued; test saving and its durable marker.
|
||||||
|
monkeypatch.setattr(index_service, 'schedule_workspace_rebuild', lambda: None)
|
||||||
|
query = asyncio.create_task(routed_vectors.search_remote('apple orchard', top_k=2, strict=True))
|
||||||
|
save = None
|
||||||
|
try:
|
||||||
|
assert await asyncio.to_thread(entered.wait, 5)
|
||||||
|
if cancel_search:
|
||||||
|
query.cancel()
|
||||||
|
save = asyncio.create_task(note_service.update_note(apple.note_id, markdown='Saved during migration', defer_vectors=True))
|
||||||
|
await asyncio.sleep(0.02)
|
||||||
|
assert not save.done()
|
||||||
|
release.set()
|
||||||
|
saved = await asyncio.wait_for(save, 5)
|
||||||
|
assert saved.markdown == 'Saved during migration'
|
||||||
|
assert (await note_service.get_note(apple.note_id)).markdown == saved.markdown
|
||||||
|
assert repository.get_index_meta()[f'note_vectors_pending:{apple.note_id}'] == '1'
|
||||||
|
# Query may observe the saved revision's pending index, but saving must succeed.
|
||||||
|
result = (await asyncio.gather(query, return_exceptions=True))[0]
|
||||||
|
if cancel_search:
|
||||||
|
assert isinstance(result, asyncio.CancelledError)
|
||||||
|
finally:
|
||||||
|
release.set()
|
||||||
|
await asyncio.gather(*([query, save] if save else [query]), return_exceptions=True)
|
||||||
|
asyncio.run(scenario())
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("outcome", ["api", "api_failure", "missing_space"])
|
@pytest.mark.parametrize("outcome", ["api", "api_failure", "missing_space"])
|
||||||
def test_benchmark_reports_actual_embedding_and_fallback(runtime, outcome):
|
def test_benchmark_reports_actual_embedding_and_fallback(runtime, outcome):
|
||||||
from app.benchmarks import service
|
from app.benchmarks import service
|
||||||
|
|||||||
Reference in New Issue
Block a user