321 lines
16 KiB
Python
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()
|