# -*- coding: utf-8 -*- """目录批量抠图:路径校验 / 目录扫描 / 输出文件命名。 三条硬约束(对应需求): 1. **校验阶段只读不写**:输入目录、输出目录都必须**已经存在**,本模块 绝不 ``mkdir``。路径写错时直接报错,而不是顺手把目录树建出来 —— 这样一次误输入(或网页上的恶意路径)不会在磁盘上留下任何东西。 2. **命名固定**:``RMBG_<原主干>_.<后缀>``,原主干超过 :data:`STEM_MAXLEN` 个字符时截断。主干统一走 :func:`birefnet_web.imageops.safe_stem` 清洗,杜绝 ``..`` 之类的穿越。 3. **后缀自适应**:透明输出需要 alpha 通道,``.jpg`` / ``.bmp`` 装不下, 此时退化为 ``.png``(强行按原后缀保存会被 PIL 静默压成白底图)。 本模块不依赖 torch,可以脱离推理环境单独测试。 """ from __future__ import annotations import os from pathlib import Path from typing import Iterable, List, Optional, Sequence from .fsutil import ALPHA_SUFFIXES, IMAGE_SUFFIXES, safe_stem #: 输出文件名统一前缀(扫描输入目录时也用来自我过滤,避免把产物再抠一遍) BATCH_PREFIX = "RMBG_" #: 原文件名主干的最大长度(超过即截断) STEM_MAXLEN = 20 #: 重名避让时最多尝试的序号 _MAX_DEDUP = 999 class BatchPathError(ValueError): """路径不合规:不存在 / 不是目录 / 不可写 / 命中系统目录黑名单。""" # --------------------------------------------------------------------------- # # 路径解析与校验 # --------------------------------------------------------------------------- # def _critical_dirs() -> List[Path]: """返回当前系统上「不该往里写文件」的目录(仅用于提示级别的高危拦截)。""" roots: List[Path] = [] names = ("SystemRoot", "windir", "ProgramFiles", "ProgramFiles(x86)", "ProgramData") for name in names: raw = os.environ.get(name) if not raw: continue p = Path(raw) if p.is_absolute(): roots.append(p) if os.name != "nt": roots.extend(Path(p) for p in ("/bin", "/sbin", "/etc", "/usr", "/boot", "/dev", "/proc", "/sys")) return roots def resolve_user_path(raw: object, base: Path) -> Path: """把用户在界面上敲的一行路径规范化成绝对路径。 Args: raw: 原始输入(可能带首尾空格、成对引号、``~`` 或 ``%VAR%``)。 base: 相对路径的基准目录(本项目的 ``PROJECT_ROOT``)。 Returns: 规范化后的绝对路径(**不校验存在性**,也不做任何创建)。 Raises: BatchPathError: 输入为空。 """ text = str(raw or "").strip().strip('"').strip("'").strip() if not text: raise BatchPathError("路径不能为空") text = os.path.expandvars(os.path.expanduser(text)) path = Path(text) if not path.is_absolute(): path = base / path return Path(os.path.normpath(str(path))) def check_input_dir(raw: object, base: Path) -> Path: """校验「图片所在目录」,必须已存在且是目录。 Raises: BatchPathError: 路径为空 / 不存在 / 不是目录。 """ path = resolve_user_path(raw, base) if not path.exists(): raise BatchPathError(f"输入目录不存在:{path}(本工具不会自动创建目录,请先确认路径)") if not path.is_dir(): raise BatchPathError(f"输入路径不是目录:{path}") return path def check_output_dir(raw: object, base: Path, *, allow_system_dir: bool = False) -> Path: """校验「结果输出目录」,必须已存在、是目录且可写。 Args: raw: 用户输入或默认输出目录。 base: 相对路径基准。 allow_system_dir: 为 True 时跳过系统目录黑名单(测试用)。 Raises: BatchPathError: 路径为空 / 不存在 / 不是目录 / 命中黑名单 / 不可写。 """ path = resolve_user_path(raw, base) if not path.exists(): raise BatchPathError(f"输出目录不存在:{path}(本工具不会自动创建目录,请先手动建好)") if not path.is_dir(): raise BatchPathError(f"输出路径不是目录:{path}") if not allow_system_dir: if path == Path(path.anchor): # C:\ 或 / raise BatchPathError(f"拒绝把结果写进磁盘根目录:{path}") for critical in _critical_dirs(): if path == critical or critical in path.parents: raise BatchPathError(f"拒绝把结果写进系统目录:{path}") if not os.access(path, os.W_OK): raise BatchPathError(f"输出目录不可写(权限不足):{path}") return path # --------------------------------------------------------------------------- # # 目录扫描 # --------------------------------------------------------------------------- # def scan_images( root: Path, *, recursive: bool = False, skip_dirs: Sequence[Path] = (), ) -> List[Path]: """扫描目录下的图片文件。 Args: root: 已校验过的图片目录。 recursive: 是否包含子目录。 skip_dirs: 需要跳过的目录(通常是输出目录,避免读到自己的产物)。 Returns: 排序后的图片路径列表(先按目录、再按文件名,保证同一批次的处理顺序稳定)。 """ skipped = [Path(d).resolve() for d in skip_dirs] walker: Iterable[Path] = root.rglob("*") if recursive else root.glob("*") found: List[Path] = [] for path in walker: try: if not path.is_file(): continue except OSError: # 权限/坏链接 continue if path.suffix.lower() not in IMAGE_SUFFIXES: continue # 本工具自己的产物(RMBG_ 前缀)不参与批量,避免反复抠同一张图 if path.name.startswith(BATCH_PREFIX): continue if skipped: parent = path.parent.resolve() if any(parent == s or s in parent.parents or parent in s.parents for s in skipped): continue found.append(path) return sorted(found, key=lambda p: (str(p.parent).lower(), p.name.lower())) # --------------------------------------------------------------------------- # # 输出命名 # --------------------------------------------------------------------------- # def choose_suffix(orig_suffix: str, background: str) -> str: """决定输出文件后缀:透明模式下如果原后缀装不下 alpha,就退化为 ``.png``。 Args: orig_suffix: 原文件后缀(含点,如 ``.jpg``)。 background: 背景模式,``transparent`` 或其它(纯色)。 Returns: 合法的输出后缀字符串。 """ suffix = (orig_suffix or "").lower() if background == "transparent": return suffix if suffix in ALPHA_SUFFIXES else ".png" return suffix if suffix in IMAGE_SUFFIXES else ".png" def build_output_name(source: Path, timestamp: int, background: str = "transparent") -> str: """按 ``RMBG_<主干(≤20)>_.<后缀>`` 组装输出文件名。 Args: source: 原始图片路径。 timestamp: Unix 时间戳(秒),由调用方在**处理时**取,便于追溯。 background: 背景模式,决定后缀是否允许退化为 ``.png``。 Returns: 输出文件名(不含目录)。 """ stem = safe_stem(source.name, fallback="image", maxlen=STEM_MAXLEN) return f"{BATCH_PREFIX}{stem}_{int(timestamp)}{choose_suffix(source.suffix, background)}" def unique_path(directory: Path, name: str) -> Path: """在同一秒内出现同名输出时,用 ``_1`` / ``_2`` 递增避让,绝不覆盖已有文件。 Args: directory: 输出目录(已存在)。 name: 目标文件名。 Returns: 可安全写入的完整路径。 Raises: BatchPathError: 避让序号耗尽(同名文件超过 :data:`_MAX_DEDUP` 个)。 """ target = directory / name if not target.exists(): return target for i in range(1, _MAX_DEDUP + 1): candidate = directory / f"{target.stem}_{i}{target.suffix}" if not candidate.exists(): return candidate raise BatchPathError(f"输出目录中同名文件过多,请检查:{name}")