Files
BiRefNet_WebUI/birefnet_web/server.py
T
2026-10-08 09:46:47 +08:00

938 lines
36 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""HTTP 服务:静态页面 + REST API + 串行任务队列。
刻意不依赖任何 Web 框架(Gradio / FastAPI / Flask 都不需要):
* 传输层 —— 标准库 http.server.ThreadingHTTPServer
* 表单解析 —— 自研 multipart/form-data 解析(`cgi` 模块在 Python 3.13 已移除)
* 并发模型 —— 单 worker 线程串行消费任务,避免 GPU 显存被打爆
"""
from __future__ import annotations
import io
import json
import os
import queue
import threading
import time
import traceback
import uuid
import zipfile
from dataclasses import dataclass, field
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Dict, Iterable, List, Optional, Sequence, Tuple
from urllib.parse import parse_qs, unquote, urlparse
from PIL import Image
from . import APP_NAME, __version__, batch, imageops
from .engine import InferenceEngine, ModelRegistry, describe_environment
# --------------------------------------------------------------------------- #
# 常量
# --------------------------------------------------------------------------- #
MAX_UPLOAD_BYTES = 512 * 1024 * 1024 # 单次请求体上限
MIME_TYPES = {
".html": "text/html; charset=utf-8",
".css": "text/css; charset=utf-8",
".js": "application/javascript; charset=utf-8",
".json": "application/json; charset=utf-8",
".png": "image/png",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".webp": "image/webp",
".bmp": "image/bmp",
".tif": "image/tiff",
".tiff": "image/tiff",
".gif": "image/gif",
".svg": "image/svg+xml",
".ico": "image/x-icon",
".woff2": "font/woff2",
}
#: 允许上传 / 扫描的图片后缀(与批量模块共用同一份定义,避免两处不一致)
IMAGE_SUFFIXES = imageops.IMAGE_SUFFIXES
DEFAULT_OPTIONS: Dict[str, object] = {
"model": "",
"device": "auto",
"dtype": "auto",
"arch": "auto",
"resolution_mode": "square",
"width": 1024,
"height": 1024,
"longest_side": 1024,
"upscale_method": "bilinear",
"mask_threshold": 0.0,
"refine_foreground": True,
"blur_size": 90,
"blur_size_two": 6,
"background": "transparent",
"bg_color": "#ffffff",
"output_mask": False,
"final_longest_side": 0,
}
class TaskCanceled(BaseException):
"""取消信号。
故意继承 BaseException:这样它会穿过引擎里 `except Exception` 的进度回调保护,
直达任务循环,实现阶段级取消。
"""
# --------------------------------------------------------------------------- #
# multipart/form-data 解析
# --------------------------------------------------------------------------- #
@dataclass
class FormPart:
name: str
filename: Optional[str]
content_type: str
data: bytes
def _parse_disposition(value: str) -> Tuple[str, Optional[str]]:
name, filename = "", None
for chunk in value.split(";"):
chunk = chunk.strip()
low = chunk.lower()
if low.startswith("name=") and name == "":
name = chunk[5:].strip().strip('"')
elif low.startswith("filename="):
raw = chunk[9:].strip().strip('"')
if raw:
filename = raw
return name, filename
def parse_multipart(body: bytes, boundary: bytes) -> List[FormPart]:
"""把一个 multipart/form-data 请求体拆成若干 FormPart。"""
parts: List[FormPart] = []
delim = b"--" + boundary
for raw in body.split(delim):
if not raw or raw in (b"--", b"--\r\n", b"\r\n"):
continue
if raw.startswith(b"--"): # 结束标记
break
if raw.startswith(b"\r\n"):
raw = raw[2:]
elif raw.startswith(b"\n"):
raw = raw[1:]
head, sep, data = raw.partition(b"\r\n\r\n")
if not sep:
head, sep, data = raw.partition(b"\n\n")
if not sep:
continue
if data.endswith(b"\r\n"):
data = data[:-2]
elif data.endswith(b"\n"):
data = data[:-1]
headers: Dict[str, str] = {}
for line in head.split(b"\r\n"):
if b":" not in line:
continue
key, _, val = line.partition(b":")
headers[key.strip().lower().decode("latin-1")] = (
val.strip().decode("utf-8", "replace")
)
name, filename = _parse_disposition(headers.get("content-disposition", ""))
if not name:
continue
parts.append(
FormPart(
name=name,
filename=filename,
content_type=headers.get("content-type", "application/octet-stream"),
data=data,
)
)
return parts
# --------------------------------------------------------------------------- #
# 任务模型
# --------------------------------------------------------------------------- #
@dataclass
class ImageJob:
index: int
filename: str
label: str
input_path: str
state: str = "pending" # pending | running | done | failed | canceled
stage: str = "等待中"
progress: float = 0.0
error: Optional[str] = None
width: int = 0
height: int = 0
source_width: int = 0
source_height: int = 0
coverage: float = 0.0
elapsed: float = 0.0
input_size: Tuple[int, int] = (0, 0)
warning: Optional[str] = None
outputs: Dict[str, str] = field(default_factory=dict)
#: 批量模式:写进用户输出目录的文件名(RMBG_<原主干>_<时间戳>.<后缀>)
output_name: Optional[str] = None
def to_dict(self, task_id: str) -> dict:
urls = {
kind: f"/api/tasks/{task_id}/file/{self.index}/{kind}" for kind in self.outputs
}
return {
"index": self.index,
"name": self.filename,
"label": self.label,
"state": self.state,
"stage": self.stage,
"progress": round(self.progress, 3),
"error": self.error,
"width": self.width,
"height": self.height,
"source_width": self.source_width,
"source_height": self.source_height,
"input_size": list(self.input_size),
"coverage": self.coverage,
"elapsed": self.elapsed,
"warning": self.warning,
"output_name": self.output_name,
"urls": urls,
}
@dataclass
class Task:
id: str
created: float
options: dict
images: List[ImageJob]
output_dir: str
#: upload = 网页上传单图(结果落在 outputs/<task_id>);batch = 目录批量(结果落在用户输出目录)
mode: str = "upload"
input_dir: Optional[str] = None
recursive: bool = False
state: str = "queued" # queued | running | done | partial | failed | canceled
stage: str = "排队中"
progress: float = 0.0
error: Optional[str] = None
finished_at: Optional[float] = None
cancel_requested: bool = False
def to_dict(self, include_options: bool = False, tail: int = 0) -> dict:
"""序列化任务。
Args:
include_options: 是否带上参数快照。
tail: 大于 0 时只回传最后 N 张图片的状态。
批量任务动辄几百张,前端轮询必须靠它把载荷压住
(处理是按序串行的,所以窗口外的一定已经是终态,不会漏更新)。
"""
total = len(self.images)
start = max(0, total - int(tail)) if tail and tail > 0 else 0
payload = {
"id": self.id,
"created": self.created,
"mode": self.mode,
"input_dir": self.input_dir,
"output_dir": self.output_dir,
"recursive": self.recursive,
"state": self.state,
"stage": self.stage,
"progress": round(self.progress, 3),
"error": self.error,
"elapsed": round((self.finished_at or time.time()) - self.created, 2),
"counts": self._counts(),
"images_total": total,
"images_from": start,
"images": [img.to_dict(self.id) for img in self.images[start:]],
}
if include_options:
payload["options"] = self.options
return payload
def _counts(self) -> dict:
out = {"total": len(self.images), "done": 0, "failed": 0, "pending": 0}
for img in self.images:
if img.state == "done":
out["done"] += 1
elif img.state == "failed":
out["failed"] += 1
elif img.state != "canceled":
out["pending"] += 1
return out
def refresh_progress(self) -> None:
if not self.images:
self.progress = 0.0
else:
self.progress = sum(i.progress for i in self.images) / len(self.images)
class TaskManager:
"""单 worker 串行任务队列。"""
def __init__(
self,
engine: InferenceEngine,
output_root: Path,
max_tasks: int = 50,
logger=None,
) -> None:
self.engine = engine
self.output_root = Path(output_root)
self.output_root.mkdir(parents=True, exist_ok=True)
self.max_tasks = max(1, int(max_tasks))
self.log = logger or (lambda msg: None)
self._tasks: Dict[str, Task] = {}
self._lock = threading.RLock()
self._queue: "queue.Queue[Optional[Task]]" = queue.Queue()
self._worker = threading.Thread(target=self._work_loop, name="birefnet-worker", daemon=True)
self._worker.start()
# -------------------------- 对外接口 -------------------------- #
def submit(self, uploads: Sequence[Tuple[str, bytes]], options: dict) -> Task:
if not uploads:
raise ValueError("没有收到任何图片")
task_id = time.strftime("%Y%m%d-%H%M%S-") + uuid.uuid4().hex[:4]
output_dir = self.output_root / task_id
input_dir = output_dir / "input"
input_dir.mkdir(parents=True, exist_ok=True)
jobs: List[ImageJob] = []
for idx, (filename, data) in enumerate(uploads):
label = imageops.safe_stem(filename, fallback=f"image_{idx + 1}")
suffix = Path(filename or "").suffix.lower()
if suffix not in IMAGE_SUFFIXES:
suffix = ".png"
stored = input_dir / f"{idx + 1:03d}_{label}{suffix}"
stored.write_bytes(data)
jobs.append(ImageJob(index=idx, filename=filename or stored.name, label=label, input_path=str(stored)))
task = Task(id=task_id, created=time.time(), options=options, images=jobs, output_dir=str(output_dir))
with self._lock:
self._tasks[task_id] = task
self._prune_locked()
self._queue.put(task)
self.log(f"[任务 {task_id}] 已提交,共 {len(jobs)} 张图片")
return task
def submit_batch(
self,
input_dir: Path,
output_dir: Path,
options: dict,
*,
recursive: bool = False,
) -> Task:
"""按目录批量建任务:原图不拷贝,结果直接写进用户指定的输出目录。
Args:
input_dir: 已校验存在的图片目录。
output_dir: 已校验存在的输出目录(本方法不创建任何目录)。
options: 与单图模式完全相同的参数集合。
recursive: 是否包含子目录。
Returns:
已入队的新任务。
Raises:
ValueError: 目录下没有可处理的图片。
"""
# 输入 == 输出时不能把输出目录当跳过项,否则一张都扫不到
skip = [] if Path(input_dir) == Path(output_dir) else [output_dir]
files = batch.scan_images(input_dir, recursive=recursive, skip_dirs=skip)
if not files:
raise ValueError(f"目录下没有找到可处理的图片:{input_dir}")
task_id = time.strftime("%Y%m%d-%H%M%S-") + uuid.uuid4().hex[:4]
jobs = [
ImageJob(
index=idx,
filename=path.name,
label=imageops.safe_stem(path.name, fallback=f"image_{idx + 1}"),
input_path=str(path),
)
for idx, path in enumerate(files)
]
task = Task(
id=task_id,
created=time.time(),
options=options,
images=jobs,
output_dir=str(output_dir),
mode="batch",
input_dir=str(input_dir),
recursive=recursive,
)
with self._lock:
self._tasks[task_id] = task
self._prune_locked()
self._queue.put(task)
self.log(f"[批量 {task_id}] 已提交,{len(jobs)} 张图片 → 输出目录 {output_dir}")
return task
def get(self, task_id: str) -> Optional[Task]:
with self._lock:
return self._tasks.get(task_id)
def list_tasks(self, limit: int = 20) -> List[dict]:
"""任务列表(含各自使用的参数,供前端恢复结果历史时直接显示参数)。"""
with self._lock:
tasks = sorted(self._tasks.values(), key=lambda t: t.created, reverse=True)[:limit]
return [t.to_dict(include_options=True) for t in tasks]
def cancel(self, task_id: str) -> bool:
task = self.get(task_id)
if task is None or task.finished_at is not None:
return False
task.cancel_requested = True
if task.state == "queued":
task.state = "canceled"
task.stage = "已取消"
task.finished_at = time.time()
self.log(f"[任务 {task_id}] 收到取消请求")
return True
def _prune_locked(self) -> None:
if len(self._tasks) <= self.max_tasks:
return
ordered = sorted(self._tasks.values(), key=lambda t: t.created)
for stale in ordered[: len(self._tasks) - self.max_tasks]:
if stale.finished_at is not None:
self._tasks.pop(stale.id, None)
# -------------------------- worker -------------------------- #
def _work_loop(self) -> None:
while True:
task = self._queue.get()
if task is None:
return
if task.cancel_requested:
continue
try:
self._process(task)
except TaskCanceled: # 正常取消,不需要堆栈
task.state = "canceled"
task.stage = "已取消"
task.finished_at = time.time()
except BaseException as exc: # noqa: BLE001 - worker 必须兜住一切
self.log(f"[任务 {task.id}] 异常终止:{exc}")
traceback.print_exc()
task.state = "failed"
task.error = str(exc)
task.stage = "失败"
task.finished_at = time.time()
def _process(self, task: Task) -> None:
task.state = "running"
task.stage = "开始处理"
self.log(f"[任务 {task.id}] 开始处理,参数:{json.dumps(task.options, ensure_ascii=False)}")
for job in task.images:
if task.cancel_requested:
if job.state in ("pending", "running"):
job.state = "canceled"
job.stage = "已取消"
job.progress = 1.0
continue
try:
self._process_one(task, job)
except TaskCanceled:
# 当前图片已被标记为 canceled;剩余图片由上面的分支收尾
continue
counts = task._counts()
task.refresh_progress()
task.progress = 1.0
task.finished_at = time.time()
if task.cancel_requested:
task.state = "canceled"
task.stage = "已取消"
elif counts["failed"] and counts["done"]:
task.state = "partial"
task.stage = "部分完成"
elif counts["failed"]:
task.state = "failed"
task.stage = "失败"
else:
task.state = "done"
task.stage = "完成"
self.log(
f"[任务 {task.id}] 结束:{task.state},成功 {counts['done']} / 失败 {counts['failed']},"
f"耗时 {task.finished_at - task.created:.1f}s"
)
def _process_one(self, task: Task, job: ImageJob) -> None:
job.state = "running"
job.progress = 0.02
task.refresh_progress()
def on_progress(stage: str, frac: float) -> None:
if task.cancel_requested:
raise TaskCanceled()
job.stage = stage
job.progress = 0.05 + 0.9 * float(frac)
task.stage = f"{job.label} · {stage}"
task.refresh_progress()
try:
with Image.open(job.input_path) as im:
im.load()
pil = im.copy()
job.source_width, job.source_height = pil.size
result = self.engine.remove_background(pil, task.options, on_progress)
job.outputs = self._store_outputs(task, job, result)
job.width = result["width"]
job.height = result["height"]
job.coverage = result["coverage"]
job.elapsed = result["elapsed"]
job.input_size = tuple(result["input_size"]) # type: ignore[assignment]
job.warning = result["warning"]
job.state = "done"
job.stage = "完成"
job.progress = 1.0
target = f" → {job.output_name}" if job.output_name else ""
self.log(
f"[{'批量' if task.mode == 'batch' else '任务'} {task.id}] {job.label}{target} 完成:"
f"{job.width}x{job.height},前景占比 {job.coverage:.1%},耗时 {job.elapsed}s"
)
except TaskCanceled:
job.state = "canceled"
job.stage = "已取消"
job.progress = 1.0
raise
except Exception as exc: # 单张失败不影响其余图片
job.state = "failed"
job.stage = "失败"
job.progress = 1.0
job.error = f"{type(exc).__name__}: {exc}"
self.log(f"[任务 {task.id}] {job.label} 处理失败:{job.error}")
traceback.print_exc()
finally:
task.refresh_progress()
def _store_outputs(self, task: Task, job: ImageJob, result: dict) -> Dict[str, str]:
"""把单张结果落盘,返回 ``kind -> 路径``。
* 单图模式:沿用 ``outputs/<task_id>/<序号>_<主干>_cutout.png``。
* 批量模式:写进用户指定的输出目录,文件名按需求固定为
``RMBG_<原主干(≤20)>_<Unix 时间戳>.<原后缀>``;同名时递增 ``_1``
避让而不是覆盖(见 :mod:`birefnet_web.batch`)。
"""
outputs: Dict[str, str] = {}
if task.mode == "batch":
out_dir = Path(task.output_dir)
name = batch.build_output_name(
Path(job.input_path),
int(time.time()),
str(task.options.get("background") or "transparent"),
)
target = batch.unique_path(out_dir, name)
imageops.save_image(result["cutout"], str(target))
job.output_name = target.name
outputs["cutout"] = str(target)
if result["mask"] is not None:
mask_target = batch.unique_path(out_dir, f"{target.stem}_mask.png")
imageops.save_image(result["mask"], str(mask_target))
outputs["mask"] = str(mask_target)
else:
cutout_png = os.path.join(task.output_dir, f"{job.index + 1:03d}_{job.label}_cutout.png")
imageops.save_image(result["cutout"], cutout_png)
outputs["cutout"] = cutout_png
if result["mask"] is not None:
mask_png = os.path.join(task.output_dir, f"{job.index + 1:03d}_{job.label}_mask.png")
imageops.save_image(result["mask"], mask_png)
outputs["mask"] = mask_png
outputs["original"] = job.input_path
return outputs
# --------------------------------------------------------------------------- #
# HTTP 处理器
# --------------------------------------------------------------------------- #
class WebUIHandler(BaseHTTPRequestHandler):
server_version = f"BiRefNetWebUI/{__version__}"
protocol_version = "HTTP/1.1"
# ------------------------- 基础工具 ------------------------- #
@property
def app(self) -> "WebUIServer": # type: ignore[override]
return self.server # type: ignore[return-value]
def log_message(self, fmt: str, *args) -> None: # noqa: A003
path = getattr(self, "path", "")
# 轮询与静态资源太吵,不打印
if path.startswith("/api/tasks/") or path.startswith("/static/") or path == "/favicon.ico":
return
super().log_message(fmt, *args)
def _send(self, status: int, body: bytes, content_type: str, extra: Optional[dict] = None) -> None:
try:
self.send_response(status)
self.send_header("Content-Type", content_type)
self.send_header("Content-Length", str(len(body)))
self.send_header("Cache-Control", "no-store")
for k, v in (extra or {}).items():
self.send_header(k, v)
self.end_headers()
if self.command != "HEAD":
self.wfile.write(body)
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
pass
def _send_json(self, payload: object, status: int = 200) -> None:
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
self._send(status, body, "application/json; charset=utf-8")
def _error(self, status: int, message: str) -> None:
self._send_json({"ok": False, "error": message}, status=status)
def _read_body(self) -> bytes:
try:
length = int(self.headers.get("Content-Length") or 0)
except ValueError:
length = 0
if length <= 0:
return b""
if length > MAX_UPLOAD_BYTES:
raise ValueError(
f"请求体过大({length / 1024 / 1024:.1f} MB),上限 {MAX_UPLOAD_BYTES / 1024 / 1024:.0f} MB"
)
return self.rfile.read(length)
# ------------------------- 路由 ------------------------- #
def do_GET(self) -> None: # noqa: N802
try:
parsed = urlparse(self.path)
path = unquote(parsed.path)
query = parse_qs(parsed.query)
if path in ("/", "/index.html"):
return self._serve_static("index.html")
if path.startswith("/static/"):
return self._serve_static(path[len("/static/"):])
if path == "/favicon.ico":
return self._send(204, b"", "image/x-icon")
if path == "/api/health":
return self._send_json({"ok": True, "app": APP_NAME, "version": __version__})
if path == "/api/state":
return self._api_state()
if path == "/api/models":
return self._send_json({"ok": True, "models": [m.to_dict() for m in self.app.registry.list()]})
if path == "/api/tasks":
return self._send_json({"ok": True, "tasks": self.app.tasks.list_tasks()})
parts = path.strip("/").split("/")
# /api/tasks/<id>[/...]
if len(parts) >= 3 and parts[0] == "api" and parts[1] == "tasks":
task_id = parts[2]
if len(parts) == 3:
task = self.app.tasks.get(task_id)
if task is None:
return self._error(404, "任务不存在")
return self._send_json({"ok": True, "task": self._task_payload(task, query)})
if len(parts) == 6 and parts[3] == "file":
return self._serve_result(task_id, parts[4], parts[5])
if len(parts) == 4 and parts[3] == "zip":
return self._serve_zip(task_id, query)
return self._error(404, f"未知路径 {path}")
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
pass
except Exception as exc: # pragma: no cover
traceback.print_exc()
self._error(500, f"{type(exc).__name__}: {exc}")
def do_HEAD(self) -> None: # noqa: N802
self.do_GET()
def do_POST(self) -> None: # noqa: N802
try:
parsed = urlparse(self.path)
path = unquote(parsed.path)
if path == "/api/models/reload":
models = self.app.registry.scan()
return self._send_json({"ok": True, "models": [m.to_dict() for m in models]})
if path == "/api/engine/unload":
self.app.engine.unload()
return self._send_json({"ok": True})
if path == "/api/shutdown":
self._send_json({"ok": True, "message": "服务正在关闭"})
threading.Thread(target=self.app.stop_soon, daemon=True).start()
return
if path == "/api/tasks":
return self._create_task()
if path == "/api/batch":
return self._create_batch_task()
parts = path.strip("/").split("/")
if len(parts) == 4 and parts[0] == "api" and parts[1] == "tasks" and parts[3] == "cancel":
ok = self.app.tasks.cancel(parts[2])
return self._send_json({"ok": ok})
return self._error(404, f"未知路径 {path}")
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
pass
except Exception as exc: # pragma: no cover
traceback.print_exc()
self._error(500, f"{type(exc).__name__}: {exc}")
# ------------------------- 具体处理 ------------------------- #
@staticmethod
def _int_param(query: dict, key: str, default: int) -> int:
"""从 query(parse_qs 的结果)里安全地取一个整数参数。"""
try:
return int(str((query.get(key) or [default])[0]))
except (TypeError, ValueError):
return default
def _task_payload(self, task: Task, query: dict) -> dict:
"""任务详情;``?tail=N`` 只回传最后 N 张(批量任务轮询用,压住载荷)。"""
return task.to_dict(include_options=True, tail=max(0, self._int_param(query, "tail", 0)))
def _merge_options(self, incoming: object) -> dict:
"""把前端传来的参数合并到默认参数上(只接受白名单键)。"""
options = dict(DEFAULT_OPTIONS)
if isinstance(incoming, dict):
options.update({k: v for k, v in incoming.items() if k in DEFAULT_OPTIONS})
return options
def _pin_model(self, options: dict) -> Optional[str]:
"""确保 ``options['model']`` 指向一个真实存在的权重。
Returns:
出错时返回给人看的错误信息;成功返回 None。
"""
models = self.app.registry.list()
if not models:
return "模型目录中没有找到任何权重文件(*.safetensors / *.pth)"
options["model"] = options.get("model") or self.app.default_model_key(
[m.to_dict() for m in models]
)
try:
self.app.registry.get(str(options["model"]))
except KeyError:
options["model"] = models[0].key
return None
def _api_state(self) -> None:
app = self.app
models = [m.to_dict() for m in app.registry.list()]
default_model = app.default_model_key(models)
self._send_json(
{
"ok": True,
"app": APP_NAME,
"version": __version__,
"environment": describe_environment(app.engine),
"node_dir": str(app.node_dir),
"model_dirs": [str(d) for d in app.registry.model_dirs],
"output_dir": str(app.output_root),
"project_root": str(app.project_root),
"models": models,
"defaults": {**DEFAULT_OPTIONS, "model": default_model},
}
)
def _create_task(self) -> None:
content_type = self.headers.get("Content-Type", "")
if "multipart/form-data" not in content_type:
return self._error(400, "Content-Type 必须是 multipart/form-data")
boundary = None
for chunk in content_type.split(";"):
chunk = chunk.strip()
if chunk.lower().startswith("boundary="):
boundary = chunk[9:].strip().strip('"').encode("latin-1")
if not boundary:
return self._error(400, "缺少 multipart boundary")
try:
body = self._read_body()
except ValueError as exc:
return self._error(413, str(exc))
options = dict(DEFAULT_OPTIONS)
uploads: List[Tuple[str, bytes]] = []
for part in parse_multipart(body, boundary):
if part.name == "options":
try:
options = self._merge_options(json.loads(part.data.decode("utf-8") or "{}"))
except json.JSONDecodeError:
return self._error(400, "options 字段不是合法 JSON")
elif part.name in ("files", "images", "file"):
if part.data:
uploads.append((part.filename or f"upload_{len(uploads) + 1}.png", part.data))
if not uploads:
return self._error(400, "没有收到任何图片,请选择 PNG / JPG / WebP 文件")
error = self._pin_model(options)
if error:
return self._error(400, error)
try:
task = self.app.tasks.submit(uploads, options)
except ValueError as exc:
return self._error(400, str(exc))
self._send_json({"ok": True, "task_id": task.id, "count": len(uploads)})
def _create_batch_task(self) -> None:
"""POST /api/batch —— 目录批量抠图(JSON 传路径,不上传文件)。
请求体:
{
"input_dir": "F:\\\\photos\\\\待抠图", # 必填,图片所在目录
"output_dir": "F:\\\\photos\\\\results", # 可空 → 用项目 outputs/
"recursive": false, # 是否含子目录
"options": { …与单图模式相同的参数… }
}
两个目录都必须**已经存在**:路径不对就报错,本接口绝不会创建目录。
"""
try:
payload = json.loads(self._read_body().decode("utf-8") or "{}")
except (ValueError, UnicodeDecodeError):
return self._error(400, "请求体不是合法 JSON")
if not isinstance(payload, dict):
return self._error(400, "请求体必须是 JSON 对象")
base = self.app.project_root
raw_out = str(payload.get("output_dir") or "").strip() or str(self.app.output_root)
try:
input_dir = batch.check_input_dir(payload.get("input_dir"), base)
output_dir = batch.check_output_dir(raw_out, base)
except batch.BatchPathError as exc:
return self._error(400, str(exc))
options = self._merge_options(payload.get("options"))
error = self._pin_model(options)
if error:
return self._error(400, error)
try:
task = self.app.tasks.submit_batch(
input_dir, output_dir, options, recursive=bool(payload.get("recursive"))
)
except ValueError as exc:
return self._error(400, str(exc))
self._send_json(
{
"ok": True,
"task_id": task.id,
"mode": "batch",
"count": len(task.images),
"input_dir": str(input_dir),
"output_dir": str(output_dir),
}
)
def _task_file(self, task_id: str, index: int, kind: str) -> Optional[str]:
task = self.app.tasks.get(task_id)
if task is None:
return None
for job in task.images:
if job.index == index:
return job.outputs.get(kind)
return None
def _serve_result(self, task_id: str, idx_part: str, kind: str) -> None:
try:
index = int(idx_part)
except ValueError:
return self._error(400, "图片序号非法")
if kind not in ("cutout", "mask", "original"):
return self._error(400, f"未知的结果类型 {kind!r}")
path = self._task_file(task_id, index, kind)
if not path or not os.path.isfile(path):
return self._error(404, "结果文件不存在")
suffix = Path(path).suffix.lower()
with open(path, "rb") as fh:
data = fh.read()
self._send(200, data, MIME_TYPES.get(suffix, "application/octet-stream"))
def _serve_zip(self, task_id: str, query: dict) -> None:
task = self.app.tasks.get(task_id)
if task is None:
return self._error(404, "任务不存在")
kinds = (query.get("kind") or ["cutout"])[0].split(",")
want = [k for k in kinds if k in ("cutout", "mask", "original")] or ["cutout"]
buffer = io.BytesIO()
added = 0
with zipfile.ZipFile(buffer, "w", zipfile.ZIP_DEFLATED) as zf:
for job in task.images:
for kind in want:
path = job.outputs.get(kind)
if path and os.path.isfile(path):
zf.write(path, arcname=os.path.basename(path))
added += 1
if added == 0:
return self._error(404, "没有可打包的结果文件")
self._send(
200,
buffer.getvalue(),
"application/zip",
extra={"Content-Disposition": f'attachment; filename="birefnet_{task_id}.zip"'},
)
def _serve_static(self, rel_path: str) -> None:
web_dir = self.app.web_dir.resolve()
target = (web_dir / rel_path.replace("\\", "/")).resolve()
try:
target.relative_to(web_dir) # 防路径穿越
except ValueError:
return self._error(403, "非法路径")
if not target.is_file():
return self._error(404, f"资源不存在:{rel_path}")
suffix = target.suffix.lower()
with open(target, "rb") as fh:
data = fh.read()
cache = "no-cache" if suffix in (".html", ".js", ".css") else "public, max-age=3600"
self._send(200, data, MIME_TYPES.get(suffix, "application/octet-stream"), {"Cache-Control": cache})
# --------------------------------------------------------------------------- #
# 服务容器
# --------------------------------------------------------------------------- #
class WebUIServer(ThreadingHTTPServer):
daemon_threads = True
allow_reuse_address = True
def __init__(
self,
address: Tuple[str, int],
registry: ModelRegistry,
engine: InferenceEngine,
tasks: TaskManager,
web_dir: Path,
output_root: Path,
node_dir: Path,
logger=None,
project_root: Optional[Path] = None,
) -> None:
super().__init__(address, WebUIHandler)
self.registry = registry
self.engine = engine
self.tasks = tasks
self.web_dir = Path(web_dir)
self.output_root = Path(output_root)
self.node_dir = Path(node_dir)
#: 相对路径型用户输入(批量处理的目录)以此为准
self.project_root = Path(project_root) if project_root else self.web_dir.parent
self.log = logger or (lambda msg: None)
def default_model_key(self, models: Optional[Iterable[dict]] = None) -> str:
models = list(models if models is not None else (m.to_dict() for m in self.registry.list()))
if not models:
return ""
for preferred in ("Portrait", "General", "General-HR"):
for m in models:
if m["name"] == preferred:
return m["key"]
return models[0]["key"]
def stop_soon(self) -> None:
time.sleep(0.4)
self.log("收到关闭指令,服务即将退出")
self.shutdown()