Files
mdns-discovery-proxy/proxy.py
T
franck e45560f763 Fix SERVFAIL crash on TXT records with valueless keys
zeroconf represents a TXT attribute with no "=" (a bare boolean flag,
valid per RFC 6763 §6.4) or an empty value as properties[key] = None.
Reconstructing the TXT record blindly did b"%s=%s" % (key, None),
which raises TypeError since bytes %-formatting rejects None. The
exception propagated out of the deferToThread call and Twisted
answered with SERVFAIL instead of the TXT record, silently breaking
resolution for any service whose TXT record includes a bare key
(e.g. some _rfb._tcp/screen-sharing advertisements).

Found by diagnosing a real screen-sharing "impossible de resoudre"
failure over a WireGuard-tunneled deployment: dig showed SRV
resolving fine but TXT returning SERVFAIL for the same instance,
which is why dns-sd -L (needs both) never completed even after
flushing the client-side mDNSResponder cache.
2026-08-05 12:46:32 +02:00

267 lines
11 KiB
Python

#!/usr/bin/env python3
from zeroconf import Zeroconf, ServiceBrowser, ServiceInfo, DNSQuestion, DNSOutgoing, IPVersion, current_time_millis
from zeroconf.const import _TYPE_A, _TYPE_AAAA, _TYPE_PTR, _TYPE_SRV, _TYPE_TXT, _CLASS_IN, _FLAGS_QR_QUERY
import sys
import time
import threading
import heapq
import queue
from twisted.internet import reactor, defer, threads
from twisted.names import client, dns, error, server
from twisted.python import log
import socket
import ipaddress
domain = sys.argv[1]
port = int(sys.argv[2])
ttl = 120
timeout = 2
negative_ttl = 10
refresh_interval = 90
refresh_pool_size = 4
ipv6_ula_network = ipaddress.ip_network('fc00::/7')
def soa_record():
return dns.RRHeader(name=domain, type=dns.SOA, ttl=negative_ttl, payload=dns.Record_SOA(
mname=domain, rname="hostmaster." + domain,
serial=1, refresh=1200, retry=180, expire=1209600, minimum=negative_ttl
))
class DynamicResolver(object):
def __init__(self):
self.zeroconf = Zeroconf(ip_version=IPVersion.All)
self.lock = threading.Lock()
self.browser = None
self.browser_types = set()
self.interests = set()
self.refresh_heap = []
self.refresh_queue = queue.Queue()
for _ in range(refresh_pool_size):
threading.Thread(target=self._refresh_worker, daemon=True).start()
threading.Thread(target=self._refresh_scheduler, daemon=True).start()
def _ensure_browser(self, localname):
"""Make sure a persistent ServiceBrowser covers this service type.
A single ServiceBrowser (and thus a single thread) tracks every
type seen so far; when a genuinely new type shows up, it is
recreated with the enlarged type set rather than starting a
second one, so the thread count stays independent of how many
distinct service types get queried. Returns True the first time
this service type is seen."""
with self.lock:
if localname in self.browser_types:
return False
self.browser_types.add(localname)
if self.browser is not None:
self.browser.cancel()
self.browser = ServiceBrowser(self.zeroconf, list(self.browser_types), [lambda *a, **kw: None])
return True
def _keep_warm(self, key, refresh):
"""Register (once) a periodic `refresh()` job for `key`, so the
zeroconf cache stays populated between DNS queries instead of
going cold and forcing a fresh mDNS round trip on every lookup.
Jobs are executed by a small fixed-size worker pool rather than
one thread per key. Returns True the first time this key is
seen."""
with self.lock:
if key in self.interests:
return False
self.interests.add(key)
heapq.heappush(self.refresh_heap, (time.time() + refresh_interval, key, refresh))
return True
def _refresh_scheduler(self):
"""Single background thread: wakes up due refresh jobs and hands
them to the worker pool via self.refresh_queue."""
while True:
job = None
with self.lock:
if self.refresh_heap and self.refresh_heap[0][0] <= time.time():
_, key, refresh = heapq.heappop(self.refresh_heap)
job = (key, refresh)
if job:
self.refresh_queue.put(job)
else:
time.sleep(1)
def _refresh_worker(self):
"""Fixed pool of worker threads (refresh_pool_size of them) that
execute due refresh jobs and reschedule them for their next run."""
while True:
key, refresh = self.refresh_queue.get()
try:
refresh()
except Exception:
pass
with self.lock:
heapq.heappush(self.refresh_heap, (time.time() + refresh_interval, key, refresh))
def _dynamicResponseRequired(self, query):
if str(query.name).endswith(domain):
return True
return False
def _doDynamicResponse(self, query):
if query.type == dns.SOA:
return defer.succeed(([soa_record()], [], []))
localname = str(query.name)[:-len(domain)] + "local."
def browse(localname):
if self._ensure_browser(localname):
time.sleep(timeout)
now = current_time_millis()
records = [r for r in self.zeroconf.cache.get_all_by_details(localname, _TYPE_PTR, _CLASS_IN) if not r.is_expired(now)]
answers = [dns.RRHeader(name=localname[:-6] + domain, ttl=int(r.get_remaining_ttl(now)), type=dns.PTR, payload=dns.Record_PTR(
name=r.alias[:-6] + domain
)) for r in records]
return answers, ([] if answers else [soa_record()]), []
def txt(localname):
self._keep_warm(('info', localname), lambda: self.zeroconf.get_service_info(localname, localname, timeout*1000))
if localname.endswith('._device-info._tcp.local.'):
info = ServiceInfo(localname, localname)
info.request(self.zeroconf, timeout*1000)
if not info.text:
return [], [soa_record()], []
else:
info = self.zeroconf.get_service_info(localname, localname, timeout*1000)
if info is None:
return [], [soa_record()], []
order = []
i = 0
while i < len(info.text):
length = info.text[i]
i += 1
kv = info.text[i : i + length].split(b'=')
order.append(kv[0])
i += length
data = [(b"%s=%s" % (p, v)) if (v := info.properties[p]) is not None else p
for p in sorted(info.properties, key=lambda k: order.index(k) if k in order else 1000)]
now = current_time_millis()
cached = self.zeroconf.cache.get_by_details(localname, _TYPE_TXT, _CLASS_IN)
record_ttl = int(cached.get_remaining_ttl(now)) if cached else ttl
answers = [dns.RRHeader(name=localname[:-6] + domain, ttl=record_ttl, type=dns.TXT, payload=dns.Record_TXT(
*data
))]
return answers, [], []
def srv(localname):
self._keep_warm(('info', localname), lambda: self.zeroconf.get_service_info(localname, localname, timeout*1000))
info = self.zeroconf.get_service_info(localname, localname, timeout*1000)
if info is None:
return [], [soa_record()], []
now = current_time_millis()
cached = self.zeroconf.cache.get_by_details(localname, _TYPE_SRV, _CLASS_IN)
srv_ttl = int(cached.get_remaining_ttl(now)) if cached else ttl
answers = [dns.RRHeader(name=localname[:-6] + domain, ttl=srv_ttl, type=dns.SRV, payload=dns.Record_SRV(
info.priority, info.weight, info.port, info.server[:-6] + domain
))]
target = info.server[:-6] + domain
additional = [dns.RRHeader(name=target, ttl=ttl, type=dns.A, payload=dns.Record_A(
str(addr)
)) for addr in info.ip_addresses_by_version(IPVersion.V4Only)]
additional += [dns.RRHeader(name=target, ttl=ttl, type=dns.AAAA, payload=dns.Record_AAAA(
str(addr)
)) for addr in info.ip_addresses_by_version(IPVersion.V6Only) if addr in ipv6_ula_network]
return answers, [], additional
def host(localname, qtype):
if qtype == dns.AAAA:
mdns_type, parse_addr = _TYPE_AAAA, ipaddress.IPv6Address
keep_addr = lambda addr: addr in ipv6_ula_network
else:
mdns_type, parse_addr = _TYPE_A, ipaddress.IPv4Address
keep_addr = lambda addr: not addr.is_link_local
record_cls = dns.Record_AAAA if qtype == dns.AAAA else dns.Record_A
def send_question():
q = DNSQuestion(localname, mdns_type, _CLASS_IN)
out = DNSOutgoing(_FLAGS_QR_QUERY)
out.add_question(q)
self.zeroconf.send(out)
def answers_from_cache():
now = current_time_millis()
result = []
for r in self.zeroconf.cache.get_all_by_details(localname, mdns_type, _CLASS_IN):
if r.is_expired(now):
continue
addr = parse_addr(r.address)
if keep_addr(addr):
result.append(dns.RRHeader(name=query.name.name, ttl=int(r.get_remaining_ttl(now)), type=qtype, payload=record_cls(str(addr))))
return result
is_new = self._keep_warm((localname, mdns_type), send_question)
answers = answers_from_cache()
if not answers and is_new:
send_question()
deadline = time.time() + timeout
while not answers and time.time() < deadline:
time.sleep(0.1)
answers = answers_from_cache()
return answers, ([] if answers else [soa_record()]), []
d = defer.Deferred()
if query.type == dns.PTR:
d = threads.deferToThread(browse, localname)
return d
elif query.type == dns.TXT:
d = threads.deferToThread(txt, localname)
return d
elif query.type == dns.SRV:
d = threads.deferToThread(srv, localname)
return d
elif query.type in (dns.A, dns.AAAA):
d = threads.deferToThread(host, localname, query.type)
return d
else:
print("Unsupported request", query)
d.callback(([], [soa_record()], []))
return d
def query(self, query, timeout=None):
if self._dynamicResponseRequired(query):
return self._doDynamicResponse(query)
else:
return defer.fail(error.DomainError())
class TruncatingDNSDatagramProtocol(dns.DNSDatagramProtocol):
def writeMessage(self, message, address):
if type(message) is dns.Message and len(message.toStr()) > 512:
message.additional = []
if len(message.toStr()) > 512:
message.trunc = 1
message.answers = []
dns.DNSDatagramProtocol.writeMessage(self, message, address)
def main():
factory = server.DNSServerFactory(
clients=[DynamicResolver()],
verbose=0
)
protocol = TruncatingDNSDatagramProtocol(controller=factory)
reactor.listenUDP(port, protocol)
reactor.listenTCP(port, factory)
log.startLogging(sys.stdout)
reactor.run()
if __name__ == '__main__':
raise SystemExit(main())