"""SQLAlchemy implementation of the local Wiki index store.""" from __future__ import annotations from collections.abc import AsyncIterator from datetime import UTC, datetime, timedelta from typing import Any from uuid import uuid4 from sqlalchemy import case, delete, func, insert, or_, select, update from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.llmwiki_index.base import LlmWikiIndexStore from deerflow.persistence.llmwiki_index.model import LlmWikiWikiPageRow, LlmWikiWikiSyncStateRow, LlmWikiWikiVectorRow def _now() -> datetime: return datetime.now(UTC) def _dict(row) -> dict[str, Any]: value = row.to_dict() for key, item in tuple(value.items()): if isinstance(item, datetime): value[key] = item.isoformat() return value class SqlLlmWikiIndexStore(LlmWikiIndexStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession], *, write_batch_size: int = 50) -> None: self._sf = session_factory self.write_batch_size = max(1, min(int(write_batch_size), 500)) async def _ensure_state(self, session: AsyncSession, mapping_id: str) -> LlmWikiWikiSyncStateRow: row = await session.get(LlmWikiWikiSyncStateRow, mapping_id) if row is None: row = LlmWikiWikiSyncStateRow(knowledge_base_mapping_id=mapping_id, state="idle", index_revision=0, updated_at=_now()) session.add(row) await session.flush() return row async def get_page_by_slug_hash(self, mapping_id: str, slug_hash: str) -> dict[str, Any] | None: count = select(func.count(LlmWikiWikiVectorRow.id)).where(LlmWikiWikiVectorRow.page_id == LlmWikiWikiPageRow.id).correlate(LlmWikiWikiPageRow).scalar_subquery() stmt = select(LlmWikiWikiPageRow, count).where(LlmWikiWikiPageRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiPageRow.slug_hash == slug_hash) async with self._sf() as session: result = (await session.execute(stmt)).one_or_none() return {**_dict(result[0]), "vector_count": result[1]} if result else None async def get_page(self, mapping_id: str, slug: str) -> dict[str, Any] | None: stmt = select(LlmWikiWikiPageRow).where(LlmWikiWikiPageRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiPageRow.slug == slug) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() return _dict(row) if row else None async def touch_page(self, page_id: str, sync_id: str) -> None: async with self._sf() as session: await session.execute( update(LlmWikiWikiPageRow) .where(LlmWikiWikiPageRow.id == page_id) .values(last_seen_sync_id=sync_id, last_error=None, updated_at=_now()) ) await session.commit() async def replace_page_vectors(self, *, page: dict[str, Any], vectors: list[dict[str, Any]], fingerprint: str, sync_id: str) -> dict[str, Any]: mapping_id = str(page["knowledge_base_mapping_id"]) async with self._sf() as session: stmt = select(LlmWikiWikiPageRow).where(LlmWikiWikiPageRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiPageRow.slug_hash == page["slug_hash"]).with_for_update() row = (await session.execute(stmt)).scalar_one_or_none() if row is not None and row.slug != page["slug"]: raise ValueError("Wiki slug hash collision detected") now = _now() if row is None: row = LlmWikiWikiPageRow(id=str(uuid4()), created_at=now, **page) session.add(row) else: for key, value in page.items(): setattr(row, key, value) row.index_status = "ready" row.indexed_fingerprint = fingerprint row.indexed_at = now row.last_error = None row.is_remote_deleted = False row.last_seen_sync_id = sync_id row.updated_at = now await session.flush() await session.execute(delete(LlmWikiWikiVectorRow).where(LlmWikiWikiVectorRow.page_id == row.id)) vector_rows = [ { "id": str(uuid4()), "page_id": row.id, "knowledge_base_mapping_id": mapping_id, "section_index": int(item["section_index"]), "heading": item.get("heading"), "section_content": str(item["section_content"]), "content_hash": str(item["content_hash"]), "embedding_fingerprint": fingerprint, "embedding_dimensions": int(item["embedding_dimensions"]), "vector_blob": bytes(item["vector_blob"]), "created_at": now, } for item in vectors ] # Keep the page replacement atomic, but cap each SQL statement so # large Wiki pages cannot flood the database with one huge insert. for start in range(0, len(vector_rows), self.write_batch_size): await session.execute( insert(LlmWikiWikiVectorRow), vector_rows[start : start + self.write_batch_size], ) state = await self._ensure_state(session, mapping_id) state.index_revision = int(state.index_revision or 0) + 1 state.embedding_fingerprint = fingerprint state.updated_at = now await session.commit() return _dict(row) async def record_page_failure(self, *, page: dict[str, Any], sync_id: str, error: str) -> None: mapping_id = str(page["knowledge_base_mapping_id"]) safe_error = error[:4000] async with self._sf() as session: stmt = select(LlmWikiWikiPageRow).where(LlmWikiWikiPageRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiPageRow.slug_hash == page["slug_hash"]) row = (await session.execute(stmt)).scalar_one_or_none() now = _now() if row is None: row = LlmWikiWikiPageRow(id=str(uuid4()), created_at=now, **page) session.add(row) row.index_status = "failed" else: vector_count = int((await session.execute(select(func.count()).select_from(LlmWikiWikiVectorRow).where(LlmWikiWikiVectorRow.page_id == row.id))).scalar_one()) # Preserve the last atomically completed generation while the # newer remote version remains pending after an embedding error. row.index_status = "ready" if vector_count else "failed" row.last_seen_sync_id = sync_id row.last_error = safe_error row.updated_at = now await self._ensure_state(session, mapping_id) await session.commit() async def mark_missing_pages_deleted(self, mapping_id: str, sync_id: str) -> int: async with self._sf() as session: stmt = ( select(LlmWikiWikiPageRow) .where( LlmWikiWikiPageRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiPageRow.is_remote_deleted.is_(False), or_(LlmWikiWikiPageRow.last_seen_sync_id.is_(None), LlmWikiWikiPageRow.last_seen_sync_id != sync_id), ) .with_for_update() ) rows = list((await session.execute(stmt)).scalars()) if not rows: return 0 now = _now() page_ids = [row.id for row in rows] await session.execute(delete(LlmWikiWikiVectorRow).where(LlmWikiWikiVectorRow.page_id.in_(page_ids))) for row in rows: row.is_remote_deleted = True row.index_status = "deleted" row.updated_at = now state = await self._ensure_state(session, mapping_id) state.index_revision = int(state.index_revision or 0) + 1 state.updated_at = now await session.commit() return len(rows) async def load_vector_snapshot(self, mapping_id: str, fingerprint: str) -> dict[str, Any]: stmt = ( select(LlmWikiWikiVectorRow, LlmWikiWikiPageRow) .join(LlmWikiWikiPageRow, LlmWikiWikiPageRow.id == LlmWikiWikiVectorRow.page_id) .where( LlmWikiWikiVectorRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiVectorRow.embedding_fingerprint == fingerprint, LlmWikiWikiPageRow.index_status == "ready", LlmWikiWikiPageRow.is_remote_deleted.is_(False), LlmWikiWikiPageRow.status.in_(["published", "draft"]), ) .order_by(LlmWikiWikiVectorRow.page_id, LlmWikiWikiVectorRow.section_index) ) async with self._sf() as session: state = await session.get(LlmWikiWikiSyncStateRow, mapping_id) pairs = list((await session.execute(stmt)).all()) rows = [{**_dict(vector), "page": _dict(page)} for vector, page in pairs] return {"revision": int(state.index_revision if state else 0), "fingerprint": fingerprint, "rows": rows} async def iter_vector_snapshot_rows( self, mapping_id: str, fingerprint: str, *, batch_size: int = 1000, ) -> AsyncIterator[dict[str, Any]]: batch = max(1, min(int(batch_size), 5000)) offset = 0 while True: stmt = ( select(LlmWikiWikiVectorRow, LlmWikiWikiPageRow) .join(LlmWikiWikiPageRow, LlmWikiWikiPageRow.id == LlmWikiWikiVectorRow.page_id) .where( LlmWikiWikiVectorRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiVectorRow.embedding_fingerprint == fingerprint, LlmWikiWikiPageRow.index_status == "ready", LlmWikiWikiPageRow.is_remote_deleted.is_(False), LlmWikiWikiPageRow.status.in_(["published", "draft"]), ) .order_by(LlmWikiWikiVectorRow.page_id, LlmWikiWikiVectorRow.section_index) .limit(batch) .offset(offset) ) async with self._sf() as session: pairs = list((await session.execute(stmt)).all()) rows = [{**_dict(vector), "page": _dict(page)} for vector, page in pairs] if not rows: break for row in rows: yield row if len(rows) < batch: break offset += len(rows) async def get_index_revision(self, mapping_id: str) -> tuple[int, str | None, str]: async with self._sf() as session: state = await session.get(LlmWikiWikiSyncStateRow, mapping_id) if state is None: return 0, None, "not_synced" return int(state.index_revision or 0), state.embedding_fingerprint, state.state async def try_acquire_sync_lease(self, mapping_id: str, *, owner: str, sync_id: str, lease_seconds: int, fingerprint: str) -> bool: now = _now() async with self._sf() as session: try: await self._ensure_state(session, mapping_id) await session.commit() except IntegrityError: # Another worker created the one-per-KB state row between our # read and insert. Roll back this isolated bootstrap session; # the conditional UPDATE below is the actual lease election. await session.rollback() async with self._sf() as session: stmt = ( update(LlmWikiWikiSyncStateRow) .where( LlmWikiWikiSyncStateRow.knowledge_base_mapping_id == mapping_id, or_( LlmWikiWikiSyncStateRow.lease_owner.is_(None), LlmWikiWikiSyncStateRow.lease_expires_at.is_(None), LlmWikiWikiSyncStateRow.lease_expires_at < now, ), ) .values( state=case((LlmWikiWikiSyncStateRow.state == "rebuild_required", "rebuilding"), else_="running"), active_sync_id=sync_id, lease_owner=owner, lease_expires_at=now + timedelta(seconds=lease_seconds), last_started_at=now, last_error=None, updated_at=now, ) ) result = await session.execute(stmt) await session.commit() return bool(result.rowcount) async def renew_sync_lease(self, mapping_id: str, *, owner: str, lease_seconds: int) -> bool: now = _now() async with self._sf() as session: result = await session.execute( update(LlmWikiWikiSyncStateRow) .where(LlmWikiWikiSyncStateRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiSyncStateRow.lease_owner == owner) .values(lease_expires_at=now + timedelta(seconds=lease_seconds), updated_at=now) ) await session.commit() return bool(result.rowcount) async def update_sync_progress( self, mapping_id: str, *, owner: str, sync_id: str, remote_page_count: int, ) -> bool: async with self._sf() as session: result = await session.execute( update(LlmWikiWikiSyncStateRow) .where( LlmWikiWikiSyncStateRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiSyncStateRow.lease_owner == owner, LlmWikiWikiSyncStateRow.active_sync_id == sync_id, ) .values(remote_page_count=max(0, remote_page_count), updated_at=_now()) ) await session.commit() return bool(result.rowcount) async def finish_sync(self, mapping_id: str, *, owner: str, sync_id: str, success: bool, remote_page_count: int, error: str | None = None) -> None: async with self._sf() as session: state = await session.get(LlmWikiWikiSyncStateRow, mapping_id) if state is None or state.lease_owner != owner or state.active_sync_id != sync_id: return visible = [LlmWikiWikiPageRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiPageRow.is_remote_deleted.is_(False)] local_count = int((await session.execute(select(func.count()).select_from(LlmWikiWikiPageRow).where(*visible))).scalar_one()) ready_count = int((await session.execute(select(func.count()).select_from(LlmWikiWikiPageRow).where(*visible, LlmWikiWikiPageRow.index_status == "ready"))).scalar_one()) failed_count = int( ( await session.execute( select(func.count()) .select_from(LlmWikiWikiPageRow) .where( *visible, or_( LlmWikiWikiPageRow.index_status.in_(["failed", "stale"]), LlmWikiWikiPageRow.last_error.is_not(None), ), ) ) ).scalar_one() ) vector_count = int((await session.execute(select(func.count()).select_from(LlmWikiWikiVectorRow).where(LlmWikiWikiVectorRow.knowledge_base_mapping_id == mapping_id))).scalar_one()) now = _now() was_rebuilding = state.state == "rebuilding" state.state = "partial" if success and error else "idle" if success else "rebuild_required" if was_rebuilding else "failed" state.active_sync_id = None state.lease_owner = None state.lease_expires_at = None state.last_error = (error or "")[:4000] or None state.remote_page_count = remote_page_count state.local_page_count = local_count state.ready_page_count = ready_count state.failed_page_count = failed_count state.vector_count = vector_count if success and not error: state.last_completed_at = now if success: state.last_success_at = now state.updated_at = now await session.commit() async def list_index_status(self, mapping_ids: list[str] | None = None) -> list[dict[str, Any]]: stmt = select(LlmWikiWikiSyncStateRow) if mapping_ids: stmt = stmt.where(LlmWikiWikiSyncStateRow.knowledge_base_mapping_id.in_(mapping_ids)) stmt = stmt.order_by(LlmWikiWikiSyncStateRow.updated_at.desc()) async with self._sf() as session: return [_dict(row) for row in (await session.execute(stmt)).scalars()] async def list_page_index_status(self, mapping_id: str) -> list[dict[str, Any]]: vector_counts = ( select( LlmWikiWikiVectorRow.page_id.label("page_id"), func.count(LlmWikiWikiVectorRow.id).label("vector_count"), ) .where(LlmWikiWikiVectorRow.knowledge_base_mapping_id == mapping_id) .group_by(LlmWikiWikiVectorRow.page_id) .subquery() ) stmt = ( select(LlmWikiWikiPageRow, func.coalesce(vector_counts.c.vector_count, 0)) .outerjoin(vector_counts, vector_counts.c.page_id == LlmWikiWikiPageRow.id) .where( LlmWikiWikiPageRow.knowledge_base_mapping_id == mapping_id, LlmWikiWikiPageRow.is_remote_deleted.is_(False), ) .order_by(LlmWikiWikiPageRow.wiki_path, LlmWikiWikiPageRow.slug) ) async with self._sf() as session: rows = (await session.execute(stmt)).all() result: list[dict[str, Any]] = [] for row, vector_count in rows: page = _dict(row) result.append( { "id": page["id"], "slug": page["slug"], "title": page["title"], "page_type": page["page_type"], "status": page["status"], "index_status": page["index_status"], "vector_count": int(vector_count or 0), "indexed_at": page.get("indexed_at"), "last_error": page.get("last_error"), "last_seen_sync_id": page.get("last_seen_sync_id"), "updated_at": page.get("updated_at"), } ) return result async def mark_rebuild_required(self, mapping_ids: list[str], fingerprint: str) -> int: if not mapping_ids: return 0 now = _now() async with self._sf() as session: for mapping_id in mapping_ids: state = await self._ensure_state(session, mapping_id) state.state = "rebuild_required" state.embedding_fingerprint = fingerprint state.updated_at = now await session.execute(update(LlmWikiWikiPageRow).where(LlmWikiWikiPageRow.knowledge_base_mapping_id.in_(mapping_ids), LlmWikiWikiPageRow.is_remote_deleted.is_(False)).values(index_status="stale", updated_at=now)) await session.commit() return len(mapping_ids)