"""Escaneo de red: ping sweep + ARP + hostname + puertos + clasificación."""
from __future__ import annotations

import socket
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import Callable, Dict, List, Optional

from .arp_watch import ArpWatcher
from .models import Device, DeviceStatus, cidr_hosts, classify_device, load_config, now_str, status_from_stats
from .oui import lookup_vendor
from .pinger import ping_host
from .ports import scan_ports


ProgressCb = Callable[[int, int, str], None]  # done, total, message
DeviceCb = Callable[[Device], None]


def resolve_hostname(ip: str, timeout: float = 0.4) -> str:
    old = socket.getdefaulttimeout()
    try:
        socket.setdefaulttimeout(timeout)
        name, _, _ = socket.gethostbyaddr(ip)
        return name
    except Exception:
        return ""
    finally:
        socket.setdefaulttimeout(old)


class NetworkScanner:
    def __init__(self, arp: Optional[ArpWatcher] = None):
        self.cfg = load_config()
        self.arp = arp or ArpWatcher()
        self.cancel = False
        self.devices: Dict[str, Device] = {}

    def stop(self) -> None:
        self.cancel = True

    def scan(
        self,
        cidr: str,
        ports: Optional[List[int]] = None,
        on_progress: Optional[ProgressCb] = None,
        on_device: Optional[DeviceCb] = None,
        resolve_names: bool = True,
        scan_ports_flag: bool = True,
        workers: int = 64,
    ) -> List[Device]:
        self.cancel = False
        hosts = cidr_hosts(cidr)
        total = len(hosts)
        timeout_ms = int(self.cfg.get("scan_timeout_ms", 800))
        ports = ports or list(self.cfg.get("default_ports", [80, 443, 554, 37777]))
        # refrescar ARP previa
        self.arp.refresh()
        found: Dict[str, Device] = {}

        def probe(ip: str) -> Optional[Device]:
            if self.cancel:
                return None
            pr = ping_host(ip, timeout_ms=timeout_ms, count=1)
            if not pr.ok:
                return None
            # ARP
            table = ArpWatcher.read_arp_table()
            mac = table.get(ip).mac if ip in table else self.arp.known_mac(ip)
            if mac:
                self.arp.record_observation(ip, mac)
            vendor = lookup_vendor(mac) if mac else ""
            hostname = resolve_hostname(ip) if resolve_names else ""
            open_ports: Dict[int, bool] = {}
            if scan_ports_flag:
                # puertos clave rápidos en scan
                key_ports = [p for p in ports if p in (80, 443, 554, 8000, 37777, 8899)]
                open_ports = scan_ports(ip, key_ports, timeout=float(self.cfg.get("port_timeout_sec", 1.0)))
            dtype = classify_device(vendor, open_ports, hostname)
            conflict = any(c.ip == ip for c in self.arp.conflicts())
            status = status_from_stats(True, 0.0, pr.rtt_ms, conflict)
            dev = Device(
                ip=ip,
                mac=mac or "",
                vendor=vendor,
                hostname=hostname,
                device_type=dtype,
                status=status,
                latency_ms=pr.rtt_ms,
                loss_pct=0.0,
                last_seen=now_str(),
                ports=open_ports,
                sent=pr.sent,
                recv=pr.recv,
                rtt_avg=pr.rtt_ms,
                rtt_min=pr.rtt_ms,
                rtt_max=pr.rtt_ms,
            )
            if conflict:
                hist = self.arp.history_for(ip)
                if len(hist) >= 2:
                    dev.notes = f"CONFLICTO IP: MAC {hist[-2][1]} -> {hist[-1][1]}"
                    dev.mac_history = [{"ts": t, "mac": m} for t, m in hist]
            return dev

        done = 0
        with ThreadPoolExecutor(max_workers=workers) as ex:
            futs = {ex.submit(probe, ip): ip for ip in hosts}
            for fut in as_completed(futs):
                if self.cancel:
                    break
                done += 1
                ip = futs[fut]
                if on_progress:
                    on_progress(done, total, f"Escaneando {ip}…")
                try:
                    dev = fut.result()
                except Exception:
                    continue
                if dev:
                    found[dev.ip] = dev
                    if on_device:
                        on_device(dev)

        # segunda pasada ARP por si ping pobló la tabla
        self.arp.refresh()
        for ip, entry in ArpWatcher.read_arp_table().items():
            if ip in found and not found[ip].mac:
                found[ip].mac = entry.mac
                found[ip].vendor = lookup_vendor(entry.mac) or found[ip].vendor
                found[ip].device_type = classify_device(
                    found[ip].vendor, found[ip].ports, found[ip].hostname
                )

        self.devices = found
        return sorted(found.values(), key=lambda d: tuple(int(x) for x in d.ip.split(".")))
