"""How a deployment restarts itself after a release is switched.

Pointing `current` at a new release changes nothing on its own — a running
process is still executing the code it started with. Something has to restart
it, and what that something is depends entirely on how the deployment is run:
systemd on a VPS, the orchestrator under Docker, or a person during a manual
first install.

A protocol rather than a branch on an environment variable, so a deployment
model this was not written for can be supported without touching the installer.
The installer only knows that it asked for a restart and was told whether it
worked.
"""

from __future__ import annotations

import logging
import shutil
import subprocess
from dataclasses import dataclass
from typing import Protocol, runtime_checkable

logger = logging.getLogger(__name__)

#: Long enough for a service that has to close a database pool and finish an
#: in-flight document; short enough that a wedged restart is reported rather
#: than hanging the install.
DEFAULT_TIMEOUT_SECONDS = 60.0


@dataclass(frozen=True, slots=True)
class RestartResult:
    ok: bool
    detail: str = ""


@runtime_checkable
class RestartStrategy(Protocol):
    """Restarts the service. Never raises — a failed restart is a result."""

    def restart(self) -> RestartResult: ...


class NoRestart:
    """Does nothing, and says so.

    The default, deliberately. An installer that guesses how to restart a
    service it does not manage will eventually guess wrong on someone's
    production box; being told "the release is live, restart it yourself" is a
    worse experience and a much better failure mode.
    """

    def restart(self) -> RestartResult:
        return RestartResult(True, "no restart configured; restart the service to pick up the release")


class CommandRestart:
    """Runs a command and treats a zero exit as success."""

    def __init__(self, argv: list[str], timeout: float = DEFAULT_TIMEOUT_SECONDS) -> None:
        if not argv:
            raise ValueError("A restart command needs at least a program name.")

        self._argv = list(argv)
        self._timeout = timeout

    def restart(self) -> RestartResult:
        if shutil.which(self._argv[0]) is None:
            return RestartResult(False, f"{self._argv[0]} is not on PATH")

        try:
            completed = subprocess.run(  # noqa: S603 - argv is operator-configured, never parsed from input
                self._argv,
                capture_output=True,
                text=True,
                timeout=self._timeout,
                check=False,
            )
        except subprocess.TimeoutExpired:
            return RestartResult(False, f"restart timed out after {self._timeout:.0f}s")
        except OSError as exc:
            return RestartResult(False, f"could not run the restart command: {exc}")

        if completed.returncode == 0:
            return RestartResult(True, "restarted")

        # Trimmed: this ends up in state.json and in an operator's terminal, and
        # a page of systemd output buries the line that matters.
        detail = (completed.stderr or completed.stdout or "").strip().splitlines()
        tail = detail[-1] if detail else f"exit {completed.returncode}"

        return RestartResult(False, f"restart failed: {tail}")

    def __repr__(self) -> str:  # pragma: no cover - diagnostics only
        return f"CommandRestart({self._argv!r})"


def systemd(unit: str, timeout: float = DEFAULT_TIMEOUT_SECONDS) -> CommandRestart:
    return CommandRestart(["systemctl", "restart", unit], timeout=timeout)


def from_env(env: dict[str, str]) -> RestartStrategy:
    """Build the configured strategy.

    `TAXPILOT_RESTART_UNIT`     a systemd unit to restart
    `TAXPILOT_RESTART_COMMAND`  an explicit command, split on spaces

    Neither set means NoRestart — the safe default, not an error. A first
    install on a box where the service is not yet managed is a normal thing to
    do, and refusing it would just mean the operator sets the variable to
    `true` to get past the check.
    """
    unit = (env.get("TAXPILOT_RESTART_UNIT") or "").strip()

    if unit:
        return systemd(unit)

    command = (env.get("TAXPILOT_RESTART_COMMAND") or "").strip()

    if command:
        return CommandRestart(command.split())

    return NoRestart()
