userver: /data/code/userver/testsuite/pytest_plugins/pytest_userver/chaos.py Source File
Loading...
Searching...
No Matches
chaos.py
1# pylint: disable=too-many-lines
2"""
3Python module that provides testsuite support for
4chaos tests; see
5@ref scripts/docs/en/userver/chaos_testing.md for an introduction.
6
7@ingroup userver_testsuite
8"""
9
10from __future__ import annotations
11
12import asyncio
13from collections.abc import Callable
14from collections.abc import Coroutine
15import dataclasses
16import functools
17import io
18import logging
19import random
20import re
21import socket
22import time
23from typing import Any
24from typing import TypeAlias
25
26import pytest
27
28from testsuite import asyncio_socket
29from testsuite.utils import callinfo
30
31
32class BaseError(Exception):
33 pass
34
35
37 pass
38
39
40@dataclasses.dataclass(frozen=True)
42 """
43 Class that describes the route for TcpGate or UdpGate.
44
45 Use `port_for_client == 0` to bind to some unused port. In that case the
46 actual address could be retrieved via BaseGate.get_sockname_for_clients().
47
48 @ingroup userver_testsuite
49 """
50
51 name: str
52 host_to_server: str
53 port_to_server: int
54 host_for_client: str = '127.0.0.1'
55 port_for_client: int = 0
56
57
58# @cond
59
60# https://docs.python.org/3/library/socket.html#socket.socket.recv
61RECV_MAX_SIZE = 4096
62MAX_DELAY = 60.0
63
64
65logger = logging.getLogger(__name__)
66
67
68Address: TypeAlias = tuple[str, int]
69EvLoop: TypeAlias = Any
70Socket: TypeAlias = socket.socket
71Interceptor: TypeAlias = Callable[[EvLoop, Socket, Socket], Coroutine[Any, Any, None]]
72
73
74class GateException(Exception):
75 pass
76
77
78class GateInterceptException(Exception):
79 pass
80
81
82async def _intercept_ok(
83 loop: EvLoop,
84 socket_from: Socket,
85 socket_to: Socket,
86) -> None:
87 data = await loop.sock_recv(socket_from, RECV_MAX_SIZE)
88 if not data:
89 raise ConnectionClosedError()
90 await loop.sock_sendall(socket_to, data)
91
92
93async def _intercept_drop(
94 loop: EvLoop,
95 socket_from: Socket,
96 socket_to: Socket,
97) -> None:
98 data = await loop.sock_recv(socket_from, RECV_MAX_SIZE)
99 if not data:
100 raise ConnectionClosedError()
101
102
103async def _intercept_delay(
104 delay: float,
105 loop: EvLoop,
106 socket_from: Socket,
107 socket_to: Socket,
108) -> None:
109 data = await loop.sock_recv(socket_from, RECV_MAX_SIZE)
110 if not data:
111 raise ConnectionClosedError()
112 await asyncio.sleep(delay)
113 await loop.sock_sendall(socket_to, data)
114
115
116async def _intercept_close_on_data(
117 loop: EvLoop,
118 socket_from: Socket,
119 socket_to: Socket,
120) -> None:
121 data = await loop.sock_recv(socket_from, 1)
122 if not data:
123 raise ConnectionClosedError()
124 raise GateInterceptException('Closing socket on data')
125
126
127async def _intercept_corrupt(
128 loop: EvLoop,
129 socket_from: Socket,
130 socket_to: Socket,
131) -> None:
132 data = await loop.sock_recv(socket_from, RECV_MAX_SIZE)
133 if not data:
134 raise ConnectionClosedError()
135 await loop.sock_sendall(socket_to, bytearray([not x for x in data]))
136
137
138class _InterceptBpsLimit:
139 def __init__(self, bytes_per_second: float):
140 assert bytes_per_second >= 1
141 self._bytes_per_second = bytes_per_second
142 self._time_last_added = 0.0
143 self._bytes_left = self._bytes_per_second
144
145 def _update_limit(self) -> None:
146 current_time = time.monotonic()
147 elapsed = current_time - self._time_last_added
148 bytes_addition = self._bytes_per_second * elapsed
149 if bytes_addition > 0:
150 self._bytes_left += bytes_addition
151 self._time_last_added = current_time
152
153 if self._bytes_left > self._bytes_per_second:
154 self._bytes_left = self._bytes_per_second
155
156 async def __call__(
157 self,
158 loop: EvLoop,
159 socket_from: Socket,
160 socket_to: Socket,
161 ) -> None:
162 self._update_limit()
163
164 bytes_to_recv = min(int(self._bytes_left), RECV_MAX_SIZE)
165 if bytes_to_recv > 0:
166 data = await loop.sock_recv(socket_from, bytes_to_recv)
167 if not data:
168 raise ConnectionClosedError()
169 self._bytes_left -= len(data)
170
171 await loop.sock_sendall(socket_to, data)
172 else:
173 logger.info('Socket hits the bytes per second limit')
174 await asyncio.sleep(1.0 / self._bytes_per_second)
175
176
177class _InterceptTimeLimit:
178 def __init__(self, timeout: float, jitter: float):
179 self._sockets: dict[Socket, float] = {}
180 assert timeout >= 0.0
181 self._timeout = timeout
182 assert jitter >= 0.0
183 self._jitter = jitter
184
185 def raise_if_timed_out(self, socket_from: Socket) -> None:
186 if socket_from not in self._sockets:
187 jitter = self._jitter * random.random()
188 expire_at = time.monotonic() + self._timeout + jitter
189 self._sockets[socket_from] = expire_at
190
191 if self._sockets[socket_from] <= time.monotonic():
192 del self._sockets[socket_from]
193 raise GateInterceptException('Socket hits the time limit')
194
195 async def __call__(
196 self,
197 loop: EvLoop,
198 socket_from: Socket,
199 socket_to: Socket,
200 ) -> None:
201 self.raise_if_timed_out(socket_from)
202 await _intercept_ok(loop, socket_from, socket_to)
203
204
205class _InterceptSmallerParts:
206 def __init__(self, max_size: int, sleep_per_packet: float):
207 assert max_size > 0
208 self._max_size = max_size
209 self._sleep_per_packet = sleep_per_packet
210
211 async def __call__(
212 self,
213 loop: EvLoop,
214 socket_from: Socket,
215 socket_to: Socket,
216 ) -> None:
217 data = await loop.sock_recv(socket_from, self._max_size)
218 if not data:
219 raise ConnectionClosedError()
220 await asyncio.sleep(self._sleep_per_packet)
221 await loop.sock_sendall(socket_to, data)
222
223
224class _InterceptConcatPackets:
225 def __init__(self, packet_size: int):
226 assert packet_size >= 0
227 self._packet_size = packet_size
228 self._expire_at: float | None = None
229 self._buf = io.BytesIO()
230
231 async def __call__(
232 self,
233 loop: EvLoop,
234 socket_from: Socket,
235 socket_to: Socket,
236 ) -> None:
237 if self._expire_at is None:
238 self._expire_at = time.monotonic() + MAX_DELAY
239
240 if self._expire_at <= time.monotonic():
241 pytest.fail(
242 f'Failed to make a packet of sufficient size in {MAX_DELAY} '
243 'seconds. Check the test logic, it should end with checking '
244 'that the data was sent and by calling TcpGate function '
245 'to_client_pass() to pass the remaining packets.',
246 )
247
248 data = await loop.sock_recv(socket_from, RECV_MAX_SIZE)
249 if not data:
250 raise ConnectionClosedError()
251 self._buf.write(data)
252 if self._buf.tell() >= self._packet_size:
253 await loop.sock_sendall(socket_to, self._buf.getvalue())
254 self._buf = io.BytesIO()
255 self._expire_at = None
256
257
258class _InterceptBytesLimit:
259 def __init__(self, bytes_limit: int, gate: BaseGate):
260 assert bytes_limit >= 0
261 self._bytes_limit = bytes_limit
262 self._bytes_remain = self._bytes_limit
263 self._gate = gate
264
265 async def __call__(
266 self,
267 loop: EvLoop,
268 socket_from: Socket,
269 socket_to: Socket,
270 ) -> None:
271 data = await loop.sock_recv(socket_from, RECV_MAX_SIZE)
272 if not data:
273 raise ConnectionClosedError()
274 if self._bytes_remain <= len(data):
275 await loop.sock_sendall(socket_to, data[0 : self._bytes_remain])
276 await self._gate.sockets_close()
277 self._bytes_remain = self._bytes_limit
278 raise GateInterceptException('Data transmission limit reached')
279 self._bytes_remain -= len(data)
280 await loop.sock_sendall(socket_to, data)
281
282
283class _InterceptSubstitute:
284 def __init__(self, pattern: str, repl: str, encoding='utf-8'):
285 self._pattern = re.compile(pattern)
286 self._repl = repl
287 self._encoding = encoding
288
289 async def __call__(
290 self,
291 loop: EvLoop,
292 socket_from: Socket,
293 socket_to: Socket,
294 ) -> None:
295 data = await loop.sock_recv(socket_from, RECV_MAX_SIZE)
296 if not data:
297 raise ConnectionClosedError()
298 try:
299 res = self._pattern.sub(self._repl, data.decode(self._encoding))
300 data = res.encode(self._encoding)
301 except UnicodeError:
302 pass
303 await loop.sock_sendall(socket_to, data)
304
305
306async def _cancel_and_join(task: asyncio.Task | None) -> None:
307 if not task or task.cancelled():
308 return
309
310 try:
311 task.cancel()
312 await task
313 except asyncio.CancelledError:
314 return
315 except Exception: # pylint: disable=broad-except
316 logger.exception('Exception in _cancel_and_join')
317
318
319def _make_socket_nonblocking(sock: Socket) -> None:
320 sock.setblocking(False)
321 if sock.type == socket.SOCK_STREAM:
322 sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
323
324
325class _UdpDemuxSocketMock:
326 """
327 Emulates a point-to-point connection over UDP socket
328 with a non-blocking socket interface
329 """
330
331 def gettimeout(self):
332 return self._sock.gettimeout()
333
334 def __init__(self, sock: Socket, peer_address: Address):
335 self._sock: Socket = sock
336 self._peeraddr: Address = peer_address
337
338 sockpair = socket.socketpair(type=socket.SOCK_DGRAM)
339 self._demux_in: Socket = sockpair[0]
340 self._demux_out: Socket = sockpair[1]
341 _make_socket_nonblocking(self._demux_in)
342 _make_socket_nonblocking(self._demux_out)
343 self._is_active: bool = True
344
345 @property
346 def peer_address(self):
347 return self._peeraddr
348
349 async def push(self, loop: EvLoop, data: bytes):
350 return await loop.sock_sendall(self._demux_in, data)
351
352 def is_active(self):
353 return self._is_active
354
355 def close(self):
356 self._is_active = False
357 self._demux_out.close()
358 self._demux_in.close()
359
360 def recvfrom(self, bufsize: int, flags: int = 0):
361 return self._demux_out.recvfrom(bufsize, flags)
362
363 def recv(self, bufsize: int, flags: int = 0):
364 return self._demux_out.recv(bufsize, flags)
365
366 def get_demux_out(self):
367 return self._demux_out
368
369 def fileno(self):
370 return self._demux_out.fileno()
371
372 def send(self, data: bytes):
373 return self._sock.sendto(data, self._peeraddr)
374
375
376class InterceptTask:
377 def __init__(self, socket_from, socket_to, interceptor):
378 self._socket_from = socket_from
379 self._socket_to = socket_to
380 self._condition = asyncio.Condition()
381 self._interceptor = interceptor
382
383 def get_interceptor(self):
384 return self._interceptor
385
386 async def set_interceptor(self, interceptor):
387 async with self._condition:
388 self._interceptor = interceptor
389 self._condition.notify()
390
391 async def run(self):
392 loop = asyncio.get_running_loop()
393 while True:
394 # Applies new interceptors faster.
395 #
396 # To avoid long awaiting on sock_recv in an outdated
397 # interceptor we wait for data before grabbing and applying
398 # the interceptor.
399 await _wait_for_data(self._socket_from)
400
401 # Wait for interceptor attached
402 async with self._condition:
403 interceptor = await self._condition.wait_for(self.get_interceptor)
404
405 logging.trace('running interceptor: %s', interceptor)
406 await interceptor(loop, self._socket_from, self._socket_to)
407
408
409class _SocketsPaired:
410 def __init__(
411 self,
412 proxy_name: str,
413 loop: EvLoop,
414 client: socket.socket | _UdpDemuxSocketMock,
415 server: socket.socket,
416 to_server_intercept: Interceptor,
417 to_client_intercept: Interceptor,
418 ) -> None:
419 self._proxy_name = proxy_name
420
421 self._client = client
422 self._server = server
423
424 self._task_to_server = InterceptTask(client, server, to_server_intercept)
425 self._task_to_client = InterceptTask(server, client, to_client_intercept)
426
427 self._task = asyncio.create_task(self._run())
428 self._interceptor_tasks = []
429
430 async def set_to_server_interceptor(self, interceptor):
431 await self._task_to_server.set_interceptor(interceptor)
432
433 async def set_to_client_interceptor(self, interceptor: Interceptor):
434 await self._task_to_client.set_interceptor(interceptor)
435
436 async def shutdown(self) -> None:
437 # An interceptor cannot join itself, self._run() shuts down anyway.
438 if asyncio.current_task() in self._interceptor_tasks:
439 return
440
441 # A second cancel would interrupt the join in self._run().
442 await _cancel_and_join(self._task)
443
444 def is_active(self) -> bool:
445 return not self._task.done()
446
447 def info(self) -> str:
448 if not self.is_active():
449 return '<inactive>'
450
451 return f'client fd={self._client.fileno()} <=> server fd={self._server.fileno()}'
452
453 async def _run(self):
454 self._interceptor_tasks = [
455 asyncio.create_task(obj.run()) for obj in (self._task_to_server, self._task_to_client)
456 ]
457 try:
458 done, _ = await asyncio.wait(self._interceptor_tasks, return_when=asyncio.FIRST_EXCEPTION)
459 for task in done:
460 task.result()
461 except GateInterceptException as exc:
462 logger.info('In "%s": %s', self._proxy_name, exc)
463 except OSError as exc:
464 logger.error('Exception in "%s": %s', self._proxy_name, exc)
465 except Exception:
466 logger.exception('interceptor failed')
467 finally:
468 for task in self._interceptor_tasks:
469 task.cancel()
470 try:
471 # Cancelled interceptors unregister the fds from the event loop.
472 # Closing before that leaves a stale registration, which breaks
473 # any new socket that reuses the fd.
474 await asyncio.gather(*self._interceptor_tasks, return_exceptions=True)
475 finally:
476 for sock in self._server, self._client:
477 try:
478 sock.close()
479 except OSError:
480 logger.exception('Exception in "%s" on closing %s:', self._proxy_name, sock)
481
482
483# @endcond
484
485
487 """
488 This base class maintain endpoints of two types:
489
490 Server-side endpoints to receive messages from clients. Address of this
491 endpoint is described by (host_for_client, port_for_client).
492
493 Client-side endpoints to forward messages to server. Server must listen on
494 (host_to_server, port_to_server).
495
496 Asynchronously concurrently passes data from client to server and from
497 server to client, allowing intercepting the data, injecting delays and
498 dropping connections.
499
500 @warning Do not use this class itself. Use one of the specifications
501 TcpGate or UdpGate
502
503 @ingroup userver_testsuite
504
505 @see @ref scripts/docs/en/userver/chaos_testing.md
506 """
507
508 _NOT_IMPLEMENTED_MESSAGE = 'Do not use BaseGate itself, use one of specializations TcpGate or UdpGate'
509
510 def __init__(self, route: GateRoute, loop: EvLoop | None = None) -> None:
511 self._route = route
512 if loop is None:
513 loop = asyncio.get_running_loop()
514 self._loop = loop
515
516 self._to_server_intercept: Interceptor = _intercept_ok
517 self._to_client_intercept: Interceptor = _intercept_ok
518
519 self._accept_sockets: list[socket.socket] = []
520 self._accept_tasks: list[asyncio.Task[None]] = []
521
522 self._sockets: set[_SocketsPaired] = set()
523
524 async def __aenter__(self) -> BaseGate:
525 self.start()
526 return self
527
528 async def __aexit__(self, exc_type, exc_value, traceback) -> None:
529 await self.stop()
530
531 def _create_accepting_sockets(self) -> list[Socket]:
532 raise NotImplementedError(self._NOT_IMPLEMENTED_MESSAGE)
533
534 def start(self):
535 """Open the socket and start accepting tasks"""
536 if self._accept_sockets:
537 return
538
540
541 if not self._accept_sockets:
542 raise GateException(
543 f'Could not resolve hostname {self._route.host_for_client}',
544 )
545
546 if self._route.port_for_client == 0:
547 # In case of stop()+start() bind to the same port
548 self._route = GateRoute(
549 name=self._route.name,
550 host_to_server=self._route.host_to_server,
551 port_to_server=self._route.port_to_server,
552 host_for_client=self._accept_sockets[0].getsockname()[0],
553 port_for_client=self._accept_sockets[0].getsockname()[1],
554 )
555
556 self.start_accepting()
557
558 def start_accepting(self) -> None:
559 """Start accepting tasks"""
560 assert self._accept_sockets
561 if not all(tsk.done() for tsk in self._accept_tasks):
562 return
563
564 self._accept_tasks.clear()
565 for sock in self._accept_sockets:
566 self._accept_tasks.append(
567 asyncio.create_task(self._do_accept(sock)),
568 )
569
570 async def stop_accepting(self) -> None:
571 """
572 Stop accepting tasks without closing the accepting socket.
573 """
574 for tsk in self._accept_tasks:
575 await _cancel_and_join(tsk)
576 self._accept_tasks.clear()
577
578 async def stop(self) -> None:
579 """
580 Stop accepting tasks, close all the sockets
581 """
582 if not self._accept_sockets and not self._sockets:
583 return
584
585 await self.to_server_pass()
586 await self.to_client_pass()
587
588 await self.stop_accepting()
589 logger.info('Before close() %s', self.info())
590 await self.sockets_close()
591 assert not self._sockets
592
593 for sock in self._accept_sockets:
594 sock.close()
595 self._accept_sockets.clear()
596 logger.info('Stopped. %s', self.info())
597
598 async def sockets_close(
599 self,
600 *,
601 count: int | None = None,
602 ) -> None:
603 """Close all the connection going through the gate"""
604 for x in list(self._sockets)[0:count]:
605 await x.shutdown()
606 self._collect_garbage()
607
608 def get_sockname_for_clients(self) -> Address:
609 """
610 Returns the client socket address that the gate listens on.
611
612 This function allows to use 0 in GateRoute.port_for_client and retrieve
613 the actual port and host.
614 """
615 assert self._route.port_for_client != 0, ('Gate was not started and the port_for_client is still 0',)
616 return (self._route.host_for_client, self._route.port_for_client)
617
618 def info(self) -> str:
619 """Print info on open sockets"""
620 if not self._sockets:
621 return f'"{self._route.name}" no active sockets'
622
623 return f'"{self._route.name}" active sockets:\n\t' + '\n\t'.join(x.info() for x in self._sockets)
624
625 def _collect_garbage(self) -> None:
626 self._sockets = {x for x in self._sockets if x.is_active()}
627
628 async def _do_accept(self, accept_sock: Socket) -> None:
629 """
630 This task should wait for connections and create SocketPair
631 """
632 raise NotImplementedError(self._NOT_IMPLEMENTED_MESSAGE)
633
634 async def set_to_server_interceptor(self, interceptor: Interceptor) -> callinfo.AsyncCallQueue:
635 """
636 Replace existing interceptor of client to server data with a custom
637 """
638 self._to_server_intercept = _create_callqueue(interceptor)
639 for x in self._sockets:
640 await x.set_to_server_interceptor(self._to_server_intercept)
641 return self._to_server_intercept
642
643 async def set_to_client_interceptor(self, interceptor: Interceptor) -> callinfo.AsyncCallQueue:
644 """
645 Replace existing interceptor of server to client data with a custom
646
647 """
648 if interceptor is not None:
649 self._to_client_intercept = _create_callqueue(interceptor)
650 else:
651 self._to_client_intercept = None
652 for x in self._sockets:
653 await x.set_to_client_interceptor(self._to_client_intercept)
654 return self._to_client_intercept
655
656 async def to_server_pass(self) -> callinfo.AsyncCallQueue:
657 """Pass data as is"""
658 logging.trace('to_server_pass')
659 return await self.set_to_server_interceptor(_intercept_ok)
660
661 async def to_client_pass(self) -> callinfo.AsyncCallQueue:
662 """Pass data as is"""
663 logging.trace('to_client_pass')
664 return await self.set_to_client_interceptor(_intercept_ok)
665
666 async def to_server_noop(self) -> callinfo.AsyncCallQueue:
667 """Do not read data, causing client to keep multiple data"""
668 logging.trace('to_server_noop')
669 return await self.set_to_server_interceptor(None)
670
671 async def to_client_noop(self) -> callinfo.AsyncCallQueue:
672 """Do not read data, causing server to keep multiple data"""
673 logging.trace('to_client_noop')
674 return await self.set_to_client_interceptor(None)
675
676 async def to_server_drop(self) -> callinfo.AsyncCallQueue:
677 """Read and discard data"""
678 logging.trace('to_server_drop')
679 return await self.set_to_server_interceptor(_intercept_drop)
680
681 async def to_client_drop(self) -> callinfo.AsyncCallQueue:
682 """Read and discard data"""
683 logging.trace('to_client_drop')
684 return await self.set_to_client_interceptor(_intercept_drop)
685
686 async def to_server_delay(self, delay: float) -> callinfo.AsyncCallQueue:
687 """Delay data transmission"""
688 logging.trace('to_server_delay, delay: %s', delay)
689
690 async def _intercept_delay_bound(
691 loop: EvLoop,
692 socket_from: Socket,
693 socket_to: Socket,
694 ) -> None:
695 await _intercept_delay(delay, loop, socket_from, socket_to)
696
697 return await self.set_to_server_interceptor(_intercept_delay_bound)
698
699 async def to_client_delay(self, delay: float) -> callinfo.AsyncCallQueue:
700 """Delay data transmission"""
701 logging.trace('to_client_delay, delay: %s', delay)
702
703 async def _intercept_delay_bound(
704 loop: EvLoop,
705 socket_from: Socket,
706 socket_to: Socket,
707 ) -> None:
708 await _intercept_delay(delay, loop, socket_from, socket_to)
709
710 return await self.set_to_client_interceptor(_intercept_delay_bound)
711
712 async def to_server_close_on_data(self) -> callinfo.AsyncCallQueue:
713 """Close on first bytes of data from client"""
714 logging.trace('to_server_close_on_data')
715 return await self.set_to_server_interceptor(_intercept_close_on_data)
716
717 async def to_client_close_on_data(self) -> callinfo.AsyncCallQueue:
718 """Close on first bytes of data from server"""
719 logging.trace('to_client_close_on_data')
720 return await self.set_to_client_interceptor(_intercept_close_on_data)
721
722 async def to_server_corrupt_data(self) -> callinfo.AsyncCallQueue:
723 """Corrupt data received from client"""
724 logging.trace('to_server_corrupt_data')
725 return await self.set_to_server_interceptor(_intercept_corrupt)
726
727 async def to_client_corrupt_data(self) -> callinfo.AsyncCallQueue:
728 """Corrupt data received from server"""
729 logging.trace('to_client_corrupt_data')
730 return await self.set_to_client_interceptor(_intercept_corrupt)
731
732 async def to_server_limit_bps(self, bytes_per_second: float) -> callinfo.AsyncCallQueue:
733 """Limit bytes per second transmission by network from client"""
734 logging.trace(
735 'to_server_limit_bps, bytes_per_second: %s',
736 bytes_per_second,
737 )
738 return await self.set_to_server_interceptor(_InterceptBpsLimit(bytes_per_second))
739
740 async def to_client_limit_bps(self, bytes_per_second: float) -> callinfo.AsyncCallQueue:
741 """Limit bytes per second transmission by network from server"""
742 logging.trace(
743 'to_client_limit_bps, bytes_per_second: %s',
744 bytes_per_second,
745 )
746 return await self.set_to_client_interceptor(_InterceptBpsLimit(bytes_per_second))
747
748 async def to_server_limit_time(self, timeout: float, jitter: float) -> callinfo.AsyncCallQueue:
749 """Limit connection lifetime on receive of first bytes from client"""
750 logging.trace(
751 'to_server_limit_time, timeout: %s, jitter: %s',
752 timeout,
753 jitter,
754 )
755 return await self.set_to_server_interceptor(_InterceptTimeLimit(timeout, jitter))
756
757 async def to_client_limit_time(self, timeout: float, jitter: float) -> callinfo.AsyncCallQueue:
758 """Limit connection lifetime on receive of first bytes from server"""
759 logging.trace(
760 'to_client_limit_time, timeout: %s, jitter: %s',
761 timeout,
762 jitter,
763 )
764 return await self.set_to_client_interceptor(_InterceptTimeLimit(timeout, jitter))
765
766 async def to_server_smaller_parts(
767 self,
768 max_size: int,
769 *,
770 sleep_per_packet: float = 0,
771 ) -> callinfo.AsyncCallQueue:
772 """
773 Pass data to server in smaller parts
774
775 @param max_size Max packet size to send to server
776 @param sleep_per_packet Optional sleep interval per packet, seconds
777 """
778 logging.trace('to_server_smaller_parts, max_size: %s', max_size)
779 return await self.set_to_server_interceptor(
780 _InterceptSmallerParts(max_size, sleep_per_packet),
781 )
782
783 async def to_client_smaller_parts(
784 self,
785 max_size: int,
786 *,
787 sleep_per_packet: float = 0,
788 ) -> callinfo.AsyncCallQueue:
789 """
790 Pass data to client in smaller parts
791
792 @param max_size Max packet size to send to client
793 @param sleep_per_packet Optional sleep interval per packet, seconds
794 """
795 logging.trace('to_client_smaller_parts, max_size: %s', max_size)
796 return await self.set_to_client_interceptor(
797 _InterceptSmallerParts(max_size, sleep_per_packet),
798 )
799
800 async def to_server_concat_packets(self, packet_size: int) -> callinfo.AsyncCallQueue:
801 """
802 Pass data in bigger parts
803 @param packet_size minimal size of the resulting packet
804 """
805 logging.trace('to_server_concat_packets, packet_size: %s', packet_size)
806 return await self.set_to_server_interceptor(_InterceptConcatPackets(packet_size))
807
808 async def to_client_concat_packets(self, packet_size: int) -> callinfo.AsyncCallQueue:
809 """
810 Pass data in bigger parts
811 @param packet_size minimal size of the resulting packet
812 """
813 logging.trace('to_client_concat_packets, packet_size: %s', packet_size)
814 return await self.set_to_client_interceptor(_InterceptConcatPackets(packet_size))
815
816 async def to_server_limit_bytes(self, bytes_limit: int) -> callinfo.AsyncCallQueue:
817 """Drop all connections each `bytes_limit` of data sent by network"""
818 logging.trace('to_server_limit_bytes, bytes_limit: %s', bytes_limit)
819 return await self.set_to_server_interceptor(_InterceptBytesLimit(bytes_limit, self))
820
821 async def to_client_limit_bytes(self, bytes_limit: int) -> callinfo.AsyncCallQueue:
822 """Drop all connections each `bytes_limit` of data sent by network"""
823 logging.trace('to_client_limit_bytes, bytes_limit: %s', bytes_limit)
824 return await self.set_to_client_interceptor(_InterceptBytesLimit(bytes_limit, self))
825
826 async def to_server_substitute(self, pattern: str, repl: str) -> callinfo.AsyncCallQueue:
827 """Apply regex substitution to data from client"""
828 logging.trace(
829 'to_server_substitute, pattern: %s, repl: %s',
830 pattern,
831 repl,
832 )
833 return await self.set_to_server_interceptor(_InterceptSubstitute(pattern, repl))
834
835 async def to_client_substitute(self, pattern: str, repl: str) -> callinfo.AsyncCallQueue:
836 """Apply regex substitution to data from server"""
837 logging.trace(
838 'to_client_substitute, pattern: %s, repl: %s',
839 pattern,
840 repl,
841 )
842 return await self.set_to_client_interceptor(_InterceptSubstitute(pattern, repl))
843
844
845class TcpGate(BaseGate):
846 """
847 Implements TCP chaos-proxy logic such as accepting incoming tcp client
848 connections. On each new connection new tcp client connects to server
849 (host_to_server, port_to_server).
850
851 @ingroup userver_testsuite
852
853 @see @ref scripts/docs/en/userver/chaos_testing.md
854 """
855
856 def __init__(self, route: GateRoute, loop: EvLoop | None = None) -> None:
857 self._connected_event = asyncio.Event()
858 super().__init__(route, loop)
859
860 def connections_count(self) -> int:
861 """
862 Returns maximal amount of connections going through the gate at
863 the moment.
864
865 @warning Some of the connections could be closing, or could be opened
866 right before the function starts. Use with caution!
867 """
868 return len(self._sockets)
869
870 async def wait_for_connections(self, *, count=1, timeout=0.0) -> None:
871 """
872 Wait for at least `count` connections going through the gate.
873
874 @throws asyncio.TimeoutError exception if failed to get the
875 required amount of connections in time.
876 """
877 if timeout <= 0.0:
878 while self.connections_count() < count:
879 await self._connected_event.wait()
880 self._connected_event.clear()
881 return
882
883 deadline = time.monotonic() + timeout
884 while self.connections_count() < count:
885 time_left = deadline - time.monotonic()
886 await asyncio.wait_for(
887 self._connected_event.wait(),
888 timeout=time_left,
889 )
890 self._connected_event.clear()
891
892 def _create_accepting_sockets(self) -> list[Socket]:
893 res: list[Socket] = []
894 for addr in socket.getaddrinfo(
895 self._route.host_for_client,
896 self._route.port_for_client,
897 type=socket.SOCK_STREAM,
898 ):
899 sock = Socket(addr[0], addr[1])
900 _make_socket_nonblocking(sock)
901 sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
902 sock.bind(addr[4])
903 sock.listen()
904 logger.debug(
905 f'Accepting connections on {sock.getsockname()}, fd={sock.fileno()}',
906 )
907 res.append(sock)
908
909 return res
910
911 async def _connect_to_server(self):
912 addrs = await self._loop.getaddrinfo(
913 self._route.host_to_server,
914 self._route.port_to_server,
915 type=socket.SOCK_STREAM,
916 )
917 for addr in addrs:
918 server = Socket(addr[0], addr[1])
919 _make_socket_nonblocking(server)
920 try:
921 await self._loop.sock_connect(server, addr[4])
922 logging.trace('Connected to %s', addr[4])
923 return server
924 except Exception as exc: # pylint: disable=broad-except
925 server.close()
926 logging.warning('Could not connect to %s: %s', addr[4], exc)
927
928 async def _do_accept(self, accept_sock: Socket) -> None:
929 while True:
930 client, _ = await self._loop.sock_accept(accept_sock)
931 _make_socket_nonblocking(client)
932
933 server = await self._connect_to_server()
934 if server:
935 self._sockets.add(
936 _SocketsPaired(
937 self._route.name,
938 self._loop,
939 client,
940 server,
941 self._to_server_intercept,
942 self._to_client_intercept,
943 ),
944 )
945 self._connected_event.set()
946 else:
947 client.close()
948
949 self._collect_garbage()
950
951
952class UdpGate(BaseGate):
953 """
954 Implements UDP chaos-proxy logic such as demuxing incoming datagrams
955 from different clients.
956 Separate connections to server are made for each new client.
957
958 @ingroup userver_testsuite
959
960 @see @ref scripts/docs/en/userver/chaos_testing.md
961 """
962
963 def __init__(self, route: GateRoute, loop: EvLoop | None = None):
964 self._clients: set[_UdpDemuxSocketMock] = set()
965 super().__init__(route, loop)
966
967 def is_connected(self) -> bool:
968 """
969 Returns True if there is active pair of sockets ready to transfer data
970 at the moment.
971 """
972 return len(self._sockets) > 0
973
974 def _create_accepting_sockets(self) -> list[Socket]:
975 res: list[Socket] = []
976 for addr in socket.getaddrinfo(
977 self._route.host_for_client,
978 self._route.port_for_client,
979 type=socket.SOCK_DGRAM,
980 ):
981 sock = socket.socket(addr[0], addr[1])
982 _make_socket_nonblocking(sock)
983 sock.bind(addr[4])
984 logger.debug(f'Accepting connections on {sock.getsockname()}')
985 res.append(sock)
986
987 return res
988
989 async def _connect_to_server(self):
990 addrs = await self._loop.getaddrinfo(
991 self._route.host_to_server,
992 self._route.port_to_server,
993 type=socket.SOCK_DGRAM,
994 )
995 for addr in addrs:
996 server = Socket(addr[0], addr[1])
997 try:
998 _make_socket_nonblocking(server)
999 await self._loop.sock_connect(server, addr[4])
1000 logging.trace('Connected to %s', addr[4])
1001 return server
1002 except Exception as exc: # pylint: disable=broad-except
1003 logging.warning('Could not connect to %s: %s', addr[4], exc)
1004
1005 def _collect_garbage(self) -> None:
1006 super()._collect_garbage()
1007 self._clients = {c for c in self._clients if c.is_active()}
1008
1009 async def _do_accept(self, accept_sock: Socket):
1010 sock = asyncio_socket.from_socket(accept_sock)
1011 while True:
1012 data, addr = await sock.recvfrom(RECV_MAX_SIZE, timeout=60.0)
1013
1014 client: _UdpDemuxSocketMock | None = None
1015 for known_clients in self._clients:
1016 if addr == known_clients.peer_address:
1017 client = known_clients
1018 break
1019
1020 if client is None:
1021 server = await self._connect_to_server()
1022 if not server:
1023 accept_sock.close()
1024 break
1025
1026 client = _UdpDemuxSocketMock(accept_sock, addr)
1027 self._clients.add(client)
1028
1029 self._sockets.add(
1030 _SocketsPaired(
1031 self._route.name,
1032 self._loop,
1033 client,
1034 server,
1035 self._to_server_intercept,
1036 self._to_client_intercept,
1037 ),
1038 )
1039
1040 await client.push(self._loop, data)
1041 self._collect_garbage()
1042
1043 async def to_server_concat_packets(self, packet_size: int) -> None:
1044 raise NotImplementedError('Udp packets cannot be concatenated')
1045
1046 async def to_client_concat_packets(self, packet_size: int) -> None:
1047 raise NotImplementedError('Udp packets cannot be concatenated')
1048
1049 async def to_server_smaller_parts(
1050 self,
1051 max_size: int,
1052 *,
1053 sleep_per_packet: float = 0,
1054 ) -> None:
1055 raise NotImplementedError('Udp packets cannot be split')
1056
1057 async def to_client_smaller_parts(
1058 self,
1059 max_size: int,
1060 *,
1061 sleep_per_packet: float = 0,
1062 ) -> None:
1063 raise NotImplementedError('Udp packets cannot be split')
1064
1065
1066def _create_callqueue(obj):
1067 if obj is None:
1068 return None
1069
1070 # workaround testsuite acallqueue that does not work with instances
1071 if isinstance(obj, callinfo.AsyncCallQueue):
1072 return obj
1073 if hasattr(obj, '__name__'):
1074 return callinfo.acallqueue(obj)
1075
1076 @functools.wraps(obj)
1077 async def wrapper(*args, **kwargs):
1078 return await obj(*args, **kwargs)
1079
1080 return callinfo.acallqueue(wrapper)
1081
1082
1083async def _wait_for_data(sock, timeout=60.0):
1084 if isinstance(sock, _UdpDemuxSocketMock):
1085 sock = sock.get_demux_out()
1086 sock = asyncio_socket.from_socket(sock)
1087 await sock.wait_for_data(timeout=timeout)