import asyncio import copy from contextlib import closing import hashlib import io import json from pathlib import Path import sqlite3 import tempfile import time from types import SimpleNamespace import unittest from unittest.mock import AsyncMock, patch import httpx import ijson from fastapi import FastAPI from starlette.requests import ClientDisconnect, Request from app.api import graph_file_import as file_api from app.api import projects as projects_api from app.auth import get_current_user from app.graph_import_files import CheckedReader, graph_records, import_staged_graph, validate_graph_file from app.project_lifecycle import GraphImportError, ProjectValidationError, normalize_provision_payload from tests.test_project_lifecycle import valid_payload class FileValidationTests(unittest.TestCase): def validate(self, body, graph=None): source = io.BytesIO(json.dumps(graph if graph is not None else body["graph_data"], ensure_ascii=False).encode()) directory = tempfile.TemporaryDirectory(prefix="test-graph-file-") self.addCleanup(directory.cleanup) result = validate_graph_file(source, Path(directory.name) / "index.sqlite", body) self.assertFalse(source.closed) return result def test_stream_and_inline_validation_preserve_same_records(self): body = valid_payload() body["graph_data"]["nodes"][0]["properties"]["details"] = {"中文": [1, True, "原始内容"]} body["graph_data"]["nodes"][0]["labels"] = ["Place"] result = self.validate(body) expected = normalize_provision_payload(body) with closing(sqlite3.connect(result["graph_data"]["staging_path"])) as connection: nodes = [json.loads(row[0]) for row in connection.execute("SELECT payload FROM nodes ORDER BY id")] relations = [json.loads(row[0]) for row in connection.execute("SELECT payload FROM relations ORDER BY seq")] self.assertEqual(nodes, expected["graph_data"]["nodes"]) self.assertEqual(relations, expected["graph_data"]["relations"]) self.assertEqual(result["counts"]["nodes"], 2) self.assertEqual(result["counts"]["relations"], 1) def test_reads_bounded_blocks_and_hashes_original_bytes(self): data = json.dumps(valid_payload()["graph_data"], ensure_ascii=False).encode() class Bounded(io.BytesIO): def read(self, size=-1): assert 0 <= size <= 65536, size return super().read(size) with tempfile.TemporaryDirectory() as directory: result = validate_graph_file(Bounded(data), Path(directory) / "index.sqlite", valid_payload()) self.assertEqual(result["file_sha256"], hashlib.sha256(data).hexdigest()) def test_relations_can_precede_nodes(self): body = valid_payload() result = self.validate(body, {"relations": body["graph_data"]["relations"], "nodes": body["graph_data"]["nodes"]}) self.assertEqual(result["counts"]["relations"], 1) def test_invalid_records_fail_before_import(self): for field, value in [("nodes", []), ("nodes", {}), ("relations", {}), ("nodes", ["bad"])]: with self.subTest(field=field, value=value): body = valid_payload() body["graph_data"][field] = value with self.assertRaises(ProjectValidationError): self.validate(body) def test_duplicates_unknown_types_and_missing_endpoints_fail(self): bodies = [] body = valid_payload() body["graph_data"]["nodes"].append(copy.deepcopy(body["graph_data"]["nodes"][0])) bodies.append(body) body = valid_payload() body["graph_data"]["nodes"][0]["type"] = "NotInSchema" bodies.append(body) body = valid_payload() body["graph_data"]["relations"][0]["target"] = "missing" bodies.append(body) for body in bodies: with self.subTest(body=body), self.assertRaises(ProjectValidationError): self.validate(body) def test_duplicate_json_fields_truncation_and_wrong_root_fail(self): for raw in [b'{"nodes":[],"nodes":[]}', b'{"nodes":[{}]', b'[]', b'{"schema":{}}', b'{"nodes":[{"id":"a","id":"b"}]}']: with self.subTest(raw=raw), self.assertRaises((ValueError, ijson.JSONError)): list(graph_records(CheckedReader(io.BytesIO(raw)))) def test_deeply_nested_or_oversized_record_rejected(self): with self.assertRaisesRegex(ValueError, "64"): list(graph_records(CheckedReader(io.BytesIO(b'[' * 65 + b']' * 65)))) with patch("app.graph_import_files.MAX_RECORD_CHARS", 8): with self.assertRaisesRegex(ValueError, "单条"): list(graph_records(CheckedReader(io.BytesIO(b'{"nodes":[{"id":"long-string"}]}')))) def test_map_validation_and_helper_exclusion_match_import_rules(self): body = valid_payload() body["spatial_map"] = {"enabled": True, "scope": "guizhou", "region_name": "贵阳市", "region_adcode": "520100", "region_level": "city"} for index, node in enumerate(body["graph_data"]["nodes"]): node["properties"].update(lng=106.7, lat=26.6, adcode="520102", place_type="Hotel", map_poi=index == 0) result = self.validate(body) self.assertEqual(result["counts"]["map_pois"], 1) body["graph_data"]["nodes"][0]["properties"]["adcode"] = "522722" with self.assertRaisesRegex(ProjectValidationError, "所选区域"): self.validate(body) def test_existing_staging_is_never_overwritten(self): with tempfile.TemporaryDirectory() as directory: staging = Path(directory) / "index.sqlite" staging.touch() with self.assertRaises(FileExistsError): validate_graph_file(io.BytesIO(b'{}'), staging, valid_payload()) class FileImportTests(unittest.TestCase): validate = FileValidationTests.validate def test_batches_preserve_labels_parallel_edges_and_counts(self): body = valid_payload() body["schema"]["entity_types"]["Hotel"] = {"fields": {"name": "string"}} body["graph_data"]["nodes"] = [ {"id": str(i), "type": "Place", "labels": ["Place", "Hotel"], "properties": {"name": str(i)}} for i in range(503) ] body["graph_data"]["relations"] = [ {"type": "RELATED_TO", "source": "0", "target": "1", "properties": {"weight": i}} for i in range(502) ] result = self.validate(body) calls = [] def query(cypher, params=None): calls.append((cypher, params)) if cypher == "MATCH (n) RETURN count(n)": return SimpleNamespace(result_set=[[503]]) return SimpleNamespace(result_set=[[502]]) graph = SimpleNamespace(query=query, delete=unittest.mock.Mock()) client = SimpleNamespace(list_graphs=lambda: [], select_graph=lambda _: graph, close=unittest.mock.Mock()) with patch("app.graph_import_files._falkor_client", return_value=client): self.assertEqual(import_staged_graph("test_only", result["graph_data"]), {"nodes": 503, "relations": 502}) node_calls = [call for call in calls if "CREATE (n:" in call[0]] edge_calls = [call for call in calls if "CREATE (s)-[r:" in call[0]] self.assertEqual([len(call[1]["rows"]) for call in node_calls], [250, 250, 3]) self.assertEqual([len(call[1]["rows"]) for call in edge_calls], [250, 250, 2]) self.assertIn("n:Place:Hotel", node_calls[0][0]) graph.delete.assert_not_called() client.close.assert_called_once() def test_never_deletes_existing_graph_on_conflict(self): result = self.validate(valid_payload()) delete = unittest.mock.Mock() client = SimpleNamespace(list_graphs=lambda: [b"existing"], select_graph=lambda _: SimpleNamespace(delete=delete), close=unittest.mock.Mock()) with patch("app.graph_import_files._falkor_client", return_value=client), self.assertRaises(FileExistsError): import_staged_graph("existing", result["graph_data"]) delete.assert_not_called() client.close.assert_called_once() def test_failed_write_cleans_only_new_graph(self): result = self.validate(valid_payload()) graph = SimpleNamespace(query=unittest.mock.Mock(side_effect=RuntimeError("write failed")), delete=unittest.mock.Mock()) client = SimpleNamespace(list_graphs=lambda: [], select_graph=lambda _: graph, close=unittest.mock.Mock()) with patch("app.graph_import_files._falkor_client", return_value=client), self.assertRaises(GraphImportError): import_staged_graph("new_only", result["graph_data"]) graph.delete.assert_called_once() class FileRouteTests(unittest.IsolatedAsyncioTestCase): async def asyncSetUp(self): self.app = FastAPI() self.app.include_router(file_api.router, prefix="/v1/admin") self.app.include_router(projects_api.router, prefix="/v1/admin") self.app.dependency_overrides[get_current_user] = lambda: {"username": "test_admin", "roles": ["admin"]} self.client = httpx.AsyncClient(transport=httpx.ASGITransport(app=self.app), base_url="http://testserver") self.addAsyncCleanup(self.client.aclose) self.body = valid_payload() def upload(self, data=None, **kwargs): return self.client.post("/v1/admin/projects/provision-file", data={"metadata": json.dumps(self.body)}, files={"file": ("graph.json", data or json.dumps(self.body["graph_data"]).encode(), "application/json")}, **kwargs) async def test_import_options_exposes_five_gib_default(self): response = await self.client.get("/v1/admin/projects/import-options") self.assertEqual(response.status_code, 200) self.assertEqual(response.json()["max_file_bytes"], 5 * 1024 * 1024 * 1024) async def test_valid_file_streams_counts_and_result_and_releases_lock(self): seen = {} async def publish(payload, user, **kwargs): seen.update(payload) self.assertTrue(Path(payload["graph_data"]["staging_path"]).is_file()) self.assertEqual(user["username"], "test_admin") return {"project": {"project_id": "test"}, "counts": payload["counts"]} with patch("app.api.projects.provision_validated_project", side_effect=publish): response = await self.upload() self.assertEqual(response.status_code, 200) messages = [json.loads(line) for line in response.text.splitlines()] self.assertEqual(messages[-1]["event"], "result") self.assertEqual(messages[-1]["data"]["counts"]["nodes"], 2) self.assertFalse(Path(seen["graph_data"]["staging_path"]).exists()) self.assertFalse(file_api._import_lock.locked()) async def test_invalid_json_reports_error_and_does_not_publish(self): with patch("app.api.projects.provision_validated_project", new_callable=AsyncMock) as publish: response = await self.upload(b'{"nodes": [') self.assertEqual(response.status_code, 200) messages = [json.loads(line) for line in response.text.splitlines()] self.assertEqual(messages[-1]["event"], "error") publish.assert_not_called() self.assertFalse(file_api._import_lock.locked()) async def test_viewer_cannot_upload_and_body_is_not_read(self): self.app.dependency_overrides[get_current_user] = lambda: {"username": "viewer", "roles": ["operator"]} with patch.object(file_api, "read_import_form", new_callable=AsyncMock) as read: response = await self.upload() self.assertEqual(response.status_code, 403) read.assert_not_called() async def test_oversize_and_low_disk_fail_clearly(self): with patch.object(file_api.settings, "graph_import_max_bytes", 1): self.assertEqual((await self.upload()).status_code, 413) with patch.object(file_api.shutil, "disk_usage", return_value=SimpleNamespace(free=0)): self.assertEqual((await self.upload()).status_code, 507) self.assertFalse(file_api._import_lock.locked()) async def test_cpu_validation_does_not_block_health_request(self): @self.app.get("/ping") async def ping(): return {"ok": True} original = file_api.validate_graph_file started = asyncio.Event() loop = asyncio.get_running_loop() def slower(*args): loop.call_soon_threadsafe(started.set) time.sleep(0.25) return original(*args) with patch.object(file_api, "validate_graph_file", side_effect=slower), patch( "app.api.projects.provision_validated_project", new_callable=AsyncMock, return_value={"project": {"project_id": "test"}}, ): task = asyncio.create_task(self.upload()) await asyncio.wait_for(started.wait(), 2) response = await asyncio.wait_for(self.client.get("/ping"), 0.15) self.assertEqual(response.status_code, 200) second = await self.upload() self.assertEqual(second.status_code, 429) await task async def test_disconnect_closes_partial_uploaded_file(self): original = tempfile.SpooledTemporaryFile files = [] def capture(*args, **kwargs): result = original(*args, **kwargs) files.append(result) return result incoming = iter([ {"type": "http.request", "body": b'--x\r\nContent-Disposition: form-data; name="file"; filename="a.json"\r\n\r\npartial', "more_body": True}, {"type": "http.disconnect"}, ]) async def receive(): return next(incoming) request = Request({"type": "http", "headers": [(b"content-type", b"multipart/form-data; boundary=x")]}, receive) with patch("starlette.formparsers.SpooledTemporaryFile", side_effect=capture), self.assertRaises(ClientDisconnect): await file_api.read_import_form(request) self.assertEqual(len(files), 1) self.assertTrue(files[0].closed) async def test_small_inline_json_remains_compatible(self): with patch.object(projects_api, "provision_validated_project", new_callable=AsyncMock, return_value={"project": {"project_id": "test"}}) as publish: response = await self.client.post("/v1/admin/projects/provision", json=self.body) self.assertEqual(response.status_code, 200) self.assertEqual(publish.call_args.args[0]["counts"]["nodes"], 2) async def test_inline_limit_applies_before_json_parsing(self): with patch.object(projects_api, "normalize_provision_payload") as normalize: response = await self.client.post("/v1/admin/projects/provision", content=b'{}', headers={"Content-Length": str(8 * 1024 * 1024 + 1)}) self.assertEqual(response.status_code, 413) normalize.assert_not_called() async def test_chunked_inline_request_cannot_bypass_limit(self): async def content(): for _ in range(9): yield b' ' * (1024 * 1024) with patch.object(projects_api, "normalize_provision_payload") as normalize: response = await self.client.post("/v1/admin/projects/provision", content=content()) self.assertEqual(response.status_code, 413) normalize.assert_not_called() async def test_inline_invalid_json_and_non_admin_fail(self): response = await self.client.post("/v1/admin/projects/provision", content=b'{') self.assertEqual(response.status_code, 422) self.app.dependency_overrides[get_current_user] = lambda: {"username": "viewer", "roles": ["operator"]} with patch.object(projects_api, "normalize_provision_payload") as normalize: response = await self.client.post("/v1/admin/projects/provision", json=self.body) self.assertEqual(response.status_code, 403) normalize.assert_not_called() if __name__ == "__main__": unittest.main()