userver: en/testsuite/utils/net.py Source File
Loading...
Searching...
No Matches
net.py
1import asyncio
3import contextlib
4import errno
5import pathlib
6import socket
7
8DEFAULT_BACKLOG = 50
9
10
11class MultipleSocketServer(asyncio.events.AbstractServer):
12 def __init__(self, servers):
13 self._servers = servers
14
15 def get_loop(self):
16 return self._servers[0].get_loop()
17
18 def start_serving(self):
19 for server in self._servers:
20 server.start_serving()
21
22 async def serve_forever(self):
23 raise NotImplementedError
24
25 def is_serving(self):
26 raise NotImplementedError
27
28 @property
29 def sockets(self):
30 sockets = []
31 for server in self._servers:
32 sockets.extend(server.sockets)
33 return sockets
34
35 def close(self):
36 for server in self._servers:
37 server.close()
38
39 async def wait_closed(self):
40 for server in self._servers:
41 server.close()
42
43
44def create_tcp_server(
45 factory,
46 *,
47 loop=None,
48 host='localhost',
49 port=0,
50 sock=None,
51 **kwargs,
52):
53 if sock is None:
54 sock = bind_socket(host, port)
55 return _create_server(factory, loop=loop, sock=sock, **kwargs)
56
57
58def create_unix_server(
59 factory,
60 path: pathlib.Path,
61 *,
62 loop=None,
63 sock=None,
64 **kwargs,
65):
66 return _create_unix_server(
67 factory, loop=loop, path=path, sock=sock, **kwargs
68 )
69
70
71@contextlib.asynccontextmanager
72async def create_server_multiple(factory, sockets, *, loop=None, **kwargs):
73 assert sockets
74 if loop is None:
75 loop = asyncio.get_running_loop()
76 servers = []
77 for sock in sockets:
78 server = await loop.create_server(factory, sock=sock, **kwargs)
79 servers.append(server)
80
81 multi = MultipleSocketServer(servers)
82 try:
83 yield multi
84 finally:
85 multi.close()
86
87
88async def start_multiple_servers(
89 client_connected_cb, sockets, *, loop=None, **kwargs
90) -> MultipleSocketServer:
91 assert sockets
92
93 servers = []
94 for sock in sockets:
95 server = await asyncio.start_server(
96 client_connected_cb, sock=sock, **kwargs
97 )
98 servers.append(server)
99 return MultipleSocketServer(servers)
100
101
102def bind_socket_multiple(
103 hostname='localhost',
104 port=0,
105 family=socket.AF_UNSPEC,
106 type=socket.SOCK_STREAM,
107 backlog=DEFAULT_BACKLOG,
108 retries=15,
109):
110 """
111 Bind multiple sockets for both IPv4 and IPv6 addresses.
112
113 If `port` is zero tries to bind the same port for all addresses,
114 `retries` times.
115 """
116
117 def bind():
118 return _bind_socket_multiple(
119 hostname,
120 port,
121 family=family,
122 type=type,
123 backlog=backlog,
124 )
125
126 if retries and not port:
127 for _ in range(retries):
128 try:
129 return bind()
130 except socket.error as err:
131 if err.errno == errno.EADDRINUSE:
132 continue
133 return bind()
134
135
136def bind_socket(
137 hostname='localhost',
138 port=0,
139 family=socket.AF_INET,
140 type=socket.SOCK_STREAM,
141 proto=-1,
142 backlog=DEFAULT_BACKLOG,
143):
144 sock = socket.socket(family, type, proto)
145 sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
146 sock.bind((hostname, port))
147 sock.listen(backlog)
148 return sock
149
150
151def bind_unix_socket(
152 socket_path,
153 backlog=DEFAULT_BACKLOG,
154):
155 sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
156 sock.bind(str(socket_path))
157 sock.listen(backlog)
158 return sock
159
160
161@contextlib.contextmanager
162def closing_sockets(sockets):
163 try:
164 yield sockets
165 finally:
166 for sock in sockets:
167 sock.close()
168
169
170@contextlib.contextmanager
171def _close_sockets_on_error(sockets):
172 try:
173 yield sockets
174 except:
175 for sock in sockets:
176 sock.close()
177 raise
178
179
180@contextlib.asynccontextmanager
181async def _create_server(factory, *, loop=None, **kwargs):
182 if loop is None:
183 loop = asyncio.get_running_loop()
184 server = await loop.create_server(factory, **kwargs)
185 try:
186 yield server
187 finally:
188 server.close()
189
190
191@contextlib.asynccontextmanager
192async def _create_unix_server(factory, *, loop=None, **kwargs):
193 if loop is None:
194 loop = asyncio.get_running_loop()
195 server = await loop.create_unix_server(factory, **kwargs)
196 try:
197 yield server
198 finally:
199 server.close()
200
201
202def _bind_socket_multiple(
203 hostname,
204 port,
205 *,
206 family,
207 type,
208 backlog,
209):
210 """
211 Bind multiple sockets for both IPv4 and IPv6 addresses.
212 """
213 infos = socket.getaddrinfo(
214 hostname, port, family=family, type=type, flags=socket.AI_PASSIVE
215 )
216 with _close_sockets_on_error([]) as sockets:
217 for af, socktype, proto, canonname, sa in infos:
218 addr = sa[0]
219 sock = bind_socket(
220 addr, port, family=af, type=socktype, proto=proto
221 )
222 sock_port = sock.getsockname()[1]
223 if port == 0:
224 port = sock_port
225 assert port == sock_port
226 sockets.append(sock)
227 return sockets