"""Request-scoped file import: spool, validate, batch-write, report progress. No daemon, persisted job queue or startup changes. The normal project transaction/compensation boundary publishes only a completely imported graph. """ from __future__ import annotations import asyncio from functools import partial import json import logging from pathlib import Path import shutil import tempfile from fastapi import APIRouter, HTTPException, Request from fastapi.responses import StreamingResponse from starlette.datastructures import FormData, UploadFile from starlette.formparsers import MultiPartException, MultiPartParser from app.auth import AdminUser from app.config import settings from app.graph_import_files import import_staged_graph, validate_graph_file from app.project_lifecycle import ProjectValidationError, normalize_provision_header router = APIRouter() logger = logging.getLogger(__name__) SCHEMA_BYTES = 2 * 1024 * 1024 MIN_FREE_BYTES = 512 * 1024 * 1024 # Bound concurrent upload/index allocations within each existing API process. _import_lock = asyncio.Lock() def event(kind: str, data: dict) -> str: return json.dumps({"event": kind, "data": data}, ensure_ascii=False, default=str) + "\n" async def read_import_form(request: Request) -> FormData: limit = settings.graph_import_max_bytes + SCHEMA_BYTES + 65536 length = request.headers.get("content-length", "") if length.isdigit(): if int(length) > limit: raise HTTPException(413, "图谱文件超过服务器文件大小限制") free = shutil.disk_usage(tempfile.gettempdir()).free needed = int(length) * 3 + MIN_FREE_BYTES if free < needed: raise HTTPException(507, f"服务器临时磁盘空间不足(可用 {free // 1048576} MiB,预计至少需要 {needed // 1048576} MiB),未接收文件") if not request.headers.get("content-type", "").lower().startswith("multipart/form-data"): raise HTTPException(415, "请使用文件上传,而不是把整个图谱放入 JSON 请求体") status = 400 async def bounded_stream(): nonlocal status received = 0 async for block in request.stream(): received += len(block) if received > limit: status = 413 raise MultiPartException("图谱文件超过服务器文件大小限制") if shutil.disk_usage(tempfile.gettempdir()).free < len(block) + MIN_FREE_BYTES: status = 507 raise MultiPartException("服务器临时磁盘空间不足,上传已停止") yield block parser = MultiPartParser(request.headers, bounded_stream(), max_files=1, max_fields=1, max_part_size=SCHEMA_BYTES) try: return await parser.parse() except BaseException as exc: # Starlette closes partial files on MultiPartException, but a network # disconnect/cancellation can also interrupt parsing. Close those too. for stream in parser._files_to_close_on_error: stream.close() if isinstance(exc, MultiPartException): raise HTTPException(status, exc.message) from exc raise @router.get("/projects/import-options") async def import_options(_user: AdminUser): return {"max_file_bytes": settings.graph_import_max_bytes, "editor_max_bytes": SCHEMA_BYTES} @router.post("/projects/provision-file") async def provision_file(request: Request, user: AdminUser): # Admin authentication finishes before reading the potentially large body. if _import_lock.locked(): raise HTTPException(429, "已有文件正在上传或导入,请等待完成后再试") await _import_lock.acquire() form = None handed_to_response = False try: form = await read_import_form(request) source, metadata = form.get("file"), form.get("metadata") if not isinstance(source, UploadFile) or not str(source.filename or "").lower().endswith(".json"): raise HTTPException(422, "请选择 .json 图谱数据文件") if not source.size or source.size > settings.graph_import_max_bytes: raise HTTPException(413, "图谱文件为空或超过服务器文件大小限制") if not isinstance(metadata, str): raise HTTPException(422, "缺少项目与 Schema 信息") try: body = json.loads(metadata) if not isinstance(body, dict): raise ValueError("项目信息必须是对象") payload = await asyncio.to_thread(normalize_provision_header, body) except (ValueError, ProjectValidationError) as exc: raise HTTPException(422, {"message": "项目或 Schema 校验失败", "errors": getattr(exc, "errors", [str(exc)])[:100]}) from exc async def response_stream(): from app.api.projects import _run_async_cleanup_to_completion, provision_validated_project queue: asyncio.Queue[dict] = asyncio.Queue(maxsize=1) loop = asyncio.get_running_loop() latest = {"phase": "服务器校验 JSON", "nodes": 0, "relations": 0} def progress(phase: str, nodes: int, relations: int): free = shutil.disk_usage(tempfile.gettempdir()).free if free < MIN_FREE_BYTES: logger.warning("graph import disk reserve reached: project=%s phase=%s free_bytes=%s", payload["project_id"], phase, free) raise RuntimeError(f"服务器临时磁盘空间不足({phase}阶段可用 {free // 1048576} MiB,须保留 512 MiB),未发布图谱") def publish(): if queue.full(): queue.get_nowait() queue.put_nowait({"phase": phase, "nodes": nodes, "relations": relations}) loop.call_soon_threadsafe(publish) async def execute(): # A private, request-owned index validates references before # any graph write. No uploaded path can select server files. with tempfile.TemporaryDirectory(prefix="kg-graph-import-") as directory: from app.api.projects import _run_blocking_to_completion prepared = await _run_blocking_to_completion( validate_graph_file, source.file, Path(directory) / "index.sqlite", payload, progress, ) return await provision_validated_project( prepared, user, graph_importer=partial(import_staged_graph, progress=progress), ) operation = asyncio.create_task(execute()) try: yield event("progress", latest) while not operation.done(): try: latest = await asyncio.wait_for(queue.get(), 1) except TimeoutError: pass yield event("progress", latest) result = await operation logger.info("graph file import completed: project=%s counts=%s", payload["project_id"], result.get("counts")) yield event("result", result) except Exception as exc: # Do not log uploaded records, credentials or the JSON body. logger.warning("graph file import failed: project=%s phase=%s error_type=%s", payload["project_id"], latest["phase"], type(exc).__name__) if isinstance(exc, HTTPException): detail = exc.detail if isinstance(exc.detail, dict) else {"message": str(exc.detail)} else: detail = {"message": "图谱导入失败", "errors": getattr(exc, "errors", [str(exc)])[:100]} yield event("error", detail) finally: # Do not delete a spool while a DB worker still reads it. An # interrupted browser must check the project list before retry. async def cleanup(): await asyncio.gather(operation, return_exceptions=True) try: await form.close() finally: _import_lock.release() await _run_async_cleanup_to_completion(cleanup()) response = StreamingResponse(response_stream(), media_type="application/x-ndjson", headers={"Cache-Control": "no-store", "X-Accel-Buffering": "no"}) handed_to_response = True return response finally: if not handed_to_response: try: if form is not None: await form.close() finally: _import_lock.release()