"""
Pipeline alert reporting — persist errors locally and push to Ops / webhook.

Every scrape/push failure is:
  1. Logged at ERROR level
  2. Appended to data/pipeline_alerts.jsonl (local audit trail)
  3. POSTed to Ops /ingest/v1/alerts (console + email from Ops)
  4. Optionally POSTed to PIPELINE_ALERT_WEBHOOK_URL (direct webhook)
"""
from __future__ import annotations

import json
import logging
import os
import socket
import threading
import traceback
from contextlib import contextmanager
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Iterator

import requests

import ops_config as cfg

LOG = logging.getLogger("baskit.pipeline.alerts")

_HOST = os.environ.get("HOSTNAME") or socket.gethostname()

# Live parallel jobs for the Ops heartbeat message (host has one heartbeat row).
_jobs_lock = threading.Lock()
_active_jobs: dict[str, dict[str, Any]] = {}


def _alert_log_path() -> Path:
    cfg.base_config.ensure_dirs()
    return cfg.base_config.DATA_DIR / "pipeline_alerts.jsonl"


def _now_iso() -> str:
    return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")


def report(
    message: str,
    *,
    severity: str = "error",
    stage: str | None = None,
    retailer: str | None = None,
    detail: str | None = None,
    run_id: int | None = None,
    exc: BaseException | None = None,
) -> None:
    """Record and dispatch a pipeline alert (never raises)."""
    if exc is not None and not detail:
        detail = "".join(traceback.format_exception(type(exc), exc, exc.__traceback__))

    payload: dict[str, Any] = {
        "source": "worker",
        "severity": severity,
        "stage": stage,
        "retailer": retailer,
        "message": message[:512],
        "detail": (detail or "")[:8000] or None,
        "run_id": run_id,
        "host": _HOST,
        "ts": _now_iso(),
    }

    LOG.error(
        "[%s] %s retailer=%s stage=%s: %s",
        severity, retailer or "—", retailer or "—", stage or "—", message,
    )
    _append_local(payload)
    _post_ops(payload)
    _post_webhook(payload)


def report_success(retailer: str, stage: str, message: str, **extra: Any) -> None:
    """Lightweight info log for completed retailer cycles (no outbound alert)."""
    LOG.info("[%s] %s retailer=%s: %s", stage, retailer, retailer, message)
    payload = {
        "source": "worker",
        "severity": "info",
        "stage": stage,
        "retailer": retailer,
        "message": message[:512],
        "detail": json.dumps(extra, default=str)[:2000] if extra else None,
        "host": _HOST,
        "ts": _now_iso(),
    }
    _append_local(payload)


def _append_local(payload: dict[str, Any]) -> None:
    try:
        path = _alert_log_path()
        with path.open("a", encoding="utf-8") as fh:
            fh.write(json.dumps(payload, ensure_ascii=False) + "\n")
    except OSError as exc:
        LOG.warning("could not write local alert log: %s", exc)


def _post_ops(payload: dict[str, Any]) -> None:
    url = f"{cfg.BASE_URL}/ingest/v1/alerts"
    body = {k: v for k, v in payload.items() if k != "ts" and v is not None}
    try:
        key = cfg.SERVICE_KEY or cfg.require_service_key()
        resp = requests.post(
            url,
            headers={
                "X-Service-Key": key,
                "Content-Type": "application/json",
                "Accept": "application/json",
            },
            json=body,
            timeout=15,
            verify=cfg.VERIFY_TLS,
        )
        if resp.status_code >= 400:
            LOG.warning("Ops alert POST failed HTTP %s: %s", resp.status_code, resp.text[:200])
    except Exception as exc:
        LOG.warning("Ops alert POST failed: %s", exc)


def _post_webhook(payload: dict[str, Any]) -> None:
    url = os.environ.get("PIPELINE_ALERT_WEBHOOK_URL", "").strip()
    if not url or payload.get("severity") != "error":
        return
    body = {
        "project_name": "baskit_worker",
        "message": f"[{payload.get('retailer') or 'pipeline'}] {payload.get('stage')}: {payload.get('message')}",
        "retailer": payload.get("retailer"),
        "stage": payload.get("stage"),
        "severity": payload.get("severity"),
        "host": payload.get("host"),
    }
    try:
        resp = requests.post(url, json=body, timeout=10)
        if resp.status_code >= 400:
            LOG.warning("Webhook alert failed HTTP %s", resp.status_code)
    except Exception as exc:
        LOG.warning("Webhook alert failed: %s", exc)


def _jobs_snapshot() -> list[dict[str, Any]]:
    with _jobs_lock:
        return list(_active_jobs.values())


def _compose_status_message(fallback: str | None = None) -> str | None:
    jobs = _jobs_snapshot()
    if not jobs:
        return (fallback or "")[:255] or None
    parts = []
    for j in jobs:
        label = j.get("retailer") or "?"
        stage = j.get("stage") or "?"
        msg = (j.get("message") or "").strip()
        parts.append(f"{label}:{stage}" + (f" ({msg})" if msg else ""))
    return "; ".join(parts)[:255]


def track_job(
    retailer: str,
    stage: str,
    *,
    message: str | None = None,
    run_id: int | None = None,
) -> None:
    """Register/update an in-flight retailer job and refresh the Ops heartbeat."""
    with _jobs_lock:
        _active_jobs[retailer] = {
            "retailer": retailer,
            "stage": stage,
            "message": message,
            "run_id": run_id,
        }
    set_status(stage, retailer=retailer, message=message, run_id=run_id)


def clear_job(retailer: str) -> None:
    """Remove a finished retailer job and refresh the Ops heartbeat."""
    with _jobs_lock:
        _active_jobs.pop(retailer, None)
    jobs = _jobs_snapshot()
    if not jobs:
        set_status("idle", message=f"Finished {retailer}")
        return
    lead = jobs[0]
    set_status(
        str(lead.get("stage") or "scrape"),
        retailer=str(lead.get("retailer") or ""),
        message=lead.get("message"),
        run_id=lead.get("run_id"),
    )


def set_status(
    stage: str,
    *,
    retailer: str | None = None,
    message: str | None = None,
    detail: str | None = None,
    run_id: int | None = None,
) -> None:
    """Update the live Ops heartbeat (never raises). Shown on Scrape runs."""
    jobs = _jobs_snapshot()
    # Prefer multi-job summary so parallel rotate stays visible in the console.
    summary = _compose_status_message(message)
    lead_retailer = retailer
    lead_stage = stage
    lead_run = run_id
    if jobs:
        lead = jobs[0]
        if len(jobs) > 1:
            lead_stage = "scrape" if any(j.get("stage") == "scrape" for j in jobs) else str(lead.get("stage") or stage)
            lead_retailer = str(lead.get("retailer") or retailer or "")
            lead_run = lead.get("run_id") if lead_run is None else lead_run
        elif retailer is None:
            lead_retailer = str(lead.get("retailer") or "")
            lead_stage = str(lead.get("stage") or stage)
            if lead_run is None:
                lead_run = lead.get("run_id")

    payload: dict[str, Any] = {
        "host": _HOST,
        "stage": lead_stage,
        "retailer": lead_retailer,
        "message": summary,
        "detail": (detail or "")[:512] or None,
        "run_id": lead_run,
    }
    LOG.info("[status] stage=%s retailer=%s %s", lead_stage, lead_retailer or "—", summary or "")
    try:
        key = cfg.SERVICE_KEY or cfg.require_service_key()
        resp = requests.post(
            f"{cfg.BASE_URL}/ingest/v1/status",
            headers={
                "X-Service-Key": key,
                "Content-Type": "application/json",
                "Accept": "application/json",
            },
            json={k: v for k, v in payload.items() if v is not None},
            timeout=10,
            verify=cfg.VERIFY_TLS,
        )
        if resp.status_code >= 400:
            LOG.warning("Ops status POST failed HTTP %s: %s", resp.status_code, resp.text[:200])
    except Exception as exc:
        LOG.warning("Ops status POST failed: %s", exc)


@contextmanager
def heartbeat_keepalive(
    interval_sec: float = 300.0,
    *,
    retailer: str | None = None,
    stage: str = "scrape",
    message: str | None = None,
    run_id: int | None = None,
) -> Iterator[None]:
    """Pulse /ingest/v1/status every interval_sec during long scrape/push work."""
    stop = threading.Event()

    def _pulse() -> None:
        while not stop.wait(interval_sec):
            set_status(stage, retailer=retailer, message=message, run_id=run_id)

    t = threading.Thread(target=_pulse, name=f"hb-{retailer or stage}", daemon=True)
    t.start()
    try:
        yield
    finally:
        stop.set()
        t.join(timeout=2)


def recent_local(limit: int = 20) -> list[dict[str, Any]]:
    """Read the tail of the local JSONL alert log."""
    path = _alert_log_path()
    if not path.is_file():
        return []
    lines = path.read_text(encoding="utf-8").splitlines()
    out: list[dict[str, Any]] = []
    for line in lines[-limit:]:
        try:
            out.append(json.loads(line))
        except json.JSONDecodeError:
            continue
    return out
