#!/usr/bin/env python3
import argparse
import csv
import os
import re
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime
from typing import Dict, Iterable, List, Optional, Sequence, Tuple

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


TS_RE = re.compile(r"(\d{2}-\d{2}\s\d{2}:\d{2}:\d{2}\.\d{3})")
N1_REG_RE = re.compile(r"regState\s*[:=]\s*([A-Z_]+)")
N1_RAT_RE = re.compile(r"\brat\b\s*[:=]\s*([A-Z0-9_]+)")
XIAOMI_DATA_REG_RE = re.compile(r"mDataRegState=(\d+)\(([^)]+)\)")
XIAOMI_DATA_RAT_RE = re.compile(r"getRilDataRadioTechnology=\d+\(([^)]+)\)")
XIAOMI_PS_WWAN_MM_RE = re.compile(
    r"domain=PS transportType=WWAN.*?cellIdentity=.*?mMcc\s*=\s*([0-9-]+)\s*mMnc\s*=\s*([0-9-]+)"
)

INVALID_VALUES = {"", "-1", "2147483647", "9223372036854775807", "null", "NULL"}


@dataclass
class Event:
    ts: datetime
    normalized_state: str
    raw_state: str
    rat: str
    source_line: str


@dataclass
class Segment:
    device: str
    start: datetime
    end: datetime
    normalized_state: str
    raw_states: Tuple[str, ...]
    rats: Tuple[str, ...]
    observation_count: int

    @property
    def duration_s(self) -> float:
        return max((self.end - self.start).total_seconds(), 0.0)


@dataclass
class CellEvent:
    ts: datetime
    in_service: bool
    mcc: Optional[str]
    mnc: Optional[str]
    rat: Optional[str]
    band: Optional[str]
    freq: Optional[str]
    bandwidth: Optional[str]
    pci: Optional[str]
    rsrp: Optional[float]
    snr: Optional[float]
    cell_id: Optional[str]
    tac: Optional[str]
    source: str


@dataclass
class CellSegment:
    device: str
    start: datetime
    end: datetime
    mcc: Optional[str]
    mnc: Optional[str]
    rat: Optional[str]
    band: Optional[str]
    freq: Optional[str]
    bandwidth: Optional[str]
    pci: Optional[str]
    cell_id: Optional[str]
    tac: Optional[str]
    rsrp_avg: Optional[float]
    snr_avg: Optional[float]
    rsrp_last: Optional[float]
    snr_last: Optional[float]
    sample_count: int
    source: str

    @property
    def duration_s(self) -> float:
        return max((self.end - self.start).total_seconds(), 0.0)


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Compare slot1 data registration state and serving cell parameters across two AP logs."
    )
    parser.add_argument("--n1-input", required=True, help="Path to N1 reg_sim1_data_reg_rat.txt")
    parser.add_argument("--xiaomi-input", required=True, help="Path to Xiaomi reg_sim1.txt")
    parser.add_argument("--outdir", required=True, help="Directory for HTML and CSV outputs")
    parser.add_argument("--year", type=int, default=datetime.now().year, help="Year for log timestamps")
    parser.add_argument(
        "--window-end",
        help="Optional compare window end time, supports 'MM-DD HH:MM:SS[.mmm]' or full datetime.",
    )
    parser.add_argument(
        "--output-prefix",
        default="did_liyang_slot1_datareg_duration_compare",
        help="Output file prefix",
    )
    return parser.parse_args()


def parse_ts(line: str, year: int) -> Optional[datetime]:
    match = TS_RE.search(line)
    if not match:
        return None
    return datetime.strptime(f"{year}-{match.group(1)}", "%Y-%m-%d %H:%M:%S.%f")


def parse_bound_time(value: Optional[str], year: int) -> Optional[datetime]:
    if not value:
        return None
    for fmt in (
        "%Y-%m-%d %H:%M:%S.%f",
        "%Y-%m-%d %H:%M:%S",
        "%Y-%m-%d %H:%M",
        "%m-%d %H:%M:%S.%f",
        "%m-%d %H:%M:%S",
        "%m-%d %H:%M",
    ):
        try:
            parsed = datetime.strptime(value, fmt)
            if fmt.startswith("%m"):
                parsed = parsed.replace(year=year)
            return parsed
        except ValueError:
            continue
    raise SystemExit(f"Unsupported time format: {value}")


def normalize_n1_state(raw_state: str) -> str:
    if raw_state in {"REG_HOME", "REG_ROAMING"}:
        return "IN_SERVICE"
    return "OUT_OF_SERVICE"


def normalize_xiaomi_state(raw_code: str, raw_name: str) -> str:
    if raw_code == "0" or raw_name == "IN_SERVICE":
        return "IN_SERVICE"
    if raw_code == "1" or raw_name == "OUT_OF_SERVICE":
        return "OUT_OF_SERVICE"
    return raw_name or raw_code or "UNKNOWN"


def clean_value(value: Optional[str]) -> Optional[str]:
    if value is None:
        return None
    cleaned = value.strip().strip(",").strip("}")
    if cleaned in INVALID_VALUES:
        return None
    return cleaned


def clean_numeric_string(value: Optional[str]) -> Optional[str]:
    return clean_value(value)


def to_float(value: Optional[str]) -> Optional[float]:
    cleaned = clean_value(value)
    if cleaned is None:
        return None
    try:
        return float(cleaned)
    except ValueError:
        return None


def format_plmn_component(value: Optional[str], width: int) -> str:
    if value is None:
        return "NA"
    if value.isdigit():
        return value.zfill(width)
    return value


def normalize_rat(raw_rat: Optional[str]) -> Optional[str]:
    cleaned = clean_value(raw_rat)
    if cleaned is None:
        return None
    upper = cleaned.upper()
    mapping = {
        "EUTRAN": "LTE",
        "LTE": "LTE",
        "LTE_CA": "LTE",
        "NGRAN": "NR",
        "NR": "NR",
        "NR_SA": "NR",
        "UNKNOWN": "UNKNOWN",
    }
    return mapping.get(upper, upper)


def extract_field(line: str, field: str) -> Optional[str]:
    match = re.search(rf"{re.escape(field)}=([^,}}]+)", line)
    if not match:
        return None
    return clean_value(match.group(1))


def axis_refs_for_row(row_index: int, secondary_y: bool = False) -> Tuple[str, str]:
    xref_map = {
        1: "x",
        2: "x2",
        3: "x3",
        4: "x4",
        5: "x5",
        6: "x6",
    }
    yref_primary_map = {
        1: "y",
        2: "y2",
        3: "y3",
        4: "y4",
        5: "y5",
        6: "y7",
    }
    yref_secondary_map = {
        5: "y6",
        6: "y8",
    }
    xref = xref_map[row_index]
    yref = yref_secondary_map[row_index] if secondary_y else yref_primary_map[row_index]
    return xref, yref


def parse_n1_events(path: str, year: int) -> List[Event]:
    events: List[Event] = []
    with open(path, "r", encoding="utf-8", errors="ignore") as f:
        for line in f:
            if "DATA_REGISTRATION_STATE" not in line:
                continue
            ts = parse_ts(line, year)
            reg_match = N1_REG_RE.search(line)
            if ts is None or reg_match is None:
                continue
            raw_state = reg_match.group(1)
            rat_match = N1_RAT_RE.search(line)
            rat = normalize_rat(rat_match.group(1) if rat_match else None) or ""
            events.append(
                Event(
                    ts=ts,
                    normalized_state=normalize_n1_state(raw_state),
                    raw_state=raw_state,
                    rat=rat,
                    source_line=line.rstrip("\n"),
                )
            )
    return sorted(events, key=lambda item: item.ts)


def parse_xiaomi_events(path: str, year: int) -> List[Event]:
    events: List[Event] = []
    with open(path, "r", encoding="utf-8", errors="ignore") as f:
        for line in f:
            if "ServiceState from" not in line or "mDataRegState=" not in line:
                continue
            ts = parse_ts(line, year)
            reg_match = XIAOMI_DATA_REG_RE.search(line)
            if ts is None or reg_match is None:
                continue
            raw_code, raw_name = reg_match.groups()
            raw_state = f"{raw_code}({raw_name})"
            rat_match = XIAOMI_DATA_RAT_RE.search(line)
            events.append(
                Event(
                    ts=ts,
                    normalized_state=normalize_xiaomi_state(raw_code, raw_name),
                    raw_state=raw_state,
                    rat=normalize_rat(rat_match.group(1) if rat_match else None) or "",
                    source_line=line.rstrip("\n"),
                )
            )
    return sorted(events, key=lambda item: item.ts)


def parse_n1_cell_events(serv_cell_path: str, year: int) -> List[CellEvent]:
    events: List[CellEvent] = []
    if not os.path.exists(serv_cell_path):
        return events
    with open(serv_cell_path, "r", encoding="utf-8", errors="ignore") as f:
        for line in f:
            if "mServingCellInfo update to CellInfo{" not in line:
                continue
            ts = parse_ts(line, year)
            if ts is None:
                continue
            data_reg_state = extract_field(line, "dataRegState")
            rat_match = re.search(r"ratType=([^ ,]+)", line)
            events.append(
                CellEvent(
                    ts=ts,
                    in_service=(data_reg_state == "0"),
                    mcc=clean_numeric_string(extract_field(line, "mcc")),
                    mnc=clean_numeric_string(extract_field(line, "mnc")),
                    rat=normalize_rat(rat_match.group(1) if rat_match else None),
                    band=clean_numeric_string(extract_field(line, "band")),
                    freq=clean_numeric_string(extract_field(line, "freq")),
                    bandwidth=clean_numeric_string(extract_field(line, "bandwidth")),
                    pci=clean_numeric_string(extract_field(line, "pci")),
                    rsrp=to_float(extract_field(line, "rsrp")),
                    snr=to_float(extract_field(line, "snr")),
                    cell_id=clean_numeric_string(extract_field(line, "cellId")),
                    tac=clean_numeric_string(extract_field(line, "tac")),
                    source="N1 reg_data_sub_serv_cell.txt",
                )
            )
    return sorted(events, key=lambda item: item.ts)


def parse_xiaomi_cell_events(path: str, year: int) -> List[CellEvent]:
    events: List[CellEvent] = []
    with open(path, "r", encoding="utf-8", errors="ignore") as f:
        for line in f:
            if "ServiceState from" not in line or "mDataRegState=" not in line:
                continue
            ts = parse_ts(line, year)
            reg_match = XIAOMI_DATA_REG_RE.search(line)
            if ts is None or reg_match is None:
                continue
            raw_code, _ = reg_match.groups()

            rat_match = XIAOMI_DATA_RAT_RE.search(line)
            rat = normalize_rat(rat_match.group(1) if rat_match else None)

            channel_number = clean_numeric_string(extract_field(line, "mChannelNumber"))
            bandwidth = None
            bandwidth_match = re.search(r"mCellBandwidths=\[([^\]]*)\]", line)
            if bandwidth_match:
                items = [item.strip() for item in bandwidth_match.group(1).split(",") if item.strip()]
                bandwidth = clean_numeric_string(items[0] if items else None)
            if bandwidth is None:
                bandwidth = clean_numeric_string(extract_field(line, "mBandwidth"))

            mm_match = XIAOMI_PS_WWAN_MM_RE.search(line)
            mcc = clean_numeric_string(mm_match.group(1)) if mm_match else None
            mnc = clean_numeric_string(mm_match.group(2)) if mm_match else None

            events.append(
                CellEvent(
                    ts=ts,
                    in_service=(raw_code == "0"),
                    mcc=mcc,
                    mnc=mnc,
                    rat=rat,
                    band=None,
                    freq=channel_number,
                    bandwidth=bandwidth,
                    pci=None,
                    rsrp=None,
                    snr=None,
                    cell_id=None,
                    tac=None,
                    source="Xiaomi reg_sim1.txt ServiceState",
                )
            )
    return sorted(events, key=lambda item: item.ts)


def build_segments(events: Sequence[Event], device: str) -> List[Segment]:
    if not events:
        return []

    segments: List[Segment] = []
    start = events[0].ts
    current_state = events[0].normalized_state
    raw_states: List[str] = [events[0].raw_state]
    rats: List[str] = [events[0].rat] if events[0].rat else []
    observation_count = 1
    last_seen = events[0].ts

    for event in events[1:]:
        if event.normalized_state == current_state:
            last_seen = event.ts
            observation_count += 1
            if event.raw_state not in raw_states:
                raw_states.append(event.raw_state)
            if event.rat and event.rat not in rats:
                rats.append(event.rat)
            continue

        segments.append(
            Segment(
                device=device,
                start=start,
                end=event.ts,
                normalized_state=current_state,
                raw_states=tuple(raw_states),
                rats=tuple(rats),
                observation_count=observation_count,
            )
        )
        start = event.ts
        current_state = event.normalized_state
        raw_states = [event.raw_state]
        rats = [event.rat] if event.rat else []
        observation_count = 1
        last_seen = event.ts

    segments.append(
        Segment(
            device=device,
            start=start,
            end=last_seen,
            normalized_state=current_state,
            raw_states=tuple(raw_states),
            rats=tuple(rats),
            observation_count=observation_count,
        )
    )
    return segments


def clip_segments(segments: Iterable[Segment], window_start: datetime, window_end: datetime) -> List[Segment]:
    clipped: List[Segment] = []
    for segment in segments:
        start = max(segment.start, window_start)
        end = min(segment.end, window_end)
        if end <= start:
            continue
        clipped.append(
            Segment(
                device=segment.device,
                start=start,
                end=end,
                normalized_state=segment.normalized_state,
                raw_states=segment.raw_states,
                rats=segment.rats,
                observation_count=segment.observation_count,
            )
        )
    return clipped


def build_cell_segments(events: Sequence[CellEvent], device: str) -> List[CellSegment]:
    if not events:
        return []

    segments: List[CellSegment] = []
    current_start: Optional[datetime] = None
    current_key: Optional[Tuple[Optional[str], ...]] = None
    current_info: Dict[str, Optional[str]] = {}
    rsrp_values: List[float] = []
    snr_values: List[float] = []
    last_seen: Optional[datetime] = None
    sample_count = 0

    def flush(end_ts: Optional[datetime]) -> None:
        nonlocal current_start, current_key, current_info, rsrp_values, snr_values, last_seen, sample_count
        if current_start is None or current_key is None or end_ts is None or end_ts <= current_start:
            current_start = None
            current_key = None
            current_info = {}
            rsrp_values = []
            snr_values = []
            last_seen = None
            sample_count = 0
            return

        segments.append(
            CellSegment(
                device=device,
                start=current_start,
                end=end_ts,
                mcc=current_info.get("mcc"),
                mnc=current_info.get("mnc"),
                rat=current_info.get("rat"),
                band=current_info.get("band"),
                freq=current_info.get("freq"),
                bandwidth=current_info.get("bandwidth"),
                pci=current_info.get("pci"),
                cell_id=current_info.get("cell_id"),
                tac=current_info.get("tac"),
                rsrp_avg=(sum(rsrp_values) / len(rsrp_values)) if rsrp_values else None,
                snr_avg=(sum(snr_values) / len(snr_values)) if snr_values else None,
                rsrp_last=rsrp_values[-1] if rsrp_values else None,
                snr_last=snr_values[-1] if snr_values else None,
                sample_count=sample_count,
                source=current_info.get("source") or "",
            )
        )
        current_start = None
        current_key = None
        current_info = {}
        rsrp_values = []
        snr_values = []
        last_seen = None
        sample_count = 0

    for event in events:
        event_key = (
            event.mcc,
            event.mnc,
            event.rat,
            event.band,
            event.freq,
            event.bandwidth,
            event.pci,
            event.cell_id,
            event.tac,
        )

        if not event.in_service:
            flush(event.ts)
            continue

        if current_key is None:
            current_start = event.ts
            current_key = event_key
            current_info = {
                "mcc": event.mcc,
                "mnc": event.mnc,
                "rat": event.rat,
                "band": event.band,
                "freq": event.freq,
                "bandwidth": event.bandwidth,
                "pci": event.pci,
                "cell_id": event.cell_id,
                "tac": event.tac,
                "source": event.source,
            }
            rsrp_values = [event.rsrp] if event.rsrp is not None else []
            snr_values = [event.snr] if event.snr is not None else []
            last_seen = event.ts
            sample_count = 1
            continue

        if event_key == current_key:
            if event.rsrp is not None:
                rsrp_values.append(event.rsrp)
            if event.snr is not None:
                snr_values.append(event.snr)
            last_seen = event.ts
            sample_count += 1
            continue

        flush(event.ts)
        current_start = event.ts
        current_key = event_key
        current_info = {
            "mcc": event.mcc,
            "mnc": event.mnc,
            "rat": event.rat,
            "band": event.band,
            "freq": event.freq,
            "bandwidth": event.bandwidth,
            "pci": event.pci,
            "cell_id": event.cell_id,
            "tac": event.tac,
            "source": event.source,
        }
        rsrp_values = [event.rsrp] if event.rsrp is not None else []
        snr_values = [event.snr] if event.snr is not None else []
        last_seen = event.ts
        sample_count = 1

    flush(last_seen)
    return segments


def clip_cell_segments(
    segments: Iterable[CellSegment], window_start: datetime, window_end: datetime
) -> List[CellSegment]:
    clipped: List[CellSegment] = []
    for segment in segments:
        start = max(segment.start, window_start)
        end = min(segment.end, window_end)
        if end <= start:
            continue
        clipped.append(
            CellSegment(
                device=segment.device,
                start=start,
                end=end,
                mcc=segment.mcc,
                mnc=segment.mnc,
                rat=segment.rat,
                band=segment.band,
                freq=segment.freq,
                bandwidth=segment.bandwidth,
                pci=segment.pci,
                cell_id=segment.cell_id,
                tac=segment.tac,
                rsrp_avg=segment.rsrp_avg,
                snr_avg=segment.snr_avg,
                rsrp_last=segment.rsrp_last,
                snr_last=segment.snr_last,
                sample_count=segment.sample_count,
                source=segment.source,
            )
        )
    return clipped


def fmt_duration(seconds: float) -> str:
    total_ms = int(round(seconds * 1000))
    ms = total_ms % 1000
    total_seconds = total_ms // 1000
    hours = total_seconds // 3600
    minutes = (total_seconds % 3600) // 60
    secs = total_seconds % 60
    if hours:
        return f"{hours}h {minutes}m {secs}s"
    if minutes:
        return f"{minutes}m {secs}s"
    if secs:
        return f"{secs}.{ms:03d}s" if ms else f"{secs}s"
    return f"{total_ms}ms"


def fmt_optional_float(value: Optional[float]) -> str:
    if value is None:
        return "N/A"
    return f"{value:.1f}"


def to_int_value(value: Optional[str]) -> Optional[int]:
    cleaned = clean_value(value)
    if cleaned is None:
        return None
    try:
        return int(cleaned)
    except ValueError:
        return None


def should_annotate_text(
    text: str,
    duration_s: float,
    total_window_s: float,
    plot_width_px: int,
    min_duration_s: float = 120.0,
) -> bool:
    if duration_s < min_duration_s or total_window_s <= 0:
        return False
    estimated_label_width_px = len(text) * 7 + 18
    estimated_bar_width_px = (duration_s / total_window_s) * plot_width_px
    return estimated_bar_width_px >= estimated_label_width_px + 8


def make_cell_label(segment: CellSegment) -> str:
    parts = [
        f"{format_plmn_component(segment.mcc, 3)}-{format_plmn_component(segment.mnc, 2)}",
        segment.rat or "UNK",
    ]
    if segment.band:
        parts.append(f"B{segment.band}")
    if segment.freq:
        parts.append(f"F{segment.freq}")
    if segment.pci:
        parts.append(f"PCI{segment.pci}")
    return " ".join(parts)


def add_plmn_compare(
    fig: go.Figure,
    row_index: int,
    field_name: str,
    title_text: str,
    n1_events: Sequence[CellEvent],
    xiaomi_events: Sequence[CellEvent],
) -> None:
    device_specs = [
        ("N1", n1_events, "#1f77b4"),
        ("Xiaomi", xiaomi_events, "#d62728"),
    ]
    plotted_values: List[int] = []
    for device, events, color in device_specs:
        filtered_events = [event for event in events if to_int_value(getattr(event, field_name)) is not None]
        if not filtered_events:
            continue
        x_values = [event.ts for event in filtered_events]
        y_values = [to_int_value(getattr(event, field_name)) for event in filtered_events]
        plotted_values.extend(value for value in y_values if value is not None)
        fig.add_trace(
            go.Scatter(
                x=x_values,
                y=y_values,
                mode="lines+markers",
                line={"shape": "hv", "width": 1.8, "color": color},
                marker={"size": 4},
                name=f"{device} {title_text}",
                hovertemplate=(
                    f"Device: {device}<br>"
                    f"{title_text}: %{{y}}<br>"
                    "Time: %{x}<extra></extra>"
                ),
            ),
            row=row_index,
            col=1,
        )
    if plotted_values:
        y_min = min(plotted_values)
        y_max = max(plotted_values)
        pad = max((y_max - y_min) * 0.12, 1)
        fig.update_yaxes(
            row=row_index,
            col=1,
            title_text=title_text,
            range=[y_min - pad, y_max + pad],
        )
    else:
        fig.update_yaxes(row=row_index, col=1, title_text=title_text)


def write_segments_csv(path: str, segments: Sequence[Segment]) -> None:
    with open(path, "w", newline="", encoding="utf-8") as f:
        writer = csv.writer(f)
        writer.writerow(
            [
                "device",
                "start",
                "end",
                "duration_s",
                "duration_label",
                "normalized_state",
                "raw_states",
                "rats",
                "observation_count",
            ]
        )
        for segment in segments:
            writer.writerow(
                [
                    segment.device,
                    segment.start.isoformat(sep=" "),
                    segment.end.isoformat(sep=" "),
                    f"{segment.duration_s:.3f}",
                    fmt_duration(segment.duration_s),
                    segment.normalized_state,
                    ";".join(segment.raw_states),
                    ";".join(segment.rats),
                    segment.observation_count,
                ]
            )


def write_summary_csv(
    path: str, segments_by_device: Dict[str, Sequence[Segment]], window_start: datetime, window_end: datetime
) -> None:
    with open(path, "w", newline="", encoding="utf-8") as f:
        writer = csv.writer(f)
        writer.writerow(
            [
                "device",
                "window_start",
                "window_end",
                "state",
                "segment_count",
                "total_duration_s",
                "total_duration_label",
                "longest_segment_s",
                "longest_segment_label",
            ]
        )
        for device, segments in segments_by_device.items():
            grouped: Dict[str, List[Segment]] = defaultdict(list)
            for segment in segments:
                grouped[segment.normalized_state].append(segment)
            for state in sorted(grouped):
                state_segments = grouped[state]
                total_duration = sum(item.duration_s for item in state_segments)
                longest = max((item.duration_s for item in state_segments), default=0.0)
                writer.writerow(
                    [
                        device,
                        window_start.isoformat(sep=" "),
                        window_end.isoformat(sep=" "),
                        state,
                        len(state_segments),
                        f"{total_duration:.3f}",
                        fmt_duration(total_duration),
                        f"{longest:.3f}",
                        fmt_duration(longest),
                    ]
                )


def write_cell_segments_csv(path: str, segments: Sequence[CellSegment]) -> None:
    with open(path, "w", newline="", encoding="utf-8") as f:
        writer = csv.writer(f)
        writer.writerow(
            [
                "device",
                "start",
                "end",
                "duration_s",
                "duration_label",
                "mcc",
                "mnc",
                "rat",
                "band",
                "freq",
                "bandwidth",
                "pci",
                "cell_id",
                "tac",
                "rsrp_avg",
                "snr_avg",
                "rsrp_last",
                "snr_last",
                "sample_count",
                "source",
            ]
        )
        for segment in segments:
            writer.writerow(
                [
                    segment.device,
                    segment.start.isoformat(sep=" "),
                    segment.end.isoformat(sep=" "),
                    f"{segment.duration_s:.3f}",
                    fmt_duration(segment.duration_s),
                    segment.mcc or "",
                    segment.mnc or "",
                    segment.rat or "",
                    segment.band or "",
                    segment.freq or "",
                    segment.bandwidth or "",
                    segment.pci or "",
                    segment.cell_id or "",
                    segment.tac or "",
                    fmt_optional_float(segment.rsrp_avg),
                    fmt_optional_float(segment.snr_avg),
                    fmt_optional_float(segment.rsrp_last),
                    fmt_optional_float(segment.snr_last),
                    segment.sample_count,
                    segment.source,
                ]
            )


def add_state_segments(
    fig: go.Figure,
    row_index: int,
    segments: Sequence[Segment],
    total_window_s: float,
    plot_width_px: int,
    color_map: Dict[str, str],
) -> None:
    y_map = {"OUT_OF_SERVICE": 0, "IN_SERVICE": 1}
    for segment in segments:
        y_center = y_map.get(segment.normalized_state, 2)
        color = color_map.get(segment.normalized_state, "#7f7f7f")
        fig.add_shape(
            type="rect",
            x0=segment.start,
            x1=segment.end,
            y0=y_center - 0.35,
            y1=y_center + 0.35,
            xref=f"x{row_index}",
            yref=f"y{row_index}",
            line={"width": 0},
            fillcolor=color,
            opacity=0.9,
        )

        center = segment.start + (segment.end - segment.start) / 2
        fig.add_trace(
            go.Scatter(
                x=[center],
                y=[y_center],
                mode="markers",
                marker={"size": 10, "color": color, "opacity": 0.0},
                showlegend=False,
                customdata=[
                    [
                        segment.device,
                        segment.normalized_state,
                        segment.start.strftime("%m-%d %H:%M:%S.%f")[:-3],
                        segment.end.strftime("%m-%d %H:%M:%S.%f")[:-3],
                        f"{segment.duration_s:.3f}",
                        fmt_duration(segment.duration_s),
                        ", ".join(segment.raw_states),
                        ", ".join(segment.rats) if segment.rats else "-",
                        str(segment.observation_count),
                    ]
                ],
                hovertemplate=(
                    "Device: %{customdata[0]}<br>"
                    "dataRegState: %{customdata[1]}<br>"
                    "Start: %{customdata[2]}<br>"
                    "End: %{customdata[3]}<br>"
                    "Duration: %{customdata[5]} (%{customdata[4]} s)<br>"
                    "Raw states: %{customdata[6]}<br>"
                    "RATs: %{customdata[7]}<br>"
                    "Observations: %{customdata[8]}<extra></extra>"
                ),
            ),
            row=row_index,
            col=1,
        )

        label = fmt_duration(segment.duration_s)
        if should_annotate_text(label, segment.duration_s, total_window_s, plot_width_px):
            fig.add_annotation(
                x=center,
                y=y_center,
                xref=f"x{row_index}",
                yref=f"y{row_index}",
                text=label,
                showarrow=False,
                font={"size": 9, "color": "#ffffff"},
                bgcolor="rgba(0,0,0,0.15)",
            )

    fig.update_yaxes(
        row=row_index,
        col=1,
        tickmode="array",
        tickvals=[0, 1],
        ticktext=["OUT_OF_SERVICE", "IN_SERVICE"],
        range=[-0.5, 1.5],
        title_text="dataRegState",
    )


def add_cell_segments(
    fig: go.Figure,
    row_index: int,
    segments: Sequence[CellSegment],
    total_window_s: float,
    plot_width_px: int,
    rat_color_map: Dict[str, str],
    signal_events: Sequence[CellEvent],
    missing_fields_note: Optional[str] = None,
) -> None:
    xref, yref = axis_refs_for_row(row_index, secondary_y=False)
    for segment in segments:
        rat = segment.rat or "UNKNOWN"
        color = rat_color_map.get(rat, rat_color_map["UNKNOWN"])
        fig.add_shape(
            type="rect",
            x0=segment.start,
            x1=segment.end,
            y0=0.2,
            y1=0.8,
            xref=xref,
            yref=yref,
            line={"width": 0},
            fillcolor=color,
            opacity=0.88,
        )

        center = segment.start + (segment.end - segment.start) / 2
        fig.add_trace(
            go.Scatter(
                x=[center],
                y=[0.5],
                mode="markers",
                marker={"size": 10, "color": color, "opacity": 0.0},
                showlegend=False,
                customdata=[
                    [
                        segment.device,
                        segment.start.strftime("%m-%d %H:%M:%S.%f")[:-3],
                        segment.end.strftime("%m-%d %H:%M:%S.%f")[:-3],
                        fmt_duration(segment.duration_s),
                        format_plmn_component(segment.mcc, 3),
                        format_plmn_component(segment.mnc, 2),
                        segment.rat or "N/A",
                        segment.band or "N/A",
                        segment.freq or "N/A",
                        segment.bandwidth or "N/A",
                        segment.pci or "N/A",
                        segment.cell_id or "N/A",
                        segment.tac or "N/A",
                        fmt_optional_float(segment.rsrp_avg),
                        fmt_optional_float(segment.snr_avg),
                        segment.source,
                        str(segment.sample_count),
                    ]
                ],
                hovertemplate=(
                    "Device: %{customdata[0]}<br>"
                    "Start: %{customdata[1]}<br>"
                    "End: %{customdata[2]}<br>"
                    "Duration: %{customdata[3]}<br>"
                    "MCC/MNC: %{customdata[4]}-%{customdata[5]}<br>"
                    "RAT: %{customdata[6]}<br>"
                    "Band: %{customdata[7]}<br>"
                    "Freq: %{customdata[8]}<br>"
                    "Bandwidth: %{customdata[9]}<br>"
                    "PCI: %{customdata[10]}<br>"
                    "CellId: %{customdata[11]}<br>"
                    "TAC: %{customdata[12]}<br>"
                    "RSRP(avg): %{customdata[13]}<br>"
                    "SNR(avg): %{customdata[14]}<br>"
                    "Source: %{customdata[15]}<br>"
                    "Samples: %{customdata[16]}<extra></extra>"
                ),
            ),
            row=row_index,
            col=1,
            secondary_y=False,
        )

        label = make_cell_label(segment)
        if should_annotate_text(label, segment.duration_s, total_window_s, plot_width_px, min_duration_s=180.0):
            fig.add_annotation(
                x=center,
                y=0.5,
                xref=xref,
                yref=yref,
                text=label,
                showarrow=False,
                font={"size": 9, "color": "#ffffff"},
                bgcolor="rgba(0,0,0,0.18)",
            )

    fig.update_yaxes(
        row=row_index,
        col=1,
        secondary_y=False,
        range=[0, 1],
        tickvals=[0.5],
        ticktext=["servingCell"],
        title_text="servingCell",
    )

    signal_rsrp = [event for event in signal_events if event.in_service and event.rsrp is not None]
    signal_snr = [event for event in signal_events if event.in_service and event.snr is not None]
    has_signal = False

    if signal_rsrp:
        has_signal = True
        fig.add_trace(
            go.Scatter(
                x=[event.ts for event in signal_rsrp],
                y=[event.rsrp for event in signal_rsrp],
                mode="lines+markers",
                name=f"{segments[0].device} RSRP" if segments else "RSRP",
                line={"color": "#ff7f0e", "width": 1.6},
                marker={"size": 4},
                hovertemplate="RSRP: %{y:.1f} dBm<br>Time: %{x}<extra></extra>",
            ),
            row=row_index,
            col=1,
            secondary_y=True,
        )
    if signal_snr:
        has_signal = True
        fig.add_trace(
            go.Scatter(
                x=[event.ts for event in signal_snr],
                y=[event.snr for event in signal_snr],
                mode="lines+markers",
                name=f"{segments[0].device} SNR" if segments else "SNR",
                line={"color": "#8c564b", "width": 1.3, "dash": "dot"},
                marker={"size": 3},
                hovertemplate="SNR: %{y:.1f} dB<br>Time: %{x}<extra></extra>",
            ),
            row=row_index,
            col=1,
            secondary_y=True,
        )

    fig.update_yaxes(
        row=row_index,
        col=1,
        secondary_y=True,
        title_text="Signal" if has_signal else "",
        showgrid=False,
        showticklabels=has_signal,
    )

    if missing_fields_note:
        fig.add_annotation(
            x=1.0,
            y=1.08,
            xref=f"{xref} domain",
            yref=f"{yref} domain",
            text=missing_fields_note,
            showarrow=False,
            xanchor="right",
            font={"size": 10, "color": "#555555"},
            bgcolor="rgba(255,255,255,0.82)",
            bordercolor="rgba(0,0,0,0.15)",
            borderwidth=1,
        )


def build_html(
    html_path: str,
    state_segments_by_device: Dict[str, Sequence[Segment]],
    cell_segments_by_device: Dict[str, Sequence[CellSegment]],
    cell_events_by_device: Dict[str, Sequence[CellEvent]],
    window_start: datetime,
    window_end: datetime,
) -> None:
    total_window_s = max((window_end - window_start).total_seconds(), 1.0)
    plot_width_px = 1500
    fig = make_subplots(
        rows=6,
        cols=1,
        shared_xaxes=True,
        vertical_spacing=0.05,
        row_heights=[0.12, 0.12, 0.12, 0.12, 0.26, 0.26],
        subplot_titles=(
            "N1 slot1 dataRegState",
            "Xiaomi slot1 dataRegState",
            "MCC compare",
            "MNC compare",
            "N1 serving cell during IN_SERVICE",
            "Xiaomi serving cell during IN_SERVICE",
        ),
        specs=[
            [{}],
            [{}],
            [{}],
            [{}],
            [{"secondary_y": True}],
            [{"secondary_y": True}],
        ],
    )

    state_color_map = {
        "IN_SERVICE": "#2ca02c",
        "OUT_OF_SERVICE": "#d62728",
    }
    rat_color_map = {
        "LTE": "#1f77b4",
        "NR": "#9467bd",
        "UNKNOWN": "#7f7f7f",
    }

    add_state_segments(fig, 1, state_segments_by_device["N1"], total_window_s, plot_width_px, state_color_map)
    add_state_segments(fig, 2, state_segments_by_device["Xiaomi"], total_window_s, plot_width_px, state_color_map)
    add_plmn_compare(fig, 3, "mcc", "MCC", cell_events_by_device["N1"], cell_events_by_device["Xiaomi"])
    add_plmn_compare(fig, 4, "mnc", "MNC", cell_events_by_device["N1"], cell_events_by_device["Xiaomi"])
    add_cell_segments(
        fig,
        5,
        cell_segments_by_device["N1"],
        total_window_s,
        plot_width_px,
        rat_color_map,
        cell_events_by_device["N1"],
    )
    add_cell_segments(
        fig,
        6,
        cell_segments_by_device["Xiaomi"],
        total_window_s,
        plot_width_px,
        rat_color_map,
        cell_events_by_device["Xiaomi"],
        missing_fields_note="Current Xiaomi source does not expose PCI/Band/RSRP/SNR in this log.",
    )

    fig.add_trace(
        go.Scatter(
            x=[None],
            y=[None],
            mode="markers",
            marker={"size": 11, "color": state_color_map["IN_SERVICE"]},
            name="dataReg IN_SERVICE",
        ),
        row=1,
        col=1,
    )
    fig.add_trace(
        go.Scatter(
            x=[None],
            y=[None],
            mode="markers",
            marker={"size": 11, "color": state_color_map["OUT_OF_SERVICE"]},
            name="dataReg OUT_OF_SERVICE",
        ),
        row=1,
        col=1,
    )
    for rat_name in ["LTE", "NR", "UNKNOWN"]:
        fig.add_trace(
            go.Scatter(
                x=[None],
                y=[None],
                mode="markers",
                marker={"size": 11, "color": rat_color_map[rat_name]},
                name=f"servingCell {rat_name}",
            ),
            row=5,
            col=1,
        )

    n1_segments = state_segments_by_device["N1"]
    xiaomi_segments = state_segments_by_device["Xiaomi"]
    n1_oos = sum(item.duration_s for item in n1_segments if item.normalized_state == "OUT_OF_SERVICE")
    xiaomi_oos = sum(item.duration_s for item in xiaomi_segments if item.normalized_state == "OUT_OF_SERVICE")
    summary_lines = [
        f"Compare window: {window_start.strftime('%m-%d %H:%M:%S.%f')[:-3]} -> {window_end.strftime('%m-%d %H:%M:%S.%f')[:-3]}",
        f"N1 OUT_OF_SERVICE total: {fmt_duration(n1_oos)}",
        f"Xiaomi OUT_OF_SERVICE total: {fmt_duration(xiaomi_oos)}",
    ]

    fig.update_layout(
        title=(
            "did_liyang slot1 dataRegState + serving cell compare"
            f"<br><sup>{summary_lines[0]}</sup>"
            f"<br><sup>{summary_lines[1]} | {summary_lines[2]}</sup>"
        ),
        title_x=0.5,
        title_xanchor="center",
        width=1650,
        height=1880,
        template="plotly_white",
        hovermode="closest",
        legend={
            "orientation": "v",
            "x": 1.02,
            "xanchor": "left",
            "y": 1.0,
            "yanchor": "top",
        },
        margin={"l": 90, "r": 240, "t": 180, "b": 90},
    )

    for row_index in range(1, 7):
        fig.update_xaxes(
            range=[window_start, window_end],
            showgrid=True,
            gridcolor="rgba(0,0,0,0.08)",
            row=row_index,
            col=1,
        )
    fig.update_xaxes(title_text="Time", row=6, col=1)

    fig.write_html(html_path, include_plotlyjs="cdn")


def ensure_dir(path: str) -> None:
    os.makedirs(path, exist_ok=True)


def main() -> None:
    args = parse_args()
    ensure_dir(args.outdir)

    n1_events = parse_n1_events(args.n1_input, args.year)
    xiaomi_events = parse_xiaomi_events(args.xiaomi_input, args.year)
    if not n1_events:
        raise SystemExit(f"No N1 events parsed from: {args.n1_input}")
    if not xiaomi_events:
        raise SystemExit(f"No Xiaomi events parsed from: {args.xiaomi_input}")

    n1_segments_all = build_segments(n1_events, "N1")
    xiaomi_segments_all = build_segments(xiaomi_events, "Xiaomi")

    window_start = max(n1_events[0].ts, xiaomi_events[0].ts)
    explicit_window_end = parse_bound_time(args.window_end, args.year)
    window_end = explicit_window_end or min(n1_events[-1].ts, xiaomi_events[-1].ts)
    if window_end <= window_start:
        raise SystemExit("The two logs do not have an overlapping compare window.")

    n1_segments = clip_segments(n1_segments_all, window_start, window_end)
    xiaomi_segments = clip_segments(xiaomi_segments_all, window_start, window_end)

    n1_serv_cell_path = os.path.join(os.path.dirname(args.n1_input), "reg_data_sub_serv_cell.txt")
    n1_cell_events_all = parse_n1_cell_events(n1_serv_cell_path, args.year)
    xiaomi_cell_events_all = parse_xiaomi_cell_events(args.xiaomi_input, args.year)

    n1_cell_segments = clip_cell_segments(
        build_cell_segments(n1_cell_events_all, "N1"),
        window_start,
        window_end,
    )
    xiaomi_cell_segments = clip_cell_segments(
        build_cell_segments(xiaomi_cell_events_all, "Xiaomi"),
        window_start,
        window_end,
    )

    n1_cell_events = [event for event in n1_cell_events_all if window_start <= event.ts <= window_end]
    xiaomi_cell_events = [event for event in xiaomi_cell_events_all if window_start <= event.ts <= window_end]

    output_base = os.path.join(args.outdir, args.output_prefix)
    html_path = output_base + ".html"
    segments_csv_path = output_base + "_segments.csv"
    summary_csv_path = output_base + "_summary.csv"
    cell_segments_csv_path = output_base + "_serving_cell_segments.csv"

    merged_segments: List[Segment] = sorted(
        list(n1_segments) + list(xiaomi_segments),
        key=lambda item: (item.device, item.start),
    )
    merged_cell_segments: List[CellSegment] = sorted(
        list(n1_cell_segments) + list(xiaomi_cell_segments),
        key=lambda item: (item.device, item.start),
    )

    write_segments_csv(segments_csv_path, merged_segments)
    write_summary_csv(
        summary_csv_path,
        {"N1": n1_segments, "Xiaomi": xiaomi_segments},
        window_start,
        window_end,
    )
    write_cell_segments_csv(cell_segments_csv_path, merged_cell_segments)
    build_html(
        html_path,
        {"N1": n1_segments, "Xiaomi": xiaomi_segments},
        {"N1": n1_cell_segments, "Xiaomi": xiaomi_cell_segments},
        {"N1": n1_cell_events, "Xiaomi": xiaomi_cell_events},
        window_start,
        window_end,
    )

    print(f"HTML: {html_path}")
    print(f"Segments CSV: {segments_csv_path}")
    print(f"Summary CSV: {summary_csv_path}")
    print(f"Serving Cell CSV: {cell_segments_csv_path}")
    print(
        "Compare window: "
        f"{window_start.strftime('%m-%d %H:%M:%S.%f')[:-3]} -> "
        f"{window_end.strftime('%m-%d %H:%M:%S.%f')[:-3]}"
    )


if __name__ == "__main__":
    main()
