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