819 lines
32 KiB
Python
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()
|