#!/usr/bin/env python3
"""USM system-package-manager helper for DNF-based systems.

Implements the USM SPM contract (see usm/README.md, "System package manager
integration") for Fedora-style systems:

    usm-spm-dnf query <usm-ref>...      -> contract JSON on STDOUT
    usm-spm-dnf install [--test] <name>... -> contract JSONL events on STDOUT

API choice: Fedora 43 ships dnf5 as the CLI, but this machine exposes the
DNF4 Python bindings (`python3-dnf`, `import dnf` succeeds) while the dnf5
bindings (`python3-libdnf5`, `import libdnf5`) are not installed. This helper
is therefore implemented against the DNF4 Python API, which resolves against
the same repositories and rpmdb as dnf5. If python3-libdnf5 becomes the only
option, this file is the one to port.

The query subcommand never modifies system state (at most it refreshes
package-metadata caches). The install subcommand must run as root; --test
resolves and downloads but runs the rpm transaction in test mode only.
"""

import argparse
import json
import os
import sys

import dnf
import dnf.callback
import dnf.transaction
import dnf.yum.rpmtrans
import hawkey

BASE_ARCH = hawkey.detect_arch()

INSTALL_ACTIONS = frozenset([
    dnf.transaction.PKG_INSTALL,
    dnf.transaction.PKG_UPGRADE,
    dnf.transaction.PKG_DOWNGRADE,
    dnf.transaction.PKG_REINSTALL,
])

EXIT_OK = 0
EXIT_FAILURE = 1
EXIT_RESOLVE = 3
EXIT_DOWNLOAD = 4
EXIT_TRANSACTION = 5


class HelperError(Exception):
    """Fatal helper failure; message is emitted as a contract error event."""

    def __init__(self, message, exit_code=EXIT_FAILURE):
        super(HelperError, self).__init__(message)
        self.exit_code = exit_code


def emit_event(event):
    """Write one contract JSONL event to STDOUT and flush it."""
    sys.stdout.write(json.dumps(event) + "\n")
    sys.stdout.flush()


def fail(message, exit_code=EXIT_FAILURE):
    """Emit the terminal error event and exit non-zero."""
    emit_event({"type": "error", "message": message})
    sys.exit(exit_code)


def fail_plain(message, exit_code=EXIT_FAILURE):
    """Report a query failure on STDERR only, keeping STDOUT contract-clean."""
    sys.stderr.write("usm-spm-dnf: %s\n" % message)
    sys.exit(exit_code)


def make_base(load_system_repo, test=False):
    """Build a filled dnf.Base, using a user cachedir when unprivileged.

    Filelists metadata is enabled in code before the sack load: the DNF4
    Python API downloads only primary metadata by default (bin:/lib:
    provides queries) and silently ignores optional_metadata_types set via
    dnf.conf, so pc:/vapi:/file-path queries would return not-found on any
    machine whose cache lacks filelists (fresh containers, for instance).
    """
    base = dnf.Base()
    if "filelists" not in base.conf.optional_metadata_types:
        base.conf.optional_metadata_types.append("filelists")
    if os.geteuid() != 0:
        base.conf.cachedir = os.path.expanduser("~/.cache/dnf")
        base.conf.history_record = False
    if test:
        base.conf.tsflags.append("test")
    base.read_all_repos()
    base.fill_sack(
        load_system_repo=load_system_repo,
        load_available_repos=True,
    )
    return base


def make_closure_base():
    """A repo-only sack restricted to the running arch for goal closures.

    dependency-count models "a solo install on this machine", so packages
    for foreign architectures are excluded the way a real single-arch
    install would never pull them.
    """
    base = make_base(load_system_repo=False)
    foreign = base.sack.query().filter(
        arch=[a for a in ("i686", "i386", "armv7hl", "ppc64le", "s390x")
              if a != BASE_ARCH])
    base.sack.add_excludes(foreign)
    return base


def translate_ref(ref):
    """Map a USM resource ref to (filter_kind, value) sack query descriptors.

    Returns None for refs whose type has no file/provides translation.
    """
    prefix, sep, resource = ref.partition(":")
    if not sep or not resource:
        return None
    file_roots = {
        "bin": "/usr/bin",
        "sbin": "/usr/sbin",
        "libexec": "/usr/libexec",
        "gir": "/usr/share/gir-1.0",
        "typelib": "/usr/lib64/girepository-1.0",
        "res": "/usr/share",
        "cfg": "/etc",
        "man": "/usr/share/man",
        "info": "/usr/share/info",
        "locale": "/usr/share/locale",
    }
    if prefix in file_roots:
        return [("file", file_roots[prefix] + "/" + resource)]
    if prefix == "vapi":
        return [
            ("file", "/usr/share/vala/vapi/" + resource),
            ("file__glob", "/usr/share/vala-*/vapi/" + resource),
        ]
    if prefix == "lib":
        return [
            ("provides", resource),
            ("file", "/usr/lib64/" + resource),
        ]
    if prefix == "gio":
        return [("file", "/usr/lib64/gio/modules/" + resource)]
    if prefix == "pc":
        return [
            ("file", "/usr/share/pkgconfig/" + resource),
            ("file", "/usr/lib64/pkgconfig/" + resource),
        ]
    if prefix == "inc":
        return [("file__glob", "/usr/include/" + resource + "*")]
    if prefix == "rootpath":
        return [("file", "/" + resource.lstrip("/"))]
    if prefix == "tag":
        return [("file", "/usr/share/usm-tags/" + resource.replace(".", "/") + ".tag")]
    return None


def arch_rank(package):
    """Sort key preferring the running arch, then noarch, then others."""
    if package.arch == BASE_ARCH:
        return (0, package.arch)
    if package.arch == "noarch":
        return (1, package.arch)
    return (2, package.arch)


def find_candidates(union, descriptors):
    """Best package per name matching any descriptor, latest version per arch."""
    by_name = {}
    for kind, value in descriptors:
        matches = union.filter(**{kind: value}).filter(latest_per_arch=True)
        for package in matches:
            current = by_name.get(package.name)
            if current is None or arch_rank(package) < arch_rank(current):
                by_name[package.name] = package
    return by_name


def closure_counts(closure_base, installed_keys, package, closure_cache):
    """(dependency-count, installed-dependency-count) for a solo install.

    The full closure (including the package itself) comes from a goal run
    against a sack with no @System, so every dependency resolves as an
    install; the overlap with @System is then counted by (name, arch).
    """
    cache_key = (package.name, package.arch)
    if cache_key in closure_cache:
        return closure_cache[cache_key]

    target_query = closure_base.sack.query().filter(
        name=package.name,
        arch=package.arch,
        epoch=package.epoch,
        version=package.version,
        release=package.release,
    )
    target = next(iter(target_query), None)
    if target is None:
        target = next(iter(
            closure_base.sack.query().filter(
                name=package.name, arch=package.arch)
            .filter(latest_per_arch=True)), None)

    counts = (1, 1 if cache_key in installed_keys else 0)
    if target is not None:
        goal = dnf.goal.Goal(closure_base.sack)
        goal.install(target)
        if goal.run():
            transaction = (goal.list_installs() + goal.list_upgrades()
                           + goal.list_downgrades() + goal.list_reinstalls())
            counts = (
                len(transaction),
                sum(1 for member in transaction
                    if (member.name, member.arch) in installed_keys),
            )
        else:
            sys.stderr.write(
                "usm-spm-dnf: could not resolve solo install of %s: %s\n"
                % (package, goal.problem_string() if hasattr(goal, "problem_string") else "unresolved dependency"))
    else:
        sys.stderr.write(
            "usm-spm-dnf: %s not found in enabled repositories, "
            "estimating dependency counts as installed-only\n" % package)
    closure_cache[cache_key] = counts
    return counts


def cmd_query(args):
    try:
        return run_query(args)
    except SystemExit:
        raise
    except Exception as e:
        fail_plain("query failed: %s" % e, EXIT_RESOLVE)


def run_query(args):
    base = make_base(load_system_repo=True)
    installed_keys = set(
        (package.name, package.arch)
        for package in base.sack.query().installed())
    union = base.sack.query().available().union(base.sack.query().installed())

    refs = list(dict.fromkeys(args.refs))
    not_found = []
    candidates = {}
    for ref in refs:
        descriptors = translate_ref(ref)
        if descriptors is None:
            sys.stderr.write(
                "usm-spm-dnf: resource type of \"%s\" has no "
                "system-package-manager translation\n" % ref)
            not_found.append(ref)
            continue
        matches = find_candidates(union, descriptors)
        if not matches:
            not_found.append(ref)
            continue
        for name, package in matches.items():
            entry = candidates.get(name)
            if entry is None:
                entry = {"package": package, "resources": []}
                candidates[name] = entry
            entry["resources"].append(ref)

    closure_base = None
    closure_cache = {}
    packages = []
    for name in sorted(candidates):
        entry = candidates[name]
        if closure_base is None:
            closure_base = make_closure_base()
        dependency_count, installed_dependency_count = closure_counts(
            closure_base, installed_keys, entry["package"], closure_cache)
        packages.append({
            "name": name,
            "resources": entry["resources"],
            "dependency-count": dependency_count,
            "installed-dependency-count": installed_dependency_count,
        })

    sys.stdout.write(json.dumps({
        "not-found": not_found,
        "packages": packages,
    }) + "\n")
    sys.stdout.flush()
    return EXIT_OK


class ContractDownloadProgress(dnf.callback.DownloadProgress):
    """Maps dnf package downloads to contract `package` events."""

    def __init__(self):
        self.total = 0
        self.index = 0
        self.last = None

    def start(self, total_files, total_size, total_drpms=0):
        self.total = total_files
        self.index = 0
        self.last = None

    def progress(self, payload, done):
        size = payload.pkg.downloadsize or 0
        fraction = (done / size) if size else 0.0
        event = (
            "package", payload.pkg.name, self.index + 1, self.total,
            round(min(fraction, 1.0), 4))
        if event == self.last:
            return
        self.last = event
        emit_event({
            "type": "package",
            "name": payload.pkg.name,
            "current": self.index + 1,
            "total": self.total,
            "progress": min(fraction, 1.0),
        })

    def end(self, payload, status, msg):
        if status == dnf.callback.STATUS_FAILED:
            raise HelperError(
                "failed to download %s: %s" % (payload.pkg.name, msg or "unknown error"),
                EXIT_DOWNLOAD)
        self.index += 1

    def message(self, msg):
        sys.stderr.write("usm-spm-dnf: %s\n" % msg)


class ContractTransactionDisplay(dnf.yum.rpmtrans.TransactionDisplay):
    """Maps the rpm transaction to contract package/complete events."""

    def __init__(self):
        super(ContractTransactionDisplay, self).__init__()
        self.installed = 0
        self.last = None

    def progress(self, package, action, ti_done, ti_total, ts_done, ts_total):
        if package is None:
            return
        fraction = (float(ti_done) / float(ti_total)) if ti_total else 0.0
        total = ts_total or 1
        current = min(ts_done + 1, total) if ts_total else 1
        event = ("package", package.name, current, total,
                 round(min(fraction, 1.0), 4))
        if event == self.last:
            return
        self.last = event
        emit_event({
            "type": "package",
            "name": package.name,
            "current": current,
            "total": total,
            "progress": min(fraction, 1.0),
        })

    def filelog(self, package, action):
        if package is None or action not in INSTALL_ACTIONS:
            return
        self.installed += 1
        emit_event({"type": "package-complete", "name": package.name})

    def scriptout(self, msgs):
        if msgs:
            sys.stderr.write(msgs if msgs.endswith("\n") else msgs + "\n")

    def error(self, message):
        raise HelperError(
            "transaction failed: %s" % (message or "unknown rpm error"),
            EXIT_TRANSACTION)


def cmd_install(args):
    base = make_base(load_system_repo=True, test=args.test)
    for name in args.names:
        try:
            base.install(name)
        except Exception as e:
            fail("could not mark \"%s\" for install: %s" % (name, e),
                 EXIT_RESOLVE)
    try:
        base.resolve()
    except Exception as e:
        fail("dependency resolution failed: %s" % e, EXIT_RESOLVE)

    install_set = base.transaction.install_set
    emit_event({"type": "begin", "total": len(install_set)})

    progress = ContractDownloadProgress()
    try:
        base.download_packages(install_set, progress=progress)
    except HelperError:
        raise
    except Exception as e:
        fail("failed to download packages: %s" % e, EXIT_DOWNLOAD)

    display = ContractTransactionDisplay()
    try:
        base.do_transaction(display=display)
    except HelperError:
        raise
    except Exception as e:
        fail("transaction failed: %s" % e, EXIT_TRANSACTION)

    emit_event({
        "type": "complete",
        "status": "ok",
        "installed": display.installed,
    })
    return EXIT_OK


def main():
    parser = argparse.ArgumentParser(
        prog="usm-spm-dnf",
        description="USM system-package-manager helper for DNF")
    subparsers = parser.add_subparsers(dest="command", required=True)

    query_parser = subparsers.add_parser(
        "query", help="resolve USM resource refs to system packages")
    query_parser.add_argument(
        "refs", nargs="+", metavar="USM-REF",
        help="resource ref, e.g. bin:valac or lib:libglib-2.0.so.0")
    query_parser.set_defaults(handler=cmd_query)

    install_parser = subparsers.add_parser(
        "install", help="install system packages, streaming progress events")
    install_parser.add_argument(
        "--test", action="store_true",
        help="resolve and download only; run the rpm transaction in test mode")
    install_parser.add_argument(
        "names", nargs="+", metavar="NAME",
        help="native system package name")
    install_parser.set_defaults(handler=cmd_install)

    args = parser.parse_args()
    try:
        return args.handler(args)
    except HelperError as e:
        fail(str(e), e.exit_code)
    except SystemExit:
        raise
    except Exception as e:
        fail("unexpected failure: %s" % e, EXIT_FAILURE)


if __name__ == "__main__":
    sys.exit(main())
