Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
dns_resolver.py183 linesDownload Raw Back to addons
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 
codekingpro/portable-devtools · Team Ai