80 lines
4.0 KiB
Python
80 lines
4.0 KiB
Python
"""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)
|