#!/usr/bin/env python3
"""Shared .env loader for scan_shield_v3 workers and report scripts."""
from __future__ import annotations

import os
from email.utils import formataddr, parseaddr
from pathlib import Path

# Keys from .env always override existing process env (fixes stale SMTP_* in shell).
OVERRIDE_KEYS = frozenset(
    {
        "SMTP_HOST",
        "SMTP_PORT",
        "SMTP_USER",
        "SMTP_PASSWORD",
        "SMTP_FROM",
        "SMTP_FROM_NAME",
        "SCAN_REPORT_SMTP_HOST",
        "SCAN_REPORT_SMTP_PORT",
        "SCAN_REPORT_SMTP_USER",
        "SCAN_REPORT_SMTP_PASSWORD",
        "SCAN_REPORT_SMTP_FROM",
        "SCAN_REPORT_SMTP_FROM_NAME",
        "SCAN_REPORT_BRAND_NAME",
        "SCAN_REPORT_EMAIL_TO",
        "SCAN_OPS_ALERT_EMAIL",
        "SUPABASE_URL",
        "SUPABASE_SERVICE_ROLE_KEY",
    }
)

# Tracks which file last set each key (for startup logging).
_loaded_from: dict[str, str] = {}


def repo_root() -> Path:
    """Directory containing generate_reports_from_update.py."""
    return Path(__file__).resolve().parent


def discover_env_files(start_path: Path | None = None) -> list[Path]:
    """Return env files in load order (later files override earlier for OVERRIDE_KEYS)."""
    root = repo_root()
    explicit = (os.environ.get("SCAN_SHIELD_ENV_FILE") or "").strip()
    if explicit:
        p = Path(explicit)
        if p.is_file():
            return [p.resolve()]

    candidates: list[Path] = [
        root / ".env",
        root / "smtp.env",
        root / "supabase.env",
    ]

    if start_path:
        sp = start_path.resolve()
        bundle = sp.parent if sp.is_file() else sp
        if (bundle / "worker_security.py").is_file():
            for name in (".env", "supabase.env"):
                extra = bundle / name
                if extra.is_file():
                    candidates.append(extra)

    seen: set[Path] = set()
    out: list[Path] = []
    for path in candidates:
        if not path.is_file():
            continue
        resolved = path.resolve()
        if resolved in seen:
            continue
        seen.add(resolved)
        out.append(resolved)
    return out


def _parse_env_line(raw: str) -> tuple[str, str] | None:
    line = raw.strip()
    if not line or line.startswith("#") or "=" not in line:
        return None
    key, _, val = line.partition("=")
    key, val = key.strip(), val.strip().strip('"').strip("'")
    if not key:
        return None
    return key, val


def load_env_file(path: Path, *, override: bool = True) -> None:
    """Load one KEY=VAL file into os.environ."""
    if not path.is_file():
        return
    label = str(path)
    for raw in path.read_text(encoding="utf-8").splitlines():
        parsed = _parse_env_line(raw)
        if not parsed:
            continue
        key, val = parsed
        if override or key in OVERRIDE_KEYS or key not in os.environ:
            os.environ[key] = val
            _loaded_from[key] = label


def load_env_files(start_path: Path | None = None, *, override: bool = True) -> None:
    """Load all discovered env files; .env values override stale process env."""
    for path in discover_env_files(start_path):
        load_env_file(path, override=override)


def _smtp_get(primary: str, fallback: str, default: str = "") -> str:
    v = os.environ.get(primary, "").strip()
    if v:
        return v
    return os.environ.get(fallback, default).strip()


def mail_from_header(default_addr: str = "security@overdrive.co.za") -> str:
    """RFC5322 From header: display name + address from env."""
    addr = _smtp_get("SCAN_REPORT_SMTP_FROM", "SMTP_FROM", default_addr) or default_addr
    name = (
        _smtp_get("SCAN_REPORT_SMTP_FROM_NAME", "SMTP_FROM_NAME")
        or (os.environ.get("SCAN_REPORT_BRAND_NAME") or "").strip()
        or "Silicon Overdrive"
    )
    return formataddr((name, addr))


def mail_from_display_name() -> str:
    """Human-readable sender name for logging."""
    name, _ = parseaddr(mail_from_header())
    return name or "(no display name)"


def log_smtp_config() -> None:
    """Print active SMTP settings (never logs password)."""
    host = _smtp_get("SCAN_REPORT_SMTP_HOST", "SMTP_HOST") or "(unset)"
    port = _smtp_get("SCAN_REPORT_SMTP_PORT", "SMTP_PORT", "465") or "465"
    user = _smtp_get("SCAN_REPORT_SMTP_USER", "SMTP_USER") or "(unset)"
    addr = _smtp_get("SCAN_REPORT_SMTP_FROM", "SMTP_FROM") or "(unset)"
    display = mail_from_display_name()
    src_key = "SCAN_REPORT_SMTP_HOST" if os.environ.get("SCAN_REPORT_SMTP_HOST") else "SMTP_HOST"
    src = _loaded_from.get(src_key) or _loaded_from.get("SMTP_HOST") or "process env"
    print(
        f"[env] SMTP host={host} port={port} user={user} "
        f"from=\"{display}\" <{addr}> (source: {src})",
        flush=True,
    )
