312 lines
14 KiB
Python
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()
|