refactor: implement SourceFileStore for managing temporary upload storage and improve upload handling

This commit is contained in:
myhloli
2026-08-26 01:49:02 +08:00
parent 41bb5ac2da
commit faca0d60c1
4 changed files with 304 additions and 48 deletions
+51 -8
View File
@@ -26,7 +26,7 @@ from .proxy import (
router_error_response,
stream_upstream,
)
from .resources import ResourceKind, ResourceRegistry, ResourceRoute
from .resources import ResourceKind, ResourceRegistry, ResourceRoute, SourceFileStore, stored_file_chunks
from .workers import RouterSettings, WorkerPool, WorkerState
@@ -187,17 +187,19 @@ def create_app(
"""创建完整代理 MinerU V1 资源面的 Router FastAPI 应用。"""
resolved_settings = settings or RouterSettings.from_env()
registry = ResourceRegistry()
source_store = SourceFileStore()
pool = WorkerPool(resolved_settings, transport=transport)
@asynccontextmanager
async def _lifespan(application: FastAPI) -> AsyncIterator[None]:
"""启动 worker pool,并在应用关闭时释放全部网络和子进程资源。"""
application.state.started_at = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
await pool.start()
try:
await pool.start()
yield
finally:
await pool.close()
source_store.close()
application = FastAPI(
title="MinerU V1 Router",
@@ -206,6 +208,7 @@ def create_app(
)
application.state.settings = resolved_settings
application.state.registry = registry
application.state.source_store = source_store
application.state.worker_pool = pool
application.state.started_at = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
@@ -282,11 +285,21 @@ def create_app(
if worker is None:
raise RouterProxyError(503, "upstream_unavailable", "No upstream accepts file uploads")
body = await request.json()
upstream = await request_upstream(pool, worker, "POST", "/v1/uploads", request=request, json_body=body)
if not isinstance(body, dict):
raise RouterProxyError(400, "invalid_request", "Upload request body must be an object")
upstream_body = dict(body)
# Router 必须收到源字节才能跨 worker 转移,因此禁止 upstream 在 PUT 前按 sha256 提前完成。
upstream_body.pop("sha256sum", None)
upstream = await request_upstream(pool, worker, "POST", "/v1/uploads", request=request, json_body=upstream_body)
result = _successful_json(upstream)
if isinstance(result, Response):
return result
return JSONResponse(rewrite_upload_payload(result, worker, registry), status_code=upstream.status_code)
rewritten = rewrite_upload_payload(result, worker, registry)
if isinstance(body.get("sha256sum"), str):
rewritten["sha256sum"] = body["sha256sum"]
route = _route_or_404(registry, "upload", str(rewritten["id"]))
route.metadata["declared"] = copy.deepcopy(body)
return JSONResponse(rewritten, status_code=upstream.status_code)
@application.get("/v1/uploads/{upload_id}")
async def get_upload(upload_id: str, request: Request) -> Response:
@@ -301,16 +314,32 @@ def create_app(
@application.put("/v1/uploads/{upload_id}/content")
async def upload_content(upload_id: str, request: Request) -> Response:
"""把上传内容写入 upload 所属 worker。"""
"""把上传内容流式暂存到 Router,再写入 upload 所属 worker。"""
route = _route_or_404(registry, "upload", upload_id)
worker = pool.get(route.worker_id)
declared = route.metadata.get("declared") if isinstance(route.metadata.get("declared"), dict) else {}
mime_type = str(declared.get("mime_type") or request.headers.get("content-type") or "application/octet-stream")
stored = await source_store.stage_upload(upload_id, request.stream(), mime_type=mime_type)
declared_bytes = declared.get("bytes")
if isinstance(declared_bytes, int) and stored.bytes != declared_bytes:
source_store.discard_upload(upload_id)
raise RouterProxyError(
400,
"upload_size_mismatch",
f"Expected {declared_bytes} upload bytes, received {stored.bytes}",
)
declared_sha256 = declared.get("sha256sum")
if isinstance(declared_sha256, str) and stored.sha256sum != declared_sha256:
source_store.discard_upload(upload_id)
raise RouterProxyError(400, "upload_sha256_mismatch", "Uploaded content SHA256 does not match request")
upstream = await request_upstream(
pool,
worker,
"PUT",
f"/v1/uploads/{route.upstream_id}/content",
request=request,
content=await request.body(),
content=stored_file_chunks(stored.path),
headers={"content-type": "application/octet-stream"},
)
return passthrough_response(upstream)
@@ -319,8 +348,14 @@ def create_app(
"""完成所属 worker 的 upload,并注册返回的 V1 file。"""
route = _route_or_404(registry, "upload", upload_id)
worker = pool.get(route.worker_id)
stored = source_store.find_upload(upload_id)
if stored is None:
raise RouterProxyError(409, "upload_not_ready", "Upload bytes have not been received by Router")
raw_body = await request.body()
body = await request.json() if raw_body else None
body = await request.json() if raw_body else {}
if not isinstance(body, dict):
raise RouterProxyError(400, "invalid_request", "Upload complete body must be an object")
body.setdefault("sha256sum", stored.sha256sum)
upstream = await request_upstream(
pool,
worker,
@@ -332,7 +367,12 @@ def create_app(
result = _successful_json(upstream)
if isinstance(result, Response):
return result
return JSONResponse(rewrite_upload_payload(result, worker, registry), status_code=upstream.status_code)
rewritten = rewrite_upload_payload(result, worker, registry)
file_payload = rewritten.get("file")
if not isinstance(file_payload, dict) or not isinstance(file_payload.get("id"), str):
raise RouterProxyError(502, "invalid_upstream_response", "Completed upload did not return a file")
source_store.bind_file(upload_id, file_payload["id"])
return JSONResponse(rewritten, status_code=upstream.status_code)
@application.post("/v1/uploads/{upload_id}/cancel")
async def cancel_upload(upload_id: str, request: Request) -> Response:
@@ -349,6 +389,7 @@ def create_app(
result = _successful_json(upstream)
if isinstance(result, Response):
return result
source_store.discard_upload(upload_id)
return JSONResponse(rewrite_upload_payload(result, worker, registry), status_code=upstream.status_code)
@application.get("/v1/files")
@@ -394,6 +435,7 @@ def create_app(
if isinstance(result, Response):
return result
registry.remove("file", file_id)
source_store.delete_file(file_id)
result["id"] = file_id
return JSONResponse(result, status_code=upstream.status_code)
@@ -448,6 +490,7 @@ def create_app(
request=request,
pool=pool,
registry=registry,
source_store=source_store,
)
copied_file_ids.append(target_file_id)
input_aliases[target_file_id] = route.public_id
+47 -26
View File
@@ -11,7 +11,7 @@ import httpx
from fastapi import Request
from fastapi.responses import JSONResponse, Response, StreamingResponse
from .resources import ResourceRegistry, ResourceRoute
from .resources import ResourceRegistry, ResourceRoute, SourceFileStore, stored_file_chunks
from .workers import WorkerPool, WorkerState
_HOP_BY_HOP_HEADERS = frozenset(
@@ -73,7 +73,7 @@ async def request_upstream(
*,
request: Request | None = None,
json_body: Any = None,
content: bytes | None = None,
content: bytes | AsyncIterator[bytes] | None = None,
headers: Mapping[str, str] | None = None,
) -> httpx.Response:
"""执行普通 upstream 请求,并把连接失败与超时转成稳定 Router 错误。"""
@@ -256,32 +256,53 @@ async def copy_file_to_worker(
request: Request,
pool: WorkerPool,
registry: ResourceRegistry,
source_store: SourceFileStore,
) -> str:
"""通过 V1 Files/Uploads API 把已有输入文件复制到目标 worker。"""
"""优先使用 Router 私有暂存输入,通过 Upload API 复制到目标 worker。"""
source = pool.get(route.worker_id)
metadata_response = await request_upstream(
pool,
source,
"GET",
f"/v1/files/{route.upstream_id}",
request=request,
)
metadata = json_or_error(metadata_response)
content_response = await request_upstream(
pool,
source,
"GET",
f"/v1/files/{route.upstream_id}/content",
request=request,
)
if content_response.status_code >= 400:
json_or_error(content_response)
metadata = route.metadata.get("payload")
if not isinstance(metadata, dict):
metadata_response = await request_upstream(
pool,
source,
"GET",
f"/v1/files/{route.upstream_id}",
request=request,
)
metadata = json_or_error(metadata_response)
stored = source_store.find_file(route.public_id)
if stored is not None:
content: bytes | AsyncIterator[bytes] = stored_file_chunks(stored.path)
content_size = stored.bytes
content_sha256 = stored.sha256sum
content_type = stored.mime_type
else:
if metadata.get("purpose") != "parse_output":
raise RouterProxyError(
409,
"source_file_unavailable",
f"Router source bytes for {route.public_id} are no longer available",
)
content_response = await request_upstream(
pool,
source,
"GET",
f"/v1/files/{route.upstream_id}/content",
request=request,
)
if content_response.status_code >= 400:
json_or_error(content_response)
content = content_response.content
content_size = len(content_response.content)
content_sha256 = str(metadata.get("sha256sum") or "")
content_type = content_response.headers.get("content-type") or "application/octet-stream"
purpose = metadata.get("purpose") if metadata.get("purpose") in {"parse", "input_image"} else "parse"
create_payload = {
"filename": metadata.get("filename") or "input.bin",
"bytes": len(content_response.content),
"mime_type": content_response.headers.get("content-type") or "application/octet-stream",
"purpose": metadata.get("purpose") or "parse",
**({"sha256sum": metadata["sha256sum"]} if metadata.get("sha256sum") else {}),
"bytes": content_size,
"mime_type": content_type,
"purpose": purpose,
**({"sha256sum": content_sha256} if content_sha256 else {}),
}
upload_response = await request_upstream(
pool,
@@ -303,7 +324,7 @@ async def copy_file_to_worker(
"PUT",
f"/v1/uploads/{upload_id}/content",
request=request,
content=content_response.content,
content=content,
headers={"content-type": "application/octet-stream"},
)
if put_response.status_code >= 400:
@@ -314,7 +335,7 @@ async def copy_file_to_worker(
"POST",
f"/v1/uploads/{upload_id}/complete",
request=request,
json_body={"sha256sum": metadata.get("sha256sum")} if metadata.get("sha256sum") else {},
json_body={"sha256sum": content_sha256} if content_sha256 else {},
)
complete_payload = json_or_error(complete_response)
file_payload = complete_payload.get("file")
+104 -1
View File
@@ -3,9 +3,13 @@
from __future__ import annotations
import hashlib
import secrets
import tempfile
from collections.abc import AsyncIterator
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Literal
ResourceKind = Literal["upload", "file", "job"]
@@ -34,6 +38,97 @@ class ResourceRoute:
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(frozen=True)
class StoredSourceFile:
"""记录 Router 私有暂存输入的路径、大小、哈希和媒体类型。"""
path: Path
bytes: int
sha256sum: str
mime_type: str
class SourceFileStore:
"""在 Router 临时目录中保存可供 cross-worker 重传的输入字节。"""
def __init__(self) -> None:
"""创建由当前 Router 进程独占并在关闭时清理的临时目录。"""
self._temp_dir = tempfile.TemporaryDirectory(prefix="mineru-v1-router-sources-")
self._root = Path(self._temp_dir.name)
self._uploads: dict[str, StoredSourceFile] = {}
self._files: dict[str, StoredSourceFile] = {}
async def stage_upload(
self,
upload_id: str,
chunks: AsyncIterator[bytes],
*,
mime_type: str,
) -> StoredSourceFile:
"""流式写入一个公共 Upload 的输入字节,并计算实际大小与 SHA256。"""
path = self._root / "uploads" / upload_id
path.parent.mkdir(parents=True, exist_ok=True)
hasher = hashlib.sha256()
byte_count = 0
with path.open("wb") as output:
async for chunk in chunks:
if not chunk:
continue
output.write(chunk)
hasher.update(chunk)
byte_count += len(chunk)
stored = StoredSourceFile(
path=path,
bytes=byte_count,
sha256sum=hasher.hexdigest(),
mime_type=mime_type,
)
previous = self._uploads.get(upload_id)
self._uploads[upload_id] = stored
if previous is not None and previous.path != path:
previous.path.unlink(missing_ok=True)
return stored
def bind_file(self, upload_id: str, file_id: str) -> StoredSourceFile:
"""把已完成 Upload 的暂存输入绑定到 Router 公共 File。"""
stored = self._uploads.pop(upload_id)
self._files[file_id] = stored
return stored
def find_upload(self, upload_id: str) -> StoredSourceFile | None:
"""读取公共 Upload 对应的私有暂存输入,不存在时返回 None。"""
return self._uploads.get(upload_id)
def find_file(self, file_id: str) -> StoredSourceFile | None:
"""读取公共 File 对应的私有暂存输入,不存在时返回 None。"""
return self._files.get(file_id)
def discard_upload(self, upload_id: str) -> None:
"""删除取消或失败 Upload 的私有暂存输入。"""
stored = self._uploads.pop(upload_id, None)
if stored is not None:
stored.path.unlink(missing_ok=True)
def delete_file(self, file_id: str) -> None:
"""删除公共 File 绑定的私有暂存输入。"""
stored = self._files.pop(file_id, None)
if stored is not None:
stored.path.unlink(missing_ok=True)
def close(self) -> None:
"""清理当前 Router 进程的全部暂存输入。"""
self._uploads.clear()
self._files.clear()
self._temp_dir.cleanup()
async def stored_file_chunks(path: Path, chunk_size: int = 1024 * 1024) -> AsyncIterator[bytes]:
"""按固定块大小异步迭代一个 Router 私有暂存文件。"""
with path.open("rb") as source:
while chunk := source.read(chunk_size):
yield chunk
class ResourceRegistry:
"""维护当前 Router 进程创建或发现的 uploads、files 与 jobs。"""
@@ -108,4 +203,12 @@ class ResourceRegistry:
return _PUBLIC_ID_PREFIXES[kind] + secrets.token_hex(12)
__all__ = ["ResourceKind", "ResourceRegistry", "ResourceRoute", "utc_now_iso"]
__all__ = [
"ResourceKind",
"ResourceRegistry",
"ResourceRoute",
"SourceFileStore",
"StoredSourceFile",
"stored_file_chunks",
"utc_now_iso",
]
+102 -13
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import asyncio
import hashlib
import json
import tomllib
from dataclasses import dataclass, field
@@ -12,6 +13,7 @@ from fastapi.testclient import TestClient
from mineru.kit.router import RouterSettings, create_app
from mineru.kit.router.workers import ManagedLocalWorker
from mineru.parser.api_server import create_app as create_api_server_app
@dataclass
@@ -27,8 +29,9 @@ class _FakeV1Upstream:
files: dict[str, dict[str, Any]] = field(default_factory=dict)
jobs: dict[str, dict[str, Any]] = field(default_factory=dict)
fail_jobs: bool = False
source_download_attempts: int = 0
def handle(self, request: httpx.Request) -> httpx.Response:
async def handle(self, request: httpx.Request) -> httpx.Response:
"""按请求 path/method 返回真实 V1 shape 或测试错误。"""
path = request.url.path
method = request.method
@@ -67,20 +70,20 @@ class _FakeV1Upstream:
if path == "/v1/uploads" and method == "POST":
self.upload_counter += 1
upload_id = f"upload_{self.name}_{self.upload_counter}"
body = self._json_body(request)
body = await self._json_body(request)
self.uploads[upload_id] = {**body, "content": b"", "status": "pending"}
return self._response(request, 200, self._upload_payload(upload_id))
if path.startswith("/v1/uploads/"):
return self._handle_upload(request)
return await self._handle_upload(request)
if path.startswith("/v1/files/"):
return self._handle_file(request)
if path == "/v1/parse/jobs" and method == "POST":
return self._create_job(request)
return await self._create_job(request)
if path.startswith("/v1/parse/jobs/"):
return self._handle_job(request)
return self._response(request, 404, self._error("not_found", path))
def _handle_upload(self, request: httpx.Request) -> httpx.Response:
async def _handle_upload(self, request: httpx.Request) -> httpx.Response:
"""处理 upload 查询、内容写入、完成和取消。"""
parts = request.url.path.split("/")
upload_id = parts[3]
@@ -89,7 +92,7 @@ class _FakeV1Upstream:
return self._response(request, 404, self._error("upload_not_found", upload_id))
action = parts[4] if len(parts) > 4 else ""
if request.method == "PUT" and action == "content":
upload["content"] = request.read()
upload["content"] = await request.aread()
return httpx.Response(200, request=request)
if request.method == "POST" and action == "complete":
upload["status"] = "completed"
@@ -120,6 +123,13 @@ class _FakeV1Upstream:
if file_record is None:
return self._response(request, 404, self._error("file_not_found", file_id))
if len(parts) > 4 and parts[4] == "content":
if file_record["purpose"] != "parse_output":
self.source_download_attempts += 1
return self._response(
request,
403,
self._error("feature_requires_api_key", "Source files cannot be downloaded"),
)
return httpx.Response(
200,
content=file_record["content"],
@@ -131,11 +141,11 @@ class _FakeV1Upstream:
return self._response(request, 200, {"id": file_id, "object": "file", "deleted": True})
return self._response(request, 200, {key: value for key, value in file_record.items() if key != "content"})
def _create_job(self, request: httpx.Request) -> httpx.Response:
async def _create_job(self, request: httpx.Request) -> httpx.Response:
"""创建 queued job,并校验请求中的 file_id 已复制到当前 upstream。"""
if self.fail_jobs:
raise httpx.ConnectError("simulated disconnect", request=request)
body = self._json_body(request)
body = await self._json_body(request)
for entry in body.get("files") or []:
source = entry.get("source") or {}
if source.get("type") == "file_id" and source.get("file_id") not in self.files:
@@ -226,9 +236,9 @@ class _FakeV1Upstream:
return payload
@staticmethod
def _json_body(request: httpx.Request) -> dict[str, Any]:
async def _json_body(request: httpx.Request) -> dict[str, Any]:
"""读取 MockTransport 请求中的 JSON object body。"""
body = request.read()
body = await request.aread()
return json.loads(body.decode("utf-8")) if body else {}
@staticmethod
@@ -248,9 +258,9 @@ def _make_router(
"""创建使用 host 分发 MockTransport 的 Router 应用。"""
by_host = {upstream.name: upstream for upstream in upstreams}
def _handler(request: httpx.Request) -> httpx.Response:
async def _handler(request: httpx.Request) -> httpx.Response:
"""把请求交给 URL host 对应的 fake upstream。"""
return by_host[request.url.host].handle(request)
return await by_host[request.url.host].handle(request)
settings = RouterSettings(
upstream_urls=tuple(f"http://{upstream.name}" for upstream in upstreams),
@@ -263,13 +273,22 @@ def _make_router(
def _upload_file(client: TestClient, *, token: str, content: bytes) -> str:
"""通过 Router 完成一次 V1 upload 并返回公共 file ID。"""
headers = {"authorization": f"Bearer {token}"}
sha256sum = hashlib.sha256(content).hexdigest()
create = client.post(
"/v1/uploads",
headers=headers,
json={"filename": f"{token}.pdf", "bytes": len(content), "mime_type": "application/pdf", "purpose": "parse"},
json={
"filename": f"{token}.pdf",
"bytes": len(content),
"mime_type": "application/pdf",
"purpose": "parse",
"sha256sum": sha256sum,
},
)
assert create.status_code == 200, create.text
upload_id = create.json()["id"]
assert create.json()["status"] == "pending"
assert create.json()["sha256sum"] == sha256sum
assert create.json()["upload_url"] == f"/v1/uploads/{upload_id}/content"
assert client.put(f"/v1/uploads/{upload_id}/content", headers=headers, content=content).status_code == 200
complete = client.post(f"/v1/uploads/{upload_id}/complete", headers=headers, json={})
@@ -282,6 +301,7 @@ def test_v1_router_aggregates_capabilities_and_routes_cross_worker_job() -> None
first = _FakeV1Upstream("worker-a", ("basic", "standard"))
second = _FakeV1Upstream("worker-b", ("standard", "advanced"))
router_app, _ = _make_router(first, second)
staged_path: Path | None = None
with TestClient(router_app) as client:
health = client.get("/v1/health")
@@ -303,6 +323,10 @@ def test_v1_router_aggregates_capabilities_and_routes_cross_worker_job() -> None
break
assert len(files_by_worker) == 2
public_file_ids = list(files_by_worker.values())
stored = router_app.state.source_store.find_file(public_file_ids[0])
assert stored is not None
staged_path = stored.path
assert staged_path.is_file()
created = client.post(
"/v1/parse/jobs",
@@ -336,6 +360,71 @@ def test_v1_router_aggregates_capabilities_and_routes_cross_worker_job() -> None
submitted_files = [entry["source"]["file_id"] for entry in next(iter(selected_job.values()))["body"]["files"]]
selected_files = first.files if first.jobs else second.files
assert all(file_id in selected_files for file_id in submitted_files)
assert first.source_download_attempts == 0
assert second.source_download_attempts == 0
assert staged_path is not None
assert not staged_path.exists()
def test_real_v1_api_rejects_public_download_of_parse_source(tmp_path: Path) -> None:
"""验证真实 api-server 保持普通 parse 输入不可公开下载的安全边界。"""
data = b"%PDF-1.7"
api_app = create_api_server_app(upload_dir=tmp_path, tier="flash")
with TestClient(api_app) as client:
created = client.post(
"/v1/uploads",
json={
"filename": "input.pdf",
"bytes": len(data),
"mime_type": "application/pdf",
"purpose": "parse",
},
)
upload_id = created.json()["id"]
uploaded = client.put(
f"/v1/uploads/{upload_id}/content",
content=data,
headers={"content-type": "application/octet-stream"},
)
completed = client.post(f"/v1/uploads/{upload_id}/complete", json={})
file_id = completed.json()["file"]["id"]
downloaded = client.get(f"/v1/files/{file_id}/content")
assert uploaded.status_code == 200
assert completed.json()["file"]["purpose"] == "parse"
assert downloaded.status_code == 403
assert downloaded.json()["error"]["message"] == "Source files cannot be downloaded"
def test_router_discards_staged_upload_with_wrong_sha256() -> None:
"""验证 Router 在 SHA256 不匹配时删除暂存输入且不转发字节。"""
upstream = _FakeV1Upstream("worker-a", ("standard",))
router_app, _ = _make_router(upstream)
expected_sha256 = hashlib.sha256(b"ok").hexdigest()
with TestClient(router_app) as client:
created = client.post(
"/v1/uploads",
json={
"filename": "input.pdf",
"bytes": 2,
"mime_type": "application/pdf",
"purpose": "parse",
"sha256sum": expected_sha256,
},
)
upload_id = created.json()["id"]
uploaded = client.put(
f"/v1/uploads/{upload_id}/content",
content=b"no",
headers={"content-type": "application/octet-stream"},
)
assert uploaded.status_code == 400
assert uploaded.json()["error"]["code"] == "upload_sha256_mismatch"
assert router_app.state.source_store.find_upload(upload_id) is None
assert next(iter(upstream.uploads.values()))["content"] == b""
def test_v1_router_reports_capability_transport_and_resource_errors() -> None: