#!/usr/bin/env python3
"""Atlas reachability helper — always find a working Atlas base URL.

The Protect box must reach the Atlas server for heartbeats, owner-login,
portal self-updates and the chat archive. Relying on one fixed URL is
fragile: many customer routers / ISPs fail to resolve the funnel name
(*.ts.net) over plain DNS even though the internet works, and the box
must never depend on Tailscale MagicDNS being enabled.

This module probes a list of candidate base URLs in order and caches the
first one that responds, so boxes keep working through:
  - env ATLAS_URL / acp.env ATLAS_URL      (funnel, public internet path)
  - env ATLAS_LAN_URL                      (optional LAN https URL, e.g. https://192.168.1.10:9443)
  - env ATLAS_URLS                         (optional comma-separated extra URLs)
  - https://atlas-server.taile9cc75.ts.net (funnel default)
  - https://atlas-server                   (MagicDNS short name, if enabled)
  - https://<tailscale-ip of atlas-server> (resolved via `tailscale ip -4
                                             atlas-server` — works even with
                                             MagicDNS off and DNS broken)

Any HTTP response (even 401/404) counts as "reachable" — only connection /
DNS errors count as failure. The winner is cached in state.json so the
agent + portal share one URL and re-probe at most every 10 minutes.
"""
from __future__ import annotations

import json
import os
import ssl
import subprocess
import time
import urllib.request
from pathlib import Path

DATA = Path("/var/lib/atlas-cyber-protect")
# Own file on purpose: the agent's save_state() rewrites state.json every loop
# and would wipe this cache. A separate file keeps the reachable URL sticky.
REACH_FILE = DATA / "atlas-reach.json"
ENV_FILE = Path("/etc/atlas-cyber-protect/acp.env")
FUNNEL = "https://atlas-server.taile9cc75.ts.net"
PROBE_PATH = "/api/cyber-protect/dashboard"
CACHE_TTL = 600  # re-probe at most every 10 min when healthy
FAIL_TTL = 45    # retry soon after everything failed


def _read_env() -> dict[str, str]:
    out: dict[str, str] = {}
    try:
        for line in ENV_FILE.read_text(encoding="utf-8").splitlines():
            line = line.strip()
            if not line or line.startswith("#") or "=" not in line:
                continue
            k, v = line.split("=", 1)
            out[k.strip()] = v.strip().strip('"').strip("'")
    except Exception:
        pass
    return out


def candidates() -> list[str]:
    """Ordered, de-duplicated list of Atlas base URLs to try."""
    env = _read_env()
    seen: set[str] = set()
    out: list[str] = []

    def add(u: str | None) -> None:
        u = (u or "").strip().rstrip("/")
        if not u or u in seen:
            return
        seen.add(u)
        out.append(u)

    add(os.environ.get("ATLAS_URL") or env.get("ATLAS_URL"))
    add(env.get("ATLAS_LAN_URL"))
    add(os.environ.get("ATLAS_LAN_URL"))
    for extra in str(os.environ.get("ATLAS_URLS") or env.get("ATLAS_URLS") or "").split(","):
        add(extra.strip())
    add(FUNNEL)
    add("https://atlas-server")
    # Tailnet IP resolved via the Tailscale CLI — no DNS resolver needed.
    try:
        ip = subprocess.check_output(
            ["tailscale", "ip", "-4", "atlas-server"],
            text=True,
            stderr=subprocess.DEVNULL,
            timeout=5,
        ).strip().splitlines()
        if ip:
            add(f"https://{ip[0].strip()}")
    except Exception:
        pass
    return out


def _probe(base: str) -> bool:
    import urllib.parse

    try:
        ctx = ssl.create_default_context()
        ctx.check_hostname = False
        ctx.verify_mode = ssl.CERT_NONE
        # When the candidate is a bare IP (tailnet), send the funnel Host header
        # so the Atlas vhost matches — same as `curl --resolve <funnel>:443:<ip>`.
        parsed = urllib.parse.urlparse(base)
        host = parsed.hostname or ""
        headers = {
            "User-Agent": "atlas-cyber-protect-box",
            "Connection": "close",
        }
        if host.replace(".", "").isdigit():
            headers["Host"] = urllib.parse.urlparse(FUNNEL).hostname or ""
        req = urllib.request.Request(
            f"{base}{PROBE_PATH}",
            headers=headers,
        )
        with urllib.request.urlopen(req, timeout=7, context=ctx) as resp:
            resp.read(128)
            return True
    except urllib.error.HTTPError as he:
        # Any HTTP answer (401 auth required, 403, 404, 5xx) means the host is
        # reachable — only connection/DNS failures mean the URL is unusable.
        try:
            he.read(32)
        except Exception:
            pass
        return True
    except Exception:
        return False


def _cache() -> dict:
    try:
        if REACH_FILE.is_file():
            c = json.loads(REACH_FILE.read_text(encoding="utf-8"))
            if isinstance(c, dict):
                return c
    except Exception:
        pass
    return {}


def _store(base: str, ok: bool) -> None:
    try:
        DATA.mkdir(parents=True, exist_ok=True)
        REACH_FILE.write_text(
            json.dumps(
                {
                    "base": base,
                    "ok": ok,
                    "checked_at": time.time(),
                    "candidates": candidates(),
                },
                indent=2,
            ),
            encoding="utf-8",
        )
    except Exception:
        pass


def atlas_base(force: bool = False) -> str:
    """Return the best currently-working Atlas base URL (probes when stale)."""
    cached = _cache()
    base = str(cached.get("base") or "")
    age = time.time() - float(cached.get("checked_at") or 0)
    ttl = FAIL_TTL if cached.get("ok") is False else CACHE_TTL
    if base and not force and age < ttl:
        return base
    tried = []
    for cand in candidates():
        tried.append(cand)
        if _probe(cand):
            _store(cand, True)
            if base != cand:
                print(f"atlas_reach: using {cand}", flush=True)
            return cand
    # Everything unreachable right now — keep the last known-good URL so the
    # caller still tries (and so nothing breaks when the internet blips).
    if base:
        _store(base, False)
        return base
    return FUNNEL


def reachable(force: bool = False) -> bool:
    """Quick health check: can we reach Atlas right now?"""
    return _probe(atlas_base(force=force))
