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

819 lines
32 KiB
Python

"""Validation and FalkorDB helpers for graph-project lifecycle operations.
This module deliberately has no dependency on ``app.data_platform``. Graph
projects and relational data-center databases are separate resources even when
they happen to use the same project id.
"""
from __future__ import annotations
import json
import math
import re
from collections import defaultdict
from collections.abc import Mapping
from typing import Any
from falkordb import FalkorDB
from app.config import settings
SAFE_RESOURCE_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.-]{0,127}$")
SAFE_SCHEMA_IDENTIFIER = re.compile(r"^[A-Za-z_][A-Za-z0-9_]{0,127}$")
SAFE_SCHEMA_VERSION = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,63}$")
SAFE_ADCODE = re.compile(r"^\d{6}$")
TYPE_ALIASES = {
"str": "string",
"text": "string",
"varchar": "string",
"int": "integer",
"long": "integer",
"float": "number",
"double": "number",
"decimal": "number",
"bool": "boolean",
"json": "object",
"dict": "object",
"list": "array",
"timestamp": "datetime",
}
SUPPORTED_FIELD_TYPES = {
"string",
"integer",
"number",
"boolean",
"object",
"array",
"date",
"datetime",
"any",
}
GRAPH_IMPORT_BATCH_THRESHOLD = 500
GRAPH_IMPORT_BATCH_SIZE = 250
SPATIAL_REGION_LEVELS = {"province", "city", "district"}
MAP_ENTITY_TYPES = {
"FoodPlace",
"Hotel",
"ScenicSpot",
"TransitFacility",
"BusStop",
}
class ProjectValidationError(ValueError):
"""Raised when a provision payload is invalid before any write occurs."""
def __init__(self, errors: list[str]):
super().__init__(errors[0] if errors else "项目数据校验失败")
self.errors = errors
class GraphImportError(RuntimeError):
"""FalkorDB import failed, optionally leaving a graph needing cleanup."""
def __init__(self, message: str, *, cleanup_error: Exception | None = None):
super().__init__(message)
self.cleanup_error = cleanup_error
self.cleanup_required = cleanup_error is not None
def _adcode_belongs_to_region(adcode: str, region_adcode: str, level: str) -> bool:
if not SAFE_ADCODE.fullmatch(adcode) or not SAFE_ADCODE.fullmatch(region_adcode):
return False
if level == "province":
return adcode[:2] == region_adcode[:2]
if level == "city":
return adcode[:4] == region_adcode[:4]
return adcode == region_adcode
def _region_names_match(value: Any, region_name: str) -> bool:
candidate = str(value or "").strip()
expected = str(region_name or "").strip()
return bool(candidate and expected and (candidate in expected or expected in candidate))
def _is_spatial_map_node(node: Mapping[str, Any]) -> bool:
properties = node.get("properties")
if not isinstance(properties, Mapping):
return False
# A semantic graph can contain coordinates without verified map-region
# metadata. Keep those nodes in the graph without adding them to the POI
# layer. Old bundles without this optional flag retain their behavior.
if properties.get("map_poi") is False:
return False
has_coordinates = (
properties.get("lng") is not None
and properties.get("lat") is not None
) or (
properties.get("longitude") is not None
and properties.get("latitude") is not None
)
if not has_coordinates:
return False
labels = {
str(label)
for label in (node.get("labels") or [])
if str(label).strip()
}
labels.add(str(node.get("type") or ""))
return bool(
labels & MAP_ENTITY_TYPES
or str(properties.get("type_label") or "").strip()
or str(properties.get("place_type") or "").strip()
)
def _as_mapping(value: Any) -> dict[str, Any]:
return dict(value) if isinstance(value, Mapping) else {}
def _normalize_named_mapping(
value: Any,
*,
kind: str,
name_keys: tuple[str, ...],
errors: list[str],
) -> dict[str, dict[str, Any]]:
if isinstance(value, Mapping):
result: dict[str, dict[str, Any]] = {}
for raw_name, raw_meta in value.items():
name = str(raw_name or "").strip()
if kind == "关系类型" and isinstance(raw_meta, (list, tuple)):
if len(raw_meta) < 2:
errors.append(f"关系类型 {name or '<未命名>'} 至少需要 from 和 to")
continue
raw_meta = {
"from": raw_meta[0],
"to": raw_meta[1],
"description": raw_meta[2] if len(raw_meta) > 2 else None,
"properties": {},
}
if not isinstance(raw_meta, Mapping):
errors.append(f"{kind} {name or '<未命名>'} 的定义必须是 JSON 对象")
continue
result[name] = dict(raw_meta)
return result
if isinstance(value, list):
result = {}
for index, raw_meta in enumerate(value):
if not isinstance(raw_meta, Mapping):
errors.append(f"{kind}列表第 {index + 1} 项必须是 JSON 对象")
continue
meta = dict(raw_meta)
name = next(
(str(meta.get(key) or "").strip() for key in name_keys if meta.get(key)),
"",
)
if not name:
errors.append(f"{kind}列表第 {index + 1} 项缺少名称")
continue
result[name] = meta
return result
errors.append(f"{kind}必须是 JSON 对象或数组")
return {}
def _normalize_field_definitions(
value: Any,
*,
path: str,
errors: list[str],
) -> dict[str, dict[str, Any]]:
if value is None:
return {}
if isinstance(value, Mapping):
raw_items = list(value.items())
elif isinstance(value, list):
raw_items = []
for index, item in enumerate(value):
if isinstance(item, str):
raw_items.append((item, {"type": "string"}))
elif isinstance(item, Mapping):
name = str(item.get("name") or item.get("field_name") or "").strip()
if not name:
errors.append(f"{path} 第 {index + 1} 个字段缺少名称")
continue
raw_items.append((name, item))
else:
errors.append(f"{path} 第 {index + 1} 个字段定义无效")
else:
errors.append(f"{path} 必须是 JSON 对象或数组")
return {}
normalized: dict[str, dict[str, Any]] = {}
for raw_name, raw_definition in raw_items:
name = str(raw_name or "").strip()
if not name:
errors.append(f"{path} 包含空字段名")
continue
if isinstance(raw_definition, str):
definition = {"type": raw_definition}
elif isinstance(raw_definition, Mapping):
definition = dict(raw_definition)
else:
errors.append(f"{path}.{name} 的字段定义必须是对象或类型字符串")
continue
raw_type = str(definition.get("type") or definition.get("value_type") or "string").lower().strip()
value_type = TYPE_ALIASES.get(raw_type, raw_type)
if value_type not in SUPPORTED_FIELD_TYPES:
errors.append(f"{path}.{name} 使用了不支持的类型 {raw_type!r}")
continue
definition["type"] = value_type
definition["required"] = bool(definition.get("required", False))
normalized[name] = definition
return normalized
def _matches_field_type(value: Any, expected: str) -> bool:
if expected == "any":
return True
if expected in {"string", "date", "datetime"}:
return isinstance(value, str)
if expected == "integer":
return isinstance(value, int) and not isinstance(value, bool)
if expected == "number":
return isinstance(value, (int, float)) and not isinstance(value, bool)
if expected == "boolean":
return isinstance(value, bool)
if expected == "object":
return isinstance(value, Mapping)
if expected == "array":
return isinstance(value, list)
return False
def _endpoint_types(value: Any) -> list[str]:
return [part.strip() for part in str(value or "").split("|") if part.strip()]
def _validate_properties(
properties: dict[str, Any],
definitions: dict[str, dict[str, Any]],
*,
path: str,
errors: list[str],
) -> None:
for field_name, definition in definitions.items():
value_present = field_name in properties and properties[field_name] is not None
if definition.get("required") and not value_present:
errors.append(f"{path} 缺少必填属性 {field_name}")
continue
if not value_present:
continue
expected = str(definition.get("type") or "string")
if not _matches_field_type(properties[field_name], expected):
errors.append(f"{path}.{field_name} 应为 {expected} 类型")
def normalize_provision_header(body: Any) -> dict[str, Any]:
"""Validate project metadata and Schema before reading graph records.
The function is intentionally pure so validation can be completed before
PostgreSQL or FalkorDB is touched.
"""
if not isinstance(body, Mapping):
raise ProjectValidationError(["请求体必须是 JSON 对象"])
errors: list[str] = []
project_id = str(body.get("project_id") or "").strip()
explicit_tenant_id = str(body.get("tenant_id") or "").strip()
explicit_graph_name = str(body.get("graph_name") or "").strip()
tenant_id = explicit_tenant_id or project_id
display_name = str(body.get("display_name") or "").strip()
graph_name = explicit_graph_name or project_id
raw_spatial_map = body.get("spatial_map", {})
if raw_spatial_map is None:
raw_spatial_map = {}
if not isinstance(raw_spatial_map, Mapping):
errors.append("spatial_map 必须是 JSON 对象")
raw_spatial_map = {}
raw_map_enabled = raw_spatial_map.get("enabled", False)
if not isinstance(raw_map_enabled, bool):
errors.append("spatial_map.enabled 必须是布尔值")
map_enabled = False
else:
map_enabled = raw_map_enabled
map_scope = str(raw_spatial_map.get("scope") or "libo").strip().lower()
if map_scope not in {"libo", "guizhou"}:
errors.append("spatial_map.scope 仅支持 libo 或 guizhou")
map_region_name = str(raw_spatial_map.get("region_name") or "").strip()
map_region_adcode = str(raw_spatial_map.get("region_adcode") or "").strip()
map_region_level = str(raw_spatial_map.get("region_level") or "").strip().lower()
if map_enabled:
if not map_region_name:
errors.append("启用知识地图时必须选择地图区域")
if not SAFE_ADCODE.fullmatch(map_region_adcode):
errors.append("spatial_map.region_adcode 必须是 6 位行政区划代码")
if map_region_level not in SPATIAL_REGION_LEVELS:
errors.append("spatial_map.region_level 仅支持 province、city 或 district")
if SAFE_ADCODE.fullmatch(map_region_adcode):
expected_scope = "libo" if map_region_adcode == "522722" else "guizhou"
if map_scope != expected_scope:
errors.append("spatial_map.scope 与所选地图区域不一致")
if not project_id:
errors.append("项目英文标识不能为空")
elif not SAFE_RESOURCE_ID.fullmatch(project_id):
errors.append("项目英文标识只能包含字母、数字、点、下划线和连字符,且必须以字母或数字开头")
if explicit_tenant_id and not SAFE_RESOURCE_ID.fullmatch(explicit_tenant_id):
errors.append("项目隔离标识只能包含安全字符")
if not display_name:
errors.append("display_name 不能为空")
if explicit_graph_name and not SAFE_RESOURCE_ID.fullmatch(explicit_graph_name):
errors.append("图谱资源标识只能包含字母、数字、点、下划线和连字符,且必须以字母或数字开头")
raw_schema = body.get("schema")
if not isinstance(raw_schema, Mapping):
errors.append("schema 必须是 JSON 对象")
raw_schema = {}
schema = dict(raw_schema)
namespace = str(schema.get("namespace") or project_id).strip()
version = str(schema.get("version") or "").strip()
if not namespace:
errors.append("schema.namespace 不能为空")
elif not SAFE_RESOURCE_ID.fullmatch(namespace):
errors.append("schema.namespace 只能包含安全字符")
if not version:
errors.append("schema.version 不能为空")
elif not SAFE_SCHEMA_VERSION.fullmatch(version):
errors.append("schema.version 只能包含字母、数字、点、下划线和连字符,最长 64 个字符")
entity_types = _normalize_named_mapping(
schema.get("entity_types"),
kind="实体类型",
name_keys=("name", "entity_type"),
errors=errors,
)
if not entity_types:
errors.append("schema.entity_types 至少需要一个实体类型")
normalized_entities: dict[str, dict[str, Any]] = {}
for entity_name, raw_meta in entity_types.items():
if not SAFE_SCHEMA_IDENTIFIER.fullmatch(entity_name):
errors.append(f"实体类型 {entity_name!r} 不是安全的 Schema/Cypher 标识符")
continue
meta = dict(raw_meta)
meta["fields"] = _normalize_field_definitions(
meta.get("fields"), path=f"schema.entity_types.{entity_name}.fields", errors=errors
)
normalized_entities[entity_name] = meta
relation_types = _normalize_named_mapping(
schema.get("relation_types", {}),
kind="关系类型",
name_keys=("name", "relation_type"),
errors=errors,
)
normalized_relations: dict[str, dict[str, Any]] = {}
for relation_name, raw_meta in relation_types.items():
if not SAFE_SCHEMA_IDENTIFIER.fullmatch(relation_name):
errors.append(f"关系类型 {relation_name!r} 不是安全的 Schema/Cypher 标识符")
continue
meta = dict(raw_meta)
source_type = str(meta.get("from") or meta.get("source") or "").strip()
target_type = str(meta.get("to") or meta.get("target") or "").strip()
source_types = _endpoint_types(source_type)
target_types = _endpoint_types(target_type)
if not source_types or any(item not in normalized_entities for item in source_types):
errors.append(f"关系类型 {relation_name} 的 from/source 必须引用已定义实体类型")
if not target_types or any(item not in normalized_entities for item in target_types):
errors.append(f"关系类型 {relation_name} 的 to/target 必须引用已定义实体类型")
meta["from"] = source_type
meta["to"] = target_type
meta["properties"] = _normalize_field_definitions(
meta.get("properties"),
path=f"schema.relation_types.{relation_name}.properties",
errors=errors,
)
normalized_relations[relation_name] = meta
if errors:
raise ProjectValidationError(errors)
schema["namespace"] = namespace
schema["version"] = version
schema["entity_types"] = normalized_entities
schema["relation_types"] = normalized_relations
return {
"project_id": project_id,
"tenant_id": tenant_id,
"display_name": display_name,
"description": str(body.get("description") or "").strip() or None,
"graph_name": graph_name,
"spatial_map": {
"enabled": map_enabled,
"scope": map_scope,
"region_name": map_region_name,
"region_adcode": map_region_adcode,
"region_level": map_region_level or "province",
},
"schema": schema,
"graph_data": {"nodes": [], "relations": []},
"counts": {
"entity_types": len(normalized_entities),
"relation_types": len(normalized_relations),
"nodes": 0,
"relations": 0,
},
}
def normalize_graph_node(raw_node: Any, index: int, normalized_entities: Mapping,
node_types: Mapping, errors: list[str]) -> dict | None:
path = f"graph_data.nodes[{index}]"
if not isinstance(raw_node, Mapping):
errors.append(f"{path} 必须是 JSON 对象")
return None
node = dict(raw_node)
node_id = str(node.get("id") or "").strip()
entity_type = str(node.get("type") or node.get("entity_type") or "").strip()
properties = node.get("properties", {})
raw_labels = node.get("labels")
normalized_labels: list[str] | None = None
if not node_id:
errors.append(f"{path}.id 不能为空")
elif node_id in node_types:
errors.append(f"节点 ID {node_id!r} 重复")
if entity_type not in normalized_entities:
errors.append(f"{path}.type {entity_type!r} 未在 Schema 中定义")
if raw_labels is not None:
if not isinstance(raw_labels, list) or not raw_labels:
errors.append(f"{path}.labels 必须是非空数组")
else:
normalized_labels = []
for label_index, raw_label in enumerate(raw_labels):
label = str(raw_label or "").strip()
label_path = f"{path}.labels[{label_index}]"
if not label or not SAFE_SCHEMA_IDENTIFIER.fullmatch(label):
errors.append(f"{label_path} 不是安全的 Schema/Cypher 标识符")
continue
if label not in normalized_entities:
errors.append(f"{label_path} {label!r} 未在 Schema 中定义")
continue
if label not in normalized_labels:
normalized_labels.append(label)
if entity_type and entity_type not in normalized_labels:
errors.append(f"{path}.labels 必须包含主类型 {entity_type!r}")
if not isinstance(properties, Mapping):
errors.append(f"{path}.properties 必须是 JSON 对象")
properties = {}
normalized_properties = dict(properties)
if (
normalized_properties.get("map_poi") is not None
and not isinstance(normalized_properties["map_poi"], bool)
):
errors.append(f"{path}.properties.map_poi 必须是布尔值")
if entity_type in normalized_entities:
_validate_properties(
normalized_properties,
normalized_entities[entity_type].get("fields", {}),
path=path,
errors=errors,
)
normalized_node = {
"id": node_id,
"type": entity_type,
"properties": normalized_properties,
}
if normalized_labels is not None:
normalized_node["labels"] = normalized_labels
return normalized_node
def normalize_graph_relation(raw_relation: Any, index: int, normalized_relations: Mapping,
node_types: Mapping, errors: list[str]) -> dict | None:
path = f"graph_data.relations[{index}]"
if not isinstance(raw_relation, Mapping):
errors.append(f"{path} 必须是 JSON 对象")
return None
relation = dict(raw_relation)
relation_type = str(relation.get("type") or relation.get("relation_type") or "").strip()
source = str(relation.get("source") or "").strip()
target = str(relation.get("target") or "").strip()
properties = relation.get("properties", {})
definition = normalized_relations.get(relation_type)
if not definition:
errors.append(f"{path}.type {relation_type!r} 未在 Schema 中定义")
if source not in node_types:
errors.append(f"{path}.source {source!r} 不存在")
if target not in node_types:
errors.append(f"{path}.target {target!r} 不存在")
if (
definition
and source in node_types
and node_types[source] not in _endpoint_types(definition.get("from"))
):
errors.append(f"{path}.source 的节点类型必须是 {definition.get('from')}")
if (
definition
and target in node_types
and node_types[target] not in _endpoint_types(definition.get("to"))
):
errors.append(f"{path}.target 的节点类型必须是 {definition.get('to')}")
if not isinstance(properties, Mapping):
errors.append(f"{path}.properties 必须是 JSON 对象")
properties = {}
normalized_properties = dict(properties)
if definition:
_validate_properties(
normalized_properties,
definition.get("properties", {}),
path=path,
errors=errors,
)
return {
"type": relation_type,
"source": source,
"target": target,
"properties": normalized_properties,
}
class SpatialMapValidation:
"""Incremental region validation shared by inline and streamed imports."""
def __init__(self, spatial_map: Mapping):
self.config = spatial_map
self.count = 0
self.unverified = 0
self.mismatched: set[str] = set()
def add(self, node: Mapping, errors: list[str]) -> None:
if not _is_spatial_map_node(node):
return
properties = node["properties"]
try:
lng = float(properties.get("lng") if properties.get("lng") is not None else properties.get("longitude"))
lat = float(properties.get("lat") if properties.get("lat") is not None else properties.get("latitude"))
valid = math.isfinite(lng) and math.isfinite(lat) and -180 <= lng <= 180 and -90 <= lat <= 90
except (TypeError, ValueError):
valid = False
if not valid:
if self.config["enabled"]:
errors.append(f"节点 {node['id']!r} 的地图经纬度无效")
return
self.count += 1
if not self.config["enabled"]:
return
code = str(properties.get("adcode") or "").strip()
level = self.config["region_level"]
if code:
if not _adcode_belongs_to_region(code, self.config["region_adcode"], level):
if len(self.mismatched) < 5:
self.mismatched.add(code)
elif not _region_names_match(properties.get({"province": "province", "city": "city", "district": "district"}[level]), self.config["region_name"]):
self.unverified += 1
def finish(self, errors: list[str]) -> None:
if not self.config["enabled"]:
return
if not self.count:
errors.append("启用知识地图时,图谱数据至少需要一个带经纬度和业态的空间节点")
if self.unverified:
errors.append(f"知识地图空间节点有 {self.unverified} 条既无可用 adcode,也没有与所选区域匹配的省市区县字段")
if self.mismatched:
errors.append(f"知识地图数据与所选区域 {self.config['region_name']} 不一致,发现行政区划代码:{'、'.join(sorted(self.mismatched))}")
def normalize_provision_payload(body: Any) -> dict[str, Any]:
"""Validate small inline requests using the same rules as file imports."""
payload = normalize_provision_header(body)
errors: list[str] = []
data = body.get("graph_data") or {}
if not isinstance(data, Mapping):
raise ProjectValidationError(["graph_data 必须是 JSON 对象"])
raw_nodes, raw_relations = data.get("nodes", []), data.get("relations", [])
if not isinstance(raw_nodes, list):
errors.append("graph_data.nodes 必须是数组")
raw_nodes = []
if not isinstance(raw_relations, list):
errors.append("graph_data.relations 必须是数组")
raw_relations = []
if not raw_nodes:
errors.append("graph_data.nodes 至少需要一个节点,FalkorDB 不支持持久化空图")
nodes, relations, node_types = [], [], {}
spatial = SpatialMapValidation(payload["spatial_map"])
for index, raw in enumerate(raw_nodes):
node = normalize_graph_node(raw, index, payload["schema"]["entity_types"], node_types, errors)
if node is not None:
nodes.append(node)
node_types.setdefault(node["id"], node["type"])
spatial.add(node, errors)
spatial.finish(errors)
for index, raw in enumerate(raw_relations):
relation = normalize_graph_relation(raw, index, payload["schema"]["relation_types"], node_types, errors)
if relation is not None:
relations.append(relation)
if errors:
raise ProjectValidationError(errors)
payload["graph_data"] = {"nodes": nodes, "relations": relations}
payload["counts"].update(nodes=len(nodes), relations=len(relations))
return payload
def _falkor_client() -> FalkorDB:
kwargs: dict[str, Any] = {
"host": settings.falkordb_host,
"port": settings.falkordb_port,
"socket_timeout": 15,
"socket_connect_timeout": 5,
}
if settings.falkordb_password:
kwargs["password"] = settings.falkordb_password
return FalkorDB(**kwargs)
def list_falkor_graphs() -> set[str]:
db = _falkor_client()
try:
return {
name.decode("utf-8") if isinstance(name, bytes) else str(name)
for name in db.list_graphs()
}
finally:
db.close()
def _graph_safe_value(value: Any) -> Any:
if value is None:
return None
if isinstance(value, (str, bool, int, float)):
return value
if isinstance(value, list) and all(isinstance(item, (str, bool, int, float)) for item in value):
return value
return json.dumps(value, ensure_ascii=False, sort_keys=True)
def _graph_safe_properties(properties: Mapping[str, Any]) -> dict[str, Any]:
return {
str(key): safe_value
for key, value in properties.items()
if (safe_value := _graph_safe_value(value)) is not None
}
def _node_labels(node: Mapping[str, Any]) -> tuple[str, ...]:
"""Return validated FalkorDB labels, preserving optional multi-label nodes."""
primary = str(node.get("type") or "").strip()
raw_labels = node.get("labels")
labels = [str(item).strip() for item in raw_labels] if isinstance(raw_labels, list) else []
if primary and primary not in labels:
labels.insert(0, primary)
if not labels:
labels = [primary]
deduplicated = tuple(dict.fromkeys(labels))
if not all(label and SAFE_SCHEMA_IDENTIFIER.fullmatch(label) for label in deduplicated):
raise ValueError("图谱节点包含不安全的 Cypher 标签")
return deduplicated
def _batches(items: list[Any], size: int = GRAPH_IMPORT_BATCH_SIZE):
for start in range(0, len(items), size):
yield items[start:start + size]
def _import_falkor_graph_batched(graph: Any, graph_data: Mapping[str, Any]) -> None:
"""Import a validated graph with bounded UNWIND batches and lookup indexes."""
nodes = list(graph_data.get("nodes", []))
relations = list(graph_data.get("relations", []))
node_types = {str(item["id"]): str(item["type"]) for item in nodes}
nodes_by_labels: dict[tuple[str, ...], list[dict[str, Any]]] = defaultdict(list)
for item in nodes:
properties = _graph_safe_properties(item.get("properties") or {})
properties.setdefault("id", str(item["id"]))
properties["__kg_node_id"] = str(item["id"])
nodes_by_labels[_node_labels(item)].append({"properties": properties})
for labels, items in nodes_by_labels.items():
label_expression = ":".join(labels)
for batch in _batches(items):
graph.query(
f"UNWIND $rows AS row CREATE (n:{label_expression}) SET n += row.properties",
{"rows": batch},
)
for label in sorted(set(node_types.values())):
# Relation creation matches both endpoint labels and this indexed id,
# avoiding repeated full-graph scans for large JSON imports.
graph.query(f"CREATE INDEX FOR (n:{label}) ON (n.__kg_node_id)")
relations_by_signature: dict[tuple[str, str, str], list[dict[str, Any]]] = defaultdict(list)
for item in relations:
source = str(item["source"])
target = str(item["target"])
signature = (
str(item["type"]),
node_types[source],
node_types[target],
)
relations_by_signature[signature].append(
{
"source": source,
"target": target,
"properties": _graph_safe_properties(item.get("properties") or {}),
}
)
for (relation_type, source_label, target_label), items in relations_by_signature.items():
for batch in _batches(items):
graph.query(
f"""
UNWIND $rows AS row
MATCH (source:{source_label} {{__kg_node_id:row.source}})
MATCH (target:{target_label} {{__kg_node_id:row.target}})
CREATE (source)-[relation:{relation_type}]->(target)
SET relation += row.properties
""",
{"rows": batch},
)
def import_falkor_graph(graph_name: str, graph_data: Mapping[str, Any]) -> dict[str, int]:
"""Create a brand-new FalkorDB graph from already validated data."""
db = _falkor_client()
created = False
try:
existing = {
name.decode("utf-8") 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)
nodes = list(graph_data.get("nodes", []))
relations = list(graph_data.get("relations", []))
if len(nodes) + len(relations) >= GRAPH_IMPORT_BATCH_THRESHOLD:
created = True
_import_falkor_graph_batched(graph, {"nodes": nodes, "relations": relations})
else:
for node in nodes:
label_expression = ":".join(_node_labels(node))
properties = _graph_safe_properties(node.get("properties") or {})
properties.setdefault("id", str(node["id"]))
properties["__kg_node_id"] = str(node["id"])
# From this point onward the graph may exist even if the client
# raises after the server accepted the command.
created = True
graph.query(
f"CREATE (n:{label_expression}) SET n += $props",
{"props": properties},
)
for relation in relations:
relation_type = str(relation["type"])
properties = _graph_safe_properties(relation.get("properties") or {})
graph.query(
f"""
MATCH (source {{__kg_node_id:$source}}), (target {{__kg_node_id:$target}})
CREATE (source)-[relation:{relation_type}]->(target)
SET relation += $props
""",
{
"source": str(relation["source"]),
"target": str(relation["target"]),
"props": properties,
},
)
created = True
if not created:
return {"nodes": 0, "relations": 0}
node_count = int(graph.query("MATCH (n) RETURN count(n)").result_set[0][0])
relation_count = int(graph.query("MATCH ()-[r]->() RETURN count(r)").result_set[0][0])
return {"nodes": node_count, "relations": relation_count}
except FileExistsError:
raise
except Exception as exc:
cleanup_error: Exception | None = None
if created:
try:
db.select_graph(graph_name).delete()
except Exception as delete_exc:
cleanup_error = delete_exc
raise GraphImportError(str(exc), cleanup_error=cleanup_error) from exc
finally:
db.close()
def delete_falkor_graph(graph_name: str) -> bool:
"""Delete one exact graph name, returning False when it no longer exists."""
db = _falkor_client()
try:
existing = {
name.decode("utf-8") if isinstance(name, bytes) else str(name)
for name in db.list_graphs()
}
if graph_name not in existing:
return False
db.select_graph(graph_name).delete()
return True
finally:
db.close()