codekingpro/portable-devtools
114k
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 