#!/usr/bin/env python3
"""ProtectMyi localhost installer bridge.

Run this on the workstation that can reach IBM i. The ProtectMyi portal talks
only to 127.0.0.1. IBM i credentials never leave this process/workstation.
"""
from __future__ import annotations
import json
import os
from pathlib import Path
import shlex
import shutil
import subprocess
import tempfile
import urllib.request
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

VERSION = "0.5.1.1"
HOST = "127.0.0.1"
PORT = 17841
ALLOWED_ORIGINS = {
    "https://pmi-app-e84ac0.web.app",
    "https://app.protectmyi.org",
    "http://localhost",
    "http://127.0.0.1",
}


def _run_password(cmd: list[str], password: str, cwd: str | None = None) -> str:
    """Run OpenSSH using an ephemeral SSH_ASKPASS program.

    The password is passed only through process environment, never argv or disk
    outside a 0700 temporary directory which is destroyed immediately.
    """
    with tempfile.TemporaryDirectory(prefix="pmi-askpass-") as td:
        askpass = Path(td) / "askpass.sh"
        askpass.write_text('#!/bin/sh\nprintf "%s\\n" "$PMI_SSH_PASSWORD"\n', encoding="utf-8")
        os.chmod(askpass, 0o700)
        env = os.environ.copy()
        env.update({
            "SSH_ASKPASS": str(askpass),
            "SSH_ASKPASS_REQUIRE": "force",
            "DISPLAY": env.get("DISPLAY") or ":0",
            "PMI_SSH_PASSWORD": password,
        })
        cp = subprocess.run(
            cmd,
            cwd=cwd,
            env=env,
            stdin=subprocess.DEVNULL,
            stdout=subprocess.PIPE,
            stderr=subprocess.STDOUT,
            text=True,
            start_new_session=True,
            timeout=300,
        )
        if cp.returncode:
            raise RuntimeError(cp.stdout.strip() or f"Command failed ({cp.returncode})")
        return cp.stdout


def _ssh_options() -> list[str]:
    return [
        "-o", "BatchMode=no",
        "-o", "PreferredAuthentications=password,keyboard-interactive",
        "-o", "NumberOfPasswordPrompts=1",
        "-o", "StrictHostKeyChecking=accept-new",
        "-o", "ConnectTimeout=12",
    ]


def preflight(payload: dict) -> dict:
    host = str(payload.get("host") or "").strip()
    user = str(payload.get("user") or "").strip()
    password = str(payload.get("password") or "")
    if not all((host, user, password)):
        raise ValueError("Hostname/IP, user profile and password are required")
    if shutil.which("ssh") is None:
        raise RuntimeError("ssh is required on this computer")
    remote = f"{user}@{host}"
    check = "test -x /QOpenSys/usr/bin/system && command -v tar >/dev/null 2>&1 && printf PROTECTMYI_OK"
    output = _run_password(["ssh", *_ssh_options(), remote, check], password)
    if "PROTECTMYI_OK" not in output:
        raise RuntimeError("SSH connected, but required IBM i PASE tools are unavailable")
    return {"ok": True}


def install(payload: dict) -> dict:
    host = str(payload.get("host") or "").strip()
    user = str(payload.get("user") or "").strip()
    password = str(payload.get("password") or "")
    token = str(payload.get("claimToken") or "").strip()
    region = str(payload.get("dataRegion") or "US").upper()
    package_url = str(payload.get("packageUrl") or "").strip()
    claim_url = str(payload.get("claimUrl") or "").strip()
    heartbeat_url = str(payload.get("heartbeatUrl") or "").strip()
    if not all((host, user, password, token, package_url, claim_url, heartbeat_url)):
        raise ValueError("Missing installation data")
    if region not in {"US", "CA"}:
        raise ValueError("Invalid data region")
    for exe in ("ssh", "scp"):
        if shutil.which(exe) is None:
            raise RuntimeError(f"{exe} is required on this computer")

    remote = f"{user}@{host}"
    opts = _ssh_options()

    with tempfile.TemporaryDirectory(prefix="protectmyi-install-") as td:
        td_path = Path(td)
        pkg = td_path / "client.tar"
        token_file = td_path / "claim.token"
        urllib.request.urlretrieve(package_url, pkg)
        token_file.write_text(token + "\n", encoding="utf-8")
        os.chmod(token_file, 0o600)

        remote_root = f"/tmp/protectmyi-install-{os.getpid()}"
        _run_password(["ssh", *opts, remote, f"mkdir -p {shlex.quote(remote_root)} && chmod 700 {shlex.quote(remote_root)}"], password)
        _run_password(["scp", *opts, str(pkg), f"{remote}:{remote_root}/client.tar"], password)
        _run_password(["scp", *opts, str(token_file), f"{remote}:{remote_root}/claim.token"], password)
        cmd = (
            f"cd {shlex.quote(remote_root)} && chmod 600 claim.token && "
            "tar -xf client.tar && "
            "/QOpenSys/usr/bin/sh protectmyi-ibmi-client/install/install.sh "
            f"--token-file {shlex.quote(remote_root + '/claim.token')} "
            f"--region {shlex.quote(region)} "
            f"--claim-url {shlex.quote(claim_url)} "
            f"--heartbeat-url {shlex.quote(heartbeat_url)}; "
            "rc=$?; rm -f claim.token; exit $rc"
        )
        output = _run_password(["ssh", *opts, remote, cmd], password)
    return {"ok": True, "output": output[-4000:]}


class Handler(BaseHTTPRequestHandler):
    server_version = "ProtectMyiInstaller/" + VERSION

    def _origin_ok(self) -> bool:
        origin = self.headers.get("Origin", "")
        return not origin or origin in ALLOWED_ORIGINS

    def _cors(self) -> None:
        origin = self.headers.get("Origin", "")
        if origin in ALLOWED_ORIGINS:
            self.send_header("Access-Control-Allow-Origin", origin)
            self.send_header("Vary", "Origin")
        self.send_header("Access-Control-Allow-Headers", "Content-Type")
        self.send_header("Access-Control-Allow-Methods", "GET,POST,OPTIONS")
        self.send_header("Access-Control-Allow-Private-Network", "true")

    def _json(self, status: int, value: dict) -> None:
        body = json.dumps(value).encode()
        self.send_response(status)
        self._cors()
        self.send_header("Content-Type", "application/json")
        self.send_header("Content-Length", str(len(body)))
        self.end_headers()
        self.wfile.write(body)

    def do_OPTIONS(self) -> None:
        if not self._origin_ok():
            return self._json(403, {"error": "origin_not_allowed"})
        self.send_response(204)
        self._cors()
        self.end_headers()

    def do_GET(self) -> None:
        if not self._origin_ok():
            return self._json(403, {"error": "origin_not_allowed"})
        if self.path == "/health":
            return self._json(200, {"ok": True, "version": VERSION})
        return self._json(404, {"error": "not_found"})

    def do_POST(self) -> None:
        if not self._origin_ok():
            return self._json(403, {"error": "origin_not_allowed"})
        if self.path not in {"/preflight", "/install"}:
            return self._json(404, {"error": "not_found"})
        try:
            length = int(self.headers.get("Content-Length", "0"))
            if length <= 0 or length > 65536:
                raise ValueError("Invalid request size")
            payload = json.loads(self.rfile.read(length))
            result = preflight(payload) if self.path == "/preflight" else install(payload)
            return self._json(200, result)
        except subprocess.TimeoutExpired:
            return self._json(504, {"error": "SSH operation timed out"})
        except Exception as exc:
            return self._json(500, {"error": str(exc)[:1000]})

    def log_message(self, fmt: str, *args) -> None:
        print("ProtectMyi Installer:", fmt % args)


if __name__ == "__main__":
    print(f"ProtectMyi Installer {VERSION}")
    print(f"Listening only on http://{HOST}:{PORT}")
    print("Leave this window open while installing from the ProtectMyi portal.")
    ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
