#!/usr/bin/env python3

import argparse
import csv
import os
import re
from collections import defaultdict
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}")
BLOCK_RE = re.compile(r"L4NetEvaluate>?\s+(?:netifStats|uidStats):")
STREAM_FIELDS = ("uid", "stat_src_type", "if_name", "app_package_name")


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


def parse_time_arg(value, is_end=False):
    value = value.strip()
    if "." in value:
        return datetime.strptime(value, "%m-%d %H:%M:%S.%f")
    parsed = datetime.strptime(value, "%m-%d %H:%M:%S")
    if is_end:
        return parsed + timedelta(seconds=1) - timedelta(microseconds=1)
    return parsed


def parse_str_field(text, field_name, default=""):
    match = re.search(rf"{re.escape(field_name)}:([^\s]+)", text)
    return match.group(1) if match else default


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


def parse_single_block(log_dt, block_text):
    uid_match = re.search(r"\buid:(\d+)\b", block_text)
    if uid_match is None:
        return None

    send_rate_bps = parse_int_field(block_text, "sendRate(Bps)", 0)
    recv_rate_bps = parse_int_field(block_text, "recvRate(Bps)", 0)

    return {
        "log_datetime": log_dt,
        "log_timestamp": log_dt.strftime("%m-%d %H:%M:%S.%f")[:-3],
        "uid": uid_match.group(1),
        "stat_src_type": parse_str_field(block_text, "statSrcType"),
        "app_package_name": parse_str_field(block_text, "appPackageName"),
        "if_name": parse_str_field(block_text, "ifName"),
        "stream_num": parse_int_field(block_text, "streamNum", 0),
        "rtt_ave": parse_int_field(block_text, "rttAve", 0),
        "total_send_bytes": parse_int_field(block_text, "totalSendBytes", 0),
        "total_recv_bytes": parse_int_field(block_text, "totalRecvBytes", 0),
        "send_packets": parse_int_field(block_text, "sendPackets", 0),
        "recv_packets": parse_int_field(block_text, "recvPackets", 0),
        "send_rate_bps": send_rate_bps,
        "recv_rate_bps": recv_rate_bps,
        "send_rate_mbps": send_rate_bps * 8 / 1_000_000,
        "recv_rate_mbps": recv_rate_bps * 8 / 1_000_000,
    }


def parse_line_records(line):
    log_dt = parse_log_datetime(line)
    if log_dt is None:
        return []

    matches = list(BLOCK_RE.finditer(line))
    records = []
    for index, match in enumerate(matches):
        start = match.end()
        end = matches[index + 1].start() if index + 1 < len(matches) else len(line)
        block_text = line[start:end].strip()
        item = parse_single_block(log_dt, block_text)
        if item is not None:
            records.append(item)
    return records


def load_records(input_path):
    records = []
    seen = set()
    with open(input_path, "r", encoding="utf-8", errors="ignore") as handle:
        for line in handle:
            for item in parse_line_records(line):
                key = (
                    item["log_timestamp"],
                    item["uid"],
                    item["stat_src_type"],
                    item["if_name"],
                    item["total_send_bytes"],
                    item["total_recv_bytes"],
                    item["send_rate_bps"],
                    item["recv_rate_bps"],
                )
                if key in seen:
                    continue
                seen.add(key)
                records.append(item)
    records.sort(key=lambda item: (item["log_datetime"], item["uid"], item["stat_src_type"], item["if_name"]))
    return records


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:
        if start_time is not None and item["log_datetime"] < start_time:
            continue
        if end_time is not None and item["log_datetime"] > end_time:
            continue
        filtered.append(item)
    return filtered


def stream_key(item):
    return tuple(item.get(field, "") for field in STREAM_FIELDS)


def stream_id_from_key(key):
    uid, stat_src_type, if_name, app_package_name = key
    parts = [f"uid{uid}"]
    if stat_src_type:
        parts.append(stat_src_type)
    if if_name:
        parts.append(if_name)
    if app_package_name:
        parts.append(app_package_name)
    return "_".join(parts)


def stream_label_from_key(key):
    uid, stat_src_type, if_name, app_package_name = key
    label_parts = [f"uid={uid}"]
    if stat_src_type:
        label_parts.append(f"src={stat_src_type}")
    if if_name:
        label_parts.append(f"if={if_name}")
    if app_package_name:
        label_parts.append(f"app={app_package_name}")
    return ", ".join(label_parts)


def group_records_by_uid(records):
    records_by_uid = defaultdict(list)
    for item in records:
        records_by_uid[item["uid"]].append(item)
    return records_by_uid


def group_records_by_stream(records):
    records_by_stream = defaultdict(list)
    for item in records:
        records_by_stream[stream_key(item)].append(item)
    return records_by_stream


def compute_stream_summary(records):
    summary_rows = []
    for key, stream_records in group_records_by_stream(records).items():
        uid, stat_src_type, if_name, app_package_name = key
        items = sorted(stream_records, key=lambda item: item["log_datetime"])
        send_bytes = 0
        recv_bytes = 0
        prev_send = None
        prev_recv = None
        max_send_bps = 0
        max_recv_bps = 0
        for item in items:
            current_send = item["total_send_bytes"]
            current_recv = item["total_recv_bytes"]
            if prev_send is not None and current_send >= prev_send:
                send_bytes += current_send - prev_send
            if prev_recv is not None and current_recv >= prev_recv:
                recv_bytes += current_recv - prev_recv
            prev_send = current_send
            prev_recv = current_recv
            max_send_bps = max(max_send_bps, item["send_rate_bps"])
            max_recv_bps = max(max_recv_bps, item["recv_rate_bps"])

        first_dt = items[0]["log_datetime"]
        last_dt = items[-1]["log_datetime"]
        duration_seconds = max((last_dt - first_dt).total_seconds(), 0.0)
        summary_rows.append(
            {
            "uid": uid,
            "stream_id": stream_id_from_key(key),
            "stream_label": stream_label_from_key(key),
            "app_package_name": app_package_name,
            "stat_src_type": stat_src_type,
            "if_name": if_name,
            "sample_count": len(items),
            "first_timestamp": first_dt.strftime("%m-%d %H:%M:%S.%f")[:-3],
            "last_timestamp": last_dt.strftime("%m-%d %H:%M:%S.%f")[:-3],
            "duration_seconds": round(duration_seconds, 3),
            "send_mb": send_bytes / 1_000_000,
            "recv_mb": recv_bytes / 1_000_000,
            "total_mb": (send_bytes + recv_bytes) / 1_000_000,
            "max_send_mbps": max_send_bps * 8 / 1_000_000,
            "max_recv_mbps": max_recv_bps * 8 / 1_000_000,
            "avg_send_mbps": send_bytes * 8 / duration_seconds / 1_000_000 if duration_seconds > 0 else 0.0,
            "avg_recv_mbps": recv_bytes * 8 / duration_seconds / 1_000_000 if duration_seconds > 0 else 0.0,
            }
        )
    summary_rows.sort(key=lambda row: (int(row["uid"]), row["stat_src_type"], row["if_name"], row["app_package_name"]))
    return summary_rows


def build_uid_timeseries(uid_records):
    rows = []
    for key, stream_records in group_records_by_stream(uid_records).items():
        cumulative_send = 0.0
        cumulative_recv = 0.0
        prev_send = None
        prev_recv = None
        for item in sorted(stream_records, key=lambda record: record["log_datetime"]):
            send_delta = 0
            recv_delta = 0
            if prev_send is not None and item["total_send_bytes"] >= prev_send:
                send_delta = item["total_send_bytes"] - prev_send
            if prev_recv is not None and item["total_recv_bytes"] >= prev_recv:
                recv_delta = item["total_recv_bytes"] - prev_recv
            prev_send = item["total_send_bytes"]
            prev_recv = item["total_recv_bytes"]
            cumulative_send += send_delta
            cumulative_recv += recv_delta
            row = dict(item)
            row["stream_id"] = stream_id_from_key(key)
            row["stream_label"] = stream_label_from_key(key)
            row["cumulative_send_mb"] = cumulative_send / 1_000_000
            row["cumulative_recv_mb"] = cumulative_recv / 1_000_000
            row["cumulative_total_mb"] = (cumulative_send + cumulative_recv) / 1_000_000
            rows.append(row)
    rows.sort(key=lambda item: (item["log_datetime"], item["stream_id"]))
    return rows


def write_summary_csv(summary, output_csv):
    headers = [
        "uid",
        "stream_id",
        "stream_label",
        "app_package_name",
        "stat_src_type",
        "if_name",
        "sample_count",
        "first_timestamp",
        "last_timestamp",
        "duration_seconds",
        "send_mb",
        "recv_mb",
        "total_mb",
        "max_send_mbps",
        "max_recv_mbps",
        "avg_send_mbps",
        "avg_recv_mbps",
    ]
    with open(output_csv, "w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=headers)
        writer.writeheader()
        for row in summary:
            writer.writerow({key: row.get(key, "") for key in headers})


def write_uid_records_csv(records, output_csv):
    headers = [
        "log_timestamp",
        "uid",
        "stream_id",
        "stream_label",
        "app_package_name",
        "stat_src_type",
        "if_name",
        "stream_num",
        "rtt_ave",
        "total_send_bytes",
        "total_recv_bytes",
        "send_rate_bps",
        "recv_rate_bps",
        "send_rate_mbps",
        "recv_rate_mbps",
        "cumulative_send_mb",
        "cumulative_recv_mb",
        "cumulative_total_mb",
    ]
    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 format_number(value, digits=3):
    return f"{value:.{digits}f}"


def plot_uid_html(uid, summary_rows, timeseries, output_html):
    fig = make_subplots(
        rows=3,
        cols=1,
        shared_xaxes=True,
        vertical_spacing=0.06,
        row_heights=[0.28, 0.36, 0.36],
        specs=[[{"type": "table"}], [{"type": "xy"}], [{"type": "xy"}]],
        subplot_titles=(
            "Summary",
            "Uplink / Downlink Rate",
            "Cumulative Traffic",
        ),
    )

    table_rows = sorted(summary_rows, key=lambda row: (row["stat_src_type"], row["if_name"], row["app_package_name"]))
    fig.add_trace(
        go.Table(
            header=dict(
                values=[
                    "Stream",
                    "Samples",
                    "Window",
                    "UL MB",
                    "DL MB",
                    "Total MB",
                    "Max UL Mbps",
                    "Max DL Mbps",
                    "Avg UL Mbps",
                    "Avg DL Mbps",
                ],
                fill_color="#263238",
                font=dict(color="white", size=12),
                align="left",
            ),
            cells=dict(
                values=[
                    [row["stream_label"] for row in table_rows],
                    [row["sample_count"] for row in table_rows],
                    [f"{row['first_timestamp']} ~ {row['last_timestamp']}" for row in table_rows],
                    [format_number(row["send_mb"]) for row in table_rows],
                    [format_number(row["recv_mb"]) for row in table_rows],
                    [format_number(row["total_mb"]) for row in table_rows],
                    [format_number(row["max_send_mbps"]) for row in table_rows],
                    [format_number(row["max_recv_mbps"]) for row in table_rows],
                    [format_number(row["avg_send_mbps"]) for row in table_rows],
                    [format_number(row["avg_recv_mbps"]) for row in table_rows],
                ],
                fill_color="#f7f9fb",
                font=dict(size=11),
                align="left",
                height=26,
            ),
        ),
        row=1,
        col=1,
    )

    colors = ["#d62728", "#1f77b4", "#2ca02c", "#ff7f0e", "#9467bd", "#8c564b", "#17becf"]
    for index, (key, stream_records) in enumerate(group_records_by_stream(timeseries).items()):
        label = stream_label_from_key(key)
        x_values = [item["log_datetime"] for item in stream_records]
        color = colors[index % len(colors)]
        fig.add_trace(
            go.Scatter(
                x=x_values,
                y=[item["send_rate_mbps"] for item in stream_records],
                mode="lines",
                name=f"{label} UL Mbps",
                line=dict(color=color, width=1.6),
            ),
            row=2,
            col=1,
        )
        fig.add_trace(
            go.Scatter(
                x=x_values,
                y=[item["recv_rate_mbps"] for item in stream_records],
                mode="lines",
                name=f"{label} DL Mbps",
                line=dict(color=color, width=1.6, dash="dot"),
            ),
            row=2,
            col=1,
        )
        fig.add_trace(
            go.Scatter(
                x=x_values,
                y=[item["cumulative_send_mb"] for item in stream_records],
                mode="lines",
                name=f"{label} UL MB",
                line=dict(color=color, width=1.6),
            ),
            row=3,
            col=1,
        )
        fig.add_trace(
            go.Scatter(
                x=x_values,
                y=[item["cumulative_recv_mb"] for item in stream_records],
                mode="lines",
                name=f"{label} DL MB",
                line=dict(color=color, width=1.6, dash="dot"),
            ),
            row=3,
            col=1,
        )

    title = f"UID {uid} Throughput / Traffic"
    fig.update_layout(
        title=title,
        height=max(950, 760 + len(table_rows) * 32),
        hovermode="x unified",
        legend=dict(orientation="h", yanchor="bottom", y=-0.18, xanchor="left", x=0),
        margin=dict(l=70, r=40, t=95, b=150),
    )
    fig.update_yaxes(title_text="Mbps", row=2, col=1)
    fig.update_yaxes(title_text="MB", row=3, col=1)
    fig.update_xaxes(
        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=3,
        col=1,
    )
    fig.write_html(output_html, include_plotlyjs="cdn")


def safe_name(value):
    return re.sub(r"[^A-Za-z0-9_.-]+", "_", value)


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


def main():
    parser = argparse.ArgumentParser(description="Plot per-UID throughput and traffic from AP L4NetEvaluate logs.")
    parser.add_argument("--input", required=True, help="Path to data_net_netifstats.txt or data_net_uidstats.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]")
    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, is_end=True) if args.end_time.strip() else None

    records = filter_records_by_time(load_records(input_path), start_time, end_time)
    if not records:
        print("No UID throughput records found.")
        return 2

    base_name = os.path.splitext(os.path.basename(input_path))[0]
    suffix = build_suffix(start_time, end_time)
    summary = compute_stream_summary(records)
    summary_csv = os.path.join(outdir, f"{base_name}_uid_tput_summary{suffix}.csv")
    write_summary_csv(summary, summary_csv)
    print(f"Summary CSV saved: {summary_csv}")

    records_by_uid = group_records_by_uid(records)
    summary_by_uid = defaultdict(list)
    for row in summary:
        summary_by_uid[row["uid"]].append(row)

    for uid in sorted(records_by_uid, key=lambda value: int(value)):
        timeseries = build_uid_timeseries(records_by_uid[uid])
        csv_path = os.path.join(outdir, f"{base_name}_uid_{safe_name(uid)}_tput{suffix}.csv")
        html_path = os.path.join(outdir, f"{base_name}_uid_{safe_name(uid)}_tput{suffix}.html")
        write_uid_records_csv(timeseries, csv_path)
        plot_uid_html(uid, summary_by_uid[uid], timeseries, html_path)
        print(f"UID CSV saved: {csv_path}")
        print(f"UID HTML saved: {html_path}")

    return 0


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