BareGit
"""Tests for command-line parsing and dispatch."""

from pathlib import Path
import tempfile
import unittest
from unittest.mock import patch

from rmbg_mask.cli import buildParser, main


class CliTest(unittest.TestCase):
    """Verify CLI defaults, options, and inference dispatch."""

    def testDefaultProcessingResolution(self) -> None:
        """Use 1024 pixels when no processing resolution is supplied."""
        arguments = buildParser().parse_args(
            ["input.jpg", "mask.png", "--model", "model"]
        )
        self.assertEqual(arguments.processing_resolution, 1024)

    def testCustomProcessingResolution(self) -> None:
        """Accept a positive custom processing resolution."""
        arguments = buildParser().parse_args(
            [
                "input.jpg",
                "mask.png",
                "--model",
                "model",
                "--processing-resolution",
                "1536",
            ]
        )
        self.assertEqual(arguments.processing_resolution, 1536)

    def testMainDispatchesArguments(self) -> None:
        """Pass parsed paths and options to mask generation."""
        with patch("rmbg_mask.cli.generateMask") as generate_mask:
            status = main(
                [
                    "input.jpg",
                    "mask.png",
                    "--model",
                    "model",
                    "-r",
                    "512",
                    "--device",
                    "cpu",
                ]
            )

        self.assertEqual(status, 0)
        generate_mask.assert_called_once_with(
            model_path=Path("model"),
            input_path=Path("input.jpg"),
            output_path=Path("mask.png"),
            processing_resolution=512,
            device_name="cpu",
        )

    def testMainReportsExpectedErrors(self) -> None:
        """Return a failure status for an expected inference error."""
        from rmbg_mask.inference import RmbgError

        with patch(
            "rmbg_mask.cli.generateMask",
            side_effect=RmbgError("test failure"),
        ):
            status = main(
                ["input.jpg", "mask.png", "--model", "model"]
            )

        self.assertEqual(status, 1)


if __name__ == "__main__":
    unittest.main()