File size: 7,781 Bytes
3c2cd23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
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())