mdns-termux/lib/mdns_tools/reverse.py

94 lines
2.9 KiB
Python

"""Reverse mDNS lookup: IP address -> hostname (PTR record on multicast)."""
import ipaddress
import socket
import struct
import select
import time
from ._stdlib_resolve import build_query, skip_name, MDNS_ADDR4, MDNS_PORT
MDNS_ADDR6 = "ff02::fb"
QTYPE_PTR = 12
def _reverse_name(ip: str) -> tuple[str, str]:
"""Return (arpa-name, address-family) for an IP."""
addr = ipaddress.ip_address(ip)
if isinstance(addr, ipaddress.IPv4Address):
return addr.reverse_pointer, "v4"
return addr.reverse_pointer, "v6"
def _parse_ptr(data: bytes) -> str | None:
"""Return first PTR target hostname in the response."""
if len(data) < 12:
return None
qd, an = struct.unpack(">HH", data[4:8])
off = 12
for _ in range(qd):
off = skip_name(data, off) + 4
for _ in range(an):
off = skip_name(data, off)
if off + 10 > len(data):
return None
rtype, _, _, rdlen = struct.unpack(">HHIH", data[off:off + 10])
off += 10
if rtype == QTYPE_PTR:
return _read_name(data, off)
off += rdlen
return None
def _read_name(data: bytes, off: int) -> str:
parts: list[str] = []
seen: set[int] = set()
while off < len(data):
b = data[off]
if b == 0:
break
if b & 0xC0 == 0xC0:
if off in seen:
break
seen.add(off)
off = ((b & 0x3F) << 8) | data[off + 1]
continue
parts.append(data[off + 1:off + 1 + b].decode(errors="replace"))
off += b + 1
return ".".join(parts)
def reverse(ip: str, timeout: float = 1.5, retries: int = 2) -> str | None:
arpa, family = _reverse_name(ip)
af = socket.AF_INET6 if family == "v6" else socket.AF_INET
dst = MDNS_ADDR6 if family == "v6" else MDNS_ADDR4
per_try = max(0.4, timeout / max(1, retries))
for _ in range(max(1, retries)):
s = socket.socket(af, socket.SOCK_DGRAM)
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
try:
if family == "v4":
mreq = socket.inet_aton(MDNS_ADDR4) + socket.inet_aton("0.0.0.0")
s.setsockopt(socket.IPPROTO_IP, socket.IP_ADD_MEMBERSHIP, mreq)
except OSError:
pass
s.setblocking(False)
try:
s.sendto(build_query(arpa, QTYPE_PTR), (dst, MDNS_PORT))
except OSError:
s.close()
continue
end = time.time() + per_try
while time.time() < end:
r, _, _ = select.select([s], [], [], min(0.1, max(0.0, end - time.time())))
if r:
try:
data, _ = s.recvfrom(4096)
except OSError:
continue
name = _parse_ptr(data)
if name:
s.close()
return name
s.close()
return None