userver: en/testsuite/mockserver/server.py Source File
Loading...
Searching...
No Matches
server.py
1import asyncio
2import contextlib
3import itertools
4import logging
5import pathlib
6import re
7import ssl
8import time
9import typing
10import urllib.parse
11import warnings
12
13import aiohttp.web
14import yarl
15
16from testsuite import utils
17from testsuite.tracing import TraceidManager
18from testsuite.utils import (
19 cached_property,
20 callinfo,
21 http,
22 url_util,
23)
24from testsuite.utils import net as net_utils
25
26from . import classes, exceptions, magicargs
27from .exceptions import __tracebackhide__ # noqa: F401
28
29DEFAULT_TRACE_ID_HEADER = 'X-YaTraceId'
30DEFAULT_SPAN_ID_HEADER = 'X-YaSpanId'
31
32REQUEST_FROM_ANOTHER_TEST_ERROR = 'Internal error: request is from other test'
33
34_SUPPORTED_ERRORS_HEADER = 'X-Testsuite-Supported-Errors'
35_ERROR_HEADER = 'X-Testsuite-Error'
36
37_LOGGER_HEADERS = (
38 ('X-YaTraceId', 'trace_id'),
39 ('X-YaSpanId', 'span_id'),
40 ('X-YaRequestId', 'link'),
41)
42
43logger = logging.getLogger(__name__)
44
45RouteParams = dict[str, str]
46
47
48class MockserverRequest(aiohttp.web.BaseRequest):
49 # We need original path including scheme and hostname
50 def __init__(self, message, *args, **kwargs):
51 super().__init__(message, *args, **kwargs)
52 self.original_path = _path_from_message(message)
53
54
55class Handler:
56 def __init__(
57 self,
58 func: typing.Callable,
59 *,
60 raw_request: bool = False,
61 json_response: bool = False,
62 strict: bool = False,
63 ) -> None:
64 self.raw_request = raw_request
65 self.json_response = json_response
66 self.orig_func = func
67 self.strict = strict
68
69 @cached_property
70 def callqueue(self):
71 return callinfo.acallqueue(self.orig_func)
72
73 @cached_property
74 def handler_args(self):
76 self.orig_func,
77 raw_request=self.raw_request,
78 )
79
80 def __repr__(self):
81 return (
82 f'{Handler.__module__}.{Handler.__name__}({self.orig_func!r}, '
83 f'raw_request={self.handler_args.raw_request}, '
84 f'json_response={self.json_response})'
85 )
86
87 async def __call__(self, request: aiohttp.web.BaseRequest, **kwargs):
88 args, kwargs = await self.handler_args.build_args(request, kwargs)
89 response = await self.callqueue(*args, **kwargs)
90 if not self.json_response:
91 return response
92 if isinstance(response, (http.Response, aiohttp.web.Response)):
93 return response
94 return http.make_response(json=response)
95
96 def collect_calls(self) -> list[dict]:
97 if not self.strict:
98 return []
99
100 callqueue = self.callqueue
101 lost_calls = []
102 while callqueue.has_calls:
103 lost_calls.append(callqueue.next_call())
104 return lost_calls
105
106
108 handlers: dict[str, Handler]
109 prefix_handlers: list[tuple[str, Handler]]
110 regex_handlers: list[tuple[typing.Pattern, Handler]]
111
112 def __init__(
113 self,
114 *,
115 asyncexc_append,
116 traceid_manager: TraceidManager,
117 tracing_enabled=True,
118 http_proxy_enabled=False,
119 mockserver_host=None,
120 ):
121 self.traceid_manager = traceid_manager
122 self.tracing_enabled = tracing_enabled
123 self.handlers = {}
124 self.prefix_handlers = []
125 self.regex_handlers = []
126 self.http_proxy_enabled = http_proxy_enabled
127 self.mockserver_host = mockserver_host
128 self._asyncexc_append = asyncexc_append
129
130 def get_handler(self, path: str) -> tuple[Handler, RouteParams]:
131 handler = self.handlers.get(path)
132 if handler is not None:
133 return handler, {}
134 for pattern, handler in reversed(self.regex_handlers):
135 match = pattern.fullmatch(path)
136 if match:
137 return handler, match.groupdict()
138 for prefix, handler in reversed(self.prefix_handlers):
139 if path.startswith(prefix):
140 return handler, {}
143 )
144
145 def _get_handler_not_found_message(self, path: str) -> str:
146 if self.tracing_enabled:
147 tracing_state = 'enabled'
148 else:
149 tracing_state = 'disabled'
150 patterns = {regex.pattern for regex, _ in self.regex_handlers}
151 prefixes = {prefix for prefix, _ in self.prefix_handlers}
152 handlers_list = '\n'.join(
153 itertools.chain(
154 (f'- {path}' for path in self.handlers),
155 (f'- REGEX {pattern}' for pattern in patterns),
156 (f'- PREFIX {prefix}' for prefix in prefixes),
157 ),
158 )
159 return (
160 f'Mockserver handler is not installed for {path!r}.\n\n'
161 f'Perhaps you forgot to setup mockserver handler. '
162 f'Tracing: {tracing_state}. Installed handlers:\n'
163 f'{handlers_list}'
164 )
165
166 async def handle_request(
167 self,
168 request: MockserverRequest,
169 nofail_404: bool,
170 ):
171 try:
172 handler, kwargs = self._get_handler_for_request(request)
174 if not nofail_404:
175 self._asyncexc_append(exc)
176 return _internal_error(f'Internal server error: {exc!r}')
177
178 try:
179 response = await handler(request, **kwargs)
180 if isinstance(response, http.Response):
181 return response.to_aiohttp()
182 elif isinstance(
183 response, (aiohttp.web.Response, aiohttp.web.WebSocketResponse)
184 ):
185 return response
186 elif isinstance(response, http.MockedError):
187 return _mocked_error_response(request, response.error_code)
189 'http.Response or aiohttp.web.Response instance is expected '
190 f'{response!r} given',
191 )
192 except http.MockedError as exc:
193 return _mocked_error_response(request, exc.error_code)
194 except Exception as exc:
195 self._asyncexc_append(exc)
196 return _internal_error(f'Internal server error: {exc!r}')
197
198 def register_handler(
199 self,
200 path: str,
201 func,
202 *,
203 prefix: bool = False,
204 regex: bool = False,
205 ):
206 if regex:
207 if prefix:
208 raise RuntimeError(
209 'Parameter value prefix=True is not supported if regex '
210 'parameter is also True.',
211 )
212 pattern = re.compile(path)
213 self.regex_handlers.append((pattern, func))
214 else:
215 if prefix:
216 self.prefix_handlers.append((path, func))
217 else:
218 self.handlers[path] = func
219 return func
220
221 def _get_handler_for_request(
222 self,
223 request: MockserverRequest,
224 ) -> tuple[Handler, RouteParams]:
225 path = request.original_path
226 if self.http_proxy_enabled:
227 host = request.headers.get('host')
228 if host and host != self.mockserver_host:
229 return self.get_handler(f'http://{host}{path}')
230 return self.get_handler(path)
231
232 def collect_calls(self) -> list[dict]:
233 calls = []
234 for _, handler in self.regex_handlers:
235 calls.extend(handler.collect_calls())
236 for _, handler in self.prefix_handlers:
237 calls.extend(handler.collect_calls())
238 for handler in self.handlers.values():
239 calls.extend(handler.collect_calls())
240 return calls
241
242
243# pylint: disable=too-many-instance-attributes
244class Server:
245 session = None
246
247 def __init__(
248 self,
249 mockserver_info: classes.MockserverInfo,
250 *,
251 nofail=False,
252 mockserver_debug=False,
253 tracing_enabled=True,
254 trace_id_header=DEFAULT_TRACE_ID_HEADER,
255 span_id_header=DEFAULT_SPAN_ID_HEADER,
256 http_proxy_enabled=False,
257 ):
258 self._info = mockserver_info
259 self._nofail = nofail
260 self._mockserver_debug = mockserver_debug
261 self._tracing_enabled = tracing_enabled
262 self._trace_id_header = trace_id_header
263 self._span_id_header = span_id_header
264 self._http_proxy_enabled = http_proxy_enabled
265
266 @property
267 def tracing_enabled(self) -> bool:
268 if self.session is None:
269 return self._tracing_enabled
270 return self.session.tracing_enabled
271
272 @property
273 def trace_id_header(self):
274 return self._trace_id_header
275
276 @property
277 def span_id_header(self):
278 return self._span_id_header
279
280 @property
281 def http_proxy_enabled(self):
282 return self._http_proxy_enabled
283
284 @property
285 def server_info(self) -> classes.MockserverInfo:
286 return self._info
287
288 def get_debug(self) -> bool:
289 return self._mockserver_debug
290
291 def set_debug(self, enabled: bool):
292 self._mockserver_debug = enabled
293
294 @contextlib.contextmanager
295 def new_session(
296 self,
297 *,
298 asyncexc_append,
299 traceid_manager: TraceidManager,
300 ):
301 self.session = Session(
302 asyncexc_append=asyncexc_append,
303 tracing_enabled=self._tracing_enabled,
304 traceid_manager=traceid_manager,
305 http_proxy_enabled=self._http_proxy_enabled,
306 mockserver_host=self._info.get_host_header(),
307 )
308 try:
309 yield self.session
310 finally:
311 self.session = None
312
313 async def handle_request(self, request):
314 started = time.perf_counter()
315 try:
316 response = await self._handle_request(request)
317 self._log_request(started, request, response)
318 return response
319 except BaseException as exc:
320 self._log_request(started, request, exc=exc)
321 raise
322
323 def _log_request(self, started, request, response=None, exc=None):
324 if exc is None and not self._mockserver_debug:
325 return
326 fields = {
327 '_type': 'mockserver_request',
328 'timestamp': utils.utcnow(),
329 'method': request.method,
330 'url': request.rel_url,
331 }
332 for header, key in _LOGGER_HEADERS:
333 if header in request.headers:
334 fields[key] = request.headers[header]
335 delay_ms = 1000 * (time.perf_counter() - started)
336 fields['delay'] = f'{delay_ms:.3f}ms'
337 if response is not None:
338 log_level = logging.DEBUG
339 fields['meta_code'] = response.status
340 fields['status'] = 'DONE'
341 else:
342 log_level = logging.ERROR
343 fields['status'] = 'FAIL'
344 fields['exc_info'] = str(exc)
345 logger.log(log_level, 'Mockserver request', extra={'tskv': fields})
346
347 async def _handle_request(self, request: MockserverRequest):
348 if not self.session:
350 'Internal error: mockserver session was not initialized'
351 )
352
353 nofail = self._nofail
354 trace_id = request.headers.get(self.trace_id_header)
355 traceid_manager = self.session.traceid_manager
356 if self.tracing_enabled and not traceid_manager.is_testsuite(trace_id):
357 nofail = True
358
359 if self.tracing_enabled and traceid_manager.is_other_test(trace_id):
360 self._report_other_test_request(request, trace_id)
361 return _internal_error(REQUEST_FROM_ANOTHER_TEST_ERROR)
362 try:
363 return await self.session.handle_request(request, nofail_404=nofail)
365 return _internal_error(
366 'Internal error: mockserver handler not found',
367 )
368
369 def _report_other_test_request(self, request, trace_id):
370 logger.warning(
371 'Mockserver called path %s with previous test trace_id %s',
372 request.path,
373 trace_id,
374 )
375
376
378 """Mockserver handler installer fixture."""
379
380 def __init__(
381 self,
382 mockserver: Server,
383 session: Session,
384 base_prefix: str = '',
385 *,
386 strict_default: bool = False,
387 ) -> None:
388 self._server = mockserver
389 self._session = session
390 self._base_prefix = base_prefix
391 self._base_prefix_re = re.escape(base_prefix)
392 self._strict_default = strict_default
393
394 def new(self, prefix: str) -> 'MockserverFixture':
395 """Create mockserver installer with given base prefix."""
396 return MockserverFixture(
397 self._server,
398 self._session,
399 self._build_fullpath(prefix),
400 )
401
402 @property
403 def base_url(self) -> str:
404 """Mockserver base url."""
405 return self._server.server_info.base_url
406
407 @property
408 def host(self) -> str | None:
409 """Mockserver hostname."""
410 return self._server.server_info.host
411
412 @property
413 def port(self) -> int | None:
414 """Mockserver port."""
415 return self._server.server_info.port
416
417 @property
418 def trace_id_header(self) -> str:
419 return self._server.trace_id_header
420
421 @property
422 def span_id_header(self) -> str:
423 return self._server.span_id_header
424
425 @property
426 def trace_id(self) -> str:
427 return self._session.traceid_manager.trace_id
428
430 self,
431 path: str,
432 *,
433 prefix: bool = False,
434 raw_request: bool = False,
435 json_response: bool = False,
436 regex: bool = False,
437 strict: typing.Optional[bool] = None,
438 ) -> classes.GenericRequestDecorator:
439 """Register basic http handler for ``path``.
440
441 Returns decorator that registers handler ``path``. Original function is
442 wrapped with :ref:`AsyncCallQueue`.
443
444 :param path: match url by prefix if ``True`` exact match otherwise
445 :param raw_request: pass ``aiohttp.web.Response`` to handler instead of
446 ``testsuite.utils.http.Request``
447 :param regex: set True to match path as regex pattern
448 :param prefix: set True to match path prefix instead of whole path
449 :param json_response: set True to let handler return json object
450 instead of full response object
451
452 .. code-block:: python
453
454 @mockserver.handler('/service/path')
455 def handler(request: testsuite.utils.http.Request):
456 return mockserver.make_response('Hello, world!')
457 """
458
459 if raw_request:
460 warnings.warn(
461 'raw_request=True is deprecated, use aiohttp_handler() instead',
462 DeprecationWarning,
463 )
464 if json_response:
465 warnings.warn(
466 'json_response=True is deprecated, use json_handler() instead',
467 DeprecationWarning,
468 )
469
470 return self._handler_installer(
471 path,
472 prefix=prefix,
473 raw_request=raw_request,
474 json_response=json_response,
475 regex=regex,
476 strict=strict,
477 )
478
480 self,
481 path: str,
482 *,
483 prefix: bool = False,
484 raw_request: bool = False,
485 regex: bool = False,
486 strict: typing.Optional[bool] = None,
487 ) -> classes.JsonRequestDecorator:
488 """Register json http handler for ``path``.
489
490 Returns decorator that registers handler ``path``. Original function is
491 wrapped with :ref:`AsyncCallQueue`.
492
493 :param path: match url by prefix if ``True`` exact match otherwise
494 :param raw_request: pass ``aiohttp.web.Response`` to handler instead of
495 ``testsuite.utils.http.Request``
496 :param prefix: set True to match path prefix instead of whole path
497 :param regex: set True to match path as regex pattern
498
499 .. code-block:: python
500
501 @mockserver.json_handler('/service/path')
502 def handler(request: testsuite.utils.http.Request):
503 # Return JSON document
504 return {...}
505 # or call to mockserver.make_response()
506 return mockserver.make_response(...)
507 """
508 if raw_request:
509 warnings.warn(
510 'raw_request=True is deprecated, '
511 'use aiohttp_json_handler() instead',
512 DeprecationWarning,
513 )
514 return self._handler_installer(
515 path,
516 prefix=prefix,
517 raw_request=raw_request,
518 json_response=True,
519 regex=regex,
520 strict=strict,
521 )
522
523 def aiohttp_handler(
524 self,
525 path: str,
526 *,
527 prefix: bool = False,
528 regex: bool = False,
529 strict: typing.Optional[bool] = None,
530 ) -> classes.GenericRequestDecorator:
531 return self._handler_installer(
532 path,
533 prefix=prefix,
534 raw_request=True,
535 json_response=False,
536 regex=regex,
537 strict=strict,
538 )
539
540 def aiohttp_json_handler(
541 self,
542 path: str,
543 *,
544 prefix: bool = False,
545 regex: bool = False,
546 strict: typing.Optional[bool] = None,
547 ) -> classes.JsonRequestDecorator:
548 return self._handler_installer(
549 path,
550 prefix=prefix,
551 raw_request=True,
552 json_response=True,
553 regex=regex,
554 strict=strict,
555 )
556
557 def url(self, path: str) -> str:
558 """Builds mockserver url for ``path``"""
559 return url_util.join(self.base_url, path)
560
561 def url_encoded(self, path: str) -> yarl.URL:
562 """Builds mockserver url for ``path``"""
563 return yarl.URL(url_util.join(self.base_url, path), encoded=True)
564
565 def ws_url(self, path: str) -> str:
566 return self._server.server_info.ws_url(path)
567
568 def ignore_trace_id(self) -> typing.ContextManager[None]:
569 return self.tracing(False)
570
571 @contextlib.contextmanager
572 def tracing(self, value: bool = True):
573 original_value = self._session.tracing_enabled
574 try:
575 self._session.tracing_enabled = value
576 yield
577 finally:
578 self._session.tracing_enabled = original_value
579
580 def get_callqueue_for(self, path) -> callinfo.AsyncCallQueue:
581 handler, _ = self._session.get_handler(path)
582 return handler.callqueue
583
584 make_response = staticmethod(http.make_response)
585
586 TimeoutError = http.TimeoutError
587 NetworkError = http.NetworkError
588
589 def _handler_installer(
590 self,
591 path: str,
592 *,
593 strict: typing.Optional[bool],
594 prefix: bool = False,
595 raw_request: bool = False,
596 json_response: bool = False,
597 regex: bool = False,
598 ) -> typing.Callable:
599 path = self._build_fullpath(path, regex)
600 if strict is None:
601 strict = self._strict_default
602
603 def decorator(func):
604 handler = Handler(
605 func,
606 raw_request=raw_request,
607 json_response=json_response,
608 strict=strict,
609 )
610 self._session.register_handler(
611 path,
612 handler,
613 prefix=prefix,
614 regex=regex,
615 )
616 return handler.callqueue
617
618 return decorator
619
620 def _build_fullpath(self, path, regex: bool = False) -> str:
621 if regex:
622 return self._base_prefix_re + path
623 if not self._base_prefix or self._base_prefix.endswith('/'):
624 if self._server.http_proxy_enabled and path.startswith('http://'):
625 return path
626 return url_util.join(self._base_prefix, path)
627 return self._base_prefix + path
628
629
630MockserverSslFixture = MockserverFixture
631
632
633def create_server(
634 *,
635 host: str,
636 port: int,
637 pytestconfig,
638 ssl_info=None,
639 loop=None,
640):
641 warnings.warn('Use mockserver_create() fixture instead', DeprecationWarning)
642
643 mockserver_socket = _create_mockserver_socket(host=host, port=port)
644 return _create_server_from_socket(
645 mockserver_socket,
646 mockserver_config=classes.MockserverConfig(),
647 ssl_cert=ssl_info,
648 loop=loop,
649 )
650
651
652def _create_ssl_context(ssl_cert: classes.SslCertInfo) -> ssl.SSLContext:
653 ssl_context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
654 ssl_context.load_cert_chain(ssl_cert.cert_path, ssl_cert.private_key_path)
655 return ssl_context
656
657
658def _internal_error(message: str = 'Internal error') -> aiohttp.web.Response:
659 return http.make_response(message, status=500).to_aiohttp()
660
661
662def _mocked_error_response(request, error_code) -> aiohttp.web.Response:
663 if _SUPPORTED_ERRORS_HEADER not in request.headers:
665 'Service does not support mockserver errors protocol',
666 )
667 supported_errors = request.headers[_SUPPORTED_ERRORS_HEADER].split(',')
668 if error_code not in supported_errors:
670 f'Service does not support mockserver error of type {error_code}',
671 )
672 return http.make_response(
673 response='',
674 status=599,
675 headers={_ERROR_HEADER: error_code},
676 ).to_aiohttp()
677
678
679def _create_server_obj(
680 mockserver_info: classes.MockserverInfo,
681 mockserver_config: classes.MockserverConfig,
682) -> Server:
683 return Server(
684 mockserver_info,
685 nofail=mockserver_config.nofail,
686 mockserver_debug=mockserver_config.debug,
687 tracing_enabled=mockserver_config.tracing_enabled,
688 trace_id_header=mockserver_config.trace_id_header,
689 span_id_header=mockserver_config.span_id_header,
690 http_proxy_enabled=mockserver_config.http_proxy_enabled,
691 )
692
693
694def _create_web_server(server: Server, loop) -> aiohttp.web.Server:
695 def request_factory(*args):
696 return MockserverRequest(*args, loop=loop)
697
698 return aiohttp.web.Server(
699 server.handle_request,
700 request_factory=request_factory,
701 loop=loop,
702 access_log=None,
703 )
704
705
706def _create_mockserver_socket(
707 socket_path=None,
708 host='localhost',
709 port=0,
710 https=False,
711):
712 if socket_path is None:
713 sockets = net_utils.bind_socket_multiple(host, port)
714 else:
715 sockets = [net_utils.bind_unix_socket(socket_path)]
716 assert sockets
717 for sock in sockets:
718 sock.setblocking(False)
719 info = _create_mockserver_info(
720 sockets[0],
721 socket_path=socket_path,
722 host=host,
723 https=https,
724 )
725 return classes.MockserverSocket(sockets=sockets, info=info)
726
727
728@contextlib.asynccontextmanager
729async def _create_server_from_socket(
730 mockserver_socket: classes.MockserverSocket,
731 mockserver_config: classes.MockserverConfig,
732 ssl_cert: classes.SslCertInfo | None = None,
733 loop=None,
734) -> typing.AsyncGenerator[Server, None]:
735 if ssl_cert:
736 ssl_context = _create_ssl_context(ssl_cert)
737 else:
738 ssl_context = None
739
740 if loop is None:
741 loop = asyncio.get_running_loop()
742
743 server = _create_server_obj(mockserver_socket.info, mockserver_config)
744 web_server = _create_web_server(server, loop)
745
746 async with net_utils.create_server_multiple(
747 web_server,
748 sockets=mockserver_socket.sockets,
749 ssl=ssl_context,
750 ) as aio_server:
751 yield server
752
753
754def _create_mockserver_info(
755 sock,
756 socket_path,
757 host: str,
758 https: bool = False,
760 if socket_path:
761 return _create_unix_mockserver_info(socket_path)
762 port = sock.getsockname()[1]
763 schema = 'https' if https else 'http'
764 base_url = f'{schema}://{host}:{port}/'
766 host=host,
767 port=port,
768 base_url=base_url,
769 https=https,
770 )
771
772
773def _create_unix_mockserver_info(
774 socket_path: pathlib.Path,
777 socket_path=socket_path,
778 # use localhost to avoid aiohttp complains on invalid url
779 base_url='http://localhost/',
780 host='localhost',
781 port=80,
782 https=False,
783 )
784
785
786def _path_from_message(message):
787 """Returns original HTTP path without query."""
788 path = str(message.url)
789 path = path.split('?')[0]
790 path = urllib.parse.unquote(path)
791 return path