#!/usr/bin/python3
import os
import shutil
import subprocess

import gi

gi.require_version("Gtk", "3.0")
gi.require_version("Gdk", "3.0")
from gi.repository import Gdk, Gtk

DRIVERCTL = shutil.which("driverctl") or "/usr/sbin/driverctl"
BUSES = ("pci", "usb")
NO_DRIVER = ("", "(none)")

CONSOLE_CSS = b"""
.console-view, .console-view text {
    font-family: monospace;
    background-color: rgba(0, 0, 0, 0.16);
}
"""


def run(args):
    return subprocess.run(args, capture_output=True, text=True)


def pci_descriptions():
    descriptions = {}
    for line in run(["lspci", "-D"]).stdout.splitlines():
        address, _, rest = line.partition(" ")
        _, _, name = rest.partition(": ")
        descriptions[address] = name or rest
    return descriptions


def usb_descriptions():
    descriptions = {}
    for line in run(["lsusb"]).stdout.splitlines():
        parts = line.split()
        if len(parts) >= 6 and parts[0] == "Bus":
            key = "%s-%s" % (parts[1], parts[3].rstrip(":"))
            descriptions[key] = " ".join(parts[6:])
    return descriptions


def driverctl(bus, *args):
    return run([DRIVERCTL, "-b", bus, *args])


def overrides(bus, command):
    result = set()
    for line in driverctl(bus, command).stdout.splitlines():
        parts = line.split()
        if parts:
            result.add(parts[0])
    return result


def list_devices(bus):
    descriptions = pci_descriptions() if bus == "pci" else usb_descriptions()
    active = overrides(bus, "list-overrides")
    persisted = overrides(bus, "list-persisted")
    devices = []
    for line in driverctl(bus, "list-devices").stdout.splitlines():
        parts = line.split()
        if not parts:
            continue
        address = parts[0]
        driver = parts[1] if len(parts) > 1 else ""
        states = []
        if address in active:
            states.append("active")
        if address in persisted:
            states.append("persisted")
        devices.append((address, descriptions.get(address, ""), driver,
                        "+".join(states)))
    return devices


def bindable_drivers(bus):
    drivers = {"vfio-pci"}
    try:
        drivers.update(os.listdir("/sys/bus/%s/drivers" % bus))
    except OSError:
        pass
    return sorted(drivers)


def all_modules():
    drivers = {"vfio-pci"}
    root = "/lib/modules/%s" % os.uname().release
    for _, _, filenames in os.walk(root):
        for filename in filenames:
            name = filename.split(".ko", 1)[0]
            if name != filename:
                drivers.add(name.replace("-", "_"))
    return sorted(drivers)


class DriverctlWindow(Gtk.Window):
    def __init__(self):
        super().__init__(title="driverctl GUI")
        self.set_default_size(860, 600)
        self.install_css()

        box = Gtk.Box(orientation=Gtk.Orientation.VERTICAL, spacing=6)
        for setter in (box.set_margin_top, box.set_margin_bottom,
                       box.set_margin_start, box.set_margin_end):
            setter(6)
        self.add(box)

        toolbar = Gtk.Box(orientation=Gtk.Orientation.HORIZONTAL, spacing=6)
        box.pack_start(toolbar, False, False, 0)

        toolbar.pack_start(Gtk.Label(label="Bus:"), False, False, 0)
        self.bus_combo = Gtk.ComboBoxText()
        for bus in BUSES:
            self.bus_combo.append_text(bus)
        self.bus_combo.set_active(0)
        self.bus_combo.connect("changed", lambda _c: self.refresh())
        toolbar.pack_start(self.bus_combo, False, False, 0)

        self.search_entry = Gtk.SearchEntry()
        self.search_entry.set_placeholder_text("Search devices")
        self.search_entry.set_hexpand(True)
        self.search_entry.connect("search-changed", lambda _e: self.filter.refilter())
        toolbar.pack_start(self.search_entry, True, True, 0)

        self.output_toggle = Gtk.ToggleButton()
        self.output_toggle.set_image(Gtk.Image.new_from_icon_name(
            "utilities-terminal-symbolic", Gtk.IconSize.BUTTON))
        self.output_toggle.set_tooltip_text("Show command output")
        self.output_toggle.connect("toggled", self.on_output_toggled)
        toolbar.pack_start(self.output_toggle, False, False, 0)

        refresh_button = Gtk.Button.new_from_icon_name("view-refresh-symbolic",
                                                       Gtk.IconSize.BUTTON)
        refresh_button.set_tooltip_text("Refresh")
        refresh_button.connect("clicked", lambda _b: self.refresh())
        toolbar.pack_start(refresh_button, False, False, 0)

        self.stack = Gtk.Stack()
        box.pack_start(self.stack, True, True, 0)

        self.store = Gtk.ListStore(str, str, str, str)
        self.filter = self.store.filter_new()
        self.filter.set_visible_func(self.row_visible)
        self.view = Gtk.TreeView(model=self.filter)
        self.view.get_selection().connect("changed", self.on_selection_changed)
        for index, title in enumerate(("Device", "Description", "Driver", "Override")):
            renderer = Gtk.CellRendererText()
            column = Gtk.TreeViewColumn(title, renderer, text=index)
            column.set_resizable(True)
            if index == 1:
                renderer.set_property("ellipsize", 3)
                column.set_expand(True)
            self.view.append_column(column)
        scrolled = Gtk.ScrolledWindow()
        scrolled.set_vexpand(True)
        scrolled.add(self.view)
        self.stack.add_named(scrolled, "list")

        self.empty_label = Gtk.Label(label="No overridable devices on this bus.")
        self.empty_label.get_style_context().add_class("dim-label")
        self.stack.add_named(self.empty_label, "empty")

        box.pack_start(self.build_console(), False, False, 0)

        self.warning_label = Gtk.Label()
        self.warning_label.set_halign(Gtk.Align.START)
        self.warning_label.get_style_context().add_class("error")
        self.warning_label.set_no_show_all(True)
        box.pack_start(self.warning_label, False, False, 0)

        box.pack_start(self.build_controls(), False, False, 0)

        self.update_flag_state()
        self.refresh()

    def install_css(self):
        provider = Gtk.CssProvider()
        provider.load_from_data(CONSOLE_CSS)
        Gtk.StyleContext.add_provider_for_screen(
            Gdk.Screen.get_default(), provider,
            Gtk.STYLE_PROVIDER_PRIORITY_APPLICATION)

    def build_console(self):
        self.output_revealer = Gtk.Revealer()
        frame = Gtk.Frame(label="Command output")
        console = Gtk.Box(orientation=Gtk.Orientation.VERTICAL, spacing=4)
        console.set_margin_top(4)
        console.set_margin_bottom(4)
        console.set_margin_start(4)
        console.set_margin_end(4)

        header = Gtk.Box(orientation=Gtk.Orientation.HORIZONTAL)
        clear = Gtk.Button(label="Clear")
        clear.connect("clicked",
                      lambda _b: self.output_view.get_buffer().set_text(""))
        header.pack_end(clear, False, False, 0)
        console.pack_start(header, False, False, 0)

        self.output_view = Gtk.TextView(editable=False, cursor_visible=False)
        self.output_view.get_style_context().add_class("console-view")
        output_scrolled = Gtk.ScrolledWindow()
        output_scrolled.set_min_content_height(150)
        output_scrolled.add(self.output_view)
        console.pack_start(output_scrolled, True, True, 0)

        frame.add(console)
        self.output_revealer.add(frame)
        return self.output_revealer

    def build_controls(self):
        row = Gtk.Box(orientation=Gtk.Orientation.HORIZONTAL, spacing=6)

        driver_box = Gtk.Box(orientation=Gtk.Orientation.VERTICAL, spacing=2)
        driver_box.set_hexpand(True)
        row.pack_start(driver_box, True, True, 0)

        self.driver_entry = Gtk.Entry()
        self.driver_entry.set_placeholder_text("driver name")
        self.completion_model = Gtk.ListStore(str)
        completion = Gtk.EntryCompletion()
        completion.set_model(self.completion_model)
        completion.set_text_column(0)
        completion.set_minimum_key_length(0)
        completion.set_match_func(self.match_driver)
        self.driver_entry.set_completion(completion)
        self.driver_entry.connect("button-press-event",
                                  lambda _w, _e: completion.complete())
        driver_box.pack_start(self.driver_entry, False, False, 0)

        self.all_drivers_check = Gtk.CheckButton(label="Show all drivers")
        self.all_drivers_check.set_halign(Gtk.Align.START)
        self.all_drivers_check.set_tooltip_text(
            "Include every installed module, even ones not bound to this bus")
        self.all_drivers_check.connect("toggled", lambda _b: self.reload_drivers())
        driver_box.pack_start(self.all_drivers_check, False, False, 0)

        self.apply_now_check = Gtk.CheckButton(label="Apply now")
        self.apply_now_check.set_active(True)
        self.apply_now_check.set_valign(Gtk.Align.START)
        self.apply_now_check.set_tooltip_text(
            "Rebind immediately; may disrupt a device in use")
        self.apply_now_check.connect("toggled", lambda _b: self.update_flag_state())
        row.pack_start(self.apply_now_check, False, False, 0)

        self.persist_check = Gtk.CheckButton(label="Persistent")
        self.persist_check.set_active(True)
        self.persist_check.set_valign(Gtk.Align.START)
        self.persist_check.set_tooltip_text("Keep the change across reboots")
        self.persist_check.connect("toggled", lambda _b: self.update_flag_state())
        row.pack_start(self.persist_check, False, False, 0)

        self.set_button = Gtk.Button(label="Set Override")
        self.set_button.set_valign(Gtk.Align.START)
        self.set_button.connect("clicked", self.on_set)
        row.pack_start(self.set_button, False, False, 0)

        self.unset_button = Gtk.Button(label="Unset Override")
        self.unset_button.set_valign(Gtk.Align.START)
        self.unset_button.connect("clicked", self.on_unset)
        row.pack_start(self.unset_button, False, False, 0)

        self.load_button = Gtk.Button(label="Load Override")
        self.load_button.set_valign(Gtk.Align.START)
        self.load_button.set_tooltip_text(
            "Apply an override that is saved but not currently active")
        self.load_button.connect("clicked", self.on_load)
        row.pack_start(self.load_button, False, False, 0)
        return row

    @property
    def bus(self):
        return self.bus_combo.get_active_text()

    def row_visible(self, model, treeiter, _data):
        text = self.search_entry.get_text().lower()
        if not text:
            return True
        return any(text in (model[treeiter][i] or "").lower() for i in range(4))

    def match_driver(self, _completion, key, treeiter):
        return key in self.completion_model[treeiter][0]

    def on_output_toggled(self, button):
        self.output_revealer.set_reveal_child(button.get_active())

    def update_flag_state(self):
        valid = self.apply_now_check.get_active() or self.persist_check.get_active()
        if valid:
            self.warning_label.hide()
        else:
            self.warning_label.set_text("Enable 'Apply now' or 'Persistent'.")
            self.warning_label.show()
        self.set_button.set_sensitive(valid)
        self.on_selection_changed(None)

    def on_selection_changed(self, _selection):
        row = self.selected_row()
        state = row[3] if row else ""
        valid = self.apply_now_check.get_active() or self.persist_check.get_active()
        self.unset_button.set_sensitive(bool(state) and valid)
        self.load_button.set_sensitive("persisted" in state and "active" not in state)

    def refresh(self):
        self.store.clear()
        devices = list_devices(self.bus)
        for row in devices:
            self.store.append(list(row))
        self.stack.set_visible_child_name("list" if devices else "empty")
        self.reload_drivers()
        self.on_selection_changed(None)

    def reload_drivers(self):
        names = all_modules() if self.all_drivers_check.get_active() \
            else bindable_drivers(self.bus)
        self.completion_model.clear()
        for name in names:
            self.completion_model.append([name])

    def selected_row(self):
        model, treeiter = self.view.get_selection().get_selected()
        return None if treeiter is None else list(model[treeiter])

    def has_driver(self, row):
        return bool(row) and row[2] not in NO_DRIVER

    def flags(self):
        args = []
        if not self.apply_now_check.get_active():
            args.append("--noprobe")
        if not self.persist_check.get_active():
            args.append("--nosave")
        return args

    def on_set(self, _button):
        row = self.selected_row()
        driver = self.driver_entry.get_text().strip()
        if not row or not driver:
            self.warn("Select a device and choose a driver.")
            return
        if driver == "vfio-pci" and not self.apply_now_check.get_active():
            if not self.confirm(
                    "vfio-pci binds immediately",
                    "driverctl always binds vfio-pci right away, even with "
                    "'Apply now' unchecked, and briefly unbinds the current "
                    "driver. If %s is in use this may disrupt the running "
                    "system. Continue?" % row[0]):
                return
        elif self.has_driver(row) and self.apply_now_check.get_active():
            if not self.confirm(
                    "Rebind %s now?" % row[0],
                    "It is currently bound to %s. Rebinding a device that is "
                    "in use may affect the running system." % row[2]):
                return
        self.privileged(self.flags() + ["set-override", row[0], driver])

    def on_unset(self, _button):
        row = self.selected_row()
        if not row:
            self.warn("Select a device.")
            return
        if self.has_driver(row):
            if not self.confirm(
                    "Unbind %s?" % row[0],
                    "Removing the override unbinds it from %s. If the device "
                    "is in use this may affect the running system." % row[2]):
                return
        self.privileged(self.flags() + ["unset-override", row[0]])

    def on_load(self, _button):
        row = self.selected_row()
        if not row:
            self.warn("Select a device.")
            return
        self.privileged(["load-override", row[0]])

    def confirm(self, text, secondary):
        dialog = Gtk.MessageDialog(
            transient_for=self, message_type=Gtk.MessageType.QUESTION,
            buttons=Gtk.ButtonsType.OK_CANCEL, text=text)
        dialog.format_secondary_text(secondary)
        response = dialog.run()
        dialog.destroy()
        return response == Gtk.ResponseType.OK

    def privileged(self, args):
        command = ["pkexec", DRIVERCTL, "-b", self.bus] + args
        result = run(command)
        self.log(command, result)
        if result.returncode not in (0, 126):
            self.warn(result.stderr.strip() or "Operation failed.")
            self.refresh()
            return
        self.refresh()
        if result.returncode == 0:
            self.report(args)

    def report(self, args):
        if "set-override" not in args:
            return
        device = args[-2]
        active = device in overrides(self.bus, "list-overrides")
        persisted = device in overrides(self.bus, "list-persisted")
        if active and persisted:
            self.info("Override applied now and saved for future boots.")
        elif active:
            self.info("Override applied now (not saved).")
        elif persisted:
            self.info("Override saved. Reboot to apply it.")

    def log(self, command, result):
        buffer = self.output_view.get_buffer()
        buffer.insert(buffer.get_end_iter(), "$ %s\n" % " ".join(command[1:]))
        if result.stdout:
            buffer.insert(buffer.get_end_iter(), result.stdout)
        if result.stderr:
            buffer.insert(buffer.get_end_iter(), result.stderr)
        if result.returncode not in (0, 126):
            buffer.insert(buffer.get_end_iter(), "[exit %d]\n" % result.returncode)
        buffer.insert(buffer.get_end_iter(), "\n")
        self.output_view.scroll_to_iter(buffer.get_end_iter(), 0.0, False, 0, 0)

    def info(self, message):
        self._dialog(Gtk.MessageType.INFO, message)

    def warn(self, message):
        self._dialog(Gtk.MessageType.WARNING, message)

    def _dialog(self, kind, message):
        dialog = Gtk.MessageDialog(
            transient_for=self, message_type=kind,
            buttons=Gtk.ButtonsType.OK, text=message)
        dialog.run()
        dialog.destroy()


def main():
    window = DriverctlWindow()
    window.connect("destroy", Gtk.main_quit)
    window.show_all()
    Gtk.main()


if __name__ == "__main__":
    main()
