"""Email/link phishing guard for Atlas Cyber Protect.

Warn-only: never blocks the customer's network. Checks URLs from PC agents
against heuristics + cached open threat feeds so Protect can warn before
someone clicks a hacker link in email.
"""
from __future__ import annotations

import hashlib
import json
import re
import time
import urllib.error
import urllib.request
from pathlib import Path
from typing import Any
from urllib.parse import urlparse

DATA = Path("/var/lib/atlas-cyber-protect")
CACHE = DATA / "link-guard-cache.json"
LEARN = DATA / "link-guard-learning.jsonl"
HITS = DATA / "link-guard-hits.json"  # portal-owned; agent must not wipe this
FEED_URLS = (
    # Host-style blocklists (domain per line / 127.0.0.1 domain)
    "https://urlhaus.abuse.ch/downloads/hostfile/",
    "https://raw.githubusercontent.com/mitchellkrogza/Phishing.Database/master/phishing-domains-ACTIVE.txt",
)

URL_RE = re.compile(
    r"(?i)\b((?:https?://|www\.)[^\s<>\"']+|(?:[a-z0-9-]+\.)+(?:com|net|org|io|co|za|ru|cn|tk|ml|ga|cf|gq|xyz|top|click|link|info|biz)(?:/[^\s<>\"']*)?)"
)

SUSPICIOUS_TLDS = {
    "tk", "ml", "ga", "cf", "gq", "xyz", "top", "click", "link", "zip", "mov", "country", "stream",
}
BRAND_BAIT = (
    "paypal", "microsoft", "apple", "google", "amazon", "nedbank", "absa", "fnb", "standardbank",
    "capitec", "outlook", "office365", "onedrive", "dhl", "sars", "bank", "secure-login",
    "verify-account", "update-billing",
)
SAFEISH_HOSTS = {
    "google.com", "www.google.com", "microsoft.com", "www.microsoft.com", "apple.com",
    "www.apple.com", "github.com", "wikipedia.org", "cloudflare.com", "openai.com",
    "atlas-server.taile9cc75.ts.net", "redbubble.com", "www.redbubble.com",
}


def _now() -> float:
    return time.time()


def _load_cache() -> dict[str, Any]:
    try:
        if CACHE.is_file():
            raw = json.loads(CACHE.read_text(encoding="utf-8"))
            if isinstance(raw, dict):
                return raw
    except Exception:
        pass
    return {"domains": [], "updated_at": 0, "sources": []}


def _save_cache(cache: dict[str, Any]) -> None:
    try:
        DATA.mkdir(parents=True, exist_ok=True)
        CACHE.write_text(json.dumps(cache, indent=2) + "\n", encoding="utf-8")
    except Exception:
        pass


def refresh_feeds(force: bool = False) -> dict[str, Any]:
    """Pull public phishing host lists (best-effort). Cached ~12h."""
    cache = _load_cache()
    age = _now() - float(cache.get("updated_at") or 0)
    if not force and age < 12 * 3600 and cache.get("domains"):
        return {"ok": True, "cached": True, "count": len(cache.get("domains") or []), "age_s": int(age)}
    domains: set[str] = set(str(d).lower() for d in (cache.get("domains") or []) if d)
    sources: list[str] = []
    for url in FEED_URLS:
        try:
            req = urllib.request.Request(url, headers={"User-Agent": "AtlasCyberProtect/link-guard"})
            with urllib.request.urlopen(req, timeout=25) as resp:
                text = resp.read().decode("utf-8", errors="ignore")
            added = 0
            for line in text.splitlines():
                line = line.strip()
                if not line or line.startswith("#"):
                    continue
                # hostfile: 127.0.0.1 evil.com
                parts = line.split()
                host = parts[-1] if parts[0].startswith("127.") or parts[0] == "0.0.0.0" else parts[0]
                host = host.lower().strip(".")
                if host.startswith("http"):
                    try:
                        host = urlparse(host).hostname or host
                    except Exception:
                        continue
                if not host or "." not in host or len(host) > 180:
                    continue
                if host not in domains:
                    domains.add(host)
                    added += 1
            sources.append(f"{url} +{added}")
        except Exception as exc:
            sources.append(f"{url} ERR {exc}")
    # Keep cache bounded
    if len(domains) > 200000:
        domains = set(list(domains)[:200000])
    cache = {"domains": sorted(domains), "updated_at": _now(), "sources": sources}
    _save_cache(cache)
    return {"ok": True, "cached": False, "count": len(domains), "sources": sources}


def extract_urls(text: str) -> list[str]:
    found: list[str] = []
    seen: set[str] = set()
    for m in URL_RE.finditer(text or ""):
        u = m.group(1).rstrip(").,;]}>\"'")
        if u.lower().startswith("www."):
            u = "http://" + u
        if not u.lower().startswith("http"):
            u = "http://" + u
        key = u.lower()
        if key in seen:
            continue
        seen.add(key)
        found.append(u[:500])
        if len(found) >= 40:
            break
    return found


def _host(url: str) -> str:
    try:
        p = urlparse(url if "://" in url else "http://" + url)
        return (p.hostname or "").lower().strip(".")
    except Exception:
        return ""


def _apex(host: str) -> str:
    parts = [p for p in host.split(".") if p]
    if len(parts) >= 2:
        return ".".join(parts[-2:])
    return host


def check_url(url: str, *, cache: dict[str, Any] | None = None) -> dict[str, Any]:
    """Return verdict for one URL. danger|suspicious|ok|unknown."""
    cache = cache or _load_cache()
    blocked = {str(d).lower() for d in (cache.get("domains") or [])}
    raw = (url or "").strip()
    host = _host(raw)
    apex = _apex(host)
    reasons: list[str] = []
    score = 0

    if not host:
        return {
            "ok": True,
            "url": raw[:500],
            "host": "",
            "verdict": "unknown",
            "score": 0,
            "reasons": ["Could not parse URL"],
            "warn": False,
        }

    if host in SAFEISH_HOSTS or apex in { _apex(h) for h in SAFEISH_HOSTS }:
        return {
            "ok": True,
            "url": raw[:500],
            "host": host,
            "verdict": "ok",
            "score": 0,
            "reasons": ["Known common safe host"],
            "warn": False,
        }

    if host in blocked or apex in blocked:
        score += 80
        reasons.append("Listed in phishing/malware host feed")

    # IP literal in URL
    if re.fullmatch(r"\d{1,3}(\.\d{1,3}){3}", host):
        score += 35
        reasons.append("Link uses raw IP address (common in phishing)")

    # Credentials in URL userinfo
    if "@" in raw.split("://", 1)[-1].split("/", 1)[0]:
        score += 50
        reasons.append("URL hides real destination with @ (classic email trick)")

    tld = host.rsplit(".", 1)[-1] if "." in host else ""
    if tld in SUSPICIOUS_TLDS:
        score += 20
        reasons.append(f"Suspicious TLD .{tld}")

    # Brand bait on non-brand domain (paypal-secure.tk, microsoft-account-verify.xyz)
    hlow = host.lower()
    safe_apexes = {_apex(h) for h in SAFEISH_HOSTS}
    compact = hlow.replace("-", "").replace(".", "")
    for brand in BRAND_BAIT:
        b = brand.replace("-", "").replace(".", "")
        if b and b in compact and apex not in safe_apexes:
            # Real brand sites are in SAFEISH; anything else with the brand name is bait
            score += 30
            reasons.append(f"Looks like '{brand}' brand bait on odd domain")
            break

    # Punycode / homograph
    if "xn--" in host:
        score += 30
        reasons.append("Internationalized (punycode) host — check carefully")

    # Very long subdomain chains
    if host.count(".") >= 4:
        score += 15
        reasons.append("Many subdomains (often used to look official)")

    # http not https for login-ish words
    if raw.lower().startswith("http://") and any(w in raw.lower() for w in ("login", "signin", "verify", "account", "bank", "password")):
        score += 25
        reasons.append("Login/bank wording over plain HTTP")

    if score >= 70:
        verdict = "danger"
    elif score >= 30:
        verdict = "suspicious"
    elif score > 0:
        verdict = "caution"
    else:
        verdict = "ok"

    return {
        "ok": True,
        "url": raw[:500],
        "host": host,
        "apex": apex,
        "verdict": verdict,
        "score": score,
        "reasons": reasons or ["No strong phishing signals"],
        "warn": verdict in {"danger", "suspicious"},
    }


def check_text(text: str, *, source: str = "paste") -> dict[str, Any]:
    refresh_feeds(force=False)
    cache = _load_cache()
    urls = extract_urls(text)
    results = [check_url(u, cache=cache) for u in urls]
    warns = [r for r in results if r.get("warn")]
    danger = [r for r in results if r.get("verdict") == "danger"]
    out = {
        "ok": True,
        "source": source,
        "url_count": len(results),
        "warn_count": len(warns),
        "danger_count": len(danger),
        "results": results,
        "summary": (
            f"Found {len(danger)} dangerous and {len(warns)} warning link(s) in {len(results)} URL(s)."
            if results
            else "No links found in that text."
        ),
        "policy": "warn_only_never_block_network",
        "checked_at": _now(),
    }
    _learn(out, text_snip=(text or "")[:200])
    return out


def load_hits(limit: int = 50) -> list[dict[str, Any]]:
    try:
        if HITS.is_file():
            raw = json.loads(HITS.read_text(encoding="utf-8"))
            if isinstance(raw, list):
                return raw[:limit]
    except Exception:
        pass
    return []


def record_hit(row: dict[str, Any]) -> None:
    """Persist portal warnings outside state.json so the LAN agent cannot wipe them."""
    try:
        DATA.mkdir(parents=True, exist_ok=True)
        hits = load_hits(200)
        hits.insert(0, row)
        HITS.write_text(json.dumps(hits[:80], indent=2) + "\n", encoding="utf-8")
    except Exception:
        pass


def _learn(out: dict[str, Any], text_snip: str = "") -> None:
    try:
        DATA.mkdir(parents=True, exist_ok=True)
        row = {
            "ts": _now(),
            "source": out.get("source"),
            "url_count": out.get("url_count"),
            "warn_count": out.get("warn_count"),
            "danger_count": out.get("danger_count"),
            "hosts": [r.get("host") for r in (out.get("results") or []) if r.get("warn")][:20],
            "snippet": text_snip,
            "lesson": out.get("summary"),
        }
        with LEARN.open("a", encoding="utf-8") as fh:
            fh.write(json.dumps(row, ensure_ascii=False) + "\n")
    except Exception:
        pass


def check_urls(urls: list[str], *, source: str = "agent") -> dict[str, Any]:
    text = "\n".join(str(u) for u in (urls or [])[:40])
    return check_text(text, source=source)
