542 lines
20 KiB
Python
542 lines
20 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 re
|
||
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}$")
|
||
|
||
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",
|
||
}
|
||
|
||
|
||
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 _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_payload(body: Any) -> dict[str, Any]:
|
||
"""Normalize and fully validate a project provision request.
|
||
|
||
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
|
||
|
||
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
|
||
|
||
raw_graph_data = body.get("graph_data")
|
||
if raw_graph_data is None:
|
||
raw_graph_data = {"nodes": [], "relations": []}
|
||
if not isinstance(raw_graph_data, Mapping):
|
||
errors.append("graph_data 必须是 JSON 对象")
|
||
raw_graph_data = {}
|
||
raw_nodes = raw_graph_data.get("nodes", [])
|
||
raw_relations = raw_graph_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: list[dict[str, Any]] = []
|
||
node_types: dict[str, str] = {}
|
||
for index, raw_node in enumerate(raw_nodes):
|
||
path = f"graph_data.nodes[{index}]"
|
||
if not isinstance(raw_node, Mapping):
|
||
errors.append(f"{path} 必须是 JSON 对象")
|
||
continue
|
||
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", {})
|
||
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 not isinstance(properties, Mapping):
|
||
errors.append(f"{path}.properties 必须是 JSON 对象")
|
||
properties = {}
|
||
normalized_properties = dict(properties)
|
||
if entity_type in normalized_entities:
|
||
_validate_properties(
|
||
normalized_properties,
|
||
normalized_entities[entity_type].get("fields", {}),
|
||
path=path,
|
||
errors=errors,
|
||
)
|
||
if node_id and node_id not in node_types:
|
||
node_types[node_id] = entity_type
|
||
nodes.append({"id": node_id, "type": entity_type, "properties": normalized_properties})
|
||
|
||
relations: list[dict[str, Any]] = []
|
||
for index, raw_relation in enumerate(raw_relations):
|
||
path = f"graph_data.relations[{index}]"
|
||
if not isinstance(raw_relation, Mapping):
|
||
errors.append(f"{path} 必须是 JSON 对象")
|
||
continue
|
||
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,
|
||
)
|
||
relations.append(
|
||
{
|
||
"type": relation_type,
|
||
"source": source,
|
||
"target": target,
|
||
"properties": normalized_properties,
|
||
}
|
||
)
|
||
|
||
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,
|
||
"schema": schema,
|
||
"graph_data": {"nodes": nodes, "relations": relations},
|
||
"counts": {
|
||
"entity_types": len(normalized_entities),
|
||
"relation_types": len(normalized_relations),
|
||
"nodes": len(nodes),
|
||
"relations": len(relations),
|
||
},
|
||
}
|
||
|
||
|
||
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 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)
|
||
for node in graph_data.get("nodes", []):
|
||
label = str(node["type"])
|
||
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}) SET n += $props", {"props": properties})
|
||
for relation in graph_data.get("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()
|