"""ONVIF WS-Discovery (multicast) — lo que un NVR suele 'ver' al pulsar Buscar."""
from __future__ import annotations

import re
import socket
import time
import uuid
from typing import Dict, Optional, Set


WS_DISCOVERY_ADDR = "239.255.255.250"
WS_DISCOVERY_PORT = 3702

PROBE_TEMPLATE = """<?xml version="1.0" encoding="UTF-8"?>
<e:Envelope xmlns:e="http://www.w3.org/2003/05/soap-envelope"
            xmlns:w="http://schemas.xmlsoap.org/ws/2004/08/addressing"
            xmlns:d="http://schemas.xmlsoap.org/ws/2005/04/discovery"
            xmlns:dn="http://www.onvif.org/ver10/network/wsdl">
  <e:Header>
    <w:MessageID>uuid:{msgid}</w:MessageID>
    <w:To>urn:schemas-xmlsoap-org:ws:2005:04:discovery</w:To>
    <w:Action>http://schemas.xmlsoap.org/ws/2005/04/discovery/Probe</w:Action>
  </e:Header>
  <e:Body>
    <d:Probe>
      <d:Types>dn:NetworkVideoTransmitter</d:Types>
    </d:Probe>
  </e:Body>
</e:Envelope>"""


def _local_ip_for_multicast(preferred: str = "") -> str:
    if preferred and not preferred.startswith("127."):
        return preferred
    try:
        s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
        s.connect(("8.8.8.8", 80))
        ip = s.getsockname()[0]
        s.close()
        return ip
    except Exception:
        return "0.0.0.0"


def _extract_ips_from_xml(text: str) -> Set[str]:
    ips: Set[str] = set()
    for token in text.replace("<", " ").replace(">", " ").replace('"', " ").split():
        if "://" in token:
            try:
                host = token.split("://", 1)[1].split("/", 1)[0]
                host = host.split(":")[0]
                parts = host.split(".")
                if len(parts) == 4 and all(p.isdigit() and 0 <= int(p) <= 255 for p in parts):
                    ips.add(host)
            except Exception:
                pass
    for m in re.finditer(r"\b(\d{1,3}(?:\.\d{1,3}){3})\b", text):
        cand = m.group(1)
        parts = cand.split(".")
        try:
            if not all(0 <= int(p) <= 255 for p in parts):
                continue
        except ValueError:
            continue
        if cand.startswith(("239.", "224.", "169.254.")):
            continue
        if cand.startswith(("192.168.", "10.")) or cand.startswith("172."):
            ips.add(cand)
    return ips


def ws_discovery_probe(
    listen_sec: float = 4.0,
    local_ip: str = "",
    probe_types: bool = True,
) -> Dict:
    """
    Envía Probe ONVIF por multicast UDP y recoge ProbeMatches.
    Simula la búsqueda automática de un NVR en la misma LAN.
    """
    result = {
        "ok": False,
        "ips": [],
        "raw_count": 0,
        "error": "",
        "listen_sec": listen_sec,
        "local_bind": "",
        "note": (
            "WS-Discovery usa multicast 239.255.255.250:3702. "
            "Si el switch/AP filtra multicast o hay aislamiento de clientes, "
            "el NVR 'Buscar' verá pocas cámaras aunque el ping unicast funcione."
        ),
    }
    bind_ip = _local_ip_for_multicast(local_ip)
    result["local_bind"] = bind_ip
    msgid = str(uuid.uuid4())
    if probe_types:
        payload = PROBE_TEMPLATE.format(msgid=msgid).encode("utf-8")
    else:
        payload = (
            PROBE_TEMPLATE.replace(
                "<d:Types>dn:NetworkVideoTransmitter</d:Types>",
                "",
            )
            .format(msgid=msgid)
            .encode("utf-8")
        )

    found: Set[str] = set()
    sock: Optional[socket.socket] = None
    try:
        sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
        sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
        try:
            sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1)
        except (AttributeError, OSError):
            pass
        sock.bind((bind_ip if bind_ip != "0.0.0.0" else "", WS_DISCOVERY_PORT))
        try:
            mreq = socket.inet_aton(WS_DISCOVERY_ADDR) + socket.inet_aton(
                bind_ip if bind_ip and bind_ip != "0.0.0.0" else "0.0.0.0"
            )
            sock.setsockopt(socket.IPPROTO_IP, socket.IP_ADD_MEMBERSHIP, mreq)
        except OSError:
            pass
        sock.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_TTL, 2)
        sock.settimeout(0.4)

        for _ in range(2):
            sock.sendto(payload, (WS_DISCOVERY_ADDR, WS_DISCOVERY_PORT))
            time.sleep(0.15)

        deadline = time.time() + listen_sec
        raw = 0
        while time.time() < deadline:
            try:
                data, addr = sock.recvfrom(65535)
            except socket.timeout:
                continue
            except OSError:
                break
            raw += 1
            text = data.decode("utf-8", errors="replace")
            ips = _extract_ips_from_xml(text)
            if not ips and addr:
                ips.add(addr[0])
            found.update(ips)
        result["raw_count"] = raw
        result["ips"] = sorted(found, key=lambda x: tuple(int(p) for p in x.split(".")))
        result["ok"] = True
    except Exception as e:
        result["error"] = str(e)
    finally:
        if sock:
            try:
                sock.close()
            except Exception:
                pass
    return result


def ws_discovery_probe_both(local_ip: str = "", listen_sec: float = 4.0) -> Dict:
    """Probe tipado + genérico; une resultados."""
    a = ws_discovery_probe(listen_sec=listen_sec, local_ip=local_ip, probe_types=True)
    b = ws_discovery_probe(listen_sec=max(2.0, listen_sec * 0.6), local_ip=local_ip, probe_types=False)
    ips = sorted(
        set(a.get("ips", [])) | set(b.get("ips", [])),
        key=lambda x: tuple(int(p) for p in x.split(".")),
    )
    return {
        "ok": a.get("ok") or b.get("ok"),
        "ips": ips,
        "raw_count": int(a.get("raw_count", 0)) + int(b.get("raw_count", 0)),
        "error": a.get("error") or b.get("error") or "",
        "listen_sec": listen_sec,
        "local_bind": a.get("local_bind") or b.get("local_bind"),
        "note": a.get("note", ""),
        "typed": a,
        "generic": b,
    }
