99 lines
4.2 KiB
Python
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)
|