168 lines
7.9 KiB
Python
168 lines
7.9 KiB
Python
#!/usr/bin/env python3
|
||
"""Exercise a real JSON file through multipart + streaming validation.
|
||
|
||
Optional --import-falkor writes only a UUID-named local test graph and removes
|
||
it in finally. System project metadata/auth are test doubles; this is NOT a
|
||
production importer, nor evidence of a browser/proxy/production deployment test.
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import asyncio
|
||
import json
|
||
from pathlib import Path
|
||
import resource
|
||
import shutil
|
||
import sys
|
||
import tempfile
|
||
import time
|
||
import uuid
|
||
from unittest.mock import patch
|
||
|
||
import httpx
|
||
from fastapi import FastAPI
|
||
|
||
ROOT = Path(__file__).resolve().parents[1]
|
||
if str(ROOT) not in sys.path:
|
||
sys.path.insert(0, str(ROOT))
|
||
|
||
from app.api import graph_file_import as file_api
|
||
from app.api.plaza import map_pois
|
||
from app.auth import get_current_user
|
||
from app.config import settings
|
||
from app.project_context import ProjectContext
|
||
from app.project_lifecycle import _falkor_client
|
||
|
||
|
||
def graph_names(db) -> set[str]:
|
||
return {name.decode() if isinstance(name, bytes) else str(name) for name in db.list_graphs()}
|
||
|
||
|
||
async def verify(args):
|
||
if settings.falkordb_host not in {"localhost", "127.0.0.1", "::1"}:
|
||
raise RuntimeError("验证仅允许本机 FalkorDB,不连接生产服务器")
|
||
if args.report.exists():
|
||
raise FileExistsError("报告已存在,请选择新路径")
|
||
if args.schema.stat().st_size > file_api.SCHEMA_BYTES:
|
||
raise ValueError("Schema 超过大小上限")
|
||
schema = json.loads(args.schema.read_text(encoding="utf-8"))
|
||
scratch = "codex_upload_verify_" + uuid.uuid4().hex
|
||
metadata = {"project_id": scratch, "display_name": "城市知识图谱-贵州*贵阳(本地导入验证)",
|
||
"schema": schema, "spatial_map": {"enabled": True, "scope": "guizhou",
|
||
"region_name": "贵阳市", "region_adcode": "520100", "region_level": "city"}}
|
||
app = FastAPI()
|
||
app.include_router(file_api.router, prefix="/v1/admin")
|
||
# Isolated in-process test app only; never bypass the running app's login.
|
||
app.dependency_overrides[get_current_user] = lambda: {"username": "test_only", "roles": ["admin"]}
|
||
|
||
@app.get("/health")
|
||
async def health():
|
||
return {"ok": True}
|
||
|
||
report = {"file_bytes": args.data.stat().st_size, "metadata_and_auth": "isolated_test_doubles",
|
||
"transport": "HTTPX ASGI multipart", "falkordb_import": bool(args.import_falkor),
|
||
"temp_free_before_bytes": shutil.disk_usage(tempfile.gettempdir()).free}
|
||
created = False
|
||
db = _falkor_client() if args.import_falkor else None
|
||
existing_graphs = graph_names(db) if db else set()
|
||
if scratch in existing_graphs:
|
||
raise FileExistsError("临时图名冲突,未开始测试")
|
||
if db:
|
||
report["existing_graphs_before"] = len(existing_graphs)
|
||
real_validator = file_api.validate_graph_file
|
||
last_print = time.monotonic()
|
||
|
||
def validate(source, staging, body, progress):
|
||
def observed(phase, nodes, relations):
|
||
nonlocal last_print
|
||
size = staging.stat().st_size if staging.exists() else 0
|
||
report["peak_staging_bytes"] = max(report.get("peak_staging_bytes", 0), size)
|
||
report["last_progress"] = {"phase": phase, "nodes": nodes, "relations": relations,
|
||
"staging_bytes": size, "temp_free_bytes": shutil.disk_usage(tempfile.gettempdir()).free}
|
||
progress(phase, nodes, relations)
|
||
if time.monotonic() - last_print > 8:
|
||
print(json.dumps({"phase": phase, "nodes": nodes, "relations": relations}, ensure_ascii=False), flush=True)
|
||
last_print = time.monotonic()
|
||
return real_validator(source, staging, body, observed)
|
||
|
||
async def publish(payload, user, *, graph_importer):
|
||
nonlocal created
|
||
report["validated_counts"] = payload["counts"]
|
||
report["file_sha256"] = payload["file_sha256"]
|
||
if db:
|
||
created = True
|
||
counts = await asyncio.to_thread(graph_importer, scratch, payload["graph_data"])
|
||
report["actual_graph_counts"] = counts
|
||
assert counts["nodes"] == payload["counts"]["nodes"]
|
||
assert counts["relations"] == payload["counts"]["relations"]
|
||
return {"project": {"project_id": scratch}, "counts": payload["counts"]}
|
||
|
||
started = time.monotonic()
|
||
latencies = []
|
||
failure = None
|
||
try:
|
||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://local-test", timeout=None) as client:
|
||
with args.data.open("rb") as source, patch.object(file_api, "validate_graph_file", side_effect=validate), patch(
|
||
"app.api.projects.provision_validated_project", side_effect=publish,
|
||
):
|
||
upload = asyncio.create_task(client.post(
|
||
"/v1/admin/projects/provision-file", data={"metadata": json.dumps(metadata, ensure_ascii=False)},
|
||
files={"file": (args.data.name, source, "application/json")},
|
||
))
|
||
while not upload.done():
|
||
t = time.monotonic()
|
||
health = await client.get("/health")
|
||
assert health.status_code == 200
|
||
latencies.append(round((time.monotonic() - t) * 1000, 2))
|
||
await asyncio.sleep(0.5)
|
||
response = await upload
|
||
report["status_code"] = response.status_code
|
||
if response.status_code != 200:
|
||
raise RuntimeError(response.text[:1500])
|
||
messages = [json.loads(line) for line in response.text.splitlines()]
|
||
report["response_events"] = len(messages)
|
||
if not messages or messages[-1]["event"] != "result":
|
||
raise RuntimeError(str(messages[-1] if messages else "没有返回结果"))
|
||
if db:
|
||
pois = await map_pois(context=ProjectContext("local_test", scratch, scratch), _user=None)
|
||
report["map_pois"] = pois["total"]
|
||
assert pois["source"] == "falkordb"
|
||
assert pois["total"] == report["validated_counts"]["map_pois"]
|
||
del pois
|
||
report["passed"] = True
|
||
except Exception as exc:
|
||
failure = exc
|
||
report.update(passed=False, error=str(exc)[:2000])
|
||
finally:
|
||
if db:
|
||
try:
|
||
if created and scratch in graph_names(db):
|
||
db.select_graph(scratch).delete()
|
||
final_graphs = graph_names(db)
|
||
if scratch in final_graphs:
|
||
raise RuntimeError("UUID 临时图清理失败")
|
||
if not existing_graphs.issubset(final_graphs):
|
||
raise RuntimeError("验证期间发现原有图谱缺失")
|
||
report.update(test_graph_removed=True, existing_graphs_preserved=len(existing_graphs))
|
||
finally:
|
||
db.close()
|
||
report.update(seconds=round(time.monotonic() - started, 2),
|
||
peak_rss_bytes=resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * (1 if sys.platform == "darwin" else 1024),
|
||
health_checks=len(latencies), max_health_latency_ms=max(latencies, default=0),
|
||
temp_free_after_bytes=shutil.disk_usage(tempfile.gettempdir()).free)
|
||
args.report.parent.mkdir(parents=True, exist_ok=True)
|
||
with args.report.open("x", encoding="utf-8") as stream:
|
||
json.dump(report, stream, ensure_ascii=False, indent=2)
|
||
print(json.dumps(report, ensure_ascii=False, indent=2), flush=True)
|
||
if failure:
|
||
raise RuntimeError("验证未通过,原因及资源占用已写入报告") from failure
|
||
|
||
|
||
if __name__ == "__main__":
|
||
parser = argparse.ArgumentParser(description=__doc__)
|
||
parser.add_argument("--schema", type=Path, required=True)
|
||
parser.add_argument("--data", type=Path, required=True)
|
||
parser.add_argument("--report", type=Path, required=True)
|
||
parser.add_argument("--import-falkor", action="store_true")
|
||
asyncio.run(verify(parser.parse_args()))
|