#!/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, info.properties[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())