# -*- 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/);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//<序号>_<主干>_cutout.png``。 * 批量模式:写进用户指定的输出目录,文件名按需求固定为 ``RMBG_<原主干(≤20)>_.<原后缀>``;同名时递增 ``_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/[/...] 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()