diff --git a/mineru/kit/router/app.py b/mineru/kit/router/app.py index f8e69fa0..df496601 100644 --- a/mineru/kit/router/app.py +++ b/mineru/kit/router/app.py @@ -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 diff --git a/mineru/kit/router/proxy.py b/mineru/kit/router/proxy.py index 47134c24..b44f07da 100644 --- a/mineru/kit/router/proxy.py +++ b/mineru/kit/router/proxy.py @@ -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") diff --git a/mineru/kit/router/resources.py b/mineru/kit/router/resources.py index 42d44e80..6e9ba8d9 100644 --- a/mineru/kit/router/resources.py +++ b/mineru/kit/router/resources.py @@ -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", +] diff --git a/tests/unittest/test_v1_router.py b/tests/unittest/test_v1_router.py index db87f906..9b9f4614 100644 --- a/tests/unittest/test_v1_router.py +++ b/tests/unittest/test_v1_router.py @@ -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: