from __future__ import annotations import argparse import importlib.util import json import sys from pathlib import Path from .image_io import InvalidImage from .inpaint import InpaintUnavailable, install_lama, lama_installed, resolve_torch_device from .models import BitmapExtractor, CleanupMode, DEFAULT_SNAP_ANGLES, Device, ProcessOptions from .ocr import OcrUnavailable, is_ocr_available from .pipeline import export_result, process_path from .sam2 import ( ShapeExtractionUnavailable, install_sam2, sam2_dependencies_installed, sam2_installed, ) IMAGE_SUFFIXES = {".png", ".jpg", ".jpeg", ".webp"} def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="editable-image") subparsers = parser.add_subparsers(dest="command", required=True) convert = subparsers.add_parser("convert", help="convert bitmap images") convert.add_argument("inputs", nargs="+", type=Path) convert.add_argument("--output", "-o", type=Path, required=True) convert.add_argument("--recursive", action="store_true") convert.add_argument("--cleanup", choices=[item.value for item in CleanupMode], default="opencv") convert.add_argument("--device", choices=[item.value for item in Device], default="auto") convert.add_argument("--confidence", type=float, default=0.5) convert.add_argument( "--snap-angles", type=parse_angles, default=list(DEFAULT_SNAP_ANGLES), metavar="DEGREES", help="comma-separated target angles (default: 0,45,90,-45,-90)", ) convert.add_argument( "--snap-tolerance", type=float, default=6.0, metavar="DEGREES", help="maximum distance from a target angle (default: 6)", ) convert.add_argument( "--snap-font-sizes", action="store_true", help="snap estimated font sizes to peaks in the image's size distribution", ) convert.add_argument( "--embed-fonts", action="store_true", help="embed used font faces in the SVG (default: link system fonts)", ) convert.add_argument( "--vectorize-shapes", action="store_true", help="promote confident diagram primitives into editable SVG shapes", ) convert.add_argument( "--bitmap-extractor", choices=[item.value for item in BitmapExtractor], default=BitmapExtractor.OPENCV.value, help="bitmap-layer mask extractor used with --vectorize-shapes (default: opencv)", ) convert.add_argument("--overwrite", action="store_true") models = subparsers.add_parser("models", help="manage optional model files") model_commands = models.add_subparsers(dest="model_command", required=True) install = model_commands.add_parser("install") install.add_argument("model", choices=["ocr", "lama", "sam2"]) model_commands.add_parser("status") return parser def main(argv: list[str] | None = None) -> int: args = build_parser().parse_args(argv) if args.command == "models": return model_command(args) return convert_command(args) def model_command(args: argparse.Namespace) -> int: if args.model_command == "install": if args.model == "lama": path = install_lama() print(f"Installed LaMa at {path}") return 0 if args.model == "sam2": path = install_sam2() print(f"Installed SAM 2 at {path}") return 0 if not is_ocr_available(): print("Install the cpu or cuda project extra before installing OCR models", file=sys.stderr) return 2 from .ocr import RapidOcrEngine RapidOcrEngine(Device.CPU) print("OCR models are ready") return 0 print( json.dumps( { "ocr": is_ocr_available(), "lama": lama_installed(), "sam2": sam2_installed(), "sam2_dependencies": sam2_dependencies_installed(), "devices": available_devices(), }, indent=2, ) ) return 0 def convert_command(args: argparse.Namespace) -> int: if args.bitmap_extractor != BitmapExtractor.OPENCV.value and not args.vectorize_shapes: print("--bitmap-extractor requires --vectorize-shapes", file=sys.stderr) return 2 options = ProcessOptions( confidence=args.confidence, cleanup=CleanupMode(args.cleanup), device=Device(args.device), snap_angles=args.snap_angles, snap_tolerance=args.snap_tolerance, snap_font_sizes=args.snap_font_sizes, vectorize_shapes=args.vectorize_shapes, bitmap_extractor=BitmapExtractor(args.bitmap_extractor), ) inputs = collect_inputs(args.inputs, args.recursive) if not inputs: print("No supported images found", file=sys.stderr) return 2 failures = 0 for source in inputs: destination = args.output / source.stem expected = [destination / f"{source.stem}-background.png", destination / f"{source.stem}-overlay.svg"] if not args.overwrite and any(path.exists() for path in expected): failures += 1 print(f"skip {source}: output exists (use --overwrite)", file=sys.stderr) continue try: result = process_path(source, options) background, overlay, assets = export_result( result, destination, source.stem, embed_fonts=args.embed_fonts, ) print( json.dumps( { "source": str(source), "background": str(background), "overlay": str(overlay), "assets": [str(path) for path in assets], } ) ) except ( InvalidImage, OcrUnavailable, InpaintUnavailable, ShapeExtractionUnavailable, OSError, ValueError, ) as exc: failures += 1 print(f"failed {source}: {exc}", file=sys.stderr) return 1 if failures else 0 def collect_inputs(paths: list[Path], recursive: bool) -> list[Path]: output: list[Path] = [] for path in paths: if path.is_file() and path.suffix.lower() in IMAGE_SUFFIXES: output.append(path) elif path.is_dir(): iterator = path.rglob("*") if recursive else path.glob("*") output.extend(item for item in iterator if item.is_file() and item.suffix.lower() in IMAGE_SUFFIXES) return sorted(set(output)) def parse_angles(value: str) -> list[float]: try: angles = [float(item.strip()) for item in value.split(",") if item.strip()] except ValueError as exc: raise argparse.ArgumentTypeError("angles must be comma-separated numbers") from exc if not angles: raise argparse.ArgumentTypeError("at least one snap angle is required") if any(angle < -180 or angle > 180 for angle in angles): raise argparse.ArgumentTypeError("snap angles must be between -180 and 180") return angles def available_devices() -> list[str]: devices = ["cpu"] try: import onnxruntime as ort if "CUDAExecutionProvider" in ort.get_available_providers(): devices.append("cuda") except ImportError: pass if importlib.util.find_spec("torch") is not None: try: if resolve_torch_device(Device.AUTO) == "mps": devices.append("mps") except InpaintUnavailable: pass return devices if __name__ == "__main__": raise SystemExit(main())