Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/aiohttp/resolver.py: 25%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

134 statements  

1import asyncio 

2import socket 

3import sys 

4import weakref 

5from typing import Any, Final, Optional 

6 

7from .abc import AbstractResolver, ResolveResult 

8 

9__all__ = ("ThreadedResolver", "AsyncResolver", "DefaultResolver") 

10 

11 

12try: 

13 import aiodns 

14 

15 aiodns_default = hasattr(aiodns.DNSResolver, "getaddrinfo") 

16except ImportError: # pragma: no cover 

17 aiodns = None # type: ignore[assignment] 

18 aiodns_default = False 

19 

20 

21_NUMERIC_SOCKET_FLAGS = socket.AI_NUMERICHOST | socket.AI_NUMERICSERV 

22_NAME_SOCKET_FLAGS = socket.NI_NUMERICHOST | socket.NI_NUMERICSERV 

23_AI_ADDRCONFIG = socket.AI_ADDRCONFIG 

24if hasattr(socket, "AI_MASK"): 

25 _AI_ADDRCONFIG &= socket.AI_MASK 

26_IS_WINDOWS = sys.platform == "win32" 

27 

28 

29def _is_windows_localhost(host: str) -> bool: 

30 return _IS_WINDOWS and host.rstrip(".").casefold() == "localhost" 

31 

32 

33class ThreadedResolver(AbstractResolver): 

34 """Threaded resolver. 

35 

36 Uses an Executor for synchronous getaddrinfo() calls. 

37 concurrent.futures.ThreadPoolExecutor is used by default. 

38 """ 

39 

40 def __init__(self, loop: asyncio.AbstractEventLoop | None = None) -> None: 

41 self._loop = loop or asyncio.get_running_loop() 

42 

43 async def resolve( 

44 self, host: str, port: int = 0, family: socket.AddressFamily = socket.AF_INET 

45 ) -> list[ResolveResult]: 

46 try: 

47 infos = await self._loop.getaddrinfo( 

48 host, 

49 port, 

50 type=socket.SOCK_STREAM, 

51 family=family, 

52 flags=_AI_ADDRCONFIG, 

53 ) 

54 except socket.gaierror: 

55 if not _is_windows_localhost(host): 

56 raise 

57 infos = await self._loop.getaddrinfo( 

58 host, 

59 port, 

60 type=socket.SOCK_STREAM, 

61 family=family, 

62 flags=0, 

63 ) 

64 

65 hosts: list[ResolveResult] = [] 

66 for family, _, proto, _, address in infos: 

67 if family == socket.AF_INET6: 

68 if len(address) < 3: 

69 # IPv6 is not supported by Python build, 

70 # or IPv6 is not enabled in the host 

71 continue 

72 if address[3]: 

73 # This is essential for link-local IPv6 addresses. 

74 # LL IPv6 is a VERY rare case. Strictly speaking, we should use 

75 # getnameinfo() unconditionally, but performance makes sense. 

76 resolved_host, _port = await self._loop.getnameinfo( 

77 address, _NAME_SOCKET_FLAGS 

78 ) 

79 port = int(_port) 

80 else: 

81 resolved_host, port = address[:2] 

82 else: # IPv4 

83 assert family == socket.AF_INET 

84 resolved_host, port = address # type: ignore[misc] 

85 hosts.append( 

86 ResolveResult( 

87 hostname=host, 

88 host=resolved_host, 

89 port=port, 

90 family=family, 

91 proto=proto, 

92 flags=_NUMERIC_SOCKET_FLAGS, 

93 ) 

94 ) 

95 

96 return hosts 

97 

98 async def close(self) -> None: 

99 pass 

100 

101 

102class AsyncResolver(AbstractResolver): 

103 """Use the `aiodns` package to make asynchronous DNS lookups""" 

104 

105 def __init__( 

106 self, 

107 loop: asyncio.AbstractEventLoop | None = None, 

108 *args: Any, 

109 **kwargs: Any, 

110 ) -> None: 

111 if aiodns is None: 

112 raise RuntimeError("Resolver requires aiodns library") 

113 

114 self._loop = loop or asyncio.get_running_loop() 

115 self._manager: _DNSResolverManager | None = None 

116 # If custom args are provided, create a dedicated resolver instance 

117 # This means each AsyncResolver with custom args gets its own 

118 # aiodns.DNSResolver instance 

119 if args or kwargs: 

120 self._resolver = aiodns.DNSResolver(*args, **kwargs) 

121 return 

122 # Use the shared resolver from the manager for default arguments 

123 self._manager = _DNSResolverManager() 

124 self._resolver = self._manager.get_resolver(self, self._loop) 

125 

126 if not hasattr(self._resolver, "gethostbyname"): 

127 # aiodns 1.1 is not available, fallback to DNSResolver.query 

128 self.resolve = self._resolve_with_query # type: ignore 

129 

130 async def resolve( 

131 self, host: str, port: int = 0, family: socket.AddressFamily = socket.AF_INET 

132 ) -> list[ResolveResult]: 

133 try: 

134 try: 

135 resp = await self._resolver.getaddrinfo( 

136 host, 

137 port=port, 

138 type=socket.SOCK_STREAM, 

139 family=family, 

140 flags=_AI_ADDRCONFIG, 

141 ) 

142 except aiodns.error.DNSError: 

143 if not _is_windows_localhost(host): 

144 raise 

145 resp = await self._resolver.getaddrinfo( 

146 host, 

147 port=port, 

148 type=socket.SOCK_STREAM, 

149 family=family, 

150 flags=0, 

151 ) 

152 except aiodns.error.DNSError as exc: 

153 msg = exc.args[1] if len(exc.args) >= 1 else "DNS lookup failed" 

154 raise OSError(None, msg) from exc 

155 hosts: list[ResolveResult] = [] 

156 for node in resp.nodes: 

157 address: tuple[bytes, int] | tuple[bytes, int, int, int] = node.addr 

158 if node.family == socket.AF_INET6: 

159 if len(address) > 3 and address[3]: 

160 # This is essential for link-local IPv6 addresses. 

161 # LL IPv6 is a VERY rare case. Strictly speaking, we should use 

162 # getnameinfo() unconditionally, but performance makes sense. 

163 result = await self._resolver.getnameinfo( 

164 (address[0].decode("ascii"), *address[1:]), 

165 _NAME_SOCKET_FLAGS, 

166 ) 

167 resolved_host = result.node 

168 else: 

169 resolved_host = address[0].decode("ascii") 

170 port = address[1] 

171 else: # IPv4 

172 assert node.family == socket.AF_INET 

173 resolved_host = address[0].decode("ascii") 

174 port = address[1] 

175 hosts.append( 

176 ResolveResult( 

177 hostname=host, 

178 host=resolved_host, 

179 port=port, 

180 family=node.family, 

181 proto=0, 

182 flags=_NUMERIC_SOCKET_FLAGS, 

183 ) 

184 ) 

185 

186 if not hosts: 

187 raise OSError(None, "DNS lookup failed") 

188 

189 return hosts 

190 

191 async def _resolve_with_query( 

192 self, host: str, port: int = 0, family: int = socket.AF_INET 

193 ) -> list[dict[str, Any]]: 

194 qtype: Final = "AAAA" if family == socket.AF_INET6 else "A" 

195 

196 try: 

197 resp = await self._resolver.query(host, qtype) 

198 except aiodns.error.DNSError as exc: 

199 msg = exc.args[1] if len(exc.args) >= 1 else "DNS lookup failed" 

200 raise OSError(None, msg) from exc 

201 

202 hosts = [] 

203 for rr in resp: 

204 hosts.append( 

205 { 

206 "hostname": host, 

207 "host": rr.host, 

208 "port": port, 

209 "family": family, 

210 "proto": 0, 

211 "flags": socket.AI_NUMERICHOST, 

212 } 

213 ) 

214 

215 if not hosts: 

216 raise OSError(None, "DNS lookup failed") 

217 

218 return hosts 

219 

220 async def close(self) -> None: 

221 if self._manager: 

222 # Release the resolver from the manager if using the shared resolver 

223 self._manager.release_resolver(self, self._loop) 

224 self._manager = None # Clear reference to manager 

225 self._resolver = None # type: ignore[assignment] # Clear reference to resolver 

226 return 

227 # Otherwise cancel our dedicated resolver 

228 if self._resolver is not None: 

229 self._resolver.cancel() 

230 self._resolver = None # type: ignore[assignment] # Clear reference 

231 

232 

233class _DNSResolverManager: 

234 """Manager for aiodns.DNSResolver objects. 

235 

236 This class manages shared aiodns.DNSResolver instances 

237 with no custom arguments across different event loops. 

238 """ 

239 

240 _instance: Optional["_DNSResolverManager"] = None 

241 

242 def __new__(cls) -> "_DNSResolverManager": 

243 if cls._instance is None: 

244 cls._instance = super().__new__(cls) 

245 cls._instance._init() 

246 return cls._instance 

247 

248 def _init(self) -> None: 

249 # Use WeakKeyDictionary to allow event loops to be garbage collected 

250 self._loop_data: weakref.WeakKeyDictionary[ 

251 asyncio.AbstractEventLoop, 

252 tuple[aiodns.DNSResolver, weakref.WeakSet[AsyncResolver]], 

253 ] = weakref.WeakKeyDictionary() 

254 

255 def get_resolver( 

256 self, client: "AsyncResolver", loop: asyncio.AbstractEventLoop 

257 ) -> "aiodns.DNSResolver": 

258 """Get or create the shared aiodns.DNSResolver instance for a specific event loop. 

259 

260 Args: 

261 client: The AsyncResolver instance requesting the resolver. 

262 This is required to track resolver usage. 

263 loop: The event loop to use for the resolver. 

264 """ 

265 # Create a new resolver and client set for this loop if it doesn't exist 

266 if loop not in self._loop_data: 

267 resolver = aiodns.DNSResolver(loop=loop) 

268 client_set: weakref.WeakSet[AsyncResolver] = weakref.WeakSet() 

269 self._loop_data[loop] = (resolver, client_set) 

270 else: 

271 # Get the existing resolver and client set 

272 resolver, client_set = self._loop_data[loop] 

273 

274 # Register this client with the loop 

275 client_set.add(client) 

276 return resolver 

277 

278 def release_resolver( 

279 self, client: "AsyncResolver", loop: asyncio.AbstractEventLoop 

280 ) -> None: 

281 """Release the resolver for an AsyncResolver client when it's closed. 

282 

283 Args: 

284 client: The AsyncResolver instance to release. 

285 loop: The event loop the resolver was using. 

286 """ 

287 # Remove client from its loop's tracking 

288 current_loop_data = self._loop_data.get(loop) 

289 if current_loop_data is None: 

290 return 

291 resolver, client_set = current_loop_data 

292 client_set.discard(client) 

293 # If no more clients for this loop, cancel and remove its resolver 

294 if not client_set: 

295 if resolver is not None: 

296 resolver.cancel() 

297 del self._loop_data[loop] 

298 

299 

300_DefaultType = type[AsyncResolver | ThreadedResolver] 

301DefaultResolver: _DefaultType = AsyncResolver if aiodns_default else ThreadedResolver