Files
Cloud-Tour-to-Libo/tests/test_project_lifecycle.py

498 lines
18 KiB
Python
Raw 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.

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