Files

185 lines
8.7 KiB
Python

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