BareGit
"""Command-line interface for local RMBG 2.0 inference."""

import argparse
from collections.abc import Sequence
from pathlib import Path
import sys

from rmbg_mask.inference import RmbgError, generateMask


DEFAULT_PROCESSING_RESOLUTION = 1024


def positiveInt(value: str) -> int:
    """Parse a positive integer for an argparse option."""
    try:
        parsed_value = int(value)
    except ValueError as error:
        raise argparse.ArgumentTypeError(
            f"expected a positive integer, got {value!r}"
        ) from error

    if parsed_value <= 0:
        raise argparse.ArgumentTypeError(
            f"expected a positive integer, got {value!r}"
        )
    return parsed_value


def buildParser() -> argparse.ArgumentParser:
    """Build and return the command-line argument parser."""
    parser = argparse.ArgumentParser(
        prog="rmbg-mask",
        description=(
            "Generate an original-size grayscale mask using a local "
            "RMBG 2.0 model."
        ),
    )
    parser.add_argument(
        "input",
        type=Path,
        help="input image path",
    )
    parser.add_argument(
        "output",
        type=Path,
        help="output grayscale mask path (PNG recommended)",
    )
    parser.add_argument(
        "--model",
        type=Path,
        required=True,
        help="directory containing the already-downloaded model",
    )
    parser.add_argument(
        "-r",
        "--processing-resolution",
        type=positiveInt,
        default=DEFAULT_PROCESSING_RESOLUTION,
        metavar="PIXELS",
        help=(
            "square inference resolution "
            f"(default: {DEFAULT_PROCESSING_RESOLUTION})"
        ),
    )
    parser.add_argument(
        "--device",
        choices=("auto", "cpu", "cuda", "mps"),
        default="auto",
        help="inference device (default: auto)",
    )
    return parser


def main(argv: Sequence[str] | None = None) -> int:
    """Run the command-line program and return its process exit status."""
    arguments = buildParser().parse_args(argv)
    try:
        generateMask(
            model_path=arguments.model,
            input_path=arguments.input,
            output_path=arguments.output,
            processing_resolution=arguments.processing_resolution,
            device_name=arguments.device,
        )
    except RmbgError as error:
        print(f"rmbg-mask: error: {error}", file=sys.stderr)
        return 1
    return 0