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

321 lines
16 KiB
Python

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