feat: initial project setup
This commit is contained in:
@@ -0,0 +1,937 @@
|
||||
# -*- 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()
|
||||
Reference in New Issue
Block a user