mirror of
https://github.com/opendatalab/MinerU.git
synced 2026-08-30 17:12:39 +08:00
refactor: implement SourceFileStore for managing temporary upload storage and improve upload handling
This commit is contained in:
@@ -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
@@ -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")
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user