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

819 lines
32 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.

"""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()