deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/persistence/llmwiki_index/sql.py
2026-09-07 18:24:55 +08:00

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)