#!/usr/bin/env python3

import argparse
import html
import json
import math
import re
from pathlib import Path
from typing import Callable, Dict, List, Optional, Tuple

import pandas as pd


MOD_TYPE_MAP = {
    "QPSK": 2,
    "16QAM": 4,
    "16 QAM": 4,
    "64QAM": 6,
    "64 QAM": 6,
    "256QAM": 8,
    "256 QAM": 8,
}

RX_NUM_MAP = {
    "2X2_MIMO": 2,
    "4X4_MIMO": 4,
}

PATTERNS = {
    "B16E_PWR": re.compile(
        r"Slot(?P<slot>\d+) B16E avg_over_last_1s "
        r"PUSCH\(tx_avg_x10dBm=(?P<tx>-?\d+) mtpl_avg_x10dBm=(?P<mtpl>-?\d+) "
        r"tpc_avg_x10dB=(?P<tpc>-?\d+) pathloss_avg_x10dB=(?P<pathloss>-?\d+) "
        r"count=(?P<count>\d+)\)"
    ),
    "B16E_UL": re.compile(
        r"Slot(?P<slot>\d+) B16E avg_over_last_1s "
        r"UL\(tput_kbps_1s=(?P<tput>\d+) count=(?P<count>\d+) "
        r"rb_avg_x100=(?P<rb>\d+) bw_mhz_x100=(?P<bw>\d+)\)"
    ),
    "B16F_PWR": re.compile(
        r"Slot(?P<slot>\d+) B16F avg_over_last_1s "
        r"PUCCH\(tx_avg_x10dBm=(?P<tx>-?\d+) tpc_avg_x10dB=(?P<tpc>-?\d+) "
        r"pathloss_avg_x10dB=(?P<pathloss>-?\d+) count=(?P<count>\d+)\)"
    ),
    "B173_1": re.compile(
        r"Slot(?P<slot>\d+) B173 avg_over_last_1s DL"
        r"\(tput_kbps_1s=(?P<tput>\d+) count=(?P<count>\d+) "
        r"avg_layers_x100=(?P<layers>\d+) avg_mcs_x100=(?P<mcs>\d+)\)"
    ),
    "B173_2": re.compile(
        r"Slot(?P<slot>\d+) B173 avg_over_last_1s DL"
        r"\(avg_mod_x100=(?P<mod>\d+) avg_rb_x100=(?P<rb>\d+) "
        r"bw_mhz_x100=(?P<bw>\d+)\)"
    ),
    "B883_1": re.compile(
        r"Slot(?P<slot>\d+) B883 avg_over_last_1s UL"
        r"\(tput_kbps_1s=(?P<tput>\d+) count=(?P<count>\d+) "
        r"avg_layers_x100=(?P<layers>\d+) avg_mod_x100=(?P<mod>\d+)\)"
    ),
    "B883_2": re.compile(
        r"Slot(?P<slot>\d+) B883 avg_over_last_1s UL"
        r"\(avg_rb_x100=(?P<rb>\d+) avg_tx_num_x100=(?P<txnum>\d+) "
        r"bw_mhz_x100=(?P<bw>\d+)\)"
    ),
    "B887_1": re.compile(
        r"Slot(?P<slot>\d+) B887 avg_over_last_1s DL"
        r"\(tput_kbps_1s=(?P<tput>\d+) count=(?P<count>\d+) "
        r"avg_layers_x100=(?P<layers>\d+) avg_mcs_x100=(?P<mcs>\d+)\)"
    ),
    "B887_2": re.compile(
        r"Slot(?P<slot>\d+) B887 avg_over_last_1s DL"
        r"\(avg_mod_x100=(?P<mod>\d+) avg_rb_x100=(?P<rb>\d+) "
        r"avg_rx_num_x100=(?P<rxnum>\d+) bw_mhz_x100=(?P<bw>\d+)\)"
    ),
    "B884": re.compile(
        r"Slot(?P<slot>\d+) B884 avg_over_last_1s "
        r"(?P<channel>PUSCH|PUCCH|SRS)"
        r"\(tx_avg_x10dBm=(?P<tx>-?\d+) mtpl_avg_x10dBm=(?P<mtpl>-?\d+) "
        r"tpc_avg_x10dB=(?P<tpc>-?\d+) pathloss_avg_x10dB=(?P<pathloss>-?\d+) "
        r"count=(?P<count>\d+)\)"
    ),
    "B890_CFG": re.compile(
        r"Slot(?P<slot>\d+) B890 config_update DRX"
        r"\(drx_enable=(?P<drx>\d+) on_duration_ms=(?P<ond>\d+) "
        r"inactivity_timer_ms=(?P<inact>\d+) long_cycle_ms=(?P<long>\d+) "
        r"max_inactive_ms=(?P<maxi>\d+)\)"
    ),
    "B890_SUM_1": re.compile(
        r"Slot(?P<slot>\d+) B890 avg_over_last_1s CDRX"
        r"\(inactive_ms_x1000=(?P<inactive>\d+) inactive_pct_x100=(?P<pct>\d+) "
        r"ref_count_avg_x100=(?P<ref>\d+) count=(?P<count>\d+)\)"
    ),
    "B8A7": re.compile(
        r"Slot(?P<slot>\d+) B8A7 avg_over_last_1s CSF"
        r"\(avg_cqi_x100=(?P<cqi>-?\d+) avg_ri_x100=(?P<ri>-?\d+) "
        r"avg_cri_x100=(?P<cri>-?\d+) avg_pmi_wb_x1_x100=(?P<pmi_wb_x1>-?\d+) "
        r"avg_pmi_wb_x2_x100=(?P<pmi_wb_x2>-?\d+) count=(?P<count>\d+)\)"
    ),
}


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Compare offline parsed CSVs against online ICD summaries and generate CSV/HTML reports."
    )
    parser.add_argument("--offline-dir", required=True, help="Offline parsed CSV directory.")
    parser.add_argument("--online-dir", required=True, help="Online parsed CSV directory.")
    parser.add_argument("--output-dir", required=True, help="Output directory for report artifacts.")
    parser.add_argument("--title", default="ICD Offline vs Online Compare", help="HTML report title.")
    parser.add_argument(
        "--file-tag",
        default="",
        help="Optional suffix appended to generated filenames, for example 'v1.1.5'.",
    )
    return parser.parse_args()


def build_output_name(base_name: str, file_tag: str) -> str:
    if not file_tag:
        return base_name
    safe_tag = re.sub(r"[^A-Za-z0-9_.-]+", "_", str(file_tag)).strip("._")
    if not safe_tag:
        return base_name
    stem, suffix = base_name.rsplit(".", 1)
    return f"{stem}_{safe_tag}.{suffix}"


def parse_dt(date_series: pd.Series, time_series: pd.Series) -> pd.Series:
    return pd.to_datetime(date_series.astype(str) + " " + time_series.astype(str), errors="coerce")


def read_csv(csv_path: Path, low_memory: bool = True) -> pd.DataFrame:
    if not csv_path.is_file():
        raise FileNotFoundError(f"Missing CSV: {csv_path}")
    return pd.read_csv(csv_path, low_memory=low_memory)


def read_csv_optional(csv_path: Path, low_memory: bool = True) -> pd.DataFrame:
    if not csv_path.is_file():
        return pd.DataFrame()
    return pd.read_csv(csv_path, low_memory=low_memory)


def round_half_away_from_zero(value: float) -> Optional[int]:
    if pd.isna(value):
        return None
    if value >= 0:
        return int(math.floor(float(value) + 0.5))
    return int(math.ceil(float(value) - 0.5))


def window_filter(
    df: pd.DataFrame,
    ts_col: str,
    end_ts: pd.Timestamp,
    sub_col: str,
    sub_id: int,
    extra_mask: Optional[Callable[[pd.DataFrame], pd.Series]] = None,
) -> pd.DataFrame:
    mask = (
        (df[sub_col] == sub_id)
        & (df[ts_col] > end_ts - pd.Timedelta(seconds=1))
        & (df[ts_col] <= end_ts)
    )
    if extra_mask is not None:
        mask &= extra_mask(df)
    return df.loc[mask]


def to_numeric_mean_x100(series: pd.Series) -> Optional[int]:
    values = pd.to_numeric(series, errors="coerce")
    return round_half_away_from_zero(values.mean() * 100.0)


def to_numeric_mean_x10(series: pd.Series) -> Optional[int]:
    values = pd.to_numeric(series, errors="coerce")
    if values.dropna().empty:
        return 0
    return round_half_away_from_zero(values.mean() * 10.0)


def to_numeric_sum_tput_kbps(series: pd.Series) -> Optional[int]:
    values = pd.to_numeric(series, errors="coerce")
    return round_half_away_from_zero(values.sum() * 8.0 / 1000.0)


def kind_display_name(kind: str) -> str:
    mapping = {
        "B16E_PWR": "B16E LTE PUSCH power summary",
        "B16E_UL": "B16E LTE UL summary",
        "B16F_PWR": "B16F LTE PUCCH power summary",
        "B173_1": "B173 LTE PDSCH summary",
        "B173_2": "B173 LTE PDSCH summary",
        "B883_1": "B883 NR UL summary",
        "B883_2": "B883 NR UL summary",
        "B887_1": "B887 NR DL summary",
        "B887_2": "B887 NR DL summary",
        "B884": "B884 NR power summary",
        "B8A7": "B8A7 NR MAC CSF summary",
    }
    return mapping.get(kind, kind)


def field_display_name(field: str) -> str:
    mapping = {
        "tput": "tput (kbps)",
        "tx": "tx power (dBm x10)",
        "mtpl": "MTPL (dBm x10)",
        "tpc": "TPC (dB x10)",
        "pathloss": "pathloss (dB x10)",
        "bw": "bandwidth (MHz x100)",
        "rb": "RB count (x100)",
        "layers": "layers (x100)",
        "mcs": "MCS (x100)",
        "mod": "modulation order (x100)",
        "txnum": "TX num (x100)",
        "rxnum": "RX num (x100)",
        "cqi": "CQI (x100)",
        "ri": "RI (x100)",
        "cri": "CRI (x100)",
        "pmi_wb_x1": "PMI WB X1 (x100)",
        "pmi_wb_x2": "PMI WB X2 (x100)",
        "count": "sample count",
    }
    return mapping.get(field, field)


class CompareContext:
    def __init__(self, offline_dir: Path, online_dir: Path) -> None:
        self.offline_dir = offline_dir
        self.online_dir = online_dir
        self.nr_ul = read_csv(offline_dir / "NR_MAC_UL_SCHEDULE.csv", low_memory=False)
        self.nr_ul["ts"] = parse_dt(self.nr_ul["Date"], self.nr_ul["Time"])
        self.nr_tx = read_csv(offline_dir / "NR_MAC_TX.csv")
        self.nr_tx["ts"] = parse_dt(self.nr_tx["Date"], self.nr_tx["Time"])
        self.nr_dl = read_csv(offline_dir / "NR_MAC_PDSCH.csv")
        self.nr_dl["ts"] = parse_dt(self.nr_dl["Date"], self.nr_dl["Time"])
        self.lte_dl = read_csv(offline_dir / "LTE_MAC_PDSCH.csv")
        self.lte_dl["ts"] = parse_dt(self.lte_dl["Date"], self.lte_dl["Time"])
        self.lte_tx = read_csv(offline_dir / "LTE_MAC_TX.csv")
        self.lte_tx["ts"] = parse_dt(self.lte_tx["Date"], self.lte_tx["Time"])
        self.cdrx = read_csv(offline_dir / "NR5G_CDRX_Config_and_State.csv")
        self.cdrx["ts"] = parse_dt(self.cdrx["Date"], self.cdrx["Time"])
        self.rrc = read_csv(offline_dir / "NR5G_RRC_CONFIG.csv")
        self.rrc["ts"] = parse_dt(self.rrc["Date"], self.rrc["Time"])
        self.nr_csf = read_csv_optional(offline_dir / "NR_MAC_CSF.csv")
        if not self.nr_csf.empty:
            self.nr_csf["ts"] = parse_dt(self.nr_csf["Date"], self.nr_csf["Time"])


def calc_b16e_power(ctx: CompareContext, end_ts: pd.Timestamp, sub_id: int) -> Dict[str, Optional[int]]:
    win = window_filter(
        ctx.lte_tx,
        "ts",
        end_ts,
        "Sub_ID",
        sub_id,
        lambda df: df["Channel"] == "PUSCH",
    )
    return {
        "tx": to_numeric_mean_x10(win["Tx_power"]),
        "mtpl": to_numeric_mean_x10(win["MTPL"]),
        "tpc": to_numeric_mean_x10(win["TPC"]),
        "pathloss": to_numeric_mean_x10(win["PathLoss"]),
        "count": int(len(win)),
    }


def calc_b16e_ul(ctx: CompareContext, end_ts: pd.Timestamp, sub_id: int) -> Dict[str, Optional[int]]:
    win = window_filter(
        ctx.lte_tx,
        "ts",
        end_ts,
        "Sub_ID",
        sub_id,
        lambda df: df["Channel"] == "PUSCH",
    )
    num_rbs = pd.to_numeric(win["Num_RBs"], errors="coerce")
    bw_x100 = round_half_away_from_zero(((num_rbs * 15.0 * 12.0 * 100.0) / 1000.0).mean())
    return {
        "tput": to_numeric_sum_tput_kbps(win["TB_size_bytes"]),
        "count": int(len(win)),
        "rb": to_numeric_mean_x100(win["Num_RBs"]),
        "bw": bw_x100,
    }


def calc_b16f_power(ctx: CompareContext, end_ts: pd.Timestamp, sub_id: int) -> Dict[str, Optional[int]]:
    win = window_filter(
        ctx.lte_tx,
        "ts",
        end_ts,
        "Sub_ID",
        sub_id,
        lambda df: df["Channel"] == "PUCCH",
    )
    return {
        "tx": to_numeric_mean_x10(win["Tx_power"]),
        "tpc": to_numeric_mean_x10(win["TPC"]),
        "pathloss": to_numeric_mean_x10(win["PathLoss"]),
        "count": int(len(win)),
    }


def calc_b173(ctx: CompareContext, end_ts: pd.Timestamp, sub_id: int) -> Dict[str, Optional[int]]:
    win = window_filter(ctx.lte_dl, "ts", end_ts, "SubId", sub_id)
    return {
        "tput": to_numeric_sum_tput_kbps(win["TB_size_bytes"]),
        "count": int(len(win)),
        "layers": to_numeric_mean_x100(win["Layers"]),
        "mcs": to_numeric_mean_x100(win["MCS"]),
        "mod": round_half_away_from_zero(win["Mod_type"].map(MOD_TYPE_MAP).mean() * 100.0),
        "rb": to_numeric_mean_x100(win["RBs"]),
    }


def calc_b883(ctx: CompareContext, end_ts: pd.Timestamp, sub_id: int) -> Dict[str, Optional[int]]:
    win = window_filter(
        ctx.nr_ul,
        "ts",
        end_ts,
        "Sub_ID",
        sub_id,
        lambda df: df["Channel"] == "PUSCH",
    )
    return {
        "tput": to_numeric_sum_tput_kbps(win["PUSCH_TB_Size_bytes"]),
        "count": int(len(win)),
        "layers": to_numeric_mean_x100(win["UL_LAYERS_PUSCH"]),
        "mod": round_half_away_from_zero(win["PUSCH_Modulation_Order"].map(MOD_TYPE_MAP).mean() * 100.0),
        "rb": to_numeric_mean_x100(win["Num_RBs"]),
        "txnum": to_numeric_mean_x100(win["UL_TX_NUM"]),
    }


def calc_b887(ctx: CompareContext, end_ts: pd.Timestamp, sub_id: int) -> Dict[str, Optional[int]]:
    win = window_filter(ctx.nr_dl, "ts", end_ts, "Sub_ID", sub_id)
    return {
        "tput": to_numeric_sum_tput_kbps(win["TB_size"]),
        "count": int(len(win)),
        "layers": to_numeric_mean_x100(win["Layers"]),
        "mcs": to_numeric_mean_x100(win["MCS"]),
        "mod": round_half_away_from_zero(win["Mod_Type"].map(MOD_TYPE_MAP).mean() * 100.0),
        "rb": to_numeric_mean_x100(win["Rbs"]),
        "rxnum": round_half_away_from_zero(win["Num_RX"].map(RX_NUM_MAP).mean() * 100.0),
    }


def calc_b884(ctx: CompareContext, end_ts: pd.Timestamp, sub_id: int, channel: str) -> Dict[str, Optional[int]]:
    win = window_filter(
        ctx.nr_tx,
        "ts",
        end_ts,
        "Sub_ID",
        sub_id,
        lambda df: df["Channel"] == channel,
    )
    return {
        "tx": to_numeric_mean_x10(win["Tx_power"]),
        "mtpl": to_numeric_mean_x10(win["MTPL"]),
        "tpc": to_numeric_mean_x10(win["TPC"]),
        "pathloss": to_numeric_mean_x10(win["PathLoss"]),
        "count": int(len(win)),
    }


def calc_b8a7(ctx: CompareContext, end_ts: pd.Timestamp, sub_id: int) -> Dict[str, Optional[int]]:
    if ctx.nr_csf.empty:
        return {
            "cqi": None,
            "ri": None,
            "cri": None,
            "pmi_wb_x1": None,
            "pmi_wb_x2": None,
            "count": 0,
        }
    win = window_filter(ctx.nr_csf, "ts", end_ts, "SubId", sub_id)
    return {
        "cqi": to_numeric_mean_x100(win["CQI"]),
        "ri": to_numeric_mean_x100(win["RI"]),
        "cri": to_numeric_mean_x100(win["CRI"]),
        "pmi_wb_x1": to_numeric_mean_x100(win["PMI_WB_X1"]),
        "pmi_wb_x2": to_numeric_mean_x100(win["PMI_WB_X2"]),
        "count": int(len(win)),
    }


def online_summary_to_df(csv_path: Path) -> pd.DataFrame:
    df = read_csv_optional(csv_path)
    if df.empty:
        return df
    df["ts"] = parse_dt(df["Date"], df["Time"])
    return df


def compare_one_kind(
    ctx: CompareContext,
    online_csv_name: str,
    kind: str,
    calc_func: Callable[..., Dict[str, Optional[int]]],
    compare_fields: List[str],
) -> List[Dict[str, object]]:
    df = online_summary_to_df(ctx.online_dir / online_csv_name)
    rows: List[Dict[str, object]] = []
    if df.empty:
        return rows

    for _, row in df.iterrows():
        summary = str(row.get("Summary", ""))
        match = PATTERNS[kind].search(summary)
        if not match:
            continue

        slot = int(match.group("slot"))
        sub_id = slot + 1
        extra_channel = match.groupdict().get("channel")
        series_name = f"{kind_display_name(kind)} Slot{slot}"
        if extra_channel:
            series_name += f" {extra_channel}"

        if extra_channel is None:
            offline_values = calc_func(ctx, row["ts"], sub_id)
        else:
            offline_values = calc_func(ctx, row["ts"], sub_id, extra_channel)

        online_values = {
            key: int(value)
            for key, value in match.groupdict().items()
            if key not in {"slot", "channel"} and value is not None
        }

        for field in compare_fields:
            online_value = online_values.get(field)
            offline_value = offline_values.get(field)
            delta = None if offline_value is None or online_value is None else online_value - offline_value
            rows.append(
                {
                    "kind": kind,
                    "kind_display": kind_display_name(kind),
                    "series_name": series_name,
                    "slot": slot,
                    "sub_id": sub_id,
                    "channel": extra_channel or "",
                    "field": field,
                    "timestamp": row["ts"],
                    "online": online_value,
                    "offline": offline_value,
                    "delta": delta,
                    "is_match": delta == 0 if delta is not None else False,
                    "summary": summary,
                }
            )
    return rows


def build_compare_details(ctx: CompareContext) -> pd.DataFrame:
    rows: List[Dict[str, object]] = []
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B16E_SUMMARY_MSG.csv",
            "B16E_PWR",
            calc_b16e_power,
            ["tx", "mtpl", "tpc", "pathloss", "count"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B16E_SUMMARY_MSG.csv",
            "B16E_UL",
            calc_b16e_ul,
            ["tput", "count", "rb", "bw"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B16F_SUMMARY_MSG.csv",
            "B16F_PWR",
            calc_b16f_power,
            ["tx", "tpc", "pathloss", "count"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B173_SUMMARY_MSG.csv",
            "B173_1",
            calc_b173,
            ["tput", "count", "layers", "mcs"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B173_SUMMARY_MSG.csv",
            "B173_2",
            calc_b173,
            ["mod", "rb"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B883_SUMMARY_MSG.csv",
            "B883_1",
            calc_b883,
            ["tput", "count", "layers", "mod"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B883_SUMMARY_MSG.csv",
            "B883_2",
            calc_b883,
            ["rb", "txnum"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B887_SUMMARY_MSG.csv",
            "B887_1",
            calc_b887,
            ["tput", "count", "layers", "mcs"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B887_SUMMARY_MSG.csv",
            "B887_2",
            calc_b887,
            ["mod", "rb", "rxnum"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B884_SUMMARY_MSG.csv",
            "B884",
            calc_b884,
            ["tx", "mtpl", "tpc", "pathloss", "count"],
        )
    )
    rows.extend(
        compare_one_kind(
            ctx,
            "BYTEDANCE_DIAG_RF_B8A7_SUMMARY_MSG.csv",
            "B8A7",
            calc_b8a7,
            ["cqi", "ri", "cri", "pmi_wb_x1", "pmi_wb_x2", "count"],
        )
    )
    detail_columns = [
        "kind",
        "kind_display",
        "series_name",
        "slot",
        "sub_id",
        "channel",
        "field",
        "field_display",
        "timestamp",
        "online",
        "offline",
        "delta",
        "is_match",
        "summary",
    ]
    df = pd.DataFrame(rows)
    if not df.empty:
        df["field_display"] = df["field"].map(field_display_name)
        df = df.sort_values(["series_name", "field", "timestamp"]).reset_index(drop=True)
    else:
        df = pd.DataFrame(columns=detail_columns)
    return df


def build_b890_config_timeline(ctx: CompareContext) -> Tuple[pd.DataFrame, pd.DataFrame]:
    online = online_summary_to_df(ctx.online_dir / "BYTEDANCE_DIAG_RF_B890_CONFIG_MSG.csv")
    online_rows: List[Dict[str, object]] = []
    if not online.empty:
        for _, row in online.iterrows():
            match = PATTERNS["B890_CFG"].search(str(row.get("Summary", "")))
            if not match:
                continue
            slot = int(match.group("slot"))
            online_rows.append(
                {
                    "source": "online",
                    "sub_id": slot + 1,
                    "slot": slot,
                    "timestamp": row["ts"],
                    "drx": int(match.group("drx")),
                    "on_duration_ms": int(match.group("ond")),
                    "inactivity_timer_ms": int(match.group("inact")),
                    "long_cycle_ms": int(match.group("long")),
                    "max_inactive_ms": int(match.group("maxi")),
                    "summary": str(row.get("Summary", "")),
                }
            )

    offline = ctx.cdrx.copy()
    offline["drx"] = pd.to_numeric(offline["DrxEnable"], errors="coerce").fillna(0).astype(int)
    offline["on_duration_ms"] = (
        pd.to_numeric(offline["OnDuration"].astype(str).str.extract(r"(\d+)")[0], errors="coerce")
        .fillna(0)
        .astype(int)
    )
    offline["inactivity_timer_ms"] = (
        pd.to_numeric(offline["InactivityTimer"].astype(str).str.extract(r"(\d+)")[0], errors="coerce")
        .fillna(0)
        .astype(int)
    )
    offline["long_cycle_ms"] = (
        pd.to_numeric(offline["LongDrxCycle"].astype(str).str.extract(r"(\d+)")[0], errors="coerce")
        .fillna(0)
        .astype(int)
    )
    offline["max_inactive_ms"] = offline.apply(
        lambda row: 0
        if row["drx"] == 0 or row["long_cycle_ms"] <= row["on_duration_ms"]
        else row["long_cycle_ms"] - row["on_duration_ms"],
        axis=1,
    )

    offline_rows: List[Dict[str, object]] = []
    for sub_id, group in offline.sort_values("ts").groupby("SubID"):
        previous_tuple = None
        for _, row in group.iterrows():
            current_tuple = (
                int(row["drx"]),
                int(row["on_duration_ms"]),
                int(row["inactivity_timer_ms"]),
                int(row["long_cycle_ms"]),
                int(row["max_inactive_ms"]),
            )
            if current_tuple == previous_tuple:
                continue
            offline_rows.append(
                {
                    "source": "offline",
                    "sub_id": int(sub_id),
                    "slot": int(sub_id) - 1,
                    "timestamp": row["ts"],
                    "drx": current_tuple[0],
                    "on_duration_ms": current_tuple[1],
                    "inactivity_timer_ms": current_tuple[2],
                    "long_cycle_ms": current_tuple[3],
                    "max_inactive_ms": current_tuple[4],
                    "summary": "",
                }
            )
            previous_tuple = current_tuple

    timeline_df = pd.DataFrame(online_rows + offline_rows)
    if not timeline_df.empty:
        timeline_df = timeline_df.sort_values(["sub_id", "timestamp", "source"]).reset_index(drop=True)

    seq_rows: List[Dict[str, object]] = []
    sequence_columns = [
        "source",
        "sub_id",
        "slot",
        "timestamp",
        "drx",
        "on_duration_ms",
        "inactivity_timer_ms",
        "long_cycle_ms",
        "max_inactive_ms",
        "summary",
    ]
    online_seq = pd.DataFrame(online_rows, columns=sequence_columns)
    offline_seq = pd.DataFrame(offline_rows, columns=sequence_columns)
    all_sub_ids = sorted(set(online_seq.get("sub_id", pd.Series(dtype=int))) | set(offline_seq.get("sub_id", pd.Series(dtype=int))))
    for sub_id in all_sub_ids:
        on = online_seq[online_seq["sub_id"] == sub_id].reset_index(drop=True)
        off = offline_seq[offline_seq["sub_id"] == sub_id].reset_index(drop=True)
        max_len = max(len(on), len(off))
        for index in range(max_len):
            on_row = on.iloc[index] if index < len(on) else None
            off_row = off.iloc[index] if index < len(off) else None
            on_tuple = None if on_row is None else [
                int(on_row["drx"]),
                int(on_row["on_duration_ms"]),
                int(on_row["inactivity_timer_ms"]),
                int(on_row["long_cycle_ms"]),
                int(on_row["max_inactive_ms"]),
            ]
            off_tuple = None if off_row is None else [
                int(off_row["drx"]),
                int(off_row["on_duration_ms"]),
                int(off_row["inactivity_timer_ms"]),
                int(off_row["long_cycle_ms"]),
                int(off_row["max_inactive_ms"]),
            ]
            status = "match" if on_tuple == off_tuple and on_tuple is not None else "mismatch"
            seq_rows.append(
                {
                    "sub_id": sub_id,
                    "seq_index": index,
                    "status": status,
                    "online_ts": None if on_row is None else on_row["timestamp"],
                    "offline_ts": None if off_row is None else off_row["timestamp"],
                    "online_tuple": on_tuple,
                    "offline_tuple": off_tuple,
                    "delta_ms": None
                    if on_row is None or off_row is None
                    else (on_row["timestamp"] - off_row["timestamp"]).total_seconds() * 1000.0,
                }
            )
    sequence_df = pd.DataFrame(seq_rows)
    return timeline_df, sequence_df


def build_b890_invariant_issues(ctx: CompareContext) -> pd.DataFrame:
    online = online_summary_to_df(ctx.online_dir / "BYTEDANCE_DIAG_RF_B890_SUMMARY_MSG.csv")
    rows: List[Dict[str, object]] = []
    if online.empty:
        return pd.DataFrame(rows)
    for _, row in online.iterrows():
        summary = str(row.get("Summary", ""))
        match = PATTERNS["B890_SUM_1"].search(summary)
        if not match:
            continue
        issues = []
        inactive_ms = int(match.group("inactive"))
        inactive_pct_x100 = int(match.group("pct"))
        if inactive_ms > 1000:
            issues.append("inactive_ms_gt_1000")
        if inactive_pct_x100 > 10000:
            issues.append("inactive_pct_gt_100pct")
        if not issues:
            continue
        rows.append(
            {
                "timestamp": row["ts"],
                "slot": int(match.group("slot")),
                "sub_id": int(match.group("slot")) + 1,
                "inactive_ms_x1000": inactive_ms,
                "inactive_pct_x100": inactive_pct_x100,
                "ref_count_avg_x100": int(match.group("ref")),
                "count": int(match.group("count")),
                "issues": ",".join(issues),
                "summary": summary,
            }
        )
    return pd.DataFrame(rows)


def build_b825_coverage(ctx: CompareContext) -> pd.DataFrame:
    return pd.DataFrame(
        [
            {
                "offline_rows": int(len(ctx.rrc)),
                "online_config_rows": int(len(read_csv(ctx.online_dir / "BYTEDANCE_DIAG_RF_B825_CONFIG_MSG.csv"))),
                "online_debug_rows": int(len(read_csv(ctx.online_dir / "BYTEDANCE_DIAG_RF_B825_DEBUG_MSG.csv"))),
                "online_error_rows": int(len(read_csv(ctx.online_dir / "BYTEDANCE_DIAG_RF_B825_ERROR_MSG.csv"))),
            }
        ]
    )


def build_error_category_counts(ctx: CompareContext) -> pd.DataFrame:
    category_names = [
        "BYTEDANCE_DIAG_RF_B16E_ERROR_MSG",
        "BYTEDANCE_DIAG_RF_B16F_ERROR_MSG",
        "BYTEDANCE_DIAG_RF_B173_ERROR_MSG",
        "BYTEDANCE_DIAG_RF_B883_ERROR_MSG",
        "BYTEDANCE_DIAG_RF_B884_ERROR_MSG",
        "BYTEDANCE_DIAG_RF_B887_ERROR_MSG",
        "BYTEDANCE_DIAG_RF_B890_ERROR_MSG",
        "BYTEDANCE_DIAG_RF_B8A7_ERROR_MSG",
    ]
    rows: List[Dict[str, object]] = []
    for category_name in category_names:
        csv_path = ctx.online_dir / f"{category_name}.csv"
        df = read_csv_optional(csv_path)
        rows.append(
            {
                "category": category_name,
                "csv_present": csv_path.is_file(),
                "row_count": int(len(df)),
            }
        )
    return pd.DataFrame(rows)


def to_plot_value(value: object) -> object:
    if isinstance(value, pd.Timestamp):
        return value.isoformat(sep=" ")
    if pd.isna(value):
        return None
    return value


def plot_div(plot_id: str, traces: List[Dict[str, object]], layout: Dict[str, object]) -> str:
    return (
        f'<div id="{plot_id}" class="plot"></div>'
        f'<script>Plotly.newPlot("{plot_id}", {json.dumps(traces, ensure_ascii=False)}, '
        f'{json.dumps(layout, ensure_ascii=False)}, '
        '{responsive: true, displaylogo: false});</script>'
    )


def build_single_field_compare_figure(field_df: pd.DataFrame, plot_id: str, title: str, y_axis_title: str) -> str:
    field_df = field_df.sort_values("timestamp")
    traces: List[Dict[str, object]] = [
        {
            "type": "scatter",
            "mode": "lines+markers",
            "x": [to_plot_value(value) for value in field_df["timestamp"]],
            "y": [to_plot_value(value) for value in field_df["offline"]],
            "name": "offline",
            "line": {"color": "#1f77b4", "width": 2.0},
            "marker": {"size": 5},
            "hovertemplate": "offline=%{y}<br>%{x}<extra></extra>",
        },
        {
            "type": "scatter",
            "mode": "lines+markers",
            "x": [to_plot_value(value) for value in field_df["timestamp"]],
            "y": [to_plot_value(value) for value in field_df["online"]],
            "name": "online",
            "line": {"color": "#d62728", "width": 2.0},
            "marker": {"size": 5},
            "hovertemplate": "online=%{y}<br>%{x}<extra></extra>",
        },
    ]
    mismatch_df = field_df[field_df["delta"].fillna(0) != 0]
    if not mismatch_df.empty:
        traces.append(
            {
                "type": "scatter",
                "mode": "markers",
                "x": [to_plot_value(value) for value in mismatch_df["timestamp"]],
                "y": [to_plot_value(value) for value in mismatch_df["online"]],
                "name": "mismatch",
                "marker": {"size": 9, "color": "#ff7f0e", "symbol": "diamond"},
                "customdata": [[to_plot_value(value)] for value in mismatch_df["delta"]],
                "hovertemplate": "delta=%{customdata[0]}<br>%{x}<extra></extra>",
            }
        )
    layout: Dict[str, object] = {
        "template": "plotly_white",
        "height": 340,
        "margin": {"l": 70, "r": 30, "t": 50, "b": 55},
        "hovermode": "x unified",
        "legend": {"orientation": "h", "yanchor": "bottom", "y": 1.02, "xanchor": "left", "x": 0.0},
        "title": {"text": title, "x": 0.0, "xanchor": "left", "font": {"size": 14}},
        "xaxis": {"title": {"text": "summary timestamp"}, "tickformat": "%H:%M:%S"},
        "yaxis": {"title": {"text": y_axis_title}, "automargin": True},
    }
    return plot_div(plot_id, traces, layout)


def build_series_section(details_df: pd.DataFrame, series_name: str, plot_id_prefix: str) -> str:
    data = details_df[details_df["series_name"] == series_name].copy()
    fields = list(dict.fromkeys(data["field"].tolist()))
    cards = []
    for index, field in enumerate(fields, start=1):
        field_df = data[data["field"] == field].copy()
        plot_html = build_single_field_compare_figure(
            field_df,
            f"{plot_id_prefix}_{index}",
            f"{series_name} / {field_display_name(field)}",
            field_display_name(field),
        )
        cards.append(f'<div class="plot-card">{plot_html}</div>')
    return (
        f'<div class="card"><h2>{html.escape(series_name)}</h2>'
        f'<div class="plot-grid">{"".join(cards)}</div></div>'
    )


def build_overview_figure(details_df: pd.DataFrame, plot_id: str) -> str:
    summary = (
        details_df.groupby(["kind_display", "field"], as_index=False)
        .agg(
            total=("delta", "size"),
            exact=("is_match", "sum"),
            median_abs_delta=("delta", lambda series: series.abs().median()),
            max_abs_delta=("delta", lambda series: series.abs().max()),
        )
        .sort_values(["kind_display", "field"])
    )
    summary["exact_rate_pct"] = summary["exact"] * 100.0 / summary["total"]
    labels = (summary["kind_display"] + " / " + summary["field"].map(field_display_name)).tolist()
    traces = [
        {
            "type": "bar",
            "y": labels,
            "x": [to_plot_value(value) for value in summary["exact_rate_pct"]],
            "name": "exact_rate_pct",
            "marker": {"color": "#2ca02c"},
            "xaxis": "x",
            "yaxis": "y",
            "orientation": "h",
        },
        {
            "type": "bar",
            "y": labels,
            "x": [to_plot_value(value) for value in summary["median_abs_delta"]],
            "name": "median_abs_delta",
            "marker": {"color": "#1f77b4"},
            "xaxis": "x2",
            "yaxis": "y2",
            "orientation": "h",
        },
        {
            "type": "bar",
            "y": labels,
            "x": [to_plot_value(value) for value in summary["max_abs_delta"]],
            "name": "max_abs_delta",
            "marker": {"color": "#d62728"},
            "xaxis": "x2",
            "yaxis": "y2",
            "orientation": "h",
        },
    ]
    layout = {
        "template": "plotly_white",
        "height": max(700, 26 * len(labels) + 220),
        "margin": {"l": 300, "r": 30, "t": 50, "b": 40},
        "barmode": "group",
        "xaxis": {"domain": [0.0, 1.0], "anchor": "y", "title": {"text": "%"}},
        "yaxis": {"domain": [0.56, 1.0], "anchor": "x", "automargin": True},
        "xaxis2": {"domain": [0.0, 1.0], "anchor": "y2", "title": {"text": "abs delta"}},
        "yaxis2": {"domain": [0.0, 0.42], "anchor": "x2", "automargin": True},
        "annotations": [
            {
                "text": "Exact match rate by parameter",
                "xref": "paper",
                "yref": "paper",
                "x": 0.0,
                "y": 1.03,
                "showarrow": False,
                "xanchor": "left",
            },
            {
                "text": "Median / max absolute delta by parameter",
                "xref": "paper",
                "yref": "paper",
                "x": 0.0,
                "y": 0.47,
                "showarrow": False,
                "xanchor": "left",
            },
        ],
    }
    return plot_div(plot_id, traces, layout)


def build_b890_config_figure(timeline_df: pd.DataFrame, plot_id_prefix: str) -> str:
    fields = [
        "drx",
        "on_duration_ms",
        "inactivity_timer_ms",
        "long_cycle_ms",
        "max_inactive_ms",
    ]
    sub_ids = sorted(timeline_df["sub_id"].dropna().unique().tolist())
    cards = []
    for sub_id in sub_ids:
        sub_df = timeline_df[timeline_df["sub_id"] == sub_id].copy()
        for index, field in enumerate(fields, start=1):
            field_df = sub_df[["timestamp", "source", field]].copy()
            traces: List[Dict[str, object]] = []
            for source, color in [("offline", "#1f77b4"), ("online", "#d62728")]:
                source_df = field_df[field_df["source"] == source].sort_values("timestamp")
                traces.append(
                    {
                        "type": "scatter",
                        "mode": "lines+markers",
                        "x": [to_plot_value(value) for value in source_df["timestamp"]],
                        "y": [to_plot_value(value) for value in source_df[field]],
                        "name": source,
                        "line": {"color": color, "width": 2.0, "shape": "hv"},
                        "marker": {"size": 6},
                    }
                )
            layout = {
                "template": "plotly_white",
                "height": 320,
                "margin": {"l": 70, "r": 30, "t": 50, "b": 55},
                "hovermode": "x unified",
                "legend": {"orientation": "h", "yanchor": "bottom", "y": 1.02, "xanchor": "left", "x": 0.0},
                "title": {"text": f"Sub{sub_id} / {field}", "x": 0.0, "xanchor": "left", "font": {"size": 14}},
                "xaxis": {"title": {"text": "config timestamp"}, "tickformat": "%H:%M:%S"},
                "yaxis": {"title": {"text": field}, "automargin": True},
            }
            cards.append(
                f'<div class="plot-card">{plot_div(f"{plot_id_prefix}_{sub_id}_{index}", traces, layout)}</div>'
            )
    return f'<div class="plot-grid">{"".join(cards)}</div>'


def df_to_html_table(df: pd.DataFrame, max_rows: Optional[int] = None) -> str:
    if df.empty:
        return "<p>No rows.</p>"
    view = df.copy()
    if max_rows is not None:
        view = view.head(max_rows)
    headers = "".join(f"<th>{html.escape(str(col))}</th>" for col in view.columns)
    body_rows = []
    for _, row in view.iterrows():
        cells = []
        for col in view.columns:
            value = row[col]
            if isinstance(value, pd.Timestamp):
                rendered = value.isoformat(sep=" ")
            else:
                rendered = str(value)
            cells.append(f"<td>{html.escape(rendered)}</td>")
        body_rows.append("<tr>" + "".join(cells) + "</tr>")
    return "<table><thead><tr>" + headers + "</tr></thead><tbody>" + "".join(body_rows) + "</tbody></table>"


def build_html_report(
    title: str,
    details_df: pd.DataFrame,
    b890_timeline_df: pd.DataFrame,
    b890_sequence_df: pd.DataFrame,
    b890_issues_df: pd.DataFrame,
    b825_coverage_df: pd.DataFrame,
    error_counts_df: pd.DataFrame,
) -> str:
    overall_rows = []
    if details_df.empty:
        summary_df = pd.DataFrame(columns=["kind_display", "field", "total", "exact"])
    else:
        summary_df = (
            details_df.groupby(["kind_display", "field"], as_index=False)
            .agg(total=("delta", "size"), exact=("is_match", "sum"))
            .sort_values(["kind_display", "field"])
        )
    for _, row in summary_df.iterrows():
        overall_rows.append(
            {
                "Kind": row["kind_display"],
                "Field": field_display_name(row["field"]),
                "Rows": int(row["total"]),
                "Exact": int(row["exact"]),
                "ExactRatePct": round(row["exact"] * 100.0 / row["total"], 1),
            }
        )

    series_sections = []
    if not details_df.empty:
        for index, series_name in enumerate(sorted(details_df["series_name"].unique().tolist()), start=1):
            series_sections.append(build_series_section(details_df, series_name, f"series_plot_{index}"))

    overview_div = (
        build_overview_figure(details_df, "overview_plot")
        if not details_df.empty
        else "<p>No offline/online summary rows to compare.</p>"
    )
    b890_div = (
        build_b890_config_figure(b890_timeline_df, "b890_config_plot")
        if not b890_timeline_df.empty
        else "<p>No B890 config data.</p>"
    )
    overall_table_html = df_to_html_table(pd.DataFrame(overall_rows))
    b890_seq_table_html = df_to_html_table(b890_sequence_df, max_rows=40)
    b890_issue_table_html = df_to_html_table(b890_issues_df, max_rows=40)
    b825_table_html = df_to_html_table(b825_coverage_df)
    error_count_table_html = df_to_html_table(error_counts_df)
    series_count = 0 if details_df.empty else details_df["series_name"].nunique()
    b890_mismatch_count = (
        int((b890_sequence_df["status"] == "mismatch").sum())
        if not b890_sequence_df.empty and "status" in b890_sequence_df.columns
        else 0
    )

    return f"""<!DOCTYPE html>
<html lang="en">
<head>
  <meta charset="utf-8">
  <title>{html.escape(title)}</title>
  <script src="https://cdn.plot.ly/plotly-2.35.2.min.js"></script>
  <style>
    body {{
      font-family: Arial, sans-serif;
      color: #1f2937;
      background: #f8fafc;
      margin: 0;
      padding: 0;
    }}
    .page {{
      max-width: 1680px;
      margin: 0 auto;
      padding: 24px;
    }}
    .card {{
      background: #ffffff;
      border: 1px solid #e5e7eb;
      border-radius: 14px;
      padding: 20px 24px;
      margin-bottom: 20px;
      box-shadow: 0 1px 2px rgba(0, 0, 0, 0.04);
    }}
    h1, h2, h3 {{
      margin-top: 0;
    }}
    p, li {{
      line-height: 1.6;
    }}
    table {{
      width: 100%;
      border-collapse: collapse;
      font-size: 13px;
    }}
    th, td {{
      border: 1px solid #e5e7eb;
      padding: 8px 10px;
      text-align: left;
      vertical-align: top;
    }}
    th {{
      background: #f8fafc;
    }}
    .meta {{
      display: grid;
      grid-template-columns: repeat(4, minmax(200px, 1fr));
      gap: 12px;
      margin-top: 14px;
    }}
    .meta-item {{
      background: #f8fafc;
      border: 1px solid #e5e7eb;
      border-radius: 10px;
      padding: 12px;
    }}
    .meta-label {{
      color: #6b7280;
      font-size: 12px;
      margin-bottom: 4px;
    }}
    .meta-value {{
      font-weight: 600;
      font-size: 15px;
    }}
    .note {{
      color: #4b5563;
    }}
    .plot {{
      width: 100%;
    }}
    .plot-grid {{
      display: grid;
      grid-template-columns: repeat(auto-fit, minmax(520px, 1fr));
      gap: 16px;
      align-items: start;
    }}
    .plot-card {{
      border: 1px solid #e5e7eb;
      border-radius: 12px;
      padding: 12px;
      background: #fcfcfd;
    }}
  </style>
</head>
<body>
  <div class="page">
    <div class="card">
      <h1>{html.escape(title)}</h1>
      <p class="note">Each chart compares one parameter only. Blue is offline, red is online, and orange diamonds mark mismatch points.</p>
      <div class="meta">
        <div class="meta-item"><div class="meta-label">Compared rows</div><div class="meta-value">{len(details_df)}</div></div>
        <div class="meta-item"><div class="meta-label">Series</div><div class="meta-value">{series_count}</div></div>
        <div class="meta-item"><div class="meta-label">B890 config seq mismatches</div><div class="meta-value">{b890_mismatch_count}</div></div>
        <div class="meta-item"><div class="meta-label">B890 invariant issues</div><div class="meta-value">{len(b890_issues_df)}</div></div>
      </div>
    </div>

    <div class="card">
      <h2>Parameter Summary</h2>
      {overall_table_html}
    </div>

    <div class="card">
      <h2>Overview</h2>
      {overview_div}
    </div>

    <div class="card">
      <h2>B890 Config Timeline</h2>
      <p class="note">Each chart compares one config parameter only, split by SubID.</p>
      {b890_div}
    </div>

    <div class="card">
      <h2>B890 Config Sequence</h2>
      {b890_seq_table_html}
    </div>

    <div class="card">
      <h2>B890 Invariant Issues</h2>
      {b890_issue_table_html}
    </div>

    <div class="card">
      <h2>B825 Coverage</h2>
      {b825_table_html}
    </div>

    <div class="card">
      <h2>Error Category Counts</h2>
      <p class="note">Error categories are diagnostic side outputs. They are counted and linked to online parsing health, but they are not plotted against offline truth data.</p>
      {error_count_table_html}
    </div>

    {"".join(series_sections)}
  </div>
</body>
</html>
"""


def main() -> None:
    args = parse_args()
    offline_dir = Path(args.offline_dir).expanduser().resolve()
    online_dir = Path(args.online_dir).expanduser().resolve()
    output_dir = Path(args.output_dir).expanduser().resolve()
    output_dir.mkdir(parents=True, exist_ok=True)

    ctx = CompareContext(offline_dir, online_dir)
    details_df = build_compare_details(ctx)
    b890_timeline_df, b890_sequence_df = build_b890_config_timeline(ctx)
    b890_issues_df = build_b890_invariant_issues(ctx)
    b825_coverage_df = build_b825_coverage(ctx)
    error_counts_df = build_error_category_counts(ctx)

    details_csv = output_dir / build_output_name("icd_offline_online_compare_details.csv", args.file_tag)
    details_df.to_csv(details_csv, index=False)

    if not b890_timeline_df.empty:
        b890_timeline_df.to_csv(output_dir / build_output_name("b890_config_timeline.csv", args.file_tag), index=False)
    if not b890_sequence_df.empty:
        b890_sequence_df.to_csv(output_dir / build_output_name("b890_config_sequence_compare.csv", args.file_tag), index=False)
    if not b890_issues_df.empty:
        b890_issues_df.to_csv(output_dir / build_output_name("b890_summary_invariant_issues.csv", args.file_tag), index=False)
    b825_coverage_df.to_csv(output_dir / build_output_name("b825_coverage.csv", args.file_tag), index=False)
    error_counts_df.to_csv(output_dir / build_output_name("error_category_counts.csv", args.file_tag), index=False)

    html_report = build_html_report(
        args.title,
        details_df,
        b890_timeline_df,
        b890_sequence_df,
        b890_issues_df,
        b825_coverage_df,
        error_counts_df,
    )
    html_path = output_dir / build_output_name("icd_offline_online_compare.html", args.file_tag)
    html_path.write_text(html_report, encoding="utf-8")

    print(f"detail csv: {details_csv}")
    print(f"html report: {html_path}")


if __name__ == "__main__":
    main()
