Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_dns.py224 linesDownload Raw Back to psycopg
1# type: ignore  # dnspython is currently optional and mypy fails if missing2"""3DNS query support4"""5 6# Copyright (C) 2021 The Psycopg Team7 8import os9import re10import warnings11from random import randint12from typing import Any, DefaultDict, Dict, List, NamedTuple, Optional, Sequence13from typing import TYPE_CHECKING14from collections import defaultdict15 16try:17    from dns.resolver import Resolver, Cache18    from dns.asyncresolver import Resolver as AsyncResolver19    from dns.exception import DNSException20except ImportError:21    raise ImportError(22        "the module psycopg._dns requires the package 'dnspython' installed"23    )24 25from . import errors as e26from .conninfo import resolve_hostaddr_async as resolve_hostaddr_async_27 28if TYPE_CHECKING:29    from dns.rdtypes.IN.SRV import SRV30 31resolver = Resolver()32resolver.cache = Cache()33 34async_resolver = AsyncResolver()35async_resolver.cache = Cache()36 37 38async def resolve_hostaddr_async(params: Dict[str, Any]) -> Dict[str, Any]:39    """40    Perform async DNS lookup of the hosts and return a new params dict.41 42    .. deprecated:: 3.143        The use of this function is not necessary anymore, because44        `psycopg.AsyncConnection.connect()` performs non-blocking name45        resolution automatically.46    """47    warnings.warn(48        "from psycopg 3.1, resolve_hostaddr_async() is not needed anymore",49        DeprecationWarning,50    )51    return await resolve_hostaddr_async_(params)52 53 54def resolve_srv(params: Dict[str, Any]) -> Dict[str, Any]:55    """Apply SRV DNS lookup as defined in :RFC:`2782`."""56    return Rfc2782Resolver().resolve(params)57 58 59async def resolve_srv_async(params: Dict[str, Any]) -> Dict[str, Any]:60    """Async equivalent of `resolve_srv()`."""61    return await Rfc2782Resolver().resolve_async(params)62 63 64class HostPort(NamedTuple):65    host: str66    port: str67    totry: bool = False68    target: Optional[str] = None69 70 71class Rfc2782Resolver:72    """Implement SRV RR Resolution as per RFC 278273 74    The class is organised to minimise code duplication between the sync and75    the async paths.76    """77 78    re_srv_rr = re.compile(r"^(?P<service>_[^\.]+)\.(?P<proto>_[^\.]+)\.(?P<target>.+)")79 80    def resolve(self, params: Dict[str, Any]) -> Dict[str, Any]:81        """Update the parameters host and port after SRV lookup."""82        attempts = self._get_attempts(params)83        if not attempts:84            return params85 86        hps = []87        for hp in attempts:88            if hp.totry:89                hps.extend(self._resolve_srv(hp))90            else:91                hps.append(hp)92 93        return self._return_params(params, hps)94 95    async def resolve_async(self, params: Dict[str, Any]) -> Dict[str, Any]:96        """Update the parameters host and port after SRV lookup."""97        attempts = self._get_attempts(params)98        if not attempts:99            return params100 101        hps = []102        for hp in attempts:103            if hp.totry:104                hps.extend(await self._resolve_srv_async(hp))105            else:106                hps.append(hp)107 108        return self._return_params(params, hps)109 110    def _get_attempts(self, params: Dict[str, Any]) -> List[HostPort]:111        """112        Return the list of host, and for each host if SRV lookup must be tried.113 114        Return an empty list if no lookup is requested.115        """116        # If hostaddr is defined don't do any resolution.117        if params.get("hostaddr", os.environ.get("PGHOSTADDR", "")):118            return []119 120        host_arg: str = params.get("host", os.environ.get("PGHOST", ""))121        hosts_in = host_arg.split(",")122        port_arg: str = str(params.get("port", os.environ.get("PGPORT", "")))123        ports_in = port_arg.split(",")124 125        if len(ports_in) == 1:126            # If only one port is specified, it applies to all the hosts.127            ports_in *= len(hosts_in)128        if len(ports_in) != len(hosts_in):129            # ProgrammingError would have been more appropriate, but this is130            # what the raise if the libpq fails connect in the same case.131            raise e.OperationalError(132                f"cannot match {len(hosts_in)} hosts with {len(ports_in)} port numbers"133            )134 135        out = []136        srv_found = False137        for host, port in zip(hosts_in, ports_in):138            m = self.re_srv_rr.match(host)139            if m or port.lower() == "srv":140                srv_found = True141                target = m.group("target") if m else None142                hp = HostPort(host=host, port=port, totry=True, target=target)143            else:144                hp = HostPort(host=host, port=port)145            out.append(hp)146 147        return out if srv_found else []148 149    def _resolve_srv(self, hp: HostPort) -> List[HostPort]:150        try:151            ans = resolver.resolve(hp.host, "SRV")152        except DNSException:153            ans = ()154        return self._get_solved_entries(hp, ans)155 156    async def _resolve_srv_async(self, hp: HostPort) -> List[HostPort]:157        try:158            ans = await async_resolver.resolve(hp.host, "SRV")159        except DNSException:160            ans = ()161        return self._get_solved_entries(hp, ans)162 163    def _get_solved_entries(164        self, hp: HostPort, entries: "Sequence[SRV]"165    ) -> List[HostPort]:166        if not entries:167            # No SRV entry found. Delegate the libpq a QNAME=target lookup168            if hp.target and hp.port.lower() != "srv":169                return [HostPort(host=hp.target, port=hp.port)]170            else:171                return []172 173        # If there is precisely one SRV RR, and its Target is "." (the root174        # domain), abort.175        if len(entries) == 1 and str(entries[0].target) == ".":176            return []177 178        return [179            HostPort(host=str(entry.target).rstrip("."), port=str(entry.port))180            for entry in self.sort_rfc2782(entries)181        ]182 183    def _return_params(184        self, params: Dict[str, Any], hps: List[HostPort]185    ) -> Dict[str, Any]:186        if not hps:187            # Nothing found, we ended up with an empty list188            raise e.OperationalError("no host found after SRV RR lookup")189 190        out = params.copy()191        out["host"] = ",".join(hp.host for hp in hps)192        out["port"] = ",".join(str(hp.port) for hp in hps)193        return out194 195    def sort_rfc2782(self, ans: "Sequence[SRV]") -> "List[SRV]":196        """197        Implement the priority/weight ordering defined in RFC 2782.198        """199        # Divide the entries by priority:200        priorities: DefaultDict[int, "List[SRV]"] = defaultdict(list)201        out: "List[SRV]" = []202        for entry in ans:203            priorities[entry.priority].append(entry)204 205        for pri, entries in sorted(priorities.items()):206            if len(entries) == 1:207                out.append(entries[0])208                continue209 210            entries.sort(key=lambda ent: ent.weight)211            total_weight = sum(ent.weight for ent in entries)212            while entries:213                r = randint(0, total_weight)214                csum = 0215                for i, ent in enumerate(entries):216                    csum += ent.weight217                    if csum >= r:218                        break219                out.append(ent)220                total_weight -= ent.weight221                del entries[i]222 223        return out224 
codekingpro/portable-devtools · Team Ai