From af2251c2ef2056df069a50d87c54854e2a7cd54e Mon Sep 17 00:00:00 2001 From: Manuel Meitinger Date: Sun, 27 Mar 2022 15:23:59 +0200 Subject: [PATCH] [dns] minor fixes --- docs/scripts/api-events.py | 12 +++++- mitmproxy/addons/dnsserver.py | 29 +++++++------ mitmproxy/addons/dumper.py | 2 +- mitmproxy/dns.py | 14 +++--- mitmproxy/proxy/layers/dns.py | 81 +++++++++++++++++++---------------- mitmproxy/proxy/server.py | 4 +- 6 files changed, 80 insertions(+), 62 deletions(-) diff --git a/docs/scripts/api-events.py b/docs/scripts/api-events.py index eca751a79..86d9496dc 100644 --- a/docs/scripts/api-events.py +++ b/docs/scripts/api-events.py @@ -8,7 +8,7 @@ from typing import List, Type import mitmproxy.addons.next_layer # noqa from mitmproxy import hooks, log, addonmanager from mitmproxy.proxy import server_hooks, layer -from mitmproxy.proxy.layers import http, modes, tcp, tls, websocket +from mitmproxy.proxy.layers import dns, http, modes, tcp, tls, websocket known = set() @@ -106,6 +106,16 @@ with outfile.open("w") as f, contextlib.redirect_stdout(f): ] ) + category( + "DNS", + "", + [ + dns.DnsRequestHook, + dns.DnsResponseHook, + dns.DnsErrorHook, + ] + ) + category( "TCP", "", diff --git a/mitmproxy/addons/dnsserver.py b/mitmproxy/addons/dnsserver.py index 1ec9757c2..86687b148 100644 --- a/mitmproxy/addons/dnsserver.py +++ b/mitmproxy/addons/dnsserver.py @@ -61,7 +61,7 @@ class DnsServer: """Start a DNS server. Disabled by default.""" ) loader.add_option( - "dns_listen_host", Optional[str], "", + "dns_listen_host", str, "", """Address to bind DNS server to.""" ) loader.add_option( @@ -87,9 +87,7 @@ class DnsServer: async def refresh_server(self) -> None: async with self._lock: - if self.server: - await self.shutdown_server() - self.server = None + await self.shutdown_server() if ctx.options.dns_server: self.server = await udp.start_server( self.handle_connection, @@ -101,18 +99,21 @@ class DnsServer: ctx.log.info(f"DNS server listening at {' and '.join(addrs)}") async def shutdown_server(self) -> None: - ctx.log.info("Stopping server...") - self.server.close() - await self.server.wait_closed() - self.server = None + if self.server is not None: + ctx.log.info("Stopping server...") + self.server.close() + await self.server.wait_closed() + self.server = None async def handle_connection(self, r: asyncio.StreamReader, w: asyncio.StreamWriter) -> None: peername = w.get_extra_info('peername') - asyncio_utils.set_task_debug_info( - asyncio.current_task(), - name=f"DnsServer.handle_connection", - client=peername, - ) + current_task = asyncio.current_task() + if current_task is not None: + asyncio_utils.set_task_debug_info( + current_task, + name=f"DnsServer.handle_connection", + client=peername, + ) handler = DnsConnectionHandler(self.master, r, w, self.options) self._connections[peername] = handler try: @@ -123,7 +124,7 @@ class DnsServer: def server_connect(self, ctx: server_hooks.ServerConnectionHookData) -> None: assert ctx.server.address self_connect = ( - ctx.server.protocol == ConnectionProtocol.UDP + ctx.server.protocol is ConnectionProtocol.UDP and ctx.server.address[1] == self.options.dns_listen_port and diff --git a/mitmproxy/addons/dumper.py b/mitmproxy/addons/dumper.py index 8bf8d933a..1405a13db 100644 --- a/mitmproxy/addons/dumper.py +++ b/mitmproxy/addons/dumper.py @@ -347,7 +347,7 @@ class Dumper: def dns_response(self, f: dns.DNSFlow): # TODO this needs to be refined - if self.match(f) and f.response.answers: + if self.match(f) and f.response and f.response.answers: self.echo("{server} answers {client} with {name} {ttl}s {class_} {type} {value}".format( client=human.format_address(f.client_conn.peername), server=human.format_address(f.server_conn.address), diff --git a/mitmproxy/dns.py b/mitmproxy/dns.py index 1f6d71c2b..dde964f6d 100644 --- a/mitmproxy/dns.py +++ b/mitmproxy/dns.py @@ -397,24 +397,24 @@ class ResourceRecord(BypassInitStateObject): return ".".join(labels) @classmethod - def A(cls, name: str, ip: IPv4Address, *, ttl = DEFAULT_TTL) -> ResourceRecord: + def A(cls, name: str, ip: IPv4Address, *, ttl: int = DEFAULT_TTL) -> ResourceRecord: """Create an IPv4 resource record.""" return ResourceRecord(name, Type.A, Class.IN, ttl, ip.packed) @classmethod - def AAAA(cls, name: str, ip: IPv6Address, *, ttl = DEFAULT_TTL) -> ResourceRecord: + def AAAA(cls, name: str, ip: IPv6Address, *, ttl: int = DEFAULT_TTL) -> ResourceRecord: """Create an IPv6 resource record.""" return ResourceRecord(name, Type.AAAA, Class.IN, ttl, ip.packed) @classmethod - def CNAME(cls, alias: str, canonical: str, *, ttl = DEFAULT_TTL) -> ResourceRecord: + def CNAME(cls, alias: str, canonical: str, *, ttl: int = DEFAULT_TTL) -> ResourceRecord: """Create a canonical internet name resource record.""" - return ResourceRecord(alias, Type.CNAME, Class.IN, ttl, ResourceRecord.encode_domain_name(canonical)) + return ResourceRecord(alias, Type.CNAME, Class.IN, ttl, ResourceRecord.pack_domain_name(canonical)) @classmethod - def PTR(cls, inaddr: str, ptr: str, *, ttl = DEFAULT_TTL) -> ResourceRecord: + def PTR(cls, inaddr: str, ptr: str, *, ttl: int = DEFAULT_TTL) -> ResourceRecord: """Create a canonical internet name resource record.""" - return ResourceRecord(inaddr, Type.PTR, Class.IN, ttl, ResourceRecord.encode_domain_name(ptr)) + return ResourceRecord(inaddr, Type.PTR, Class.IN, ttl, ResourceRecord.pack_domain_name(ptr)) # comments are taken from rfc1035 @@ -628,7 +628,7 @@ class Message(BypassInitStateObject): except struct.error as e: raise struct.error(f"question #{i}: {str(e)}") - def unpack_rrs(section: List[ResourceRecord], section_name: str, count: int) -> int: + def unpack_rrs(section: List[ResourceRecord], section_name: str, count: int) -> None: nonlocal buffer, offset for i in range(0, count): try: diff --git a/mitmproxy/proxy/layers/dns.py b/mitmproxy/proxy/layers/dns.py index b08ec3fca..3b24c5f62 100644 --- a/mitmproxy/proxy/layers/dns.py +++ b/mitmproxy/proxy/layers/dns.py @@ -62,53 +62,60 @@ class DNSLayer(layer.Layer): @classmethod def simple_resolve(cls, questions: List[dns.Question]) -> List[dns.ResourceRecord]: - answers = [] + answers: List[dns.ResourceRecord] = [] - def resolve_by_name(family: socket.AddressFamily, ip: Callable[[str], Union[ipaddress.IPv4Address, ipaddress.IPv6Address]]) -> None: - nonlocal answers, question - try: - addrinfos = socket.getaddrinfo(host=question.name, port=0, family=family) - except socket.gaierror as e: - if e.errno == socket.EAI_NODATA: + for question in questions: + + def resolve_by_name( + family: socket.AddressFamily, + ip: Callable[[str], Union[ipaddress.IPv4Address, ipaddress.IPv6Address]] + ) -> None: + nonlocal answers, question + try: + addrinfos = socket.getaddrinfo(host=question.name, port=0, family=family) + except socket.gaierror as e: + if e.errno == socket.EAI_NODATA: + raise DnsResolveError(dns.ResponseCode.NXDOMAIN) + else: + # NOTE might fail on Windows for IPv6 queries: + # https://stackoverflow.com/questions/66755681/getaddrinfo-c-on-windows-not-handling-ipv6-correctly-returning-error-code-1 + raise DnsResolveError(dns.ResponseCode.SERVFAIL) + for addrinfo in addrinfos: + _, _, _, _, addr = addrinfo + answers.append(dns.ResourceRecord( + name=question.name, + type=question.type, + class_=question.class_, + ttl=dns.ResourceRecord.DEFAULT_TTL, + data=ip(addr[0]).packed, + )) + + def resolve_by_addr( + suffix: str, + ip: Callable[[List[str]], Union[ipaddress.IPv4Address, ipaddress.IPv6Address]] + ) -> bool: + nonlocal answers, question + if not question.name.lower().endswith(suffix.lower()): + return False + try: + addr = ip(question.name[0:-len(suffix)].split(".")[::-1]) + except ValueError: + raise DnsResolveError(dns.ResponseCode.FORMERR) + try: + name, _, _ = socket.gethostbyaddr(str(addr)) + except socket.herror: raise DnsResolveError(dns.ResponseCode.NXDOMAIN) - else: - # NOTE might fail on Windows for IPv6 queries: - # https://stackoverflow.com/questions/66755681/getaddrinfo-c-on-windows-not-handling-ipv6-correctly-returning-error-code-1 + except socket.gaierror: raise DnsResolveError(dns.ResponseCode.SERVFAIL) - for addrinfo in addrinfos: - _, _, _, _, (addr, _) = addrinfo answers.append(dns.ResourceRecord( name=question.name, type=question.type, class_=question.class_, ttl=dns.ResourceRecord.DEFAULT_TTL, - data=ip(addr).packed, + data=dns.ResourceRecord.pack_domain_name(name), )) + return True - def resolve_by_addr(suffix: str, ip: Callable[[List[str]], Union[ipaddress.IPv4Address, ipaddress.IPv6Address]]) -> bool: - nonlocal answers, question - if not question.name.lower().endswith(suffix.lower()): - return False - try: - addr = ip(question.name[0:-len(suffix)].split(".")[::-1]) - except ValueError: - raise DnsResolveError(dns.ResponseCode.FORMERR) - try: - name, _, _ = socket.gethostbyaddr(str(addr)) - except socket.herror: - raise DnsResolveError(dns.ResponseCode.NXDOMAIN) - except socket.gaierror: - raise DnsResolveError(dns.ResponseCode.SERVFAIL) - answers.append(dns.ResourceRecord( - name=question.name, - type=question.type, - class_=question.class_, - ttl=dns.ResourceRecord.DEFAULT_TTL, - data=dns.ResourceRecord.pack_domain_name(name), - )) - return True - - for question in questions: if question.class_ is not dns.Class.IN: raise DnsResolveError(dns.ResponseCode.NOTIMP) if question.type is dns.Type.A: diff --git a/mitmproxy/proxy/server.py b/mitmproxy/proxy/server.py index 12669f19f..2a47c6f3f 100644 --- a/mitmproxy/proxy/server.py +++ b/mitmproxy/proxy/server.py @@ -154,9 +154,9 @@ class ConnectionHandler(metaclass=abc.ABCMeta): try: command.connection.timestamp_start = time.time() open_connection = ( - asyncio.open_connection if command.connection.protocol == ConnectionProtocol.TCP + asyncio.open_connection if command.connection.protocol is ConnectionProtocol.TCP else - udp.open_connection if command.connection.protocol == ConnectionProtocol.UDP + udp.open_connection if command.connection.protocol is ConnectionProtocol.UDP else None )