#!/usr/bin/python3

import os
import sys
import tempfile
import unittest


sys.path.insert(
    0,
    os.path.join(os.path.dirname(__file__), "usr/lib/linuxmint/mintsysadm"),
)

from common.kernel_cleanup import (
    CleanupCandidate,
    CleanupSettings,
    get_cleanup_candidates,
    load_cleanup_settings,
    prepare_candidate,
    prune_protected_kernels,
    save_cleanup_settings,
)
from common.kernels import Series


class CleanupPlannerTests(unittest.TestCase):

    def setUp(self):
        self.packages = [
            "linux-headers-6.8.0-40",
            "linux-headers-6.8.0-40-generic",
            "linux-image-6.8.0-40-generic",
            "linux-modules-6.8.0-40-generic",
            "linux-headers-6.8.0-41",
            "linux-headers-6.8.0-41-generic",
            "linux-image-6.8.0-41-generic",
            "linux-modules-6.8.0-41-generic",
            "linux-headers-6.8.0-42",
            "linux-headers-6.8.0-42-generic",
            "linux-image-6.8.0-42-generic",
            "linux-modules-6.8.0-42-generic",
        ]

    def test_tracked_series_keeps_newest_installed_versions(self):
        series = Series(
            name="6.8",
            version="6.8",
            flavor="generic",
            tracked=True,
            installed_versions={
                "6.8.0-40",
                "6.8.0-41",
                "6.8.0-42",
            },
        )
        candidates = get_cleanup_candidates(
            [series],
            2,
            "6.8.0-42-generic",
            self.packages,
        )
        self.assertEqual(len(candidates), 1)
        self.assertEqual(candidates[0].kernel_version, "6.8.0-40")
        self.assertEqual(
            candidates[0].packages,
            {
                "linux-headers-6.8.0-40",
                "linux-headers-6.8.0-40-generic",
                "linux-image-6.8.0-40-generic",
                "linux-modules-6.8.0-40-generic",
            },
        )

    def test_untracked_series_does_not_retain_old_versions(self):
        series = Series(
            name="6.8",
            version="6.8",
            flavor="generic",
            installed_versions={"6.8.0-40", "6.8.0-41"},
        )
        candidates = get_cleanup_candidates(
            [series],
            2,
            "6.8.0-42-generic",
            self.packages,
        )
        versions = set()
        for candidate in candidates:
            versions.add(candidate.kernel_version)
        self.assertEqual(versions, {"6.8.0-40", "6.8.0-41"})

    def test_running_kernel_is_never_a_candidate(self):
        series = Series(
            name="6.8",
            version="6.8",
            flavor="generic",
            installed_versions={"6.8.0-40"},
        )
        candidates = get_cleanup_candidates(
            [series],
            2,
            "6.8.0-40-generic",
            self.packages,
        )
        self.assertEqual(candidates, [])

    def test_duplicate_tracks_share_tracked_state(self):
        hwe = Series(
            name="6.8",
            version="6.8",
            flavor="generic",
            track="hwe",
            tracked=True,
            installed_versions={"6.8.0-40", "6.8.0-41"},
        )
        edge = Series(
            name="6.8",
            version="6.8",
            flavor="generic",
            track="hwe",
            edge=True,
            installed_versions={"6.8.0-40", "6.8.0-41"},
        )
        candidates = get_cleanup_candidates(
            [hwe, edge],
            1,
            "different-kernel",
            self.packages,
        )
        self.assertEqual(len(candidates), 1)
        self.assertEqual(candidates[0].kernel_version, "6.8.0-40")

    def test_settings_round_trip(self):
        with tempfile.TemporaryDirectory() as directory:
            path = os.path.join(directory, "cleanup.conf")
            save_cleanup_settings(
                CleanupSettings(
                    enabled=True,
                    retain=3,
                    protected_kernels={"6.8.0-40-generic"},
                ),
                path,
            )
            self.assertEqual(
                load_cleanup_settings(path),
                CleanupSettings(
                    enabled=True,
                    retain=3,
                    protected_kernels={"6.8.0-40-generic"},
                ),
            )

    def test_protected_kernel_is_not_a_cleanup_candidate(self):
        series = Series(
            name="6.8",
            version="6.8",
            flavor="generic",
            installed_versions={"6.8.0-40", "6.8.0-41"},
        )
        candidates = get_cleanup_candidates(
            [series],
            1,
            "different-kernel",
            self.packages,
            protected_kernels={"6.8.0-40-generic"},
        )
        self.assertEqual(
            {candidate.kernel_version for candidate in candidates},
            {"6.8.0-41"},
        )

    def test_stale_protected_kernels_are_pruned(self):
        settings = CleanupSettings(
            protected_kernels={
                "6.8.0-40-generic",
                "6.8.0-39-generic",
            },
        )
        changed = prune_protected_kernels(
            settings,
            [
                "linux-image-6.8.0-40-generic",
                "linux-modules-6.8.0-40-generic",
            ],
        )
        self.assertTrue(changed)
        self.assertEqual(
            settings.protected_kernels,
            {"6.8.0-40-generic"},
        )

    def test_candidate_is_rejected_if_apt_would_remove_tracked_meta(self):
        class FakePackage:

            def __init__(self, name, cache):
                self.name = name
                self.cache = cache
                self.is_installed = True
                self.marked_delete = False

            def mark_delete(self, auto_fix=True, purge=True):
                self.marked_delete = True
                self.cache["linux-generic"].marked_delete = True

        class FakeCache(dict):

            def get_changes(self):
                changes = []
                for package in self.values():
                    if package.marked_delete:
                        changes.append(package)
                return changes

        cache = FakeCache()
        cache["linux-generic"] = FakePackage("linux-generic", cache)
        package_name = "linux-image-6.8.0-40-generic"
        cache[package_name] = FakePackage(package_name, cache)
        candidate = CleanupCandidate(
            series_version="6.8",
            flavor="generic",
            kernel_version="6.8.0-40",
            packages={package_name},
        )
        removals, reason = prepare_candidate(
            cache,
            candidate,
            {"linux-generic"},
        )
        self.assertIsNone(removals)
        self.assertIn("linux-generic", reason)

    def test_candidate_is_rejected_if_apt_would_remove_non_kernel_package(self):
        class FakePackage:

            def __init__(self, name, cache):
                self.name = name
                self.cache = cache
                self.is_installed = True
                self.marked_delete = False

            def mark_delete(self, auto_fix=True, purge=True):
                self.marked_delete = True
                self.cache["unrelated-package"].marked_delete = True

        class FakeCache(dict):

            def get_changes(self):
                return [
                    package
                    for package in self.values()
                    if package.marked_delete
                ]

        cache = FakeCache()
        package_name = "linux-image-6.8.0-40-generic"
        cache[package_name] = FakePackage(package_name, cache)
        cache["unrelated-package"] = FakePackage(
            "unrelated-package",
            cache,
        )
        candidate = CleanupCandidate(
            series_version="6.8",
            flavor="generic",
            kernel_version="6.8.0-40",
            packages={package_name},
        )
        removals, reason = prepare_candidate(cache, candidate, set())
        self.assertIsNone(removals)
        self.assertIn("unrelated-package", reason)

    def test_candidate_accepts_versioned_main_module_packages(self):
        class FakePackage:

            def __init__(self, name):
                self.name = name
                self.is_installed = True
                self.marked_delete = False

            def mark_delete(self, auto_fix=True, purge=True):
                self.marked_delete = True

        class FakeCache(dict):

            def get_changes(self):
                return [
                    package
                    for package in self.values()
                    if package.marked_delete
                ]

        image = "linux-image-7.0.0-1009-oem"
        modules = "linux-main-modules-zfs-7.0.0-1009-oem"
        cache = FakeCache({
            image: FakePackage(image),
            modules: FakePackage(modules),
        })
        candidate = CleanupCandidate(
            series_version="7.0",
            flavor="oem",
            kernel_version="7.0.0-1009",
            packages={image, modules},
        )
        removals, reason = prepare_candidate(cache, candidate, set())
        self.assertEqual(removals, {image, modules})
        self.assertEqual(reason, "")


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