185 lines
8.7 KiB
Python
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()
|