401 lines
19 KiB
Python
401 lines
19 KiB
Python
"""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)
|