#!/usr/bin/env python3

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

import matplotlib.dates as mdates
import matplotlib.pyplot as plt
import plotly.graph_objects as go
from plotly.subplots import make_subplots


INVALID_FIRST_RATIO = -1.0
RAW_TS_RE = re.compile(r"\d{2}-\d{2} \d{2}:\d{2}:\d{2}\.\d{3}")
PHONE_RE = re.compile(r"\[PHONE(\d+)\]")


def extract_log_datetime(text):
    match = RAW_TS_RE.search(text)
    if not match:
        return None
    return datetime.strptime(match.group(0), "%m-%d %H:%M:%S.%f")


def parse_phone_id(text):
    match = PHONE_RE.search(text)
    if not match:
        return "unknown"
    return match.group(1)


def parse_int_field(text, field_name, default=0):
    match = re.search(rf"{re.escape(field_name)}=(-?\d+)", text)
    if not match:
        return default
    return int(match.group(1))


def parse_csv_int(row, key, default=0):
    value = row.get(key, "")
    if value in ("", None):
        return default
    try:
        return int(float(value))
    except (TypeError, ValueError):
        return default


def parse_csv_float(row, key, default=0.0):
    value = row.get(key, "")
    if value in ("", None):
        return default
    try:
        return float(value)
    except (TypeError, ValueError):
        return default


def parse_tx_array(text):
    match = re.search(r"mTxTimeMs\[\]=\[([^\]]+)\]", text)
    if not match:
        return [0, 0, 0, 0, 0]
    parts = [item.strip() for item in match.group(1).replace(";", ",").split(",")]
    values = []
    for item in parts:
        if not item:
            values.append(0)
            continue
        try:
            values.append(int(item))
        except ValueError:
            values.append(0)
    while len(values) < 5:
        values.append(0)
    return values[:5]


def build_record_base(log_dt, phone_id, m_timestamp, sleep_ms, idle_ms, rx_ms, tx_values):
    return {
        "log_datetime": log_dt,
        "log_timestamp": log_dt.strftime("%m-%d %H:%M:%S.%f")[:-3],
        "phone_id": phone_id,
        "mTimestamp": m_timestamp,
        "mSleepTimeMs": sleep_ms,
        "mIdleTimeMs": idle_ms,
        "mRxTimeMs": rx_ms,
        "mRat": "UNKNOWN",
        "mFrequencyRange": "UNKNOWN",
        "tx_time_0_5": tx_values[0],
        "tx_time_5_10": tx_values[1],
        "tx_time_10_15": tx_values[2],
        "tx_time_15_20": tx_values[3],
        "tx_time_20_25": tx_values[4],
    }


def parse_raw_modem_activity_line(line):
    if "ModemActivityInfo{" not in line:
        return None
    log_dt = extract_log_datetime(line)
    if log_dt is None:
        return None
    start = line.find("ModemActivityInfo{")
    end = line.find("}]} [PHONE")
    if start == -1 or end == -1 or end <= start:
        return None
    body = line[start + len("ModemActivityInfo{"):end]
    item = build_record_base(
        log_dt=log_dt,
        phone_id=parse_phone_id(line),
        m_timestamp=parse_int_field(body, "mTimestamp"),
        sleep_ms=parse_int_field(body, "mSleepTimeMs"),
        idle_ms=parse_int_field(body, "mIdleTimeMs"),
        rx_ms=parse_int_field(body, "mRxTimeMs"),
        tx_values=parse_tx_array(body),
    )
    rat_match = re.search(r"mRat=([^,\]]+)", body)
    if rat_match:
        item["mRat"] = rat_match.group(1).strip()
    freq_match = re.search(r"mFrequencyRange=([^,\]]+)", body)
    if freq_match:
        item["mFrequencyRange"] = freq_match.group(1).strip()
    return item


def parse_filtered_modem_activity_line(line):
    if "mTimestamp=" not in line or "[PHONE" not in line:
        return None
    log_dt = extract_log_datetime(line)
    if log_dt is None:
        return None
    item = build_record_base(
        log_dt=log_dt,
        phone_id=parse_phone_id(line),
        m_timestamp=parse_int_field(line, "mTimestamp"),
        sleep_ms=parse_int_field(line, "mSleepTimeMs"),
        idle_ms=parse_int_field(line, "mIdleTimeMs"),
        rx_ms=parse_int_field(line, "mRxTimeMs"),
        tx_values=[
            parse_int_field(line, "tx_time_0-5"),
            parse_int_field(line, "tx_time_5-10"),
            parse_int_field(line, "tx_time_10-15"),
            parse_int_field(line, "tx_time_15-20"),
            parse_int_field(line, "tx_time_20-25"),
        ],
    )
    rat_match = re.search(r"mRat=([^,\]]+)", line)
    if rat_match:
        item["mRat"] = rat_match.group(1).strip()
    freq_match = re.search(r"mFrequencyRange=([^,\]]+)", line)
    if freq_match:
        item["mFrequencyRange"] = freq_match.group(1).strip()
    return item


def parse_csv_records(input_path):
    records = []
    with open(input_path, "r", encoding="utf-8", errors="ignore", newline="") as csvfile:
        reader = csv.DictReader(csvfile)
        for row in reader:
            log_ts = row.get("log_timestamp", "")
            if not log_ts:
                continue
            try:
                log_dt = datetime.strptime(log_ts, "%m-%d %H:%M:%S.%f")
            except ValueError:
                continue
            item = build_record_base(
                log_dt=log_dt,
                phone_id=str(row.get("phone_id", "0") or "0"),
                m_timestamp=parse_csv_int(row, "mTimestamp"),
                sleep_ms=parse_csv_int(row, "mSleepTimeMs"),
                idle_ms=parse_csv_int(row, "mIdleTimeMs"),
                rx_ms=parse_csv_int(row, "mRxTimeMs"),
                tx_values=[
                    parse_csv_int(row, "tx_time_0-5"),
                    parse_csv_int(row, "tx_time_5-10"),
                    parse_csv_int(row, "tx_time_10-15"),
                    parse_csv_int(row, "tx_time_15-20"),
                    parse_csv_int(row, "tx_time_20-25"),
                ],
            )
            item["mRat"] = row.get("mRat", "UNKNOWN") or "UNKNOWN"
            item["mFrequencyRange"] = row.get("mFrequencyRange", "UNKNOWN") or "UNKNOWN"
            existing_ratio = parse_csv_float(row, "modemSleepRatio", INVALID_FIRST_RATIO)
            if existing_ratio != INVALID_FIRST_RATIO:
                item["existing_sleep_ratio"] = existing_ratio
            records.append(item)
    return records


def parse_text_records(input_path):
    records = []
    with open(input_path, "r", encoding="utf-8", errors="ignore") as handle:
        for line in handle:
            item = parse_raw_modem_activity_line(line)
            if item is None:
                item = parse_filtered_modem_activity_line(line)
            if item is not None:
                records.append(item)
    return records


def load_records(input_path):
    if input_path.lower().endswith(".csv"):
        return parse_csv_records(input_path)
    return parse_text_records(input_path)


def dedupe_records(records):
    deduped = []
    seen = set()
    for item in records:
        key = (
            item["phone_id"],
            item["mTimestamp"],
            item["mSleepTimeMs"],
            item["mIdleTimeMs"],
            item["mRxTimeMs"],
            item["tx_time_0_5"],
            item["tx_time_5_10"],
            item["tx_time_10_15"],
            item["tx_time_15_20"],
            item["tx_time_20_25"],
        )
        if key in seen:
            continue
        seen.add(key)
        deduped.append(item)
    return deduped


def add_derived_metrics(records):
    grouped = defaultdict(list)
    for item in records:
        grouped[item["phone_id"]].append(item)

    tx_power_midpoints = [2.5, 7.5, 12.5, 17.5, 22.5]
    results = []

    for phone_id, items in grouped.items():
        items.sort(key=lambda x: x["log_datetime"])
        prev_timestamp = None
        base_log_datetime = None
        base_modem_timestamp = None

        for item in items:
            current_timestamp = item["mTimestamp"]
            if prev_timestamp is None or current_timestamp < prev_timestamp:
                base_log_datetime = item["log_datetime"]
                base_modem_timestamp = current_timestamp

            offset_ms = current_timestamp - base_modem_timestamp
            aligned_log_datetime = base_log_datetime + timedelta(milliseconds=offset_ms)
            item["aligned_log_datetime"] = aligned_log_datetime
            item["aligned_log_timestamp"] = aligned_log_datetime.strftime("%m-%d %H:%M:%S.%f")[:-3]
            item["log_align_delta_ms"] = int((item["log_datetime"] - aligned_log_datetime).total_seconds() * 1000)

            total_tx = (
                item["tx_time_0_5"]
                + item["tx_time_5_10"]
                + item["tx_time_10_15"]
                + item["tx_time_15_20"]
                + item["tx_time_20_25"]
            )
            item["total_tx_ms"] = total_tx

            weighted_sum = (
                item["tx_time_0_5"] * tx_power_midpoints[0]
                + item["tx_time_5_10"] * tx_power_midpoints[1]
                + item["tx_time_10_15"] * tx_power_midpoints[2]
                + item["tx_time_15_20"] * tx_power_midpoints[3]
                + item["tx_time_20_25"] * tx_power_midpoints[4]
            )
            item["avg_tx_power_dbm"] = round(weighted_sum / total_tx, 2) if total_tx > 0 else 0.0

            if prev_timestamp is None or current_timestamp < prev_timestamp:
                item["duration_ms"] = 0
                item["sleep_ratio"] = INVALID_FIRST_RATIO
                item["idle_ratio"] = INVALID_FIRST_RATIO
                item["residual_ms"] = -1
                item["interval_start"] = aligned_log_datetime
                item["interval_end"] = aligned_log_datetime
            else:
                duration = current_timestamp - prev_timestamp
                item["duration_ms"] = duration
                item["interval_end"] = aligned_log_datetime
                item["interval_start"] = aligned_log_datetime - timedelta(milliseconds=duration)
                if duration > 0:
                    item["sleep_ratio"] = round(item["mSleepTimeMs"] / duration, 4)
                    item["idle_ratio"] = round(item["mIdleTimeMs"] / duration, 4)
                    known = total_tx + item["mRxTimeMs"] + item["mSleepTimeMs"] + item["mIdleTimeMs"]
                    item["residual_ms"] = duration - known
                else:
                    item["sleep_ratio"] = 0.0
                    item["idle_ratio"] = 0.0
                    item["residual_ms"] = 0

            prev_timestamp = current_timestamp
            results.append(item)

    results.sort(key=lambda x: (x["phone_id"], x["aligned_log_datetime"]))
    return results


def parse_time_arg(value):
    return datetime.strptime(value, "%m-%d %H:%M:%S")


def filter_records_by_time(records, start_time=None, end_time=None):
    if start_time is None and end_time is None:
        return records
    filtered = []
    for item in records:
        log_dt = item.get("aligned_log_datetime") or item.get("log_datetime")
        if log_dt is None:
            continue
        if start_time is not None and log_dt < start_time:
            continue
        if end_time is not None and log_dt > end_time:
            continue
        filtered.append(item)
    return filtered


def write_csv(records, output_csv):
    headers = [
        "log_timestamp",
        "aligned_log_timestamp",
        "log_align_delta_ms",
        "phone_id",
        "mTimestamp",
        "duration_ms",
        "mRat",
        "mFrequencyRange",
        "mSleepTimeMs",
        "mIdleTimeMs",
        "mRxTimeMs",
        "sleep_ratio",
        "idle_ratio",
        "residual_ms",
        "avg_tx_power_dbm",
        "total_tx_ms",
        "tx_time_0_5",
        "tx_time_5_10",
        "tx_time_10_15",
        "tx_time_15_20",
        "tx_time_20_25",
    ]
    with open(output_csv, "w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=headers)
        writer.writeheader()
        for item in records:
            writer.writerow({key: item.get(key, "") for key in headers})


def build_valid_intervals(records):
    valid = [item for item in records if item["sleep_ratio"] != INVALID_FIRST_RATIO and item["duration_ms"] > 0]
    for item in valid:
        item["interval_center"] = item["interval_start"] + timedelta(milliseconds=item["duration_ms"] / 2.0)
    return valid


def format_interval_hover(item):
    return (
        f"start={item['interval_start'].strftime('%m-%d %H:%M:%S.%f')[:-3]}<br>"
        f"end={item['interval_end'].strftime('%m-%d %H:%M:%S.%f')[:-3]}<br>"
        f"duration={item['duration_ms']} ms"
    )


def plot_phone_records(phone_id, records, output_png):
    valid = build_valid_intervals(records)
    if not valid:
        print(f"No valid plot data for PHONE{phone_id}")
        return

    centers = [item["interval_center"] for item in valid]
    widths_days = [item["duration_ms"] / 1000.0 / 86400.0 for item in valid]
    min_width_days = 0.001 / 86400.0
    widths_days = [max(width, min_width_days) for width in widths_days]

    sleep_ratios = [item["sleep_ratio"] for item in valid]
    idle_ratios = [item["idle_ratio"] for item in valid]
    rx_times = [item["mRxTimeMs"] for item in valid]
    residual_ms = [item["residual_ms"] for item in valid]
    avg_tx_power = [item["avg_tx_power_dbm"] for item in valid]

    tx_labels = [
        ("tx_time_0_5", "TX 0-5 dBm", "tab:blue"),
        ("tx_time_5_10", "TX 5-10 dBm", "tab:green"),
        ("tx_time_10_15", "TX 10-15 dBm", "tab:orange"),
        ("tx_time_15_20", "TX 15-20 dBm", "tab:red"),
        ("tx_time_20_25", "TX 20-25 dBm", "tab:purple"),
    ]

    fig, axes = plt.subplots(6, 1, figsize=(16, 22), sharex=True)
    fig.suptitle(f"Modem Activity Interval Analysis - PHONE{phone_id}", fontsize=16, y=0.995)

    axes[0].bar(centers, sleep_ratios, width=widths_days, align="center", color="tab:blue", alpha=0.9, edgecolor="black", linewidth=0.3)
    axes[0].set_ylabel("Sleep Ratio")
    axes[0].set_ylim(0, 1)
    axes[0].set_title("Sleep Ratio By Real Interval")
    axes[0].grid(True, alpha=0.3)

    axes[1].bar(centers, idle_ratios, width=widths_days, align="center", color="tab:cyan", alpha=0.9, edgecolor="black", linewidth=0.3)
    axes[1].set_ylabel("Idle Ratio")
    axes[1].set_ylim(0, 1)
    axes[1].set_title("Idle Ratio By Real Interval")
    axes[1].grid(True, alpha=0.3)

    for idx, (field_name, label, color) in enumerate(tx_labels):
        shifted_centers = []
        shifted_widths = []
        values = []
        for item in valid:
            single_width_days = max((item["duration_ms"] / 1000.0 / 86400.0) * 0.9 / len(tx_labels), min_width_days / len(tx_labels))
            offset_days = (idx - 2) * single_width_days
            shifted_centers.append(item["interval_center"] + timedelta(days=offset_days))
            shifted_widths.append(single_width_days * 1.1)
            values.append(item[field_name])
        axes[2].bar(shifted_centers, values, width=shifted_widths, align="center", color=color, alpha=0.9, edgecolor="black", linewidth=0.3, label=label)
    axes[2].set_ylabel("TX Time (ms)")
    axes[2].set_title("TX Power Distribution")
    axes[2].grid(True, alpha=0.3, axis="y")
    axes[2].legend(ncol=3, loc="upper left")

    axes[3].plot(centers, avg_tx_power, marker="^", color="tab:red", label="Average TX Power")
    axes[3].set_ylabel("Avg TX (dBm)")
    axes[3].set_ylim(0, 25)
    axes[3].grid(True, alpha=0.3)
    axes[3].legend()

    axes[4].plot(centers, rx_times, marker="s", color="tab:green", label="RX Time")
    axes[4].set_ylabel("RX (ms)")
    axes[4].grid(True, alpha=0.3)
    axes[4].legend()

    axes[5].plot(centers, residual_ms, marker="d", color="tab:brown", label="Residual Time")
    axes[5].set_ylabel("Residual (ms)")
    axes[5].set_xlabel("Time")
    axes[5].grid(True, alpha=0.3)
    axes[5].legend()
    axes[5].xaxis.set_major_locator(mdates.AutoDateLocator(minticks=6, maxticks=20))
    axes[5].xaxis.set_major_formatter(mdates.DateFormatter("%H:%M:%S"))
    plt.setp(axes[5].get_xticklabels(), rotation=45, ha="right")

    fig.tight_layout(rect=[0, 0, 1, 0.985])
    plt.savefig(output_png, dpi=200, bbox_inches="tight")
    plt.close(fig)


def plot_phone_records_html(phone_id, records, output_html):
    valid = build_valid_intervals(records)
    if not valid:
        print(f"No valid HTML plot data for PHONE{phone_id}")
        return

    centers = [item["interval_center"] for item in valid]
    widths_ms = [max(item["duration_ms"], 1) for item in valid]
    hover_text = [format_interval_hover(item) for item in valid]

    fig = make_subplots(
        rows=6,
        cols=1,
        shared_xaxes=True,
        vertical_spacing=0.04,
        subplot_titles=(
            "Sleep Ratio",
            "Idle Ratio",
            "TX Power Distribution",
            "Average TX Power",
            "RX Time",
            "Residual Time",
        ),
    )

    fig.add_trace(
        go.Bar(
            x=centers,
            y=[item["sleep_ratio"] for item in valid],
            width=widths_ms,
            name="Sleep Ratio",
            marker_color="#1f77b4",
            customdata=hover_text,
            hovertemplate="%{customdata}<br>sleep_ratio=%{y}<extra></extra>",
        ),
        row=1,
        col=1,
    )
    fig.add_trace(
        go.Bar(
            x=centers,
            y=[item["idle_ratio"] for item in valid],
            width=widths_ms,
            name="Idle Ratio",
            marker_color="#17becf",
            customdata=hover_text,
            hovertemplate="%{customdata}<br>idle_ratio=%{y}<extra></extra>",
        ),
        row=2,
        col=1,
    )

    tx_series = [
        ("tx_time_0_5", "TX 0-5 dBm", "#1f77b4"),
        ("tx_time_5_10", "TX 5-10 dBm", "#2ca02c"),
        ("tx_time_10_15", "TX 10-15 dBm", "#ff7f0e"),
        ("tx_time_15_20", "TX 15-20 dBm", "#d62728"),
        ("tx_time_20_25", "TX 20-25 dBm", "#9467bd"),
    ]
    for field_name, label, color in tx_series:
        fig.add_trace(
            go.Bar(
                x=centers,
                y=[item[field_name] for item in valid],
                width=[max(item["duration_ms"] * 0.16, 1) for item in valid],
                name=label,
                marker_color=color,
                customdata=hover_text,
                hovertemplate="%{customdata}<br>" + label + "=%{y} ms<extra></extra>",
            ),
            row=3,
            col=1,
        )

    fig.add_trace(
        go.Scatter(
            x=centers,
            y=[item["avg_tx_power_dbm"] for item in valid],
            mode="lines+markers",
            name="Average TX Power",
            line=dict(color="#d62728"),
            customdata=hover_text,
            hovertemplate="%{customdata}<br>avg_tx=%{y} dBm<extra></extra>",
        ),
        row=4,
        col=1,
    )
    fig.add_trace(
        go.Scatter(
            x=centers,
            y=[item["mRxTimeMs"] for item in valid],
            mode="lines+markers",
            name="RX Time",
            line=dict(color="#2ca02c"),
            customdata=hover_text,
            hovertemplate="%{customdata}<br>rx=%{y} ms<extra></extra>",
        ),
        row=5,
        col=1,
    )
    fig.add_trace(
        go.Scatter(
            x=centers,
            y=[item["residual_ms"] for item in valid],
            mode="lines+markers",
            name="Residual Time",
            line=dict(color="#8c564b"),
            customdata=hover_text,
            hovertemplate="%{customdata}<br>residual=%{y} ms<extra></extra>",
        ),
        row=6,
        col=1,
    )

    fig.update_yaxes(title_text="Ratio", range=[0, 1], row=1, col=1)
    fig.update_yaxes(title_text="Ratio", range=[0, 1], row=2, col=1)
    fig.update_yaxes(title_text="TX Time (ms)", row=3, col=1)
    fig.update_yaxes(title_text="dBm", row=4, col=1)
    fig.update_yaxes(title_text="ms", row=5, col=1)
    fig.update_yaxes(title_text="ms", row=6, col=1)

    fig.update_layout(
        title=f"Modem Activity Interval Analysis - PHONE{phone_id}",
        height=1600,
        barmode="group",
        hovermode="x unified",
        legend=dict(orientation="h", yanchor="bottom", y=1.01, xanchor="left", x=0),
    )
    fig.update_xaxes(
        title_text="Time",
        rangeslider=dict(visible=True),
        rangeselector=dict(
            buttons=[
                dict(count=10, label="10m", step="minute", stepmode="backward"),
                dict(count=30, label="30m", step="minute", stepmode="backward"),
                dict(count=1, label="1h", step="hour", stepmode="backward"),
                dict(step="all", label="All"),
            ]
        ),
        row=6,
        col=1,
    )
    fig.write_html(output_html, include_plotlyjs="cdn")


def main():
    parser = argparse.ArgumentParser(description="Analyze modem activity with real interval-aware plots.")
    parser.add_argument("--input", required=True, help="Path to modem_activity_info.csv or modem activity text")
    parser.add_argument("--outdir", default="", help="Output directory, default is input file directory")
    parser.add_argument("--keep-duplicates", action="store_true", help="Keep duplicate modem activity records")
    parser.add_argument("--start-time", default="", help="Start time, format: MM-DD HH:MM:SS")
    parser.add_argument("--end-time", default="", help="End time, format: MM-DD HH:MM:SS")
    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) if args.start_time.strip() else None
    end_time = parse_time_arg(args.end_time) if args.end_time.strip() else None

    records = load_records(input_path)
    if not records:
        print("No modem activity records found.")
        return 2

    if not args.keep_duplicates:
        before = len(records)
        records = dedupe_records(records)
        print(f"Deduped records: {before} -> {len(records)}")

    records = add_derived_metrics(records)
    filtered_records = filter_records_by_time(records, start_time, end_time)
    print(f"Time filtered records: {len(records)} -> {len(filtered_records)}")
    if not filtered_records:
        print("No modem activity records in selected time range.")
        return 3

    base_name = os.path.splitext(os.path.basename(input_path))[0]
    suffix = ""
    if start_time or end_time:
        start_str = start_time.strftime("%m%d_%H%M%S") if start_time else "begin"
        end_str = end_time.strftime("%m%d_%H%M%S") if end_time else "end"
        suffix = f"_{start_str}_{end_str}"

    csv_path = os.path.join(outdir, f"{base_name}_quick_analysis{suffix}.csv")
    write_csv(filtered_records, csv_path)
    print(f"CSV saved: {csv_path}")

    grouped = defaultdict(list)
    for item in filtered_records:
        grouped[item["phone_id"]].append(item)

    for phone_id, items in grouped.items():
        png_path = os.path.join(outdir, f"{base_name}_phone{phone_id}_quick_analysis{suffix}.png")
        html_path = os.path.join(outdir, f"{base_name}_phone{phone_id}_quick_analysis{suffix}.html")
        plot_phone_records(phone_id, items, png_path)
        plot_phone_records_html(phone_id, items, html_path)
        print(f"PNG saved: {png_path}")
        print(f"HTML saved: {html_path}")

    return 0


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