"""
RedAmon - Masscan Port Scanner Module

High-speed SYN port scanning for large networks and IP ranges.
Runs as a native binary (built into the recon container) for simplicity.

Features:
- Fastest port scanner — optimized for subnet/CIDR scanning
- SYN scan only (no CONNECT fallback)
- Banner grabbing support
- NDJSON output for reliable parsing
- Normalized output compatible with Naabu's port_scan structure
"""

import json
import subprocess
import shutil
import tempfile
from pathlib import Path
from datetime import datetime
from typing import Dict, List, Set, Tuple
import sys

PROJECT_ROOT = Path(__file__).parent.parent
sys.path.insert(0, str(PROJECT_ROOT))

from helpers.iana_services import get_service_name_friendly as get_service_name

# AI surface recon — reuse the port-catalogue annotator defined in port_scan.py.
# Importing from there (instead of duplicating the helper) keeps the catalogue
# rules single-sourced and means a future change (e.g. promoting an ambiguous
# port via http-probe confirmation) only needs to land once.
from main_recon_modules.port_scan import _annotate_ai_port_catalog


# =============================================================================
# Prerequisites
# =============================================================================

def is_masscan_installed() -> bool:
    """Check if masscan binary is available."""
    return shutil.which("masscan") is not None


def _is_mock_hostname(hostname: str, ip: str) -> bool:
    """
    Detect mock hostnames generated by run_ip_recon for IPs without PTR records.
    e.g. "10-0-0-1" for IP "10.0.0.1", "2001-db8--1" for "2001:db8::1"
    """
    expected_mock = ip.replace('.', '-').replace(':', '-')
    return hostname == expected_mock


# =============================================================================
# Target Preparation
# =============================================================================

def resolve_targets_to_ips(recon_data: dict) -> Tuple[List[str], Dict[str, List[str]]]:
    """
    Extract IPs from recon data. Masscan requires IP addresses (not hostnames).

    Returns:
        Tuple of (ip_list, ip_to_hostnames mapping)
    """
    unique_ips: Set[str] = set()
    ip_to_hostnames: Dict[str, List[str]] = {}

    dns_data = recon_data.get("dns", {})

    # Root domain
    domain_dns = dns_data.get("domain", {})
    domain_name = recon_data.get("domain", "")
    if domain_name:
        for ip in domain_dns.get("ips", {}).get("ipv4", []) + domain_dns.get("ips", {}).get("ipv6", []):
            unique_ips.add(ip)
            ip_to_hostnames.setdefault(ip, [])
            if domain_name not in ip_to_hostnames[ip]:
                ip_to_hostnames[ip].append(domain_name)

    # Subdomains
    for subdomain, sub_data in dns_data.get("subdomains", {}).items():
        if not sub_data.get("has_records", False):
            continue
        for ip in sub_data.get("ips", {}).get("ipv4", []) + sub_data.get("ips", {}).get("ipv6", []):
            unique_ips.add(ip)
            ip_to_hostnames.setdefault(ip, [])
            if subdomain not in ip_to_hostnames[ip]:
                ip_to_hostnames[ip].append(subdomain)

    return list(unique_ips), ip_to_hostnames


# =============================================================================
# Command Builder
# =============================================================================

def build_masscan_command(targets_file: str, output_file: str, settings: dict) -> List[str]:
    """
    Build the masscan command with project settings.

    Args:
        targets_file: Path to file with one IP/CIDR per line
        output_file: Path for NDJSON output (one JSON object per line)
        settings: Settings dictionary from main.py

    Returns:
        List of command arguments
    """
    MASSCAN_CUSTOM_PORTS = settings.get('MASSCAN_CUSTOM_PORTS', '')
    MASSCAN_TOP_PORTS = settings.get('MASSCAN_TOP_PORTS', '1000')
    MASSCAN_RATE = settings.get('MASSCAN_RATE', 1000)
    MASSCAN_BANNERS = settings.get('MASSCAN_BANNERS', False)
    MASSCAN_WAIT = settings.get('MASSCAN_WAIT', 10)
    MASSCAN_RETRIES = settings.get('MASSCAN_RETRIES', 1)
    MASSCAN_EXCLUDE_TARGETS = settings.get('MASSCAN_EXCLUDE_TARGETS', '')

    cmd = ["masscan"]

    # Targets from file
    cmd.extend(["-iL", targets_file])

    # Port configuration
    if MASSCAN_CUSTOM_PORTS:
        cmd.extend(["-p", MASSCAN_CUSTOM_PORTS])
    elif MASSCAN_TOP_PORTS:
        top_ports_str = str(MASSCAN_TOP_PORTS)
        if top_ports_str.lower() == 'full':
            cmd.extend(["-p", "0-65535"])
        else:
            cmd.extend(["--top-ports", top_ports_str])

    # Performance
    cmd.extend(["--rate", str(MASSCAN_RATE)])
    cmd.extend(["--wait", str(MASSCAN_WAIT)])
    cmd.extend(["--retries", str(MASSCAN_RETRIES)])

    # Banner grabbing
    if MASSCAN_BANNERS:
        cmd.append("--banners")

    # Output: NDJSON (one JSON object per line, avoids -oJ trailing comma bug)
    cmd.extend(["-oD", output_file])

    # Exclude targets
    if MASSCAN_EXCLUDE_TARGETS.strip():
        exclude_file = str(Path(output_file).parent / "masscan_exclude.txt")
        with open(exclude_file, 'w') as f:
            for line in MASSCAN_EXCLUDE_TARGETS.strip().split(','):
                line = line.strip()
                if line:
                    f.write(f"{line}\n")
        cmd.extend(["--excludefile", exclude_file])

    return cmd


# =============================================================================
# Result Parsing
# =============================================================================

def parse_masscan_output(output_file: str, ip_to_hostnames: Dict[str, List[str]], settings: dict = None) -> Dict:
    """
    Parse Masscan NDJSON (-oD) output into the same structure as Naabu's parser.

    AI surface recon: when ``settings`` is provided, ports matching the AI
    port catalogue receive an ``ai_service`` annotation on their port_details
    entry (gated by MASSCAN_AI_PORT_CATALOG_ENABLED).

    Masscan NDJSON format — one record per line, one port per record:
    {"ip":"x.x.x.x","timestamp":"...","port":80,"proto":"tcp","rec_type":"status","data":{"status":"open","reason":"syn-ack","ttl":48}}

    With --banners, banner records also appear:
    {"ip":"x.x.x.x","timestamp":"...","port":80,"proto":"tcp","rec_type":"banner","data":{"banner":"..."}}

    Note: this is NOT the -oJ format which uses ports:[{...}] arrays.

    Returns:
        Dict with by_host, by_ip, all_ports, summary — identical to Naabu's output
    """
    by_host: Dict = {}
    by_ip: Dict = {}
    all_ports: Set[int] = set()

    if not Path(output_file).exists():
        return _empty_result()

    with open(output_file, 'r') as f:
        for line in f:
            line = line.strip()
            if not line or line.startswith('#'):
                continue

            try:
                entry = json.loads(line)
            except json.JSONDecodeError:
                continue

            ip = entry.get("ip", "")
            if not ip:
                continue

            # NDJSON has port/proto at top level (not in a "ports" array)
            port = entry.get("port")
            proto = entry.get("proto", "tcp")
            rec_type = entry.get("rec_type", "")
            data = entry.get("data", {})

            # Only process status records with open ports; skip banner records
            if rec_type == "banner":
                continue
            status = data.get("status", "") if isinstance(data, dict) else ""
            if port is None or status != "open":
                continue

            hostnames = ip_to_hostnames.get(ip, [])

            # Initialize IP record
            if ip not in by_ip:
                by_ip[ip] = {
                    "ip": ip,
                    "hostnames": list(hostnames),
                    "ports": [],
                    "cdn": None,
                    "is_cdn": False,
                }

            all_ports.add(port)

            # Add to by_ip
            if port not in by_ip[ip]["ports"]:
                by_ip[ip]["ports"].append(port)

            # Add to by_host for each hostname mapped to this IP
            for hostname in hostnames:
                if hostname not in by_host:
                    by_host[hostname] = {
                        "host": hostname,
                        "ip": ip,
                        "ports": [],
                        "port_details": [],
                        "cdn": None,
                        "is_cdn": False,
                    }
                if port not in by_host[hostname]["ports"]:
                    by_host[hostname]["ports"].append(port)
                    service = get_service_name(port)
                    by_host[hostname]["port_details"].append({
                        "port": port,
                        "protocol": proto,
                        "service": service,
                    })

            # If no hostname mapping, use the IP itself as the host key
            if not hostnames:
                if ip not in by_host:
                    by_host[ip] = {
                        "host": ip,
                        "ip": ip,
                        "ports": [],
                        "port_details": [],
                        "cdn": None,
                        "is_cdn": False,
                    }
                if port not in by_host[ip]["ports"]:
                    by_host[ip]["ports"].append(port)
                    service = get_service_name(port)
                    by_host[ip]["port_details"].append({
                        "port": port,
                        "protocol": proto,
                        "service": service,
                    })

    # Sort ports
    for host_data in by_host.values():
        host_data["ports"].sort()
        host_data["port_details"].sort(key=lambda x: x["port"])
    for ip_data in by_ip.values():
        ip_data["ports"].sort()

    all_ports_sorted = sorted(list(all_ports))

    # AI surface recon — annotate AI-bearing ports on each port_details entry
    ai_annotations = _annotate_ai_port_catalog(by_host, settings, detected_by="masscan-ai-port")

    summary = {
        "hosts_scanned": len(by_host),
        "ips_scanned": len(by_ip),
        "hosts_with_open_ports": len([h for h in by_host.values() if h["ports"]]),
        "total_open_ports": sum(len(h["ports"]) for h in by_host.values()),
        "unique_ports": all_ports_sorted,
        "unique_port_count": len(all_ports_sorted),
        "cdn_hosts": 0,
        "ai_ports_annotated": ai_annotations,
    }

    return {
        "by_host": by_host,
        "by_ip": by_ip,
        "all_ports": all_ports_sorted,
        "summary": summary,
    }


def _empty_result() -> Dict:
    return {
        "by_host": {},
        "by_ip": {},
        "all_ports": [],
        "summary": {
            "hosts_scanned": 0,
            "ips_scanned": 0,
            "hosts_with_open_ports": 0,
            "total_open_ports": 0,
            "unique_ports": [],
            "unique_port_count": 0,
            "cdn_hosts": 0,
            "ai_ports_annotated": 0,
        },
    }


# =============================================================================
# Main Scan Function
# =============================================================================

def run_masscan_scan(recon_data: dict, output_file: Path = None, settings: dict = None) -> dict:
    """
    Run Masscan port scan on targets from recon data.

    Args:
        recon_data: Dictionary containing DNS/subdomain data
        output_file: Path to save enriched results (optional)
        settings: Settings dictionary from main.py

    Returns:
        Enriched recon_data with "masscan_scan" section added
    """
    print("\n" + "=" * 60)
    print("[*][Masscan] PORT SCANNER")
    print("=" * 60)

    if settings is None:
        settings = {}

    if not settings.get('MASSCAN_ENABLED', True):
        print("[-][Masscan] Disabled — skipping")
        return recon_data

    from recon.helpers import print_effective_settings
    print_effective_settings(
        "Masscan",
        settings,
        keys=[
            ("MASSCAN_ENABLED", "Toggle"),
            ("MASSCAN_TOP_PORTS", "Ports"),
            ("MASSCAN_CUSTOM_PORTS", "Ports"),
            ("MASSCAN_RATE", "Performance"),
            ("MASSCAN_WAIT", "Performance"),
            ("MASSCAN_RETRIES", "Performance"),
            ("MASSCAN_BANNERS", "Features"),
            ("MASSCAN_EXCLUDE_TARGETS", "Features"),
        ],
    )

    MASSCAN_RATE = settings.get('MASSCAN_RATE', 1000)
    MASSCAN_CUSTOM_PORTS = settings.get('MASSCAN_CUSTOM_PORTS', '')
    MASSCAN_TOP_PORTS = settings.get('MASSCAN_TOP_PORTS', '1000')
    MASSCAN_BANNERS = settings.get('MASSCAN_BANNERS', False)

    if not is_masscan_installed():
        print("[!][Masscan] Binary not found. Ensure masscan is installed.")
        return recon_data

    # Extract targets — masscan needs IPs, not hostnames
    print("[*][Masscan] Extracting IP targets from recon data...")

    # In IP mode, use expanded IPs directly (may include CIDRs)
    metadata = recon_data.get("metadata", {})
    if metadata.get("ip_mode"):
        all_targets = metadata.get("expanded_ips", [])
        raw_map = metadata.get("ip_to_hostname", {})
        # Normalize: main.py stores {ip: hostname_str}, we need {ip: [hostname_str]}.
        # Mock hostnames (e.g. "10-0-0-1" for IPs without PTR) must be excluded — they're
        # not routable and would produce invalid URLs in http_probe. For IPs without
        # real PTR records, use the IP itself as the hostname (matching Naabu's behavior).
        ip_to_hostnames = {}
        for ip, val in raw_map.items():
            if isinstance(val, list):
                real = [h for h in val if not _is_mock_hostname(h, ip)]
                ip_to_hostnames[ip] = real if real else [ip]
            elif isinstance(val, str) and val and not _is_mock_hostname(val, ip):
                ip_to_hostnames[ip] = [val]
            else:
                ip_to_hostnames[ip] = [ip]
    else:
        all_targets, ip_to_hostnames = resolve_targets_to_ips(recon_data)

    if not all_targets:
        print("[!][Masscan] No IP targets found in recon data")
        return recon_data

    print(f"[*][Masscan] Total IP targets: {len(all_targets)}")

    scan_temp_dir = Path(tempfile.mkdtemp(prefix="redamon_masscan_"))

    try:
        # Write targets
        targets_file = scan_temp_dir / "targets.txt"
        with open(targets_file, 'w') as f:
            for target in all_targets:
                f.write(f"{target}\n")

        masscan_output = scan_temp_dir / "masscan_output.ndjson"

        cmd = build_masscan_command(str(targets_file), str(masscan_output), settings)

        print(f"[*][Masscan] Starting scan...")
        print(f"[*][Masscan] Ports: {MASSCAN_CUSTOM_PORTS if MASSCAN_CUSTOM_PORTS else f'top {MASSCAN_TOP_PORTS}'}")
        print(f"[*][Masscan] Rate: {MASSCAN_RATE} pps")
        print(f"[*][Masscan] Wait: {settings.get('MASSCAN_WAIT', 10)}s")
        print(f"[*][Masscan] Retries: {settings.get('MASSCAN_RETRIES', 1)}")
        print(f"[*][Masscan] Banners: {MASSCAN_BANNERS}")

        start_time = datetime.now()

        process = subprocess.Popen(
            cmd,
            stdout=subprocess.PIPE,
            stderr=subprocess.PIPE,
            text=True,
        )

        _, stderr = process.communicate(timeout=1800)

        end_time = datetime.now()
        duration = (end_time - start_time).total_seconds()

        if process.returncode != 0 and not masscan_output.exists():
            if "permission" in (stderr or "").lower() or "raw socket" in (stderr or "").lower():
                print(f"[!][Masscan] Permission denied — masscan requires root or CAP_NET_RAW")
            else:
                print(f"[!][Masscan] Scan failed: {(stderr or '')[:200]}")
            return recon_data

        # Parse results
        print(f"[*][Masscan] Parsing results...")
        results = parse_masscan_output(str(masscan_output), ip_to_hostnames, settings=settings)
        if results.get("summary", {}).get("ai_ports_annotated"):
            print(f"[+][Masscan] AI port catalog matched {results['summary']['ai_ports_annotated']} port(s)")

        masscan_results = {
            "scan_metadata": {
                "scanner": "masscan",
                "scan_timestamp": start_time.isoformat(),
                "scan_duration_seconds": round(duration, 2),
                "scan_type": "syn",
                "ports_config": MASSCAN_CUSTOM_PORTS if MASSCAN_CUSTOM_PORTS else f"top-{MASSCAN_TOP_PORTS}",
                "rate_limit": MASSCAN_RATE,
                "banners_enabled": MASSCAN_BANNERS,
                "total_targets": len(all_targets),
            },
            "by_host": results["by_host"],
            "by_ip": results["by_ip"],
            "all_ports": results["all_ports"],
            "ip_to_hostnames": ip_to_hostnames,
            "summary": results["summary"],
        }

        summary = results["summary"]
        print(f"[✓][Masscan] Scan completed in {duration:.1f} seconds")
        print(f"[+][Masscan] Hosts with open ports: {summary['hosts_with_open_ports']}")
        print(f"[+][Masscan] Total open ports found: {summary['total_open_ports']}")
        print(f"[+][Masscan] Unique ports: {summary['unique_port_count']}")

        if results["all_ports"]:
            ports_preview = ', '.join(map(str, results['all_ports'][:20]))
            extra = f"... (+{len(results['all_ports'])-20} more)" if len(results['all_ports']) > 20 else ""
            print(f"[+][Masscan] Ports discovered: {ports_preview}{extra}")

        recon_data["masscan_scan"] = masscan_results

        if output_file:
            with open(output_file, 'w') as f:
                json.dump(recon_data, f, indent=2, default=str)
            print(f"[✓][Masscan] Results saved to {output_file}")

        return recon_data

    except subprocess.TimeoutExpired:
        print("[!][Masscan] Scan timed out after 30 minutes — killing process")
        try:
            process.kill()
            process.wait(timeout=10)
        except Exception:
            pass
        return recon_data
    except Exception as e:
        print(f"[!][Masscan] Error during scan: {e}")
        return recon_data
    finally:
        try:
            if scan_temp_dir.exists():
                for f in scan_temp_dir.iterdir():
                    f.unlink()
                scan_temp_dir.rmdir()
        except Exception:
            pass


def run_masscan_scan_isolated(recon_data: dict, settings: dict = None) -> dict:
    """
    Run masscan scan and return only the 'masscan_scan' data dict.

    Thread-safe: does not mutate recon_data.

    Args:
        recon_data: The pipeline's combined result dictionary (read-only)
        settings: Settings dictionary from main.py

    Returns:
        The 'masscan_scan' data dict, or empty dict if scan produced no results.
    """
    import copy
    snapshot = copy.copy(recon_data)
    run_masscan_scan(snapshot, output_file=None, settings=settings)
    return snapshot.get("masscan_scan", {})
