#!/usr/bin/env python3
"""Generate an HTML report for modem search process from parsed mdlog CSVs."""

from __future__ import annotations

import argparse
import html
import math
import re
from pathlib import Path
from typing import Dict, Iterable, List, Tuple

import pandas as pd
import plotly.graph_objects as go
from plotly.subplots import make_subplots


FILE_NAMES = [
    "LTE_PLMN_SEARCH.csv",
    "NR5G_PLMN_SEARCH.csv",
    "LTE_BAND_SCAN.csv",
    "LTE_SYSTEM_SCAN.csv",
    "LTE_INIT_ACQ.csv",
    "NR5G_ACQ.csv",
]

SOURCE_LABELS = {
    "LTE_PLMN_SEARCH.csv": "LTE PLMN Search",
    "NR5G_PLMN_SEARCH.csv": "NR5G PLMN Search",
    "LTE_BAND_SCAN.csv": "LTE Band Scan",
    "LTE_SYSTEM_SCAN.csv": "LTE System Scan",
    "LTE_INIT_ACQ.csv": "LTE Init ACQ",
    "NR5G_ACQ.csv": "NR5G ACQ",
}

SOURCE_COLORS = {
    "LTE_PLMN_SEARCH.csv": "#1f77b4",
    "NR5G_PLMN_SEARCH.csv": "#d62728",
    "LTE_BAND_SCAN.csv": "#17becf",
    "LTE_SYSTEM_SCAN.csv": "#9467bd",
    "LTE_INIT_ACQ.csv": "#2ca02c",
    "NR5G_ACQ.csv": "#ff7f0e",
}

STATUS_COLORS = {
    "SUCCESS": "#2ca02c",
    "FAILURE": "#d62728",
    "COMPLETED": "#1f77b4",
    "ABORT CONFIRMATION": "#ff7f0e",
}


def load_csvs(base_dir: Path) -> Dict[str, pd.DataFrame]:
    frames: Dict[str, pd.DataFrame] = {}
    for name in FILE_NAMES:
        path = base_dir / name
        if not path.exists():
            frames[name] = pd.DataFrame()
            continue
        df = pd.read_csv(path)
        if "Timestamp" in df.columns:
            df["Timestamp"] = pd.to_numeric(df["Timestamp"], errors="coerce")
            df = df.dropna(subset=["Timestamp"]).copy()
            df["dt"] = pd.to_datetime(df["Timestamp"], unit="s")
        frames[name] = df
    return frames


def unique_join(values: Iterable, limit: int = 12, sep: str = ", ") -> str:
    seen: List[str] = []
    for value in values:
        if pd.isna(value):
            continue
        text = str(value).strip()
        if not text or text in seen:
            continue
        seen.append(text)
    if not seen:
        return "-"
    if len(seen) <= limit:
        return sep.join(seen)
    return sep.join(seen[:limit]) + f" ...(+{len(seen) - limit})"


def plmn_join(df: pd.DataFrame) -> str:
    if "PLMN_MCC" not in df.columns or "PLMN_MNC" not in df.columns:
        return "-"
    pairs = []
    for mcc, mnc in df[["PLMN_MCC", "PLMN_MNC"]].dropna().itertuples(index=False):
        try:
            pairs.append(f"{int(float(mcc)):03d}-{int(float(mnc)):02d}")
        except (TypeError, ValueError):
            continue
    return unique_join(pairs, limit=10)


def parse_numeric_prefix(value) -> float | None:
    if pd.isna(value):
        return None
    match = re.search(r"-?\d+(?:\.\d+)?", str(value))
    return float(match.group(0)) if match else None


def seconds_to_ms(start_ts: float, end_ts: float) -> float:
    return max((end_ts - start_ts) * 1000.0, 0.0)


def format_ts(ts: float) -> str:
    return pd.to_datetime(ts, unit="s").strftime("%Y-%m-%d %H:%M:%S.%f")[:-3]


def human_duration(ms: float) -> str:
    sec = ms / 1000.0
    if sec >= 60:
        return f"{sec / 60:.2f} min"
    return f"{sec:.3f} s"


def find_clusters(frames: Dict[str, pd.DataFrame], gap_seconds: float) -> List[dict]:
    events: List[Tuple[float, str]] = []
    for name, df in frames.items():
        if df.empty:
            continue
        events.extend((float(ts), name) for ts in df["Timestamp"].dropna().tolist())
    events.sort(key=lambda item: item[0])
    if not events:
        return []

    clusters: List[dict] = []
    cluster_id = 1
    start_ts = prev_ts = events[0][0]
    sources = [events[0][1]]
    for ts, source in events[1:]:
        if ts - prev_ts > gap_seconds:
            clusters.append(
                {
                    "cluster_id": cluster_id,
                    "start_ts": start_ts,
                    "end_ts": prev_ts,
                    "sources": sorted(set(sources)),
                }
            )
            cluster_id += 1
            start_ts = ts
            sources = []
        sources.append(source)
        prev_ts = ts

    clusters.append(
        {
            "cluster_id": cluster_id,
            "start_ts": start_ts,
            "end_ts": prev_ts,
            "sources": sorted(set(sources)),
        }
    )
    return clusters


def assign_cluster(ts: float, clusters: List[dict]) -> int | None:
    for cluster in clusters:
        if cluster["start_ts"] <= ts <= cluster["end_ts"]:
            return cluster["cluster_id"]
    return None


def filter_cluster(df: pd.DataFrame, cluster: dict) -> pd.DataFrame:
    if df.empty:
        return df
    return df[(df["Timestamp"] >= cluster["start_ts"]) & (df["Timestamp"] <= cluster["end_ts"])].copy()


def extract_plmn_stages(df: pd.DataFrame, tech: str, clusters: List[dict]) -> List[dict]:
    stages: List[dict] = []
    if df.empty or "Search_Group_Id" not in df.columns:
        return stages

    for group_id, group in df.groupby("Search_Group_Id", sort=False):
        request = group[group["Message_Type"].astype(str).str.contains("Request", case=False, na=False)]
        response = group[group["Message_Type"].astype(str).str.contains("Response", case=False, na=False)]

        start_candidates = []
        end_candidates = []
        for column in ["Search_Start_Timestamp", "Timestamp"]:
            if column in request.columns:
                start_candidates.extend(pd.to_numeric(request[column], errors="coerce").dropna().tolist())
        for column in ["Search_End_Timestamp", "Timestamp"]:
            if column in response.columns:
                end_candidates.extend(pd.to_numeric(response[column], errors="coerce").dropna().tolist())
        if not start_candidates:
            start_candidates = pd.to_numeric(group["Timestamp"], errors="coerce").dropna().tolist()
        if not end_candidates:
            end_candidates = pd.to_numeric(group["Timestamp"], errors="coerce").dropna().tolist()

        start_ts = float(min(start_candidates))
        end_ts = float(max(end_candidates))
        cluster_id = assign_cluster(start_ts, clusters)

        status = "-"
        for column in ["Search_Status", "Network_Search_Status"]:
            if column in group.columns:
                status = unique_join(group[column].dropna().astype(str).str.upper().tolist())
                if status != "-":
                    break

        requested_rat = []
        for column in ["RAT_List_RAT", "Current_Search_RAT", "Response_RAT", "Source_RAT"]:
            if column in group.columns:
                requested_rat.extend(group[column].dropna().astype(str).tolist())

        result_plmn = plmn_join(group)
        target_plmn = "-"
        if tech == "LTE" and "MCC_List" in request.columns:
            target_plmn = unique_join(request["MCC_List"].dropna().tolist())
        elif result_plmn != "-":
            target_plmn = result_plmn

        bands = "-"
        if "Band" in group.columns:
            bands = unique_join(sorted(group["Band"].dropna().astype(int).astype(str).tolist()))
        elif "Candidate_Band" in group.columns:
            bands = unique_join(sorted(group["Candidate_Band"].dropna().astype(int).astype(str).tolist()))

        freqs = "-"
        if "ARFCN" in group.columns:
            freqs = unique_join(sorted(group["ARFCN"].dropna().astype(int).astype(str).tolist()))

        stages.append(
            {
                "cluster_id": cluster_id,
                "stage_kind": "PLMN_SEARCH",
                "source": f"{tech}_PLMN",
                "group_id": str(group_id),
                "start_ts": start_ts,
                "end_ts": end_ts,
                "duration_ms": seconds_to_ms(start_ts, end_ts),
                "rat": unique_join(requested_rat),
                "target_plmn": target_plmn,
                "result_plmn": result_plmn,
                "bands": bands,
                "freqs": freqs,
                "cells": "-",
                "status": status,
                "note": unique_join(
                    [group["Search_Type"].iloc[0] if "Search_Type" in group.columns and not group.empty else None]
                ),
            }
        )

    return stages


def summarize_nr_acq(df: pd.DataFrame, clusters: List[dict]) -> Tuple[List[dict], pd.DataFrame]:
    if df.empty:
        return [], pd.DataFrame()

    group_cols = ["Timestamp", "TransId", "Raster_Band", "Raster_ARFCN", "Raster_Status"]
    existing_cols = [col for col in group_cols if col in df.columns]
    rows: List[dict] = []
    for keys, group in df.groupby(existing_cols, sort=False, dropna=False):
        if not isinstance(keys, tuple):
            keys = (keys,)
        row = dict(zip(existing_cols, keys))
        ts = float(group["Timestamp"].iloc[0])
        row["cluster_id"] = assign_cluster(ts, clusters)
        row["cell_count"] = int(pd.to_numeric(group.get("Raster_NumCell"), errors="coerce").fillna(0).max())
        row["cells"] = unique_join(
            [
                f"PCI {int(pci)} ({rsrp:.2f} dBm)"
                for pci, rsrp in group[["Cell_PCI", "Cell_RSRP"]].dropna().itertuples(index=False)
            ],
            limit=20,
        )
        row["best_rsrp"] = pd.to_numeric(group.get("Cell_RSRP"), errors="coerce").max()
        rows.append(row)

    detail_df = pd.DataFrame(rows)
    detail_df["dt"] = pd.to_datetime(detail_df["Timestamp"], unit="s")

    stages = []
    for cluster_id, group in detail_df.groupby("cluster_id", dropna=True):
        start_ts = float(group["Timestamp"].min())
        end_ts = float(group["Timestamp"].max())
        success_group = group[group["Raster_Status"].astype(str).str.upper() == "E_SUCCESS"]
        stages.append(
            {
                "cluster_id": int(cluster_id),
                "stage_kind": "ACQ_WINDOW",
                "source": "NR5G_ACQ",
                "group_id": f"NR_ACQ_C{int(cluster_id)}",
                "start_ts": start_ts,
                "end_ts": end_ts,
                "duration_ms": seconds_to_ms(start_ts, end_ts),
                "rat": "NR5G",
                "target_plmn": "-",
                "result_plmn": "-",
                "bands": unique_join(group["Raster_Band"].dropna().tolist()),
                "freqs": unique_join(group["Raster_ARFCN"].dropna().astype(int).astype(str).tolist()),
                "cells": unique_join(success_group["cells"].tolist(), limit=8, sep="<br>"),
                "status": unique_join(group["Raster_Status"].dropna().astype(str).str.upper().tolist()),
                "note": f"{len(group)} raster attempts",
            }
        )
    return stages, detail_df


def summarize_lte_init_acq(df: pd.DataFrame, clusters: List[dict]) -> Tuple[List[dict], pd.DataFrame]:
    if df.empty:
        return [], pd.DataFrame()

    detail_df = df.copy()
    detail_df["cluster_id"] = detail_df["Timestamp"].apply(lambda ts: assign_cluster(float(ts), clusters))
    detail_df["Result"] = detail_df["Result"].astype(str).str.upper()
    stages = []
    for cluster_id, group in detail_df.groupby("cluster_id", dropna=True):
        start_ts = float(group["Timestamp"].min())
        end_ts = float(group["Timestamp"].max())
        success_group = group[group["Result"] == "SUCCESS"]
        stages.append(
            {
                "cluster_id": int(cluster_id),
                "stage_kind": "ACQ_WINDOW",
                "source": "LTE_INIT_ACQ",
                "group_id": f"LTE_INIT_C{int(cluster_id)}",
                "start_ts": start_ts,
                "end_ts": end_ts,
                "duration_ms": seconds_to_ms(start_ts, end_ts),
                "rat": "LTE",
                "target_plmn": "-",
                "result_plmn": "-",
                "bands": unique_join(group["Band"].dropna().astype(int).astype(str).tolist()),
                "freqs": unique_join(group["EARFCN"].dropna().astype(int).astype(str).tolist()),
                "cells": unique_join(
                    [
                        f"EARFCN {int(row.EARFCN)} search_results={int(row.Num_Search_Results)}"
                        for row in success_group.itertuples(index=False)
                    ],
                    limit=8,
                ),
                "status": unique_join(group["Result"].dropna().tolist()),
                "note": f"{len(group)} init-acq attempts",
            }
        )
    return stages, detail_df


def summarize_window_source(
    df: pd.DataFrame,
    clusters: List[dict],
    source_name: str,
    rat: str,
    band_col: str,
    freq_col: str,
) -> List[dict]:
    if df.empty:
        return []

    temp = df.copy()
    temp["cluster_id"] = temp["Timestamp"].apply(lambda ts: assign_cluster(float(ts), clusters))
    rows: List[dict] = []
    for cluster_id, group in temp.groupby("cluster_id", dropna=True):
        rows.append(
            {
                "cluster_id": int(cluster_id),
                "stage_kind": "SCAN_WINDOW",
                "source": source_name,
                "group_id": f"{source_name}_C{int(cluster_id)}",
                "start_ts": float(group["Timestamp"].min()),
                "end_ts": float(group["Timestamp"].max()),
                "duration_ms": seconds_to_ms(float(group["Timestamp"].min()), float(group["Timestamp"].max())),
                "rat": rat,
                "target_plmn": "-",
                "result_plmn": "-",
                "bands": unique_join(group[band_col].dropna().astype(str).tolist()) if band_col in group.columns else "-",
                "freqs": unique_join(group[freq_col].dropna().astype(str).tolist()) if freq_col in group.columns else "-",
                "cells": "-",
                "status": "-",
                "note": f"{len(group)} rows",
            }
        )
    return rows


def attach_overlap_details(stages: List[dict], frames: Dict[str, pd.DataFrame]) -> None:
    band_scan = frames["LTE_BAND_SCAN.csv"]
    system_scan = frames["LTE_SYSTEM_SCAN.csv"]
    nr_acq = frames["NR5G_ACQ.csv"]

    for stage in stages:
        start_ts, end_ts = stage["start_ts"], stage["end_ts"]
        in_lte_band = filter_between(band_scan, start_ts, end_ts)
        in_lte_system = filter_between(system_scan, start_ts, end_ts)
        in_nr_acq = filter_between(nr_acq, start_ts, end_ts)

        band_parts = []
        freq_parts = []
        if stage["bands"] != "-":
            band_parts.append(stage["bands"])
        if not in_lte_band.empty:
            band_parts.append("LTE_BAND_SCAN:" + unique_join(in_lte_band["Candidate_Band"].dropna().astype(int).astype(str).tolist()))
            freq_parts.append("LTE_BAND_SCAN:" + unique_join(in_lte_band["EARFCN"].dropna().astype(int).astype(str).tolist()))
        if not in_lte_system.empty:
            band_parts.append("LTE_SYSTEM_SCAN:" + unique_join(in_lte_system["Band"].dropna().astype(int).astype(str).tolist()))
            freq_parts.append("LTE_SYSTEM_SCAN:" + unique_join(in_lte_system["EARFCN"].dropna().astype(int).astype(str).tolist()))
        if "ARFCN" in stage and stage["freqs"] != "-":
            freq_parts.append(stage["freqs"])

        if not in_nr_acq.empty:
            cell_rows = []
            for arfcn, group in in_nr_acq.groupby("Raster_ARFCN", sort=False, dropna=False):
                cells = unique_join(
                    [
                        f"PCI {int(pci)} ({float(rsrp):.2f} dBm)"
                        for pci, rsrp in group[["Cell_PCI", "Cell_RSRP"]].dropna().itertuples(index=False)
                    ],
                    limit=10,
                )
                arfcn_text = str(int(float(arfcn))) if not pd.isna(arfcn) else "NA"
                cell_rows.append(f"ARFCN {arfcn_text}: {cells}")
            stage["cells"] = unique_join(cell_rows, limit=6, sep="<br>")

        stage["bands"] = unique_join(band_parts, limit=8, sep="<br>")
        stage["freqs"] = unique_join(freq_parts, limit=8, sep="<br>")


def filter_between(df: pd.DataFrame, start_ts: float, end_ts: float) -> pd.DataFrame:
    if df.empty:
        return df
    return df[(df["Timestamp"] >= start_ts) & (df["Timestamp"] <= end_ts)].copy()


def make_overview_figure(frames: Dict[str, pd.DataFrame], clusters: List[dict]) -> str:
    fig = go.Figure()
    y_order = list(reversed(FILE_NAMES))
    y_map = {name: idx for idx, name in enumerate(y_order)}

    for name in y_order:
        df = frames[name]
        if df.empty:
            continue
        fig.add_trace(
            go.Scatter(
                x=df["dt"],
                y=[y_map[name]] * len(df),
                mode="markers",
                marker=dict(size=5, color=SOURCE_COLORS[name], opacity=0.7),
                name=SOURCE_LABELS[name],
                hovertemplate=(
                    f"{SOURCE_LABELS[name]}<br>"
                    "Time=%{x|%Y-%m-%d %H:%M:%S.%L}<br>"
                    "<extra></extra>"
                ),
            )
        )

    for cluster in clusters:
        fig.add_vrect(
            x0=pd.to_datetime(cluster["start_ts"], unit="s"),
            x1=pd.to_datetime(cluster["end_ts"], unit="s"),
            fillcolor="rgba(180, 180, 180, 0.08)",
            line_width=0,
        )

    fig.update_layout(
        template="plotly_white",
        height=420,
        margin=dict(l=110, r=220, t=30, b=50),
        legend=dict(orientation="v", x=1.02, y=1.0, xanchor="left", yanchor="top"),
        hovermode="closest",
    )
    fig.update_yaxes(
        tickmode="array",
        tickvals=list(y_map.values()),
        ticktext=[SOURCE_LABELS[name] for name in y_order],
        title_text="Packet Source",
        range=[-0.8, len(y_order) - 0.2],
    )
    fig.update_xaxes(title_text="Absolute Time")
    return fig.to_html(full_html=False, include_plotlyjs="inline")


def add_timeline_trace(fig: go.Figure, row: int, stage: dict, y_value: int, showlegend: bool) -> None:
    color = STATUS_COLORS.get(stage["status"].split(",")[0].upper(), SOURCE_COLORS.get(stage_to_filename(stage["source"]), "#444"))
    hover = (
        f"Stage={html.escape(stage['group_id'])}<br>"
        f"Source={html.escape(stage['source'])}<br>"
        f"Start={html.escape(format_ts(stage['start_ts']))}<br>"
        f"End={html.escape(format_ts(stage['end_ts']))}<br>"
        f"Duration={html.escape(human_duration(stage['duration_ms']))}<br>"
        f"RAT={stage['rat']}<br>"
        f"Target PLMN={stage['target_plmn']}<br>"
        f"Result PLMN={stage['result_plmn']}<br>"
        f"Bands={stage['bands']}<br>"
        f"Freqs={stage['freqs']}<br>"
        f"Cells={stage['cells']}<br>"
        f"Status={stage['status']}<br>"
        f"Note={stage['note']}<extra></extra>"
    )
    fig.add_trace(
        go.Scatter(
            x=[pd.to_datetime(stage["start_ts"], unit="s"), pd.to_datetime(stage["end_ts"], unit="s")],
            y=[y_value, y_value],
            mode="lines+markers",
            line=dict(color=color, width=16),
            marker=dict(size=8, color=color),
            name=stage["source"],
            legendgroup=stage["source"],
            showlegend=showlegend,
            hovertemplate=hover,
        ),
        row=row,
        col=1,
    )


def stage_to_filename(source: str) -> str:
    mapping = {
        "LTE_PLMN": "LTE_PLMN_SEARCH.csv",
        "NR5G_PLMN": "NR5G_PLMN_SEARCH.csv",
        "LTE_BAND_SCAN": "LTE_BAND_SCAN.csv",
        "LTE_SYSTEM_SCAN": "LTE_SYSTEM_SCAN.csv",
        "LTE_INIT_ACQ": "LTE_INIT_ACQ.csv",
        "NR5G_ACQ": "NR5G_ACQ.csv",
    }
    return mapping.get(source, "")


def make_cluster_figure(
    cluster: dict,
    frames: Dict[str, pd.DataFrame],
    cluster_stages: List[dict],
    nr_acq_detail: pd.DataFrame,
    lte_init_detail: pd.DataFrame,
) -> str:
    fig = make_subplots(
        rows=5,
        cols=1,
        shared_xaxes=True,
        vertical_spacing=0.035,
        row_heights=[0.18, 0.14, 0.23, 0.22, 0.23],
    )

    shown_sources = set()
    stage_order = []
    for stage in cluster_stages:
        if stage["source"] not in stage_order:
            stage_order.append(stage["source"])
    if not stage_order:
        stage_order = ["NO_STAGE"]
    y_map = {name: len(stage_order) - idx for idx, name in enumerate(stage_order)}
    for stage in cluster_stages:
        add_timeline_trace(
            fig,
            row=1,
            stage=stage,
            y_value=y_map[stage["source"]],
            showlegend=stage["source"] not in shown_sources,
        )
        shown_sources.add(stage["source"])

    lte_plmn = filter_cluster(frames["LTE_PLMN_SEARCH.csv"], cluster)
    nr_plmn = filter_cluster(frames["NR5G_PLMN_SEARCH.csv"], cluster)
    req_points = []
    if not lte_plmn.empty and "MCC_List" in lte_plmn.columns:
        requests = lte_plmn[lte_plmn["Message_Type"].astype(str).str.contains("Request", case=False, na=False)].copy()
        requests = requests.dropna(subset=["MCC_List"])
        for row in requests.itertuples(index=False):
            req_points.append(
                {
                    "dt": row.dt,
                    "label": f"REQ MCC {str(row.MCC_List)}",
                    "source": "LTE target MCC",
                }
            )
    req_df = pd.DataFrame(req_points)
    if not req_df.empty:
        fig.add_trace(
            go.Scatter(
                x=req_df["dt"],
                y=req_df["label"],
                mode="markers",
                marker=dict(symbol="triangle-right", size=10, color="#1f77b4"),
                name="LTE target MCC",
                legendgroup="plmn_req",
                hovertemplate="LTE target MCC<br>%{y}<br>%{x|%H:%M:%S.%L}<extra></extra>",
            ),
            row=2,
            col=1,
        )

    for df, source_name, symbol, color in [
        (lte_plmn, "LTE result PLMN", "circle", "#1f77b4"),
        (nr_plmn, "NR result PLMN", "diamond", "#d62728"),
    ]:
        if df.empty or "PLMN_MCC" not in df.columns or "PLMN_MNC" not in df.columns:
            continue
        result_df = df.dropna(subset=["PLMN_MCC", "PLMN_MNC"]).copy()
        if result_df.empty:
            continue
        result_df["plmn"] = result_df.apply(
            lambda row: f"{int(float(row['PLMN_MCC'])):03d}-{int(float(row['PLMN_MNC'])):02d}",
            axis=1,
        )
        fig.add_trace(
            go.Scatter(
                x=result_df["dt"],
                y=result_df["plmn"],
                mode="markers",
                marker=dict(symbol=symbol, size=9, color=color, line=dict(color="white", width=0.5)),
                name=source_name,
                legendgroup=source_name,
                hovertemplate=(
                    f"{source_name}<br>"
                    "PLMN=%{y}<br>"
                    "Time=%{x|%H:%M:%S.%L}<br>"
                    "<extra></extra>"
                ),
            ),
            row=2,
            col=1,
        )

    lte_band_scan = filter_cluster(frames["LTE_BAND_SCAN.csv"], cluster).copy()
    if not lte_band_scan.empty:
        lte_band_scan["EARFCN"] = pd.to_numeric(lte_band_scan["EARFCN"], errors="coerce")
        fig.add_trace(
            go.Scatter(
                x=lte_band_scan["dt"],
                y=lte_band_scan["EARFCN"],
                mode="markers",
                marker=dict(symbol="circle", size=5, color="#17becf", opacity=0.6),
                name="LTE Band Scan",
                legendgroup="LTE Band Scan",
                hovertemplate=(
                    "LTE Band Scan<br>"
                    "Time=%{x|%H:%M:%S.%L}<br>"
                    "EARFCN=%{y}<br>"
                    "Band=%{customdata[0]}<br>"
                    "Packet_Band=%{customdata[1]}<br>"
                    "Bandwidth=%{customdata[2]}<br>"
                    "WB_Energy=%{customdata[3]}<extra></extra>"
                ),
                customdata=lte_band_scan[["Candidate_Band", "Packet_Band", "Bandwidth", "WB_Energy"]],
            ),
            row=3,
            col=1,
        )

    lte_system_scan = filter_cluster(frames["LTE_SYSTEM_SCAN.csv"], cluster).copy()
    if not lte_system_scan.empty:
        lte_system_scan["EARFCN"] = pd.to_numeric(lte_system_scan["EARFCN"], errors="coerce")
        fig.add_trace(
            go.Scatter(
                x=lte_system_scan["dt"],
                y=lte_system_scan["EARFCN"],
                mode="markers",
                marker=dict(symbol="x", size=8, color="#9467bd", opacity=0.9),
                name="LTE System Scan",
                legendgroup="LTE System Scan",
                hovertemplate=(
                    "LTE System Scan<br>"
                    "Time=%{x|%H:%M:%S.%L}<br>"
                    "EARFCN=%{y}<br>"
                    "Band=%{customdata[0]}<br>"
                    "Candidates=%{customdata[1]}<br>"
                    "Energy=%{customdata[2]}<br>"
                    "NB_Energy=%{customdata[3]}<extra></extra>"
                ),
                customdata=lte_system_scan[["Band", "Num_Candidates", "Energy", "NB_Energy"]],
            ),
            row=3,
            col=1,
        )

    lte_init = lte_init_detail[lte_init_detail["cluster_id"] == cluster["cluster_id"]].copy()
    if not lte_init.empty:
        lte_init["EARFCN"] = pd.to_numeric(lte_init["EARFCN"], errors="coerce")
        colors = lte_init["Result"].map(lambda value: STATUS_COLORS.get(str(value).upper(), "#2ca02c"))
        fig.add_trace(
            go.Scatter(
                x=lte_init["dt"],
                y=lte_init["EARFCN"],
                mode="markers",
                marker=dict(symbol="diamond-open", size=10, color=colors, line=dict(width=1.6)),
                name="LTE Init ACQ",
                legendgroup="LTE Init ACQ",
                hovertemplate=(
                    "LTE Init ACQ<br>"
                    "Time=%{x|%H:%M:%S.%L}<br>"
                    "EARFCN=%{y}<br>"
                    "Band=%{customdata[0]}<br>"
                    "Result=%{customdata[1]}<br>"
                    "Search_Results=%{customdata[2]}<br>"
                    "PBCH_Attempt_Cells=%{customdata[3]}<extra></extra>"
                ),
                customdata=lte_init[["Band", "Result", "Num_Search_Results", "Num_PBCH_Decode_Attempt_Cells"]],
            ),
            row=3,
            col=1,
        )

    nr_plmn_freq = nr_plmn.dropna(subset=["ARFCN"]).copy()
    if not nr_plmn_freq.empty:
        nr_plmn_freq["ARFCN"] = pd.to_numeric(nr_plmn_freq["ARFCN"], errors="coerce")
        fig.add_trace(
            go.Scatter(
                x=nr_plmn_freq["dt"],
                y=nr_plmn_freq["ARFCN"],
                mode="markers",
                marker=dict(symbol="diamond", size=9, color="#d62728"),
                name="NR PLMN Result Freq",
                legendgroup="NR PLMN Result Freq",
                hovertemplate=(
                    "NR PLMN Result<br>"
                    "Time=%{x|%H:%M:%S.%L}<br>"
                    "ARFCN=%{y}<br>"
                    "Band=%{customdata[0]}<br>"
                    "PLMN=%{customdata[1]}<extra></extra>"
                ),
                customdata=nr_plmn_freq.apply(
                    lambda row: [row["Band"], f"{int(float(row['PLMN_MCC'])):03d}-{int(float(row['PLMN_MNC'])):02d}"],
                    axis=1,
                    result_type="expand",
                ),
            ),
            row=4,
            col=1,
        )

    nr_acq = nr_acq_detail[nr_acq_detail["cluster_id"] == cluster["cluster_id"]].copy()
    if not nr_acq.empty:
        nr_acq["Raster_ARFCN"] = pd.to_numeric(nr_acq["Raster_ARFCN"], errors="coerce")
        nr_colors = nr_acq["Raster_Status"].astype(str).str.upper().map(lambda value: STATUS_COLORS.get(value, "#ff7f0e"))
        fig.add_trace(
            go.Scatter(
                x=nr_acq["dt"],
                y=nr_acq["Raster_ARFCN"],
                mode="markers",
                marker=dict(symbol="triangle-up", size=10, color=nr_colors, line=dict(width=0.8, color="white")),
                name="NR5G ACQ",
                legendgroup="NR5G ACQ",
                hovertemplate=(
                    "NR5G ACQ<br>"
                    "Time=%{x|%H:%M:%S.%L}<br>"
                    "ARFCN=%{y}<br>"
                    "Band=%{customdata[0]}<br>"
                    "Status=%{customdata[1]}<br>"
                    "NumCell=%{customdata[2]}<br>"
                    "Cells=%{customdata[3]}<extra></extra>"
                ),
                customdata=nr_acq[["Raster_Band", "Raster_Status", "cell_count", "cells"]],
            ),
            row=4,
            col=1,
        )

        cell_df = filter_cluster(frames["NR5G_ACQ.csv"], cluster).dropna(subset=["Cell_PCI", "Cell_RSRP"]).copy()
        if not cell_df.empty:
            cell_df["Cell_PCI"] = pd.to_numeric(cell_df["Cell_PCI"], errors="coerce")
            cell_df["Cell_RSRP"] = pd.to_numeric(cell_df["Cell_RSRP"], errors="coerce")
            fig.add_trace(
                go.Scatter(
                    x=cell_df["dt"],
                    y=cell_df["Cell_PCI"],
                    mode="markers",
                    marker=dict(
                        size=9,
                        color=cell_df["Cell_RSRP"],
                        colorscale="RdYlGn",
                        reversescale=True,
                        colorbar=dict(title="RSRP(dBm)", x=1.12, len=0.18),
                        line=dict(width=0.6, color="white"),
                    ),
                    name="NR Cell PCI",
                    legendgroup="NR Cell PCI",
                    hovertemplate=(
                        "NR Cell Found<br>"
                        "Time=%{x|%H:%M:%S.%L}<br>"
                        "PCI=%{y}<br>"
                        "ARFCN=%{customdata[0]}<br>"
                        "Band=%{customdata[1]}<br>"
                        "RSRP=%{customdata[2]:.2f} dBm<br>"
                        "Barred=%{customdata[3]}<extra></extra>"
                    ),
                    customdata=cell_df[["Raster_ARFCN", "Raster_Band", "Cell_RSRP", "MIB_CellBarred"]],
                ),
                row=5,
                col=1,
            )

    fig.update_layout(
        template="plotly_white",
        height=1540,
        margin=dict(l=120, r=280, t=30, b=60),
        legend=dict(orientation="v", x=1.02, y=1.0, xanchor="left", yanchor="top"),
        hovermode="closest",
    )
    fig.update_xaxes(range=[pd.to_datetime(cluster["start_ts"], unit="s"), pd.to_datetime(cluster["end_ts"], unit="s")])
    fig.update_yaxes(
        row=1,
        col=1,
        tickmode="array",
        tickvals=list(y_map.values()),
        ticktext=list(y_map.keys()),
        title_text="Stage",
        automargin=True,
    )
    fig.update_yaxes(row=2, col=1, title_text="PLMN", automargin=True)
    fig.update_yaxes(row=3, col=1, title_text="LTE EARFCN", automargin=True)
    fig.update_yaxes(row=4, col=1, title_text="NR ARFCN", automargin=True)
    fig.update_yaxes(row=5, col=1, title_text="Cell PCI", automargin=True)
    fig.update_xaxes(row=5, col=1, title_text="Absolute Time")
    return fig.to_html(full_html=False, include_plotlyjs=False)


def dataframe_to_html(df: pd.DataFrame, columns: List[str]) -> str:
    if df.empty:
        return "<p class='empty'>No data in this epoch.</p>"
    safe_df = df.loc[:, [col for col in columns if col in df.columns]].copy()
    for column in safe_df.columns:
        safe_df[column] = safe_df[column].fillna("-")
    return safe_df.to_html(index=False, escape=False, classes="detail-table")


def make_stage_table(stages: List[dict], cluster_id: int) -> str:
    cluster_rows = [stage for stage in stages if stage["cluster_id"] == cluster_id]
    if not cluster_rows:
        return "<p class='empty'>No stage summary.</p>"

    table_df = pd.DataFrame(cluster_rows)
    table_df["Start"] = table_df["start_ts"].apply(format_ts)
    table_df["End"] = table_df["end_ts"].apply(format_ts)
    table_df["Duration"] = table_df["duration_ms"].apply(human_duration)
    table_df = table_df.rename(
        columns={
            "source": "Source",
            "group_id": "Stage_ID",
            "rat": "RAT",
            "target_plmn": "Target_PLMN",
            "result_plmn": "Result_PLMN",
            "bands": "Bands",
            "freqs": "Freqs",
            "cells": "Cells",
            "status": "Status",
            "note": "Note",
        }
    )
    return dataframe_to_html(
        table_df,
        ["Source", "Stage_ID", "Start", "End", "Duration", "RAT", "Target_PLMN", "Result_PLMN", "Bands", "Freqs", "Cells", "Status", "Note"],
    )


def make_cluster_tables(cluster: dict, frames: Dict[str, pd.DataFrame], nr_acq_detail: pd.DataFrame, lte_init_detail: pd.DataFrame) -> str:
    cid = cluster["cluster_id"]
    parts = []

    band_scan = filter_cluster(frames["LTE_BAND_SCAN.csv"], cluster)
    if not band_scan.empty:
        grouped = (
            band_scan.groupby(["Timestamp", "Packet_Band"], sort=False)
            .agg(
                Num_Candidates=("Num_Candidates", "max"),
                Bands=("Candidate_Band", lambda s: unique_join(s.astype(int).astype(str).tolist())),
                EARFCNs=("EARFCN", lambda s: unique_join(s.astype(int).astype(str).tolist(), limit=16)),
                Bandwidth=("Bandwidth", lambda s: unique_join(s.tolist(), limit=4)),
                WB_Energy=("WB_Energy", lambda s: unique_join(s.tolist(), limit=4)),
            )
            .reset_index()
        )
        grouped["Time"] = grouped["Timestamp"].apply(format_ts)
        parts.append(
            "<details><summary>LTE Band Scan batches</summary>"
            + dataframe_to_html(grouped, ["Time", "Packet_Band", "Num_Candidates", "Bands", "EARFCNs", "Bandwidth", "WB_Energy"])
            + "</details>"
        )

    system_scan = filter_cluster(frames["LTE_SYSTEM_SCAN.csv"], cluster)
    if not system_scan.empty:
        grouped = (
            system_scan.groupby("Timestamp", sort=False)
            .agg(
                Num_Candidates=("Num_Candidates", "max"),
                Bands=("Band", lambda s: unique_join(s.astype(int).astype(str).tolist())),
                EARFCNs=("EARFCN", lambda s: unique_join(s.astype(int).astype(str).tolist(), limit=16)),
                Energy=("Energy", lambda s: unique_join(s.tolist(), limit=8)),
            )
            .reset_index()
        )
        grouped["Time"] = grouped["Timestamp"].apply(format_ts)
        parts.append(
            "<details><summary>LTE System Scan batches</summary>"
            + dataframe_to_html(grouped, ["Time", "Num_Candidates", "Bands", "EARFCNs", "Energy"])
            + "</details>"
        )

    cluster_lte_init = lte_init_detail[lte_init_detail["cluster_id"] == cid].copy()
    if not cluster_lte_init.empty:
        cluster_lte_init["Time"] = cluster_lte_init["Timestamp"].apply(format_ts)
        parts.append(
            "<details><summary>LTE Init ACQ results</summary>"
            + dataframe_to_html(
                cluster_lte_init,
                ["Time", "EARFCN", "Band", "Result", "Num_Search_Results", "Num_PBCH_Decode_Attempt_Cells", "Num_Blocked_Cells"],
            )
            + "</details>"
        )

    cluster_nr_acq = nr_acq_detail[nr_acq_detail["cluster_id"] == cid].copy()
    if not cluster_nr_acq.empty:
        cluster_nr_acq["Time"] = cluster_nr_acq["Timestamp"].apply(format_ts)
        parts.append(
            "<details><summary>NR5G ACQ raster and cell findings</summary>"
            + dataframe_to_html(
                cluster_nr_acq,
                ["Time", "Raster_Band", "Raster_ARFCN", "Raster_Status", "cell_count", "cells", "best_rsrp"],
            )
            + "</details>"
        )

    if not parts:
        return "<p class='empty'>No detail tables in this epoch.</p>"
    return "\n".join(parts)


def build_cluster_cards(clusters: List[dict]) -> str:
    cards = []
    for cluster in clusters:
        gap_note = unique_join([SOURCE_LABELS[name] for name in cluster["sources"]], limit=6)
        cards.append(
            f"""
            <div class="card">
              <div class="card-title">Epoch {cluster['cluster_id']}</div>
              <div class="card-item"><b>Start</b>: {html.escape(format_ts(cluster['start_ts']))}</div>
              <div class="card-item"><b>End</b>: {html.escape(format_ts(cluster['end_ts']))}</div>
              <div class="card-item"><b>Duration</b>: {html.escape(human_duration(seconds_to_ms(cluster['start_ts'], cluster['end_ts'])))}</div>
              <div class="card-item"><b>Sources</b>: {html.escape(gap_note)}</div>
            </div>
            """
        )
    return "\n".join(cards)


def build_html(
    base_dir: Path,
    clusters: List[dict],
    frames: Dict[str, pd.DataFrame],
    stages: List[dict],
    nr_acq_detail: pd.DataFrame,
    lte_init_detail: pd.DataFrame,
) -> str:
    overview_html = make_overview_figure(frames, clusters)
    sections = []
    for cluster in clusters:
        cluster_stages = [stage for stage in stages if stage["cluster_id"] == cluster["cluster_id"]]
        cluster_stages.sort(key=lambda item: (item["start_ts"], item["source"], item["group_id"]))
        fig_html = make_cluster_figure(cluster, frames, cluster_stages, nr_acq_detail, lte_init_detail)
        section = f"""
        <section class="epoch">
          <h2>Epoch {cluster['cluster_id']}</h2>
          <p class="epoch-note">
            This figure keeps the real time spacing inside the epoch. Labels are moved to hover and tables to avoid overlap.
          </p>
          <div class="plot-wrap">{fig_html}</div>
          <h3>Stage Summary</h3>
          {make_stage_table(stages, cluster['cluster_id'])}
          <h3>Detail Tables</h3>
          {make_cluster_tables(cluster, frames, nr_acq_detail, lte_init_detail)}
        </section>
        """
        sections.append(section)

    start_ts = min(cluster["start_ts"] for cluster in clusters) if clusters else 0
    end_ts = max(cluster["end_ts"] for cluster in clusters) if clusters else 0
    return f"""
<!DOCTYPE html>
<html lang="en">
<head>
  <meta charset="utf-8">
  <title>Modem Search Process</title>
  <style>
    body {{
      margin: 0;
      font-family: Arial, Helvetica, sans-serif;
      color: #1f2937;
      background: #f8fafc;
    }}
    .container {{
      max-width: 1800px;
      margin: 0 auto;
      padding: 24px 28px 40px;
    }}
    h1, h2, h3 {{
      margin: 0 0 12px;
    }}
    p {{
      margin: 0 0 12px;
      line-height: 1.5;
    }}
    .subtle {{
      color: #475569;
    }}
    .card-grid {{
      display: grid;
      grid-template-columns: repeat(auto-fit, minmax(260px, 1fr));
      gap: 12px;
      margin: 18px 0 20px;
    }}
    .card {{
      background: white;
      border: 1px solid #dbeafe;
      border-radius: 12px;
      padding: 14px 16px;
      box-shadow: 0 1px 2px rgba(15, 23, 42, 0.06);
    }}
    .card-title {{
      font-size: 16px;
      font-weight: 700;
      margin-bottom: 8px;
    }}
    .card-item {{
      font-size: 13px;
      margin-top: 4px;
      word-break: break-word;
    }}
    .plot-wrap {{
      background: white;
      border-radius: 12px;
      padding: 10px 10px 0;
      border: 1px solid #e2e8f0;
      overflow-x: auto;
    }}
    .epoch {{
      margin-top: 28px;
      padding-top: 20px;
      border-top: 2px solid #dbeafe;
    }}
    .epoch-note {{
      color: #475569;
      margin-bottom: 10px;
    }}
    .detail-table {{
      width: 100%;
      border-collapse: collapse;
      margin: 10px 0 18px;
      background: white;
      table-layout: fixed;
    }}
    .detail-table th,
    .detail-table td {{
      border: 1px solid #e2e8f0;
      padding: 8px 10px;
      vertical-align: top;
      font-size: 12px;
      word-break: break-word;
    }}
    .detail-table th {{
      background: #eff6ff;
      position: sticky;
      top: 0;
      z-index: 1;
    }}
    details {{
      background: #ffffff;
      border: 1px solid #dbeafe;
      border-radius: 10px;
      padding: 8px 12px;
      margin-bottom: 12px;
    }}
    summary {{
      cursor: pointer;
      font-weight: 600;
      color: #0f172a;
    }}
    .empty {{
      color: #64748b;
      font-style: italic;
      margin: 8px 0 14px;
    }}
    .note-list {{
      margin: 12px 0 20px;
      padding-left: 18px;
    }}
    .note-list li {{
      margin-bottom: 6px;
      line-height: 1.45;
    }}
  </style>
</head>
<body>
  <div class="container">
    <h1>Modem Search Process Timeline</h1>
    <p class="subtle">
      Source directory: {html.escape(str(base_dir))}
    </p>
    <p class="subtle">
      Absolute range: {html.escape(format_ts(start_ts))} ~ {html.escape(format_ts(end_ts))}
    </p>
    <ul class="note-list">
      <li>Overview keeps the full absolute timeline, so the long idle gap between search epochs remains visible.</li>
      <li>Each epoch figure uses shared absolute time on the x-axis, so interval gaps between behaviors stay real.</li>
      <li>To avoid label pile-up, all dense details are moved into hover and the tables below each figure.</li>
      <li>Cell-level PCI/RSRP is only shown where the source packets actually provide it, mainly from <code>NR5G_ACQ.csv</code>.</li>
    </ul>
    <div class="card-grid">
      {build_cluster_cards(clusters)}
    </div>
    <section>
      <h2>Overview</h2>
      <div class="plot-wrap">{overview_html}</div>
    </section>
    {''.join(sections)}
  </div>
</body>
</html>
"""


def main() -> None:
    parser = argparse.ArgumentParser(description="Generate modem search process HTML report.")
    parser.add_argument("--input-dir", required=True, help="Directory containing parsed CSV files.")
    parser.add_argument("--output-html", required=True, help="Output HTML path.")
    parser.add_argument("--cluster-gap-seconds", type=float, default=120.0, help="Gap threshold for epoch split.")
    args = parser.parse_args()

    input_dir = Path(args.input_dir).expanduser().resolve()
    output_html = Path(args.output_html).expanduser().resolve()
    output_html.parent.mkdir(parents=True, exist_ok=True)

    frames = load_csvs(input_dir)
    clusters = find_clusters(frames, gap_seconds=args.cluster_gap_seconds)
    lte_stages = extract_plmn_stages(frames["LTE_PLMN_SEARCH.csv"], "LTE", clusters)
    nr_stages = extract_plmn_stages(frames["NR5G_PLMN_SEARCH.csv"], "NR5G", clusters)
    nr_acq_stages, nr_acq_detail = summarize_nr_acq(frames["NR5G_ACQ.csv"], clusters)
    lte_init_stages, lte_init_detail = summarize_lte_init_acq(frames["LTE_INIT_ACQ.csv"], clusters)
    lte_band_windows = summarize_window_source(frames["LTE_BAND_SCAN.csv"], clusters, "LTE_BAND_SCAN", "LTE", "Candidate_Band", "EARFCN")
    lte_system_windows = summarize_window_source(frames["LTE_SYSTEM_SCAN.csv"], clusters, "LTE_SYSTEM_SCAN", "LTE", "Band", "EARFCN")

    stages = lte_stages + nr_stages + nr_acq_stages + lte_init_stages + lte_band_windows + lte_system_windows
    attach_overlap_details(stages, frames)
    stages.sort(key=lambda item: (item["cluster_id"], item["start_ts"], item["source"], item["group_id"]))

    html_doc = build_html(input_dir, clusters, frames, stages, nr_acq_detail, lte_init_detail)
    output_html.write_text(html_doc, encoding="utf-8")
    print(f"Generated {output_html}")


if __name__ == "__main__":
    main()
