Files
Cloud-Tour-to-Libo/app/graph_import_files.py
T

312 lines
14 KiB
Python

"""Bounded-memory validation and import of large graph-data JSON files.
The browser never needs to parse these files. SQLite is a temporary index,
not another graph database: it validates cross-record references and supplies
bounded UNWIND batches. Only the completed FalkorDB graph is published.
"""
from __future__ import annotations
import hashlib
import json
import math
import sqlite3
from collections.abc import Mapping
from contextlib import nullcontext
from functools import lru_cache
from pathlib import Path
from typing import BinaryIO, Callable
import ijson
from ijson.common import ObjectBuilder
from app.project_lifecycle import (
GraphImportError, ProjectValidationError, SpatialMapValidation,
_falkor_client, _graph_safe_properties, _node_labels,
normalize_graph_node, normalize_graph_relation, normalize_provision_header,
)
Progress = Callable[[str, int, int], None]
MAX_STRING_BYTES = 4 * 1024 * 1024
MAX_RECORD_CHARS = 8 * 1024 * 1024
class CheckedReader:
"""Reject oversized tokens/deep nesting before the parser buffers them."""
def __init__(self, stream):
self.stream = stream
self.sha256 = hashlib.sha256()
self.count = 0
self.in_string = self.escaped = False
self.token_size = self.depth = 0
def read(self, size=-1):
block = self.stream.read(65536 if size < 0 else min(size, 65536))
self.sha256.update(block)
self.count += len(block)
for char in block:
if self.in_string:
self.token_size += 1
if self.token_size > MAX_STRING_BYTES:
raise ValueError("单个 JSON 字符串超过 4 MiB,请将大正文/附件拆分为记录或外部资源")
if self.escaped:
self.escaped = False
elif char == 92:
self.escaped = True
elif char == 34:
self.in_string = False
self.token_size = 0
elif char == 34:
self.in_string = True
self.token_size = 0
elif char in (123, 91):
self.depth += 1
self.token_size = 0
if self.depth > 64:
raise ValueError("JSON 嵌套不能超过 64 层")
elif char in (125, 93):
self.depth -= 1
self.token_size = 0
elif char in (9, 10, 13, 32, 44, 58):
self.token_size = 0
else:
self.token_size += 1
if self.token_size > 1024:
raise ValueError("JSON 数值或标记过长")
return block
def graph_records(reader):
"""One pass, strict root/array shape and duplicate-key detection."""
builder = None
record_kind = ""
record_depth = 0
record_chars = 0
maps: list[set[str] | None] = []
root_fields: set[str] = set()
first = True
for prefix, event, value in ijson.parse(reader, use_float=True):
if first:
first = False
if (prefix, event) != ("", "start_map"):
raise ValueError("图谱数据 JSON 顶层必须是对象")
if event == "map_key":
keys = maps[-1]
if value in keys:
raise ValueError(f"JSON 对象中存在重复字段:{str(value)[:120]}")
keys.add(value)
if len(keys) > 100000:
raise ValueError("单个 JSON 对象字段过多")
if prefix == "":
root_fields.add(value)
elif event in ("start_map", "start_array"):
maps.append(set() if event == "start_map" else None)
elif event in ("end_map", "end_array"):
maps.pop()
if prefix in ("nodes", "relations") and event not in ("start_array", "end_array", "map_key"):
raise ValueError(f"graph_data.{prefix} 必须是数组")
if builder is None and prefix in ("nodes.item", "relations.item"):
if event != "start_map":
raise ValueError(f"{prefix} 必须是 JSON 对象")
builder, record_kind, record_depth, record_chars = ObjectBuilder(), prefix.split(".")[0], 0, 0
if builder is not None:
record_chars += len(value) if isinstance(value, str) else 8
if record_chars > MAX_RECORD_CHARS:
raise ValueError("单条节点或关系超过 8 MiB,请拆分大附件")
builder.event(event, value)
if event in ("start_map", "start_array"):
record_depth += 1
elif event in ("end_map", "end_array"):
record_depth -= 1
if record_depth == 0:
yield record_kind, builder.value
builder = None
if "nodes" not in root_fields:
raise ValueError("图谱数据缺少 nodes 数组;请上传 data JSON,而不是 Schema 或 bundle 文件")
class NodeTypes(Mapping):
"""A bounded LRU over SQLite, rather than an all-nodes Python dictionary."""
def __init__(self, connection):
self.connection = connection
self.lookup = lru_cache(maxsize=4096)(self._lookup)
def _lookup(self, key):
row = self.connection.execute("SELECT type FROM nodes WHERE id=?", (key,)).fetchone()
if row is None:
raise KeyError(key)
return row[0]
def __getitem__(self, key):
return self.lookup(key)
def __iter__(self):
return (row[0] for row in self.connection.execute("SELECT id FROM nodes"))
def __len__(self):
return self.connection.execute("SELECT count(*) FROM nodes").fetchone()[0]
def _dumps(value):
return json.dumps(value, ensure_ascii=False, separators=(",", ":"), allow_nan=False)
def validate_graph_file(source: Path | BinaryIO, staging: Path, body: dict,
progress: Progress = lambda *_: None) -> dict:
"""Fully validate before any PostgreSQL project/FalkorDB write."""
payload = normalize_provision_header(body)
if staging.exists():
raise FileExistsError("校验暂存文件已存在")
connection = sqlite3.connect(staging)
node_types = NodeTypes(connection)
spatial = SpatialMapValidation(payload["spatial_map"])
counts = {"nodes": 0, "relations": 0}
try:
connection.execute("PRAGMA cache_size=-4096")
connection.execute("PRAGMA temp_store=FILE")
connection.executescript("""
CREATE TABLE nodes(id TEXT PRIMARY KEY, type TEXT NOT NULL, labels TEXT NOT NULL, payload TEXT NOT NULL);
CREATE TABLE relations(seq INTEGER PRIMARY KEY, type TEXT, source_type TEXT, target_type TEXT, payload TEXT NOT NULL);
""")
# Multipart uploads are already spooled to disk; reuse their handle
# instead of making another full-size string or file copy.
with (source.open("rb") if isinstance(source, Path) else nullcontext(source)) as stream:
stream.seek(0)
reader = CheckedReader(stream)
for kind, raw in graph_records(reader):
errors: list[str] = []
if kind == "nodes":
node = normalize_graph_node(raw, counts[kind], payload["schema"]["entity_types"], node_types, errors)
if node is not None:
spatial.add(node, errors)
if errors:
raise ProjectValidationError(errors[:100])
connection.execute(
"INSERT INTO nodes VALUES(?,?,?,?)",
(node["id"], node["type"], _dumps(_node_labels(node)), _dumps(node)),
)
else:
# Relations may appear before nodes in valid JSON.
connection.execute("INSERT INTO relations(seq,payload) VALUES(?,?)", (counts[kind], _dumps(raw)))
counts[kind] += 1
if (counts["nodes"] + counts["relations"]) % 1000 == 0:
connection.commit()
progress("校验节点" if kind == "nodes" else "读取关系", counts["nodes"], counts["relations"])
if not counts["nodes"]:
raise ProjectValidationError(["graph_data.nodes 至少需要一个节点,FalkorDB 不支持持久化空图"])
errors = []
spatial.finish(errors)
if errors:
raise ProjectValidationError(errors)
connection.commit()
for index, text in connection.execute("SELECT seq,payload FROM relations ORDER BY seq"):
errors = []
relation = normalize_graph_relation(json.loads(text), index, payload["schema"]["relation_types"], node_types, errors)
if errors:
raise ProjectValidationError(errors[:100])
connection.execute(
"UPDATE relations SET type=?, source_type=?, target_type=?, payload=? WHERE seq=?",
(relation["type"], node_types[relation["source"]], node_types[relation["target"]], _dumps(relation), index),
)
if index % 1000 == 0:
connection.commit()
progress("校验关系", counts["nodes"], index + 1)
connection.executescript("""
CREATE INDEX nodes_group ON nodes(labels);
CREATE INDEX relations_group ON relations(type, source_type, target_type);
""")
connection.commit()
payload["counts"].update(counts)
payload["counts"]["map_pois"] = spatial.count
payload["graph_data"] = {"staging_path": str(staging)}
payload["file_sha256"] = reader.sha256.hexdigest()
progress("校验完成", counts["nodes"], counts["relations"])
return payload
except (ValueError, ijson.JSONError) as exc:
if isinstance(exc, ProjectValidationError):
raise
raise ProjectValidationError([f"图谱 JSON 校验失败:{str(exc)[:1000]}"]) from exc
finally:
node_types.lookup.cache_clear()
connection.close()
def _record_batches(cursor):
batch, size = [], 0
for (text,) in cursor:
if batch and (len(batch) >= 250 or size + len(text) > 2 * 1024 * 1024):
yield batch
batch, size = [], 0
batch.append(json.loads(text))
size += len(text)
if batch:
yield batch
def import_staged_graph(graph_name: str, graph_data: Mapping,
progress: Progress = lambda *_: None) -> dict[str, int]:
"""Import only a private validated spool; never overwrite an existing graph."""
db = _falkor_client()
created = False
done_nodes = done_relations = 0
connection = None
try:
connection = sqlite3.connect(f"file:{Path(graph_data['staging_path']).as_posix()}?mode=ro", uri=True)
connection.execute("PRAGMA cache_size=-4096")
existing = {name.decode() if isinstance(name, bytes) else str(name) for name in db.list_graphs()}
if graph_name in existing:
raise FileExistsError(f"FalkorDB 图 {graph_name!r} 已存在")
graph = db.select_graph(graph_name)
for (labels_json,) in connection.execute("SELECT DISTINCT labels FROM nodes"):
labels = json.loads(labels_json)
expression = ":".join(_node_labels({"type": labels[0], "labels": labels}))
cursor = connection.execute("SELECT payload FROM nodes WHERE labels=?", (labels_json,))
for batch in _record_batches(cursor):
rows = []
for node in batch:
properties = _graph_safe_properties(node["properties"])
properties.setdefault("id", node["id"])
properties["__kg_node_id"] = node["id"]
rows.append({"properties": properties})
created = True
graph.query(f"UNWIND $rows AS row CREATE (n:{expression}) SET n += row.properties", {"rows": rows})
done_nodes += len(batch)
progress("导入节点", done_nodes, done_relations)
for (label,) in connection.execute("SELECT DISTINCT type FROM nodes"):
_node_labels({"type": label})
graph.query(f"CREATE INDEX FOR (n:{label}) ON (n.__kg_node_id)")
for signature in connection.execute("SELECT DISTINCT type,source_type,target_type FROM relations"):
rel_type, source_type, target_type = signature
for label in signature:
_node_labels({"type": label})
cursor = connection.execute(
"SELECT payload FROM relations WHERE type=? AND source_type=? AND target_type=?", signature)
for batch in _record_batches(cursor):
rows = [{"source": row["source"], "target": row["target"],
"properties": _graph_safe_properties(row["properties"])} for row in batch]
graph.query(
f"UNWIND $rows AS row MATCH (s:{source_type} {{__kg_node_id:row.source}}) "
f"MATCH (t:{target_type} {{__kg_node_id:row.target}}) "
f"CREATE (s)-[r:{rel_type}]->(t) SET r += row.properties", {"rows": rows})
done_relations += len(batch)
progress("导入关系", done_nodes, done_relations)
return {
"nodes": int(graph.query("MATCH (n) RETURN count(n)").result_set[0][0]),
"relations": int(graph.query("MATCH ()-[r]->() RETURN count(r)").result_set[0][0]),
}
except FileExistsError:
raise
except Exception as exc:
cleanup_error = None
if created:
try:
db.select_graph(graph_name).delete()
except Exception as cleanup_exc:
cleanup_error = cleanup_exc
raise GraphImportError(str(exc), cleanup_error=cleanup_error) from exc
finally:
if connection is not None:
connection.close()
db.close()