deerflow-code/offline-backend-20260512/backend/tests/test_llmwiki_local_index_repository.py
2026-09-07 18:24:55 +08:00

99 lines
4.2 KiB
Python

from __future__ import annotations
from datetime import UTC, datetime
import pytest
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from deerflow.integrations.weknora.local_index.normalize import normalize_wiki_page
from deerflow.integrations.weknora.local_index.vector_codec import decode_vector, encode_vector
from deerflow.persistence.base import Base
from deerflow.persistence.llmwiki.model import LlmWikiKnowledgeBaseRow
from deerflow.persistence.llmwiki_index.model import LlmWikiWikiPageRow, LlmWikiWikiSyncStateRow, LlmWikiWikiVectorRow
from deerflow.persistence.llmwiki_index.sql import SqlLlmWikiIndexStore
@pytest.mark.asyncio
async def test_sql_store_atomically_replaces_vectors_and_marks_missing_deleted(tmp_path) -> None:
engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'wiki.db'}")
async with engine.begin() as connection:
await connection.run_sync(
lambda sync_connection: Base.metadata.create_all(
sync_connection,
tables=[LlmWikiKnowledgeBaseRow.__table__, LlmWikiWikiPageRow.__table__, LlmWikiWikiVectorRow.__table__, LlmWikiWikiSyncStateRow.__table__],
)
)
session_factory = async_sessionmaker(engine, expire_on_commit=False)
async with session_factory() as session:
session.add(
LlmWikiKnowledgeBaseRow(
id="mapping-1",
weknora_id="remote-1",
owner_user_id="u1",
name="Wiki",
description="",
kb_type="document",
publication_status="private",
created_at=datetime.now(UTC),
updated_at=datetime.now(UTC),
)
)
await session.commit()
store = SqlLlmWikiIndexStore(session_factory, write_batch_size=1)
assert await store.try_acquire_sync_lease("mapping-1", owner="worker", sync_id="sync-1", lease_seconds=60, fingerprint="fp")
assert not await store.try_acquire_sync_lease("mapping-1", owner="worker", sync_id="sync-2", lease_seconds=60, fingerprint="fp")
page = normalize_wiki_page("mapping-1", {"id": "p1", "slug": "people/alice", "title": "Alice", "content": "Alice founded Acme.", "status": "published"})
blob, dimensions = encode_vector([3.0, 4.0], dimensions=2)
saved = await store.replace_page_vectors(
page=page,
vectors=[
{
"section_index": index,
"heading": "History",
"section_content": f"Alice founded Acme. Part {index}",
"content_hash": f"section-{index}",
"embedding_dimensions": dimensions,
"vector_blob": blob,
}
for index in range(3)
],
fingerprint="fp",
sync_id="sync-1",
)
snapshot = await store.load_vector_snapshot("mapping-1", "fp")
assert saved["slug"] == "people/alice"
assert store.write_batch_size == 1
assert len(snapshot["rows"]) == 3
assert decode_vector(snapshot["rows"][0]["vector_blob"], 2).tolist() == pytest.approx([0.6, 0.8])
streamed = [row async for row in store.iter_vector_snapshot_rows("mapping-1", "fp", batch_size=2)]
assert [row["section_index"] for row in streamed] == [0, 1, 2]
assert streamed[0]["page"]["slug"] == "people/alice"
page_statuses = await store.list_page_index_status("mapping-1")
assert len(page_statuses) == 1
assert page_statuses[0] == {
**page_statuses[0],
"id": saved["id"],
"slug": "people/alice",
"title": "Alice",
"page_type": "article",
"status": "published",
"index_status": "ready",
"vector_count": 3,
"last_error": None,
"last_seen_sync_id": "sync-1",
}
assert page_statuses[0]["indexed_at"]
assert page_statuses[0]["updated_at"]
assert await store.mark_missing_pages_deleted("mapping-1", "sync-2") == 1
assert (await store.load_vector_snapshot("mapping-1", "fp"))["rows"] == []
await engine.dispose()
def test_vector_codec_rejects_corrupt_and_non_finite_vectors() -> None:
with pytest.raises(ValueError, match="NaN"):
encode_vector([1.0, float("nan")], dimensions=2)
with pytest.raises(ValueError, match="Corrupt"):
decode_vector(b"123", 2)