userver: en/testsuite/asyncio_socket.py Source File
Loading...
Searching...
No Matches
asyncio_socket.py
1# mypy: disable-error-code=attr-defined
2
3
4import asyncio
5import select
6import socket
7import sys
8
9DEFAULT_TIMEOUT = 10.0
10_DefaultTimeout = object()
11
12_ASYNCIO_HAS_SENDTO = sys.version_info >= (3, 11)
13
14
16 def __init__(
17 self,
18 sock: socket.socket,
19 loop: asyncio.AbstractEventLoop | None = None,
20 timeout=DEFAULT_TIMEOUT,
21 ):
22 if loop is None:
23 loop = asyncio.get_running_loop()
24 self._loop: asyncio.AbstractEventLoop = loop
25 self._sock: socket.socket = sock
26 self._default_timeout = timeout
27 sock.setblocking(False)
28
29 def __repr__(self):
30 return f'<AsyncioSocket for {self._sock}>'
31
32 def __enter__(self):
33 return self
34
35 def __exit__(self, *args):
36 self.close()
37
38 @property
39 def type(self):
40 return self._sock.type
41
42 @property
43 def socket(self) -> socket.socket:
44 return self._sock
45
46 def fileno(self) -> int:
47 return self._sock.fileno()
48
49 async def connect(self, address, *, timeout=_DefaultTimeout):
50 coro = self._loop.sock_connect(self._sock, address)
51 return await self._with_timeout(coro, timeout=timeout)
52
53 async def sendto(self, *args, timeout=_DefaultTimeout):
54 # asyncio support added in version 3.11
55 if not _ASYNCIO_HAS_SENDTO:
56 return await self._sendto_legacy(*args, timeout=timeout)
57 coro = self._loop.sock_sendto(self._sock, *args)
58 try:
59 return await self._with_timeout(coro, timeout=timeout)
60 except NotImplementedError:
61 return await self._sendto_legacy(*args, timeout=timeout)
62
63 async def sendall(self, data, *, timeout=_DefaultTimeout):
64 coro = self._loop.sock_sendall(self._sock, data)
65 return await self._with_timeout(coro, timeout=timeout)
66
67 async def recv(self, size, *, timeout=_DefaultTimeout):
68 coro = self._loop.sock_recv(self._sock, size)
69 return await self._with_timeout(coro, timeout=timeout)
70
71 async def recvfrom(self, *args, timeout=_DefaultTimeout):
72 # asyncio support added in version 3.11
73 if not _ASYNCIO_HAS_SENDTO:
74 return await self._recvfrom_legacy(*args, timeout=timeout)
75 coro = self._loop.sock_recvfrom(self._sock, *args)
76 try:
77 return await self._with_timeout(coro, timeout=timeout)
78 except NotImplementedError:
79 return await self._recvfrom_legacy(*args, timeout=timeout)
80
81 async def accept(self, *, timeout=_DefaultTimeout):
82 coro = self._loop.sock_accept(self._sock)
83 conn, address = await self._with_timeout(coro, timeout=timeout)
84 return from_socket(conn), address
85
86 def bind(self, address):
87 return self._sock.bind(address)
88
89 def listen(self, *args):
90 return self._sock.listen(*args)
91
92 def getsockname(self):
93 return self._sock.getsockname()
94
95 def setsockopt(self, *args, **kwargs):
96 self._sock.setsockopt(*args, **kwargs)
97
98 def close(self):
99 self._sock.close()
100
101 def has_data(self) -> bool:
102 rlist, _, _ = select.select([self._sock], [], [], 0)
103 return bool(rlist)
104
105 def can_write(self) -> bool:
106 _, wlist, _ = select.select([], [self._sock], [], 0)
107 return bool(wlist)
108
109 async def wait_for_data(self, timeout=_DefaultTimeout):
110 if self.has_data():
111 return
112 coro = _wait_for_data(self._loop, self._sock)
113 return await self._with_timeout(coro, timeout=timeout)
114
115 async def _with_timeout(self, awaitable, timeout):
116 # TODO(python3.11): switch to `asyncio.timeout()`
117 if timeout is _DefaultTimeout:
118 timeout = self._default_timeout
119
120 return await asyncio.wait_for(awaitable, timeout=timeout)
121
122 async def _sendto_legacy(self, *args, timeout):
123 # uvloop and python < 3.11
124 try:
125 return self._sock.sendto(*args)
126 except (BlockingIOError, InterruptedError):
127 pass
128 fut = self._loop.create_future()
129 try:
130 self._loop.add_writer(
131 self.fileno(),
132 _legacy_io_handler,
133 fut,
134 self._sock.sendto,
135 *args,
136 )
137 return await self._with_timeout(fut, timeout=timeout)
138 finally:
139 self._loop.remove_writer(self.fileno())
140
141 async def _recvfrom_legacy(self, *args, timeout):
142 # uvloop and python < 3.11
143 try:
144 return self._sock.recvfrom(*args)
145 except (BlockingIOError, InterruptedError):
146 pass
147 fut = self._loop.create_future()
148 try:
149 self._loop.add_reader(
150 self.fileno(),
151 _legacy_io_handler,
152 fut,
153 self._sock.recvfrom,
154 *args,
155 )
156 return await self._with_timeout(fut, timeout=timeout)
157 finally:
158 self._loop.remove_reader(self.fileno())
159
160
162 def __init__(self, loop=None):
163 if loop is None:
164 loop = asyncio.get_running_loop()
165 self._loop = loop
166
167 def from_socket(self, sock, timeout=DEFAULT_TIMEOUT):
168 return from_socket(sock, loop=self._loop, timeout=timeout)
169
170 def socket(self, *args, timeout=DEFAULT_TIMEOUT):
171 sock = socket.socket(*args)
172 return self.from_socket(sock, timeout=timeout)
173
174 async def getaddrinfo(self, *args, timeout=DEFAULT_TIMEOUT, **kwargs):
175 coro = self._loop.getaddrinfo(*args, **kwargs)
176 return await asyncio.wait_for(coro, timeout=timeout)
177
178 def tcp(self, *, timeout=DEFAULT_TIMEOUT):
179 return self.socket(socket.AF_INET, socket.SOCK_STREAM, timeout=timeout)
180
181 def udp(self, *, timeout=DEFAULT_TIMEOUT):
182 return self.socket(socket.AF_INET, socket.SOCK_DGRAM, timeout=timeout)
183
184 def socketpair(self, *args, timeout=DEFAULT_TIMEOUT, **kwargs):
185 sock1, sock2 = socket.socketpair(*args, **kwargs)
186 return self.from_socket(sock1, timeout=timeout), self.from_socket(
187 sock2, timeout=timeout
188 )
189
190
191def from_socket(
192 sock: socket.socket | AsyncioSocket,
193 *,
194 loop=None,
195 timeout=DEFAULT_TIMEOUT,
196) -> AsyncioSocket:
197 if isinstance(sock, AsyncioSocket):
198 return sock
199 return AsyncioSocket(sock, loop=loop, timeout=timeout)
200
201
202def create_socket(*args, timeout=DEFAULT_TIMEOUT):
203 return AsyncioSocketsFactory().socket(*args, timeout=timeout)
204
205
206def create_tcp_socket(*args, timeout=DEFAULT_TIMEOUT):
207 return AsyncioSocketsFactory().tcp(timeout=timeout)
208
209
210def create_udp_socket(timeout=DEFAULT_TIMEOUT):
211 return AsyncioSocketsFactory().udp(timeout=timeout)
212
213
214def create_socketpair(*args, timeout=DEFAULT_TIMEOUT, **kwargs):
215 return AsyncioSocketsFactory().socketpair(*args, **kwargs, timeout=timeout)
216
217
218async def getaddrinfo(*args, **kwargs):
219 return await AsyncioSocketsFactory().getaddrinfo(*args, **kwargs)
220
221
222def _legacy_io_handler(fut, sock_handler, *args):
223 if fut.done():
224 return
225 try:
226 data = sock_handler(*args)
227 except (BlockingIOError, InterruptedError):
228 return # try again next time
229 except (SystemExit, KeyboardInterrupt):
230 raise
231 except BaseException as exc:
232 fut.set_exception(exc)
233 else:
234 fut.set_result(data)
235
236
237async def _wait_for_data(loop, sock):
238 fut = loop.create_future()
239 fd = sock.fileno()
240
241 def on_data():
242 if fut.done():
243 return
244 fut.set_result(None)
245
246 try:
247 loop.add_reader(fd, on_data)
248 await fut
249 finally:
250 loop.remove_reader(fd)