#!/usr/bin/env python3
"""Summarize authorized packet captures with Scapy.

This script reads PCAP/PCAPNG files only. It does not capture live traffic.
"""

from __future__ import annotations

import argparse
import csv
from collections import Counter
from pathlib import Path

from scapy.all import DNS, IP, TCP, UDP, rdpcap  # type: ignore


def packet_rows(pcap_path: Path) -> list[dict[str, str | int]]:
    packets = rdpcap(str(pcap_path))
    rows: list[dict[str, str | int]] = []
    for index, packet in enumerate(packets, start=1):
        if IP not in packet:
            continue
        proto = "tcp" if TCP in packet else "udp" if UDP in packet else str(packet[IP].proto)
        src_port = packet[TCP].sport if TCP in packet else packet[UDP].sport if UDP in packet else ""
        dst_port = packet[TCP].dport if TCP in packet else packet[UDP].dport if UDP in packet else ""
        dns_query = ""
        if DNS in packet and packet[DNS].qd:
            dns_query = packet[DNS].qd.qname.decode(errors="replace").rstrip(".")
        rows.append({
            "packet": index,
            "src": packet[IP].src,
            "dst": packet[IP].dst,
            "protocol": proto,
            "src_port": src_port,
            "dst_port": dst_port,
            "length": len(packet),
            "dns_query": dns_query,
        })
    return rows


def main() -> int:
    parser = argparse.ArgumentParser(description="Summarize a packet capture for incident response.")
    parser.add_argument("pcap", type=Path)
    parser.add_argument("--output", type=Path, default=Path("packet_summary.csv"))
    args = parser.parse_args()

    rows = packet_rows(args.pcap)
    conversations = Counter((row["src"], row["dst"], row["protocol"]) for row in rows)
    with args.output.open("w", newline="", encoding="utf-8") as handle:
        fieldnames = ["packet", "src", "dst", "protocol", "src_port", "dst_port", "length", "dns_query"]
        writer = csv.DictWriter(handle, fieldnames=fieldnames)
        writer.writeheader()
        writer.writerows(rows)
    print(f"Wrote {len(rows)} IP packets to {args.output}")
    print("Top conversations:")
    for conversation, count in conversations.most_common(10):
        print(f"  {conversation}: {count}")
    return 0


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