498 lines
18 KiB
Python
498 lines
18 KiB
Python
import asyncio
|
||
import threading
|
||
import unittest
|
||
from unittest.mock import patch
|
||
|
||
from fastapi import HTTPException
|
||
|
||
from app import db as app_db
|
||
from app.api import projects as projects_api
|
||
from app.project_lifecycle import (
|
||
GraphImportError,
|
||
ProjectValidationError,
|
||
import_falkor_graph,
|
||
normalize_provision_payload,
|
||
)
|
||
|
||
|
||
def valid_payload() -> dict:
|
||
return {
|
||
"project_id": "lifecycle_test",
|
||
"tenant_id": "test_tenant",
|
||
"display_name": "生命周期测试",
|
||
"graph_name": "lifecycle_test_graph",
|
||
"schema": {
|
||
"namespace": "lifecycle_test",
|
||
"version": "1.0.0",
|
||
"entity_types": {
|
||
"Place": {
|
||
"label": "地点",
|
||
"fields": {
|
||
"name": {"type": "string", "required": True},
|
||
"rating": {"type": "number"},
|
||
},
|
||
}
|
||
},
|
||
"relation_types": {
|
||
"RELATED_TO": {
|
||
"from": "Place",
|
||
"to": "Place",
|
||
"properties": {"weight": {"type": "number"}},
|
||
}
|
||
},
|
||
},
|
||
"graph_data": {
|
||
"nodes": [
|
||
{"id": "p1", "type": "Place", "properties": {"name": "甲", "rating": 4.8}},
|
||
{"id": "p2", "type": "Place", "properties": {"name": "乙"}},
|
||
],
|
||
"relations": [
|
||
{
|
||
"type": "RELATED_TO",
|
||
"source": "p1",
|
||
"target": "p2",
|
||
"properties": {"weight": 0.9},
|
||
}
|
||
],
|
||
},
|
||
}
|
||
|
||
|
||
class ProjectPayloadValidationTests(unittest.TestCase):
|
||
def test_valid_payload_is_normalized_and_counted(self) -> None:
|
||
payload = normalize_provision_payload(valid_payload())
|
||
|
||
self.assertEqual(payload["tenant_id"], "test_tenant")
|
||
self.assertEqual(payload["graph_name"], "lifecycle_test_graph")
|
||
self.assertEqual(payload["counts"], {
|
||
"entity_types": 1,
|
||
"relation_types": 1,
|
||
"nodes": 2,
|
||
"relations": 1,
|
||
})
|
||
self.assertEqual(
|
||
payload["schema"]["entity_types"]["Place"]["fields"]["rating"]["type"],
|
||
"number",
|
||
)
|
||
|
||
def test_hidden_resource_ids_default_to_project_id(self) -> None:
|
||
body = valid_payload()
|
||
body.pop("tenant_id")
|
||
body.pop("graph_name")
|
||
|
||
payload = normalize_provision_payload(body)
|
||
|
||
self.assertEqual(payload["tenant_id"], "lifecycle_test")
|
||
self.assertEqual(payload["graph_name"], "lifecycle_test")
|
||
|
||
def test_blank_resource_ids_also_default_to_project_id(self) -> None:
|
||
body = valid_payload()
|
||
body["tenant_id"] = " "
|
||
body["graph_name"] = " "
|
||
|
||
payload = normalize_provision_payload(body)
|
||
|
||
self.assertEqual(payload["tenant_id"], "lifecycle_test")
|
||
self.assertEqual(payload["graph_name"], "lifecycle_test")
|
||
|
||
def test_unsafe_project_and_graph_names_are_rejected(self) -> None:
|
||
body = valid_payload()
|
||
body["project_id"] = "../../biz_secret"
|
||
body["graph_name"] = "graph name"
|
||
|
||
with self.assertRaises(ProjectValidationError) as raised:
|
||
normalize_provision_payload(body)
|
||
|
||
joined = ";".join(raised.exception.errors)
|
||
self.assertIn("项目英文标识", joined)
|
||
self.assertIn("图谱资源标识", joined)
|
||
self.assertNotIn("graph_name", joined)
|
||
|
||
def test_unknown_relation_endpoint_type_is_rejected(self) -> None:
|
||
body = valid_payload()
|
||
body["schema"]["relation_types"]["RELATED_TO"]["to"] = "MissingType"
|
||
|
||
with self.assertRaises(ProjectValidationError) as raised:
|
||
normalize_provision_payload(body)
|
||
|
||
self.assertTrue(any("to/target" in error for error in raised.exception.errors))
|
||
|
||
def test_duplicate_node_and_missing_endpoint_are_rejected(self) -> None:
|
||
body = valid_payload()
|
||
body["graph_data"]["nodes"][1]["id"] = "p1"
|
||
body["graph_data"]["relations"][0]["target"] = "missing"
|
||
|
||
with self.assertRaises(ProjectValidationError) as raised:
|
||
normalize_provision_payload(body)
|
||
|
||
joined = ";".join(raised.exception.errors)
|
||
self.assertIn("重复", joined)
|
||
self.assertIn("不存在", joined)
|
||
|
||
def test_required_and_property_types_are_checked(self) -> None:
|
||
body = valid_payload()
|
||
body["graph_data"]["nodes"][0]["properties"] = {"rating": "excellent"}
|
||
|
||
with self.assertRaises(ProjectValidationError) as raised:
|
||
normalize_provision_payload(body)
|
||
|
||
joined = ";".join(raised.exception.errors)
|
||
self.assertIn("必填属性 name", joined)
|
||
self.assertIn("rating 应为 number", joined)
|
||
|
||
def test_data_center_tables_are_not_in_delete_allow_lists(self) -> None:
|
||
protected = {
|
||
"project_databases",
|
||
"project_table_definitions",
|
||
"project_table_overrides",
|
||
}
|
||
self.assertTrue(protected.isdisjoint(projects_api._PROJECT_SCOPED_TABLES))
|
||
self.assertTrue(protected.isdisjoint(projects_api._PROJECT_DELETE_ORDER))
|
||
self.assertFalse(any(name.startswith("biz_") for name in projects_api._PROJECT_DELETE_ORDER))
|
||
self.assertTrue(protected.isdisjoint(app_db._LIFECYCLE_PROJECT_GUARDED_TABLES))
|
||
self.assertFalse(any(name.startswith("biz_") for name in app_db._LIFECYCLE_PROJECT_GUARDED_TABLES))
|
||
|
||
def test_empty_graph_is_rejected(self) -> None:
|
||
body = valid_payload()
|
||
body["graph_data"] = {"nodes": [], "relations": []}
|
||
|
||
with self.assertRaises(ProjectValidationError) as raised:
|
||
normalize_provision_payload(body)
|
||
|
||
self.assertTrue(any("至少需要一个节点" in error for error in raised.exception.errors))
|
||
|
||
def test_tuple_relation_and_union_endpoints_are_normalized(self) -> None:
|
||
body = valid_payload()
|
||
body["schema"]["entity_types"]["Area"] = {
|
||
"fields": {"name": "text"},
|
||
}
|
||
body["schema"]["relation_types"]["RELATED_TO"] = [
|
||
"Place|Area",
|
||
"Place|Area",
|
||
"联合端点关系",
|
||
]
|
||
body["graph_data"]["nodes"].append(
|
||
{"id": "a1", "type": "Area", "properties": {"name": "区域"}}
|
||
)
|
||
body["graph_data"]["relations"][0]["target"] = "a1"
|
||
|
||
payload = normalize_provision_payload(body)
|
||
|
||
self.assertEqual(payload["schema"]["relation_types"]["RELATED_TO"]["from"], "Place|Area")
|
||
self.assertEqual(
|
||
payload["schema"]["entity_types"]["Area"]["fields"]["name"]["type"],
|
||
"string",
|
||
)
|
||
|
||
def test_unsafe_schema_identifier_and_version_are_rejected(self) -> None:
|
||
body = valid_payload()
|
||
body["schema"]["version"] = "../../1"
|
||
body["schema"]["entity_types"] = {"Bad Label": {"fields": {}}}
|
||
|
||
with self.assertRaises(ProjectValidationError) as raised:
|
||
normalize_provision_payload(body)
|
||
|
||
joined = ";".join(raised.exception.errors)
|
||
self.assertIn("schema.version", joined)
|
||
self.assertIn("Schema/Cypher", joined)
|
||
|
||
def test_schema_source_path_rejects_unsafe_components(self) -> None:
|
||
row = {"project_id": "../outside", "namespace": "safe", "version": "1.0.0"}
|
||
self.assertIsNone(projects_api._schema_source_from_file(row, {}))
|
||
|
||
def test_graph_release_uses_project_id_as_its_implicit_graph_name(self) -> None:
|
||
self.assertEqual(projects_api._resolve_graph_name("stable-project"), "stable-project")
|
||
self.assertEqual(projects_api._resolve_graph_name("stable-project", " "), "stable-project")
|
||
self.assertEqual(
|
||
projects_api._resolve_graph_name("stable-project", "legacy_graph"),
|
||
"legacy_graph",
|
||
)
|
||
|
||
def test_schema_response_prefers_json_semantic_version_over_legacy_revision(self) -> None:
|
||
row = {
|
||
"id": 42,
|
||
"project_id": "stable-project",
|
||
"namespace": "stable-project",
|
||
"version": 7,
|
||
"schema_jsonb": {
|
||
"namespace": "stable-project",
|
||
"version": "1.0.0",
|
||
"entity_types": {},
|
||
"relation_types": {},
|
||
},
|
||
}
|
||
|
||
with patch("app.api.projects._schema_source_from_file", return_value=None):
|
||
result = projects_api._attach_schema_source(row)
|
||
|
||
self.assertEqual(result["id"], 42)
|
||
self.assertEqual(result["version"], "1.0.0")
|
||
|
||
|
||
class _ScriptedCursor:
|
||
def __init__(self, *rows: dict):
|
||
self.rows = list(rows)
|
||
self.calls: list[tuple[str, object]] = []
|
||
|
||
async def execute(self, query: str, params=None) -> None:
|
||
self.calls.append((query, params))
|
||
|
||
async def fetchone(self):
|
||
return self.rows.pop(0) if self.rows else None
|
||
|
||
|
||
class SchemaVersionStorageTests(unittest.IsolatedAsyncioTestCase):
|
||
async def test_integer_column_allocates_revision_for_semantic_version(self) -> None:
|
||
cursor = _ScriptedCursor(
|
||
{"data_type": "integer", "udt_name": "int4"},
|
||
{"next_version": 7},
|
||
)
|
||
|
||
stored_version = await projects_api._ontology_schema_storage_version(
|
||
cursor,
|
||
tenant_id="tenant",
|
||
project_id="project",
|
||
namespace="namespace",
|
||
semantic_version="1.0.0",
|
||
)
|
||
|
||
self.assertEqual(stored_version, 7)
|
||
self.assertEqual(len(cursor.calls), 2)
|
||
self.assertIn("MAX(version)", cursor.calls[1][0])
|
||
self.assertEqual(cursor.calls[1][1], ("tenant", "project", "namespace"))
|
||
|
||
async def test_text_column_stores_semantic_version_directly(self) -> None:
|
||
cursor = _ScriptedCursor({"data_type": "text", "udt_name": "text"})
|
||
|
||
stored_version = await projects_api._ontology_schema_storage_version(
|
||
cursor,
|
||
tenant_id="tenant",
|
||
project_id="project",
|
||
namespace="namespace",
|
||
semantic_version="1.0.0",
|
||
)
|
||
|
||
self.assertEqual(stored_version, "1.0.0")
|
||
self.assertEqual(len(cursor.calls), 1)
|
||
|
||
async def test_unknown_version_column_type_fails_before_insert(self) -> None:
|
||
cursor = _ScriptedCursor({"data_type": "uuid", "udt_name": "uuid"})
|
||
|
||
with self.assertRaisesRegex(RuntimeError, "不受支持"):
|
||
await projects_api._ontology_schema_storage_version(
|
||
cursor,
|
||
tenant_id="tenant",
|
||
project_id="project",
|
||
namespace="namespace",
|
||
semantic_version="1.0.0",
|
||
)
|
||
|
||
self.assertEqual(len(cursor.calls), 1)
|
||
|
||
|
||
class _Result:
|
||
def __init__(self, count: int):
|
||
self.result_set = [[count]]
|
||
|
||
|
||
class _FakeGraph:
|
||
def __init__(self, *, fail_on_query: int | None = None, delete_fails: bool = False):
|
||
self.queries: list[tuple[str, dict | None]] = []
|
||
self.fail_on_query = fail_on_query
|
||
self.delete_fails = delete_fails
|
||
self.deleted = False
|
||
|
||
def query(self, query: str, params: dict | None = None):
|
||
self.queries.append((query, params))
|
||
if self.fail_on_query and len(self.queries) == self.fail_on_query:
|
||
raise RuntimeError("write failed")
|
||
if "MATCH (n) RETURN count(n)" in query:
|
||
return _Result(2)
|
||
if "MATCH ()-[r]->() RETURN count(r)" in query:
|
||
return _Result(1)
|
||
return _Result(0)
|
||
|
||
def delete(self):
|
||
if self.delete_fails:
|
||
raise RuntimeError("cleanup failed")
|
||
self.deleted = True
|
||
|
||
|
||
class _FakeDb:
|
||
def __init__(self, graph: _FakeGraph, existing: list[str] | None = None):
|
||
self.graph = graph
|
||
self.existing = existing or []
|
||
self.closed = False
|
||
|
||
def list_graphs(self):
|
||
return self.existing
|
||
|
||
def select_graph(self, _name: str):
|
||
return self.graph
|
||
|
||
def close(self):
|
||
self.closed = True
|
||
|
||
|
||
class FalkorLifecycleTests(unittest.TestCase):
|
||
def test_import_creates_validated_nodes_and_relations(self) -> None:
|
||
payload = normalize_provision_payload(valid_payload())
|
||
graph = _FakeGraph()
|
||
database = _FakeDb(graph)
|
||
|
||
with patch("app.project_lifecycle._falkor_client", return_value=database):
|
||
counts = import_falkor_graph("lifecycle_test_graph", payload["graph_data"])
|
||
|
||
self.assertEqual(counts, {"nodes": 2, "relations": 1})
|
||
self.assertEqual(len([query for query, _ in graph.queries if "CREATE (n:Place)" in query]), 2)
|
||
self.assertTrue(any("[relation:RELATED_TO]" in query for query, _ in graph.queries))
|
||
self.assertFalse(graph.deleted)
|
||
self.assertTrue(database.closed)
|
||
|
||
def test_import_failure_compensates_only_the_new_graph(self) -> None:
|
||
payload = normalize_provision_payload(valid_payload())
|
||
graph = _FakeGraph(fail_on_query=2)
|
||
database = _FakeDb(graph)
|
||
|
||
with patch("app.project_lifecycle._falkor_client", return_value=database):
|
||
with self.assertRaises(GraphImportError) as raised:
|
||
import_falkor_graph("lifecycle_test_graph", payload["graph_data"])
|
||
|
||
self.assertTrue(graph.deleted)
|
||
self.assertFalse(raised.exception.cleanup_required)
|
||
|
||
def test_failed_compensation_is_reported_as_retryable(self) -> None:
|
||
payload = normalize_provision_payload(valid_payload())
|
||
graph = _FakeGraph(fail_on_query=2, delete_fails=True)
|
||
database = _FakeDb(graph)
|
||
|
||
with patch("app.project_lifecycle._falkor_client", return_value=database):
|
||
with self.assertRaises(GraphImportError) as raised:
|
||
import_falkor_graph("lifecycle_test_graph", payload["graph_data"])
|
||
|
||
self.assertTrue(raised.exception.cleanup_required)
|
||
|
||
def test_existing_graph_is_never_overwritten_or_deleted(self) -> None:
|
||
payload = normalize_provision_payload(valid_payload())
|
||
graph = _FakeGraph()
|
||
database = _FakeDb(graph, existing=["lifecycle_test_graph"])
|
||
|
||
with patch("app.project_lifecycle._falkor_client", return_value=database):
|
||
with self.assertRaises(FileExistsError):
|
||
import_falkor_graph("lifecycle_test_graph", payload["graph_data"])
|
||
|
||
self.assertFalse(graph.queries)
|
||
self.assertFalse(graph.deleted)
|
||
|
||
|
||
class LifecycleWriteGuardTests(unittest.IsolatedAsyncioTestCase):
|
||
async def test_blocking_graph_worker_is_drained_before_cancellation_returns(self) -> None:
|
||
started = threading.Event()
|
||
release = threading.Event()
|
||
finished = threading.Event()
|
||
|
||
def blocking_mutation() -> str:
|
||
started.set()
|
||
release.wait(timeout=2)
|
||
finished.set()
|
||
return "done"
|
||
|
||
task = asyncio.create_task(
|
||
projects_api._run_blocking_to_completion(blocking_mutation)
|
||
)
|
||
await asyncio.to_thread(started.wait, 1)
|
||
task.cancel()
|
||
await asyncio.sleep(0.02)
|
||
self.assertFalse(task.done())
|
||
release.set()
|
||
with self.assertRaises(asyncio.CancelledError):
|
||
await task
|
||
self.assertTrue(finished.is_set())
|
||
|
||
async def test_async_lock_cleanup_is_drained_before_cancellation_returns(self) -> None:
|
||
started = asyncio.Event()
|
||
release = asyncio.Event()
|
||
finished = asyncio.Event()
|
||
|
||
async def cleanup() -> None:
|
||
started.set()
|
||
await release.wait()
|
||
finished.set()
|
||
|
||
task = asyncio.create_task(
|
||
projects_api._run_async_cleanup_to_completion(cleanup())
|
||
)
|
||
await started.wait()
|
||
task.cancel()
|
||
await asyncio.sleep(0)
|
||
self.assertFalse(task.done())
|
||
release.set()
|
||
with self.assertRaises(asyncio.CancelledError):
|
||
await task
|
||
self.assertTrue(finished.is_set())
|
||
|
||
async def test_invalid_provision_payload_performs_no_external_calls(self) -> None:
|
||
body = valid_payload()
|
||
body["project_id"] = "../unsafe"
|
||
user = {"username": "admin", "roles": ["admin"]}
|
||
|
||
with patch("app.api.projects.get_conn") as get_conn, patch(
|
||
"app.api.projects.list_falkor_graphs"
|
||
) as list_graphs:
|
||
with self.assertRaises(HTTPException) as raised:
|
||
await projects_api.provision_project(body, user)
|
||
|
||
self.assertEqual(raised.exception.status_code, 422)
|
||
get_conn.assert_not_called()
|
||
list_graphs.assert_not_called()
|
||
|
||
async def test_derived_graph_collision_uses_project_id_message_before_database_access(self) -> None:
|
||
body = valid_payload()
|
||
body.pop("tenant_id")
|
||
body.pop("graph_name")
|
||
user = {"username": "admin", "roles": ["admin"]}
|
||
|
||
with patch("app.api.projects.get_conn") as get_conn, patch(
|
||
"app.api.projects.list_falkor_graphs",
|
||
return_value={"lifecycle_test"},
|
||
):
|
||
with self.assertRaises(HTTPException) as raised:
|
||
await projects_api.provision_project(body, user)
|
||
|
||
self.assertEqual(raised.exception.status_code, 409)
|
||
self.assertIn("项目英文标识“lifecycle_test”", str(raised.exception.detail))
|
||
self.assertNotIn("graph_name", str(raised.exception.detail))
|
||
self.assertNotIn("FalkorDB", str(raised.exception.detail))
|
||
get_conn.assert_not_called()
|
||
|
||
async def test_delete_confirmation_mismatch_performs_no_write(self) -> None:
|
||
user = {"username": "admin", "roles": ["admin"]}
|
||
|
||
with patch("app.api.projects.get_conn") as get_conn, patch(
|
||
"app.api.projects.delete_falkor_graph"
|
||
) as delete_graph:
|
||
with self.assertRaises(HTTPException) as raised:
|
||
await projects_api.delete_project("expected", "wrong", user)
|
||
|
||
self.assertEqual(raised.exception.status_code, 400)
|
||
get_conn.assert_not_called()
|
||
delete_graph.assert_not_called()
|
||
|
||
async def test_rename_rejects_stable_fields_before_database_access(self) -> None:
|
||
user = {"username": "admin", "roles": ["admin"]}
|
||
|
||
with patch("app.api.projects.get_conn") as get_conn:
|
||
with self.assertRaises(HTTPException) as raised:
|
||
await projects_api.rename_project(
|
||
"stable-id",
|
||
{"display_name": "新名称", "graph_name": "forbidden"},
|
||
user,
|
||
)
|
||
|
||
self.assertEqual(raised.exception.status_code, 422)
|
||
get_conn.assert_not_called()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|