#!/usr/bin/env python3

import argparse
import csv
import os
import re
from collections import Counter
from datetime import datetime, timedelta

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


TS_RE = re.compile(r"\d{2}-\d{2} \d{2}:\d{2}:\d{2}(?:\.\d{3,6})?")
CELL_RE = re.compile(r"CellInfo\{(.*)\}")
INVALID_VALUES = {"", "-1", "2147483647", "9223372036854775807", "null", "NULL"}
CELL_FIELDS = (
    "inService",
    "isOos",
    "ratType",
    "pci",
    "freq",
    "bandwidth",
    "scs",
    "band",
    "tac",
    "rsrp",
    "snr",
    "rsrq",
    "networkType",
    "dataRegState",
    "cellId",
    "isCA",
    "isInBlackList",
    "roamingType",
    "mcc",
    "mnc",
    "channelNumber",
    "operatorNumeric",
)


def parse_log_datetime(line, year):
    match = TS_RE.search(line)
    if match is None:
        return None
    value = match.group(0)
    if "." in value:
        parsed = datetime.strptime(value, "%m-%d %H:%M:%S.%f")
    else:
        parsed = datetime.strptime(value, "%m-%d %H:%M:%S")
    return parsed.replace(year=year)


def parse_time_arg(value, year, is_end=False):
    value = value.strip()
    for fmt in ("%Y-%m-%d %H:%M:%S.%f", "%Y-%m-%d %H:%M:%S", "%m-%d %H:%M:%S.%f", "%m-%d %H:%M:%S"):
        try:
            parsed = datetime.strptime(value, fmt)
            if fmt.startswith("%m"):
                parsed = parsed.replace(year=year)
            if is_end and "." not in value:
                return parsed + timedelta(seconds=1) - timedelta(microseconds=1)
            return parsed
        except ValueError:
            continue
    raise SystemExit(f"Unsupported time format: {value}")


def in_window(dt, start_time, end_time):
    if dt is None:
        return False
    if start_time is not None and dt < start_time:
        return False
    if end_time is not None and dt > end_time:
        return False
    return True


def timestamp_text(dt):
    return dt.strftime("%m-%d %H:%M:%S.%f")[:-3]


def clean_value(value):
    if value is None:
        return ""
    cleaned = value.strip().strip(",").strip("}")
    return "" if cleaned in INVALID_VALUES else cleaned


def parse_int(value):
    cleaned = clean_value(value)
    if not cleaned:
        return None
    try:
        return int(cleaned)
    except ValueError:
        return None


def parse_bool(value):
    return 1 if str(value).strip().lower() == "true" else 0


def normalize_rat(raw_rat):
    rat = clean_value(raw_rat).upper()
    mapping = {
        "EUTRAN": "4G LTE",
        "LTE": "4G LTE",
        "NGRAN": "5G NR",
        "NR": "5G NR",
        "NR_SA": "5G NR",
        "NR_NSA": "5G NR",
    }
    return mapping.get(rat, rat or "UNKNOWN")


def rat_code(rat):
    if rat == "4G LTE":
        return 4
    if rat == "5G NR":
        return 5
    return 0


def parse_cell_fields(body):
    fields = {}
    for field in CELL_FIELDS:
        match = re.search(rf"(?<![A-Za-z0-9_]){re.escape(field)}=([^,\s}}]+)", body)
        if match:
            fields[field] = match.group(1)
    return fields


def parse_record(line, year):
    match = CELL_RE.search(line)
    if match is None:
        return None
    dt = parse_log_datetime(line, year)
    if dt is None:
        return None
    fields = parse_cell_fields(match.group(1))
    rat = normalize_rat(fields.get("ratType", ""))
    raw_freq = fields.get("freq", fields.get("channelNumber", ""))
    freq = parse_int(raw_freq)
    channel_number = parse_int(fields.get("channelNumber", raw_freq))
    item = {
        "dt": dt,
        "timestamp": timestamp_text(dt),
        "in_service": parse_bool(fields.get("inService", "false")),
        "is_oos": parse_bool(fields.get("isOos", "false")),
        "rat": rat,
        "rat_code": rat_code(rat),
        "raw_rat": clean_value(fields.get("ratType", "")),
        "pci": parse_int(fields.get("pci")),
        "freq": freq,
        "channel_number": channel_number,
        "band": parse_int(fields.get("band")),
        "rsrp": parse_int(fields.get("rsrp")),
        "snr": parse_int(fields.get("snr")),
        "rsrq": parse_int(fields.get("rsrq")),
        "bandwidth": parse_int(fields.get("bandwidth")),
        "scs": parse_int(fields.get("scs")),
        "tac": clean_value(fields.get("tac", "")),
        "cell_id": clean_value(fields.get("cellId", "")),
        "mcc": clean_value(fields.get("mcc", "")),
        "mnc": clean_value(fields.get("mnc", "")),
        "operator": clean_value(fields.get("operatorNumeric", "")),
        "data_reg_state": parse_int(fields.get("dataRegState")),
        "network_type": parse_int(fields.get("networkType")),
        "source_line": line.rstrip("\n"),
    }
    return item


def load_records(input_path, year, start_time=None, end_time=None):
    records = []
    seen = set()
    with open(input_path, "r", encoding="utf-8", errors="ignore") as handle:
        for line in handle:
            item = parse_record(line, year)
            if item is None or not in_window(item["dt"], start_time, end_time):
                continue
            key = (
                item["timestamp"],
                item["in_service"],
                item["is_oos"],
                item["rat"],
                item["pci"],
                item["freq"],
                item["band"],
                item["rsrp"],
                item["snr"],
                item["rsrq"],
                item["cell_id"],
            )
            if key in seen:
                continue
            seen.add(key)
            records.append(item)
    records.sort(key=lambda row: (row["dt"], row["rat"], row["pci"] if row["pci"] is not None else -1))
    return records


def format_nullable(value):
    return "" if value is None else value


def write_records_csv(records, output_csv):
    headers = [
        "timestamp",
        "in_service",
        "is_oos",
        "rat",
        "raw_rat",
        "pci",
        "freq",
        "channel_number",
        "band",
        "rsrp",
        "snr",
        "rsrq",
        "bandwidth",
        "scs",
        "tac",
        "cell_id",
        "mcc",
        "mnc",
        "operator",
        "data_reg_state",
        "network_type",
    ]
    with open(output_csv, "w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=headers)
        writer.writeheader()
        for row in records:
            writer.writerow({key: format_nullable(row.get(key)) for key in headers})


def count_by(records, field):
    counter = Counter(row[field] for row in records if row.get(field) not in (None, ""))
    return ", ".join(f"{key}:{value}" for key, value in counter.most_common()) or "N/A"


def summary_rows(records):
    if not records:
        return []
    window = f"{records[0]['timestamp']} ~ {records[-1]['timestamp']}"
    rows = [
        {
            "domain": "Input",
            "metric": "samples/window",
            "value": f"{len(records)} samples, {window}",
        },
        {
            "domain": "Service",
            "metric": "inService/isOos",
            "value": f"inService=1: {sum(row['in_service'] for row in records)}, isOos=1: {sum(row['is_oos'] for row in records)}",
        },
        {
            "domain": "RAT",
            "metric": "ratType",
            "value": count_by(records, "rat"),
        },
        {
            "domain": "Cell",
            "metric": "PCI",
            "value": count_by(records, "pci"),
        },
        {
            "domain": "Cell",
            "metric": "band",
            "value": count_by(records, "band"),
        },
        {
            "domain": "Cell",
            "metric": "freq/channelNumber",
            "value": count_by(records, "freq"),
        },
    ]
    return rows


def hover_text(row):
    return (
        f"{row['timestamp']}<br>"
        f"rat={row['rat']} raw={row['raw_rat']}<br>"
        f"inService={row['in_service']} isOos={row['is_oos']} dataRegState={format_nullable(row['data_reg_state'])}<br>"
        f"pci={format_nullable(row['pci'])} band={format_nullable(row['band'])} "
        f"freq={format_nullable(row['freq'])} channelNumber={format_nullable(row['channel_number'])}<br>"
        f"rsrp={format_nullable(row['rsrp'])} snr={format_nullable(row['snr'])} rsrq={format_nullable(row['rsrq'])}<br>"
        f"cellId={row['cell_id']} tac={row['tac']} operator={row['operator']}"
    )


def add_marker_trace(fig, records, y_field, name, row, color, legend="legend", marker_symbol="circle"):
    items = [item for item in records if item.get(y_field) is not None]
    if not items:
        return
    fig.add_trace(
        go.Scatter(
            x=[item["dt"] for item in items],
            y=[item[y_field] for item in items],
            mode="markers",
            name=name,
            legend=legend,
            marker=dict(color=color, size=7, opacity=0.86, symbol=marker_symbol),
            customdata=[hover_text(item) for item in items],
            hovertemplate="%{customdata}<extra></extra>",
        ),
        row=row,
        col=1,
    )


def interval_items(records, y_field):
    items = [item for item in records if item.get(y_field) is not None]
    if not items:
        return []
    if len(items) == 1:
        return [
            {
                "item": items[0],
                "center": items[0]["dt"],
                "width_ms": 1,
            }
        ]

    rows = []
    deltas = [
        max((items[index + 1]["dt"] - items[index]["dt"]).total_seconds() * 1000, 1)
        for index in range(len(items) - 1)
    ]
    fallback_width_ms = sorted(deltas)[len(deltas) // 2] if deltas else 1
    for index, item in enumerate(items):
        if index + 1 < len(items):
            width_ms = max((items[index + 1]["dt"] - item["dt"]).total_seconds() * 1000, 1)
        else:
            width_ms = fallback_width_ms
        center = item["dt"] + timedelta(milliseconds=width_ms / 2)
        rows.append(
            {
                "item": item,
                "center": center,
                "width_ms": width_ms,
            }
        )
    return rows


def add_interval_bar_trace(fig, records, y_field, name, row, color, legend="legend"):
    bars = interval_items(records, y_field)
    if not bars:
        return
    fig.add_trace(
        go.Bar(
            x=[bar["center"] for bar in bars],
            y=[bar["item"][y_field] for bar in bars],
            width=[bar["width_ms"] for bar in bars],
            name=name,
            legend=legend,
            marker_color=color,
            opacity=0.72,
            customdata=[hover_text(bar["item"]) for bar in bars],
            hovertemplate="%{customdata}<extra></extra>",
        ),
        row=row,
        col=1,
    )


def add_state_trace(fig, records, y_field, name, row, color, legend):
    fig.add_trace(
        go.Scatter(
            x=[item["dt"] for item in records],
            y=[item[y_field] for item in records],
            mode="lines+markers",
            name=name,
            legend=legend,
            line=dict(color=color, width=1.8, shape="hv"),
            marker=dict(color=color, size=6),
            customdata=[hover_text(item) for item in records],
            hovertemplate="%{customdata}<extra></extra>",
        ),
        row=row,
        col=1,
    )


def add_panel_regions(fig, panel_titles):
    shapes = []
    annotations = []
    for axis_index, panel_title in enumerate(panel_titles, start=1):
        yaxis_name = "yaxis" if axis_index == 1 else f"yaxis{axis_index}"
        y_domain = list(fig.layout[yaxis_name].domain)
        panel_y0 = max(y_domain[0] - 0.028, 0.0)
        panel_y1 = min(y_domain[1] + 0.032, 1.0)
        shapes.append(
            dict(
                type="rect",
                xref="paper",
                yref="paper",
                x0=0.0,
                x1=1.0,
                y0=panel_y0,
                y1=panel_y1,
                fillcolor="rgba(248, 250, 252, 0.72)",
                line=dict(color="rgba(148, 163, 184, 0.55)", width=1),
                layer="below",
            )
        )
        annotations.append(
            dict(
                x=0.012,
                y=panel_y1 - 0.006,
                xref="paper",
                yref="paper",
                text=f"<b>{panel_title}</b>",
                showarrow=False,
                xanchor="left",
                yanchor="top",
                align="left",
                font=dict(size=13, color="#263238"),
                bgcolor="rgba(248, 250, 252, 0.85)",
            )
        )
    fig.update_layout(shapes=shapes, annotations=annotations)


def build_plot(records, output_html, title, start_time=None, end_time=None):
    rows = summary_rows(records)
    panel_titles = (
        "inService / isOos",
        "RAT Type (4G LTE / 5G NR)",
        "RSRP by RAT",
        "SNR by RAT",
        "RSRQ by RAT",
        "PCI",
        "Band",
        "Frequency",
    )
    fig = make_subplots(
        rows=9,
        cols=1,
        shared_xaxes=False,
        vertical_spacing=0.05,
        row_heights=[0.16, 0.09, 0.09, 0.13, 0.13, 0.13, 0.09, 0.09, 0.10],
        specs=[[{"type": "table"}], [{}], [{}], [{}], [{}], [{}], [{}], [{}], [{}]],
    )
    fig.add_trace(
        go.Table(
            header=dict(
                values=["Domain", "Metric", "Value"],
                fill_color="#263238",
                font=dict(color="white", size=12),
                align="left",
            ),
            cells=dict(
                values=[
                    [row["domain"] for row in rows],
                    [row["metric"] for row in rows],
                    [row["value"] for row in rows],
                ],
                fill_color="#f7f9fb",
                font=dict(size=11),
                align="left",
                height=26,
            ),
        ),
        row=1,
        col=1,
    )

    add_state_trace(fig, records, "in_service", "inService", 2, "#2ca02c", "legend")
    add_state_trace(fig, records, "is_oos", "isOos", 2, "#d62728", "legend")
    add_interval_bar_trace(fig, records, "rat_code", "RAT code", 3, "#1f77b4", legend="legend2")

    rat_colors = {"4G LTE": "#ff7f0e", "5G NR": "#1f77b4", "UNKNOWN": "#7f7f7f"}
    lte_records = [item for item in records if item["rat"] == "4G LTE"]
    nr_records = [item for item in records if item["rat"] == "5G NR"]
    add_marker_trace(fig, lte_records, "rsrp", "4G LTE RSRP", 4, rat_colors["4G LTE"], legend="legend3", marker_symbol="circle")
    add_marker_trace(fig, nr_records, "rsrp", "5G NR RSRP", 4, rat_colors["5G NR"], legend="legend3", marker_symbol="diamond")
    add_marker_trace(fig, lte_records, "snr", "4G LTE SNR", 5, rat_colors["4G LTE"], legend="legend4", marker_symbol="circle")
    add_marker_trace(fig, nr_records, "snr", "5G NR SNR", 5, rat_colors["5G NR"], legend="legend4", marker_symbol="diamond")
    add_marker_trace(fig, lte_records, "rsrq", "4G LTE RSRQ", 6, rat_colors["4G LTE"], legend="legend5", marker_symbol="circle")
    add_marker_trace(fig, nr_records, "rsrq", "5G NR RSRQ", 6, rat_colors["5G NR"], legend="legend5", marker_symbol="diamond")

    add_interval_bar_trace(fig, records, "pci", "PCI", 7, "#9467bd", legend="legend6")
    add_interval_bar_trace(fig, records, "band", "Band", 8, "#8c564b", legend="legend7")
    add_interval_bar_trace(fig, records, "freq", "freq", 9, "#17becf", legend="legend8")

    fig.update_yaxes(title_text="0/1", range=[-0.1, 1.1], row=2, col=1)
    fig.update_yaxes(title_text="RAT", tickmode="array", tickvals=[0, 4, 5], ticktext=["UNKNOWN", "4G", "5G"], row=3, col=1)
    fig.update_yaxes(title_text="dBm", row=4, col=1)
    fig.update_yaxes(title_text="dB", row=5, col=1)
    fig.update_yaxes(title_text="dB", row=6, col=1)
    fig.update_yaxes(title_text="PCI", row=7, col=1)
    fig.update_yaxes(title_text="Band", row=8, col=1)
    fig.update_yaxes(title_text="ARFCN/EARFCN", row=9, col=1)

    x_start = start_time or records[0]["dt"]
    x_end = end_time or records[-1]["dt"]
    for row_index in range(2, 10):
        fig.update_xaxes(
            type="date",
            range=[x_start, x_end],
            domain=[0.06, 0.78],
            ticks="inside",
            ticklabelposition="inside",
            row=row_index,
            col=1,
        )
    for row_index in range(2, 9):
        fig.update_xaxes(matches="x8", row=row_index, col=1)
    fig.update_xaxes(
        rangeslider=dict(visible=True),
        row=9,
        col=1,
    )
    add_panel_regions(fig, panel_titles)

    legend_style = dict(
        x=0.795,
        xanchor="left",
        orientation="v",
        bgcolor="rgba(255,255,255,0.75)",
        bordercolor="rgba(180,180,180,0.4)",
        borderwidth=1,
        font=dict(size=10),
    )
    fig.update_layout(
        title=title,
        height=2300,
        hovermode="x unified",
        template="plotly_white",
        barmode="overlay",
        margin=dict(l=85, r=55, t=95, b=90),
        legend=dict(**legend_style, y=0.802, yanchor="top", title_text="Service"),
        legend2=dict(**legend_style, y=0.704, yanchor="top", title_text="RAT"),
        legend3=dict(**legend_style, y=0.606, yanchor="top", title_text="RSRP"),
        legend4=dict(**legend_style, y=0.465, yanchor="top", title_text="SNR"),
        legend5=dict(**legend_style, y=0.323, yanchor="top", title_text="RSRQ"),
        legend6=dict(**legend_style, y=0.189, yanchor="top", title_text="PCI"),
        legend7=dict(**legend_style, y=0.097, yanchor="top", title_text="Band"),
        legend8=dict(**legend_style, y=0.025, yanchor="top", title_text="Freq"),
    )
    fig.write_html(output_html, include_plotlyjs="cdn")


def safe_suffix(start_time=None, end_time=None):
    if start_time is None and end_time is None:
        return ""
    start_text = start_time.strftime("%m%d_%H%M%S") if start_time else "begin"
    end_text = end_time.strftime("%m%d_%H%M%S") if end_time else "end"
    return f"_{start_text}_{end_text}"


def main():
    parser = argparse.ArgumentParser(description="Plot NetworkBrain serving cell timeline from reg_data_sub_serv_cell.txt.")
    parser.add_argument("--input", required=True, help="Path to reg_data_sub_serv_cell.txt")
    parser.add_argument("--outdir", default="", help="Output directory, default is input file directory")
    parser.add_argument("--start-time", default="", help="Start time, format: MM-DD HH:MM:SS[.mmm]")
    parser.add_argument("--end-time", default="", help="End time, format: MM-DD HH:MM:SS[.mmm]")
    parser.add_argument("--year", type=int, default=datetime.now().year, help="Year for AP log timestamps")
    args = parser.parse_args()

    input_path = os.path.abspath(args.input)
    if not os.path.isfile(input_path):
        print(f"Input file not found: {input_path}")
        return 1

    outdir = args.outdir.strip() if args.outdir.strip() else os.path.dirname(input_path)
    os.makedirs(outdir, exist_ok=True)

    start_time = parse_time_arg(args.start_time, args.year) if args.start_time.strip() else None
    end_time = parse_time_arg(args.end_time, args.year, is_end=True) if args.end_time.strip() else None
    records = load_records(input_path, args.year, start_time=start_time, end_time=end_time)
    if not records:
        print("No serving cell records found.")
        return 2

    base_name = os.path.splitext(os.path.basename(input_path))[0]
    suffix = safe_suffix(start_time, end_time)
    output_csv = os.path.join(outdir, f"{base_name}_serving_cell_timeline{suffix}.csv")
    output_html = os.path.join(outdir, f"{base_name}_serving_cell_timeline{suffix}.html")
    write_records_csv(records, output_csv)
    title = f"{base_name} NetworkBrain Serving Cell Timeline"
    build_plot(records, output_html, title, start_time=start_time, end_time=end_time)
    print(f"Serving cell CSV saved: {output_csv}")
    print(f"Serving cell HTML saved: {output_html}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
