938 lines
36 KiB
Python
938 lines
36 KiB
Python
# -*- 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()
|