"""Unit and CLI tests; no ML framework, network or GPU is required."""

import subprocess
import sys
import unittest
from pathlib import Path

from qlora_memory import GIB, Projection, estimate, ideal_four_bit_bytes, lora_parameter_count

SCRIPT = Path(__file__).with_name("qlora_memory.py")


class ArithmeticTests(unittest.TestCase):
    def test_four_bit_payload_and_binary_units(self):
        self.assertEqual(ideal_four_bit_bytes(7_000_000_000), 3_500_000_000)
        self.assertAlmostEqual(3_500_000_000 / GIB, 3.259629011154175)

    def test_odd_parameter_count_rounds_up_without_floats(self):
        self.assertEqual(ideal_four_bit_bytes(3), 2)
        self.assertEqual(ideal_four_bit_bytes(2 ** 60 + 1), 2 ** 59 + 1)

    def test_one_square_adapter_has_both_matrices(self):
        self.assertEqual(lora_parameter_count(Projection(4096, 4096, 1), 16), 131_072)

    def test_rectangular_projection_uses_both_dimensions(self):
        self.assertEqual(lora_parameter_count(Projection(4096, 1024, 32), 16), 2_621_440)

    def test_two_square_projections_per_layer(self):
        result = estimate(7_000_000_000, [Projection(4096, 4096, 64)], 16)
        self.assertEqual(result.adapter_parameters, 8_388_608)
        self.assertEqual(result.adapter_weight_bytes, 33_554_432)
        self.assertEqual(result.total_bytes, 3_533_554_432)

    def test_seven_projections_per_layer(self):
        result = estimate(7_000_000_000, [Projection(4096, 4096, 128),
                                       Projection(4096, 11008, 96)], 16)
        self.assertEqual(result.adapter_parameters, 39_976_960)
        self.assertEqual(result.adapter_weight_bytes, 159_907_840)
        self.assertEqual(result.total_bytes, 3_659_907_840)

    def test_doubling_rank_only_doubles_adapter_storage(self):
        shapes = [Projection(4096, 4096, 128), Projection(4096, 11008, 96)]
        low = estimate(7_000_000_000, shapes, 16)
        high = estimate(7_000_000_000, shapes, 32)
        self.assertEqual(high.adapter_parameters, 79_953_920)
        self.assertEqual(high.adapter_weight_bytes, 319_815_680)
        self.assertEqual(high.ideal_base_bytes, low.ideal_base_bytes)
        self.assertEqual(high.adapter_weight_bytes, low.adapter_weight_bytes * 2)

    def test_explicit_storage_dtype_changes_bytes_not_parameter_count(self):
        half = estimate(7_000_000_000, [Projection(4096, 4096, 64)], 16, adapter_bytes=2)
        single = estimate(7_000_000_000, [Projection(4096, 4096, 64)], 16, adapter_bytes=4)
        self.assertEqual(half.adapter_parameters, single.adapter_parameters)
        self.assertEqual(half.adapter_weight_bytes * 2, single.adapter_weight_bytes)

    def test_small_floor_never_proves_a_training_fit(self):
        result = estimate(7_000_000_000, [Projection(4096, 4096, 64)], 16)
        self.assertEqual(result.budget_status(8 * GIB),
                         "undetermined: training overhead is not estimated")
        self.assertEqual(result.budget_status(result.total_bytes),
                         "undetermined: training overhead is not estimated")

    def test_floor_can_rule_out_resident_weights(self):
        result = estimate(7_000_000_000, [Projection(4096, 4096, 64)], 16)
        self.assertEqual(result.budget_status(3 * GIB),
                         "exceeds budget before training overhead")


class InputTests(unittest.TestCase):
    def test_parameter_count_rejects_nonpositive_and_noninteger_values(self):
        for value in (0, -1, 3.5, True, "7", None):
            with self.subTest(value=value), self.assertRaises(ValueError):
                ideal_four_bit_bytes(value)

    def test_every_projection_field_requires_a_positive_integer(self):
        for field in range(3):
            for invalid in (0, -1, True, 1.5, "4"):
                values = [4, 8, 2]
                values[field] = invalid
                with self.subTest(field=field, invalid=invalid), self.assertRaises(ValueError):
                    Projection(*values)

    def test_rank_and_storage_validation(self):
        for rank in (0, -1, True, 1.5):
            with self.subTest(rank=rank), self.assertRaises(ValueError):
                estimate(1000, [Projection(4, 8, 2)], rank)
        for size in (0, 1, 3, True, 4.0):
            with self.subTest(size=size), self.assertRaises(ValueError):
                estimate(1000, [Projection(4, 8, 2)], 2, adapter_bytes=size)

    def test_empty_or_untyped_projection_list_is_rejected(self):
        for projections in ([], [(4, 8, 2)]):
            with self.subTest(projections=projections), self.assertRaises(ValueError):
                estimate(1000, projections, 2)

    def test_targets_cannot_exceed_the_supplied_base_model(self):
        with self.assertRaisesRegex(ValueError, "more weights"):
            estimate(10, [Projection(4, 8, 2)], 2)

    def test_budget_requires_positive_integer_bytes(self):
        result = estimate(1000, [Projection(4, 8, 2)], 2)
        for budget in (0, -1, True, 10.5):
            with self.subTest(budget=budget), self.assertRaises(ValueError):
                result.budget_status(budget)


class CommandLineTests(unittest.TestCase):
    def run_cli(self, *args):
        return subprocess.run([sys.executable, "-B", str(SCRIPT), *args],
                              capture_output=True, text=True, check=False)

    def test_reader_command(self):
        result = self.run_cli("--parameters", "7000000000", "--rank", "16",
                              "--projection", "4096:4096:64", "--budget-gib", "8")
        self.assertEqual(result.returncode, 0, result.stderr)
        self.assertIn("LoRA parameters: 8,388,608", result.stdout)
        self.assertIn("Budget verdict: undetermined", result.stdout)
        self.assertIn("Excluded:", result.stdout)

    def test_invalid_cli_values_fail_without_a_traceback(self):
        valid = ["--parameters", "7000000000", "--rank", "16", "--projection", "4096:4096:64"]
        variants = [valid + ["--budget-gib", value] for value in ("nan", "inf", "0", "-1", "abc")]
        variants += [valid + ["--projection", "4096:0:32"],
                     valid + ["--projection", "4096:32"], valid + ["--rank", "0"],
                     valid + ["--parameters", "100"], valid + ["--adapter-bytes", "3"]]
        for args in variants:
            with self.subTest(args=args):
                result = self.run_cli(*args)
                self.assertEqual(result.returncode, 2)
                self.assertIn("error:", result.stderr)
                self.assertNotIn("Traceback", result.stderr)
                self.assertEqual(result.stdout, "")


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