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)