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

80 lines
4.0 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Disk-backed, replayable package records. Memory is bounded by one JSONL record.
The staging index is local SQLite, never a remote provider's database. It also
allows restoration to fetch only the vectors belonging to the current page.
"""
from __future__ import annotations
import json
import sqlite3
import tempfile
from collections.abc import Iterator
from contextlib import closing
from pathlib import Path
from typing import Any
class Records:
def __init__(self, path: Path, kind: str, slug: str | None = None):
self.path, self.kind, self.slug = path, kind, slug
def _where(self):
return ("kind=?", (self.kind,)) if self.slug is None else ("kind=? AND slug=?", (self.kind, self.slug))
def __iter__(self) -> Iterator[dict[str, Any]]:
where, args = self._where()
with closing(sqlite3.connect(self.path)) as db:
for (payload,) in db.execute(f"SELECT payload FROM records WHERE {where} ORDER BY seq", args):
yield json.loads(payload)
def __len__(self) -> int:
where, args = self._where()
with closing(sqlite3.connect(self.path)) as db:
return db.execute(f"SELECT count(*) FROM records WHERE {where}", args).fetchone()[0]
class PackageArchive(dict):
def __init__(self, path: Path):
super().__init__()
temporary = tempfile.NamedTemporaryFile(prefix="wiki-package-", suffix=".sqlite", dir=path.parent, delete=False)
self.path = Path(temporary.name)
temporary.close()
try:
with closing(sqlite3.connect(self.path)) as db, path.open(encoding="utf-8-sig") as handle:
db.execute("CREATE TABLE records(seq INTEGER PRIMARY KEY, kind TEXT, slug TEXT, payload TEXT)")
for line_no, line in enumerate(handle, 1):
if not line.strip():
continue
record = json.loads(line)
kind, payload = record.get("record_type"), record.get("payload")
if not isinstance(payload, dict):
raise ValueError(f"导出包第 {line_no} 行格式错误")
if kind == "source":
self["source"] = {**self.get("source", {}), **payload}
elif kind == "graph":
for key, subkind in (("nodes", "graph_node"), ("edges", "graph_edge")):
db.executemany("INSERT INTO records(kind, payload) VALUES (?, ?)", [(subkind, json.dumps(v, ensure_ascii=False)) for v in payload.get(key) or []])
elif kind in {"wiki_page", "document", "chunk", "vector", "entity", "relation"}:
slug = str(payload.get("wiki_slug") or payload.get("page_slug") or payload.get("slug") or "").strip("/")
db.execute("INSERT INTO records(kind, slug, payload) VALUES (?, ?, ?)", (kind, slug, json.dumps(payload, ensure_ascii=False)))
else:
raise ValueError(f"不支持的导出记录类型:{kind}")
db.execute("CREATE INDEX by_kind_slug ON records(kind, slug)")
db.commit()
source = self.get("source", {})
if source.get("package_version") != 1 or (source.get("provider") or source.get("type")) not in {"weknora", "assistant"}:
raise ValueError("仅接受 WeKnora / 助手知识库生成的兼容 Wiki 向量包(v1)")
for key, kind in (("wiki_pages", "wiki_page"), ("documents", "document"), ("chunks", "chunk"), ("vectors", "vector"), ("entities", "entity"), ("relations", "relation")):
self[key] = Records(self.path, kind)
self["graph"] = {"nodes": Records(self.path, "graph_node"), "edges": Records(self.path, "graph_edge")}
except BaseException:
self.close()
raise
def vectors_for(self, slug: str) -> Records:
return Records(self.path, "vector", slug)
def close(self):
self.path.unlink(missing_ok=True)