codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import asyncio4import ipaddress5import logging6import socket7from collections.abc import Sequence8from functools import cache9from typing import Protocol10 11import mitmproxy_rs12from mitmproxy import ctx13from mitmproxy import dns14from mitmproxy.flow import Error15from mitmproxy.proxy import mode_specs16 17logger = logging.getLogger(__name__)18 19 20class DnsResolver:21 def load(self, loader):22 loader.add_option(23 "dns_use_hosts_file",24 bool,25 True,26 "Use the hosts file for DNS lookups in regular DNS mode/wireguard mode.",27 )28 29 loader.add_option(30 "dns_name_servers",31 Sequence[str],32 [],33 "Name servers to use for lookups in regular DNS mode/wireguard mode. Default: operating system's name servers",34 )35 36 def configure(self, updated):37 if "dns_use_hosts_file" in updated or "dns_name_servers" in updated:38 self.resolver.cache_clear()39 self.name_servers.cache_clear()40 41 @cache42 def name_servers(self) -> list[str]:43 """44 Returns the operating system's name servers unless custom name servers are set.45 On error, an empty list is returned.46 """47 try:48 return (49 ctx.options.dns_name_servers50 or mitmproxy_rs.dns.get_system_dns_servers()51 )52 except RuntimeError as e:53 logger.warning(54 f"Failed to get system dns servers: {e}\n"55 f"The dns_name_servers option needs to be set manually."56 )57 return []58 59 @cache60 def resolver(self) -> Resolver:61 """62 Returns:63 The DNS resolver to use.64 Raises:65 MissingNameServers, if name servers are unknown and `dns_use_hosts_file` is disabled.66 """67 if ns := self.name_servers():68 # We always want to use our own resolver if name server info is available.69 return mitmproxy_rs.dns.DnsResolver(70 name_servers=ns,71 use_hosts_file=ctx.options.dns_use_hosts_file,72 )73 elif ctx.options.dns_use_hosts_file:74 # Fallback to getaddrinfo as hickory's resolver isn't as reliable75 # as we would like it to be (https://github.com/mitmproxy/mitmproxy/issues/7064).76 return GetaddrinfoFallbackResolver()77 else:78 raise MissingNameServers()79 80 async def dns_request(self, flow: dns.DNSFlow) -> None:81 if self._should_resolve(flow):82 all_ip_lookups = (83 flow.request.query84 and flow.request.op_code == dns.op_codes.QUERY85 and flow.request.question86 and flow.request.question.class_ == dns.classes.IN87 and flow.request.question.type in (dns.types.A, dns.types.AAAA)88 )89 if all_ip_lookups:90 try:91 flow.response = await self.resolve(flow.request)92 except MissingNameServers:93 flow.error = Error("Cannot resolve, dns_name_servers unknown.")94 elif name_servers := self.name_servers():95 # For other records, the best we can do is to forward the query96 # to an upstream server.97 flow.server_conn.address = (name_servers[0], 53)98 else:99 flow.error = Error("Cannot resolve, dns_name_servers unknown.")100 101 @staticmethod102 def _should_resolve(flow: dns.DNSFlow) -> bool:103 return (104 (105 isinstance(flow.client_conn.proxy_mode, mode_specs.DnsMode)106 or (107 isinstance(flow.client_conn.proxy_mode, mode_specs.WireGuardMode)108 and flow.server_conn.address == ("10.0.0.53", 53)109 )110 )111 and flow.live112 and not flow.response113 and not flow.error114 )115 116 async def resolve(117 self,118 message: dns.DNSMessage,119 ) -> dns.DNSMessage:120 q = message.question121 assert q122 try:123 if q.type == dns.types.A:124 ip_addrs = await self.resolver().lookup_ipv4(q.name)125 else:126 ip_addrs = await self.resolver().lookup_ipv6(q.name)127 except socket.gaierror as e:128 match e.args[0]:129 case socket.EAI_NONAME:130 return message.fail(dns.response_codes.NXDOMAIN)131 case socket.EAI_NODATA:132 ip_addrs = []133 case _:134 return message.fail(dns.response_codes.SERVFAIL)135 136 return message.succeed(137 [138 dns.ResourceRecord(139 name=q.name,140 type=q.type,141 class_=q.class_,142 ttl=dns.ResourceRecord.DEFAULT_TTL,143 data=ipaddress.ip_address(ip).packed,144 )145 for ip in ip_addrs146 ]147 )148 149 150class Resolver(Protocol):151 async def lookup_ip(self, domain: str) -> list[str]: # pragma: no cover152 ...153 154 async def lookup_ipv4(self, domain: str) -> list[str]: # pragma: no cover155 ...156 157 async def lookup_ipv6(self, domain: str) -> list[str]: # pragma: no cover158 ...159 160 161class GetaddrinfoFallbackResolver(Resolver):162 async def lookup_ip(self, domain: str) -> list[str]:163 return await self._lookup(domain, socket.AF_UNSPEC)164 165 async def lookup_ipv4(self, domain: str) -> list[str]:166 return await self._lookup(domain, socket.AF_INET)167 168 async def lookup_ipv6(self, domain: str) -> list[str]:169 return await self._lookup(domain, socket.AF_INET6)170 171 async def _lookup(self, domain: str, family: socket.AddressFamily) -> list[str]:172 addrinfos = await asyncio.get_running_loop().getaddrinfo(173 host=domain,174 port=None,175 family=family,176 type=socket.SOCK_STREAM,177 )178 return [addrinfo[4][0] for addrinfo in addrinfos]179 180 181class MissingNameServers(RuntimeError):182 pass183 