codekingpro/portable-devtools
115k
1import asyncio
2import asyncio.streams
3import sys
4import traceback
5import warnings
6from collections import deque
7from contextlib import suppress
8from html import escape as html_escape
9from http import HTTPStatus
10from logging import Logger
11from typing import (
12 TYPE_CHECKING,
13 Any,
14 Awaitable,
15 Callable,
16 Deque,
17 Optional,
18 Sequence,
19 Tuple,
20 Type,
21 Union,
22 cast,
23)
24
25import attr
26import yarl
27from propcache import under_cached_property
28
29from .abc import AbstractAccessLogger, AbstractStreamWriter
30from .base_protocol import BaseProtocol
31from .helpers import ceil_timeout
32from .http import (
33 HttpProcessingError,
34 HttpRequestParser,
35 HttpVersion10,
36 RawRequestMessage,
37 StreamWriter,
38)
39from .http_exceptions import BadHttpMethod
40from .log import access_logger, server_logger
41from .streams import EMPTY_PAYLOAD, StreamReader
42from .tcp_helpers import tcp_keepalive
43from .web_exceptions import HTTPException, HTTPInternalServerError
44from .web_log import AccessLogger
45from .web_request import BaseRequest
46from .web_response import Response, StreamResponse
47
48__all__ = ("RequestHandler", "RequestPayloadError", "PayloadAccessError")
49
50if TYPE_CHECKING:
51 import ssl
52
53 from .web_server import Server
54
55
56_RequestFactory = Callable[
57 [
58 RawRequestMessage,
59 StreamReader,
60 "RequestHandler",
61 AbstractStreamWriter,
62 "asyncio.Task[None]",
63 ],
64 BaseRequest,
65]
66
67_RequestHandler = Callable[[BaseRequest], Awaitable[StreamResponse]]
68
69ERROR = RawRequestMessage(
70 "UNKNOWN",
71 "/",
72 HttpVersion10,
73 {}, # type: ignore[arg-type]
74 {}, # type: ignore[arg-type]
75 True,
76 None,
77 False,
78 False,
79 yarl.URL("/"),
80)
81
82
83class RequestPayloadError(Exception):
84 """Payload parsing error."""
85
86
87class PayloadAccessError(Exception):
88 """Payload was accessed after response was sent."""
89
90
91_PAYLOAD_ACCESS_ERROR = PayloadAccessError()
92
93
94@attr.s(auto_attribs=True, frozen=True, slots=True)
95class _ErrInfo:
96 status: int
97 exc: BaseException
98 message: str
99
100
101_MsgType = Tuple[Union[RawRequestMessage, _ErrInfo], StreamReader]
102
103
104class RequestHandler(BaseProtocol):
105 """HTTP protocol implementation.
106
107 RequestHandler handles incoming HTTP request. It reads request line,
108 request headers and request payload and calls handle_request() method.
109 By default it always returns with 404 response.
110
111 RequestHandler handles errors in incoming request, like bad
112 status line, bad headers or incomplete payload. If any error occurs,
113 connection gets closed.
114
115 keepalive_timeout -- number of seconds before closing
116 keep-alive connection
117
118 tcp_keepalive -- TCP keep-alive is on, default is on
119
120 debug -- enable debug mode
121
122 logger -- custom logger object
123
124 access_log_class -- custom class for access_logger
125
126 access_log -- custom logging object
127
128 access_log_format -- access log format string
129
130 loop -- Optional event loop
131
132 max_line_size -- Optional maximum header line size
133
134 max_field_size -- Optional maximum header field size
135
136 max_headers -- Optional maximum header size
137
138 timeout_ceil_threshold -- Optional value to specify
139 threshold to ceil() timeout
140 values
141
142 """
143
144 __slots__ = (
145 "max_field_size",
146 "max_headers",
147 "max_line_size",
148 "_request_count",
149 "_keepalive",
150 "_manager",
151 "_request_handler",
152 "_request_factory",
153 "_tcp_keepalive",
154 "_next_keepalive_close_time",
155 "_keepalive_handle",
156 "_keepalive_timeout",
157 "_lingering_time",
158 "_messages",
159 "_message_tail",
160 "_handler_waiter",
161 "_waiter",
162 "_task_handler",
163 "_upgrade",
164 "_payload_parser",
165 "_request_parser",
166 "_reading_paused",
167 "logger",
168 "debug",
169 "access_log",
170 "access_logger",
171 "_close",
172 "_force_close",
173 "_current_request",
174 "_timeout_ceil_threshold",
175 "_request_in_progress",
176 "_logging_enabled",
177 "_cache",
178 )
179
180 def __init__(
181 self,
182 manager: "Server",
183 *,
184 loop: asyncio.AbstractEventLoop,
185 # Default should be high enough that it's likely longer than a reverse proxy.
186 keepalive_timeout: float = 3630,
187 tcp_keepalive: bool = True,
188 logger: Logger = server_logger,
189 access_log_class: Type[AbstractAccessLogger] = AccessLogger,
190 access_log: Logger = access_logger,
191 access_log_format: str = AccessLogger.LOG_FORMAT,
192 debug: bool = False,
193 max_line_size: int = 8190,
194 max_headers: int = 128,
195 max_field_size: int = 8190,
196 lingering_time: float = 10.0,
197 read_bufsize: int = 2**16,
198 auto_decompress: bool = True,
199 timeout_ceil_threshold: float = 5,
200 ):
201 super().__init__(loop)
202
203 # _request_count is the number of requests processed with the same connection.
204 self._request_count = 0
205 self._keepalive = False
206 self._current_request: Optional[BaseRequest] = None
207 self._manager: Optional[Server] = manager
208 self._request_handler: Optional[_RequestHandler] = manager.request_handler
209 self._request_factory: Optional[_RequestFactory] = manager.request_factory
210
211 self.max_line_size = max_line_size
212 self.max_headers = max_headers
213 self.max_field_size = max_field_size
214
215 self._tcp_keepalive = tcp_keepalive
216 # placeholder to be replaced on keepalive timeout setup
217 self._next_keepalive_close_time = 0.0
218 self._keepalive_handle: Optional[asyncio.Handle] = None
219 self._keepalive_timeout = keepalive_timeout
220 self._lingering_time = float(lingering_time)
221
222 self._messages: Deque[_MsgType] = deque()
223 self._message_tail = b""
224
225 self._waiter: Optional[asyncio.Future[None]] = None
226 self._handler_waiter: Optional[asyncio.Future[None]] = None
227 self._task_handler: Optional[asyncio.Task[None]] = None
228
229 self._upgrade = False
230 self._payload_parser: Any = None
231 self._request_parser: Optional[HttpRequestParser] = HttpRequestParser(
232 self,
233 loop,
234 read_bufsize,
235 max_line_size=max_line_size,
236 max_field_size=max_field_size,
237 max_headers=max_headers,
238 payload_exception=RequestPayloadError,
239 auto_decompress=auto_decompress,
240 )
241
242 self._timeout_ceil_threshold: float = 5
243 try:
244 self._timeout_ceil_threshold = float(timeout_ceil_threshold)
245 except (TypeError, ValueError):
246 pass
247
248 self.logger = logger
249 self.debug = debug
250 self.access_log = access_log
251 if access_log:
252 self.access_logger: Optional[AbstractAccessLogger] = access_log_class(
253 access_log, access_log_format
254 )
255 self._logging_enabled = self.access_logger.enabled
256 else:
257 self.access_logger = None
258 self._logging_enabled = False
259
260 self._close = False
261 self._force_close = False
262 self._request_in_progress = False
263 self._cache: dict[str, Any] = {}
264
265 def __repr__(self) -> str:
266 return "<{} {}>".format(
267 self.__class__.__name__,
268 "connected" if self.transport is not None else "disconnected",
269 )
270
271 @under_cached_property
272 def ssl_context(self) -> Optional["ssl.SSLContext"]:
273 """Return SSLContext if available."""
274 return (
275 None
276 if self.transport is None
277 else self.transport.get_extra_info("sslcontext")
278 )
279
280 @under_cached_property
281 def peername(
282 self,
283 ) -> Optional[Union[str, Tuple[str, int, int, int], Tuple[str, int]]]:
284 """Return peername if available."""
285 return (
286 None
287 if self.transport is None
288 else self.transport.get_extra_info("peername")
289 )
290
291 @property
292 def keepalive_timeout(self) -> float:
293 return self._keepalive_timeout
294
295 async def shutdown(self, timeout: Optional[float] = 15.0) -> None:
296 """Do worker process exit preparations.
297
298 We need to clean up everything and stop accepting requests.
299 It is especially important for keep-alive connections.
300 """
301 self._force_close = True
302
303 if self._keepalive_handle is not None:
304 self._keepalive_handle.cancel()
305
306 # Wait for graceful handler completion
307 if self._request_in_progress:
308 # The future is only created when we are shutting
309 # down while the handler is still processing a request
310 # to avoid creating a future for every request.
311 self._handler_waiter = self._loop.create_future()
312 try:
313 async with ceil_timeout(timeout):
314 await self._handler_waiter
315 except (asyncio.CancelledError, asyncio.TimeoutError):
316 self._handler_waiter = None
317 if (
318 sys.version_info >= (3, 11)
319 and (task := asyncio.current_task())
320 and task.cancelling()
321 ):
322 raise
323 # Then cancel handler and wait
324 try:
325 async with ceil_timeout(timeout):
326 if self._current_request is not None:
327 self._current_request._cancel(asyncio.CancelledError())
328
329 if self._task_handler is not None and not self._task_handler.done():
330 await asyncio.shield(self._task_handler)
331 except (asyncio.CancelledError, asyncio.TimeoutError):
332 if (
333 sys.version_info >= (3, 11)
334 and (task := asyncio.current_task())
335 and task.cancelling()
336 ):
337 raise
338
339 # force-close non-idle handler
340 if self._task_handler is not None:
341 self._task_handler.cancel()
342
343 self.force_close()
344
345 def connection_made(self, transport: asyncio.BaseTransport) -> None:
346 super().connection_made(transport)
347
348 real_transport = cast(asyncio.Transport, transport)
349 if self._tcp_keepalive:
350 tcp_keepalive(real_transport)
351
352 assert self._manager is not None
353 self._manager.connection_made(self, real_transport)
354
355 loop = self._loop
356 if sys.version_info >= (3, 12):
357 task = asyncio.Task(self.start(), loop=loop, eager_start=True)
358 else:
359 task = loop.create_task(self.start())
360 self._task_handler = task
361
362 def connection_lost(self, exc: Optional[BaseException]) -> None:
363 if self._manager is None:
364 return
365 self._manager.connection_lost(self, exc)
366
367 # Grab value before setting _manager to None.
368 handler_cancellation = self._manager.handler_cancellation
369
370 self.force_close()
371 super().connection_lost(exc)
372 self._manager = None
373 self._request_factory = None
374 self._request_handler = None
375 self._request_parser = None
376
377 if self._keepalive_handle is not None:
378 self._keepalive_handle.cancel()
379
380 if self._current_request is not None:
381 if exc is None:
382 exc = ConnectionResetError("Connection lost")
383 self._current_request._cancel(exc)
384
385 if handler_cancellation and self._task_handler is not None:
386 self._task_handler.cancel()
387
388 self._task_handler = None
389
390 if self._payload_parser is not None:
391 self._payload_parser.feed_eof()
392 self._payload_parser = None
393
394 def set_parser(self, parser: Any) -> None:
395 # Actual type is WebReader
396 assert self._payload_parser is None
397
398 self._payload_parser = parser
399
400 if self._message_tail:
401 self._payload_parser.feed_data(self._message_tail)
402 self._message_tail = b""
403
404 def eof_received(self) -> None:
405 pass
406
407 def data_received(self, data: bytes) -> None:
408 if self._force_close or self._close:
409 return
410 # parse http messages
411 messages: Sequence[_MsgType]
412 if self._payload_parser is None and not self._upgrade:
413 assert self._request_parser is not None
414 try:
415 messages, upgraded, tail = self._request_parser.feed_data(data)
416 except HttpProcessingError as exc:
417 messages = [
418 (_ErrInfo(status=400, exc=exc, message=exc.message), EMPTY_PAYLOAD)
419 ]
420 upgraded = False
421 tail = b""
422
423 for msg, payload in messages or ():
424 self._request_count += 1
425 self._messages.append((msg, payload))
426
427 waiter = self._waiter
428 if messages and waiter is not None and not waiter.done():
429 # don't set result twice
430 waiter.set_result(None)
431
432 self._upgrade = upgraded
433 if upgraded and tail:
434 self._message_tail = tail
435
436 # no parser, just store
437 elif self._payload_parser is None and self._upgrade and data:
438 self._message_tail += data
439
440 # feed payload
441 elif data:
442 eof, tail = self._payload_parser.feed_data(data)
443 if eof:
444 self.close()
445
446 def keep_alive(self, val: bool) -> None:
447 """Set keep-alive connection mode.
448
449 :param bool val: new state.
450 """
451 self._keepalive = val
452 if self._keepalive_handle:
453 self._keepalive_handle.cancel()
454 self._keepalive_handle = None
455
456 def close(self) -> None:
457 """Close connection.
458
459 Stop accepting new pipelining messages and close
460 connection when handlers done processing messages.
461 """
462 self._close = True
463 if self._waiter:
464 self._waiter.cancel()
465
466 def force_close(self) -> None:
467 """Forcefully close connection."""
468 self._force_close = True
469 if self._waiter:
470 self._waiter.cancel()
471 if self.transport is not None:
472 self.transport.close()
473 self.transport = None
474
475 def log_access(
476 self, request: BaseRequest, response: StreamResponse, time: Optional[float]
477 ) -> None:
478 if self._logging_enabled and self.access_logger is not None:
479 if TYPE_CHECKING:
480 assert time is not None
481 self.access_logger.log(request, response, self._loop.time() - time)
482
483 def log_debug(self, *args: Any, **kw: Any) -> None:
484 if self.debug:
485 self.logger.debug(*args, **kw)
486
487 def log_exception(self, *args: Any, **kw: Any) -> None:
488 self.logger.exception(*args, **kw)
489
490 def _process_keepalive(self) -> None:
491 self._keepalive_handle = None
492 if self._force_close or not self._keepalive:
493 return
494
495 loop = self._loop
496 now = loop.time()
497 close_time = self._next_keepalive_close_time
498 if now < close_time:
499 # Keep alive close check fired too early, reschedule
500 self._keepalive_handle = loop.call_at(close_time, self._process_keepalive)
501 return
502
503 # handler in idle state
504 if self._waiter and not self._waiter.done():
505 self.force_close()
506
507 async def _handle_request(
508 self,
509 request: BaseRequest,
510 start_time: Optional[float],
511 request_handler: Callable[[BaseRequest], Awaitable[StreamResponse]],
512 ) -> Tuple[StreamResponse, bool]:
513 self._request_in_progress = True
514 try:
515 try:
516 self._current_request = request
517 resp = await request_handler(request)
518 finally:
519 self._current_request = None
520 except HTTPException as exc:
521 resp = exc
522 resp, reset = await self.finish_response(request, resp, start_time)
523 except asyncio.CancelledError:
524 raise
525 except asyncio.TimeoutError as exc:
526 self.log_debug("Request handler timed out.", exc_info=exc)
527 resp = self.handle_error(request, 504)
528 resp, reset = await self.finish_response(request, resp, start_time)
529 except Exception as exc:
530 resp = self.handle_error(request, 500, exc)
531 resp, reset = await self.finish_response(request, resp, start_time)
532 else:
533 # Deprecation warning (See #2415)
534 if getattr(resp, "__http_exception__", False):
535 warnings.warn(
536 "returning HTTPException object is deprecated "
537 "(#2415) and will be removed, "
538 "please raise the exception instead",
539 DeprecationWarning,
540 )
541
542 resp, reset = await self.finish_response(request, resp, start_time)
543 finally:
544 self._request_in_progress = False
545 if self._handler_waiter is not None:
546 self._handler_waiter.set_result(None)
547
548 return resp, reset
549
550 async def start(self) -> None:
551 """Process incoming request.
552
553 It reads request line, request headers and request payload, then
554 calls handle_request() method. Subclass has to override
555 handle_request(). start() handles various exceptions in request
556 or response handling. Connection is being closed always unless
557 keep_alive(True) specified.
558 """
559 loop = self._loop
560 manager = self._manager
561 assert manager is not None
562 keepalive_timeout = self._keepalive_timeout
563 resp = None
564 assert self._request_factory is not None
565 assert self._request_handler is not None
566
567 while not self._force_close:
568 if not self._messages:
569 try:
570 # wait for next request
571 self._waiter = loop.create_future()
572 await self._waiter
573 finally:
574 self._waiter = None
575
576 message, payload = self._messages.popleft()
577
578 # time is only fetched if logging is enabled as otherwise
579 # its thrown away and never used.
580 start = loop.time() if self._logging_enabled else None
581
582 manager.requests_count += 1
583 writer = StreamWriter(self, loop)
584 if isinstance(message, _ErrInfo):
585 # make request_factory work
586 request_handler = self._make_error_handler(message)
587 message = ERROR
588 else:
589 request_handler = self._request_handler
590
591 # Important don't hold a reference to the current task
592 # as on traceback it will prevent the task from being
593 # collected and will cause a memory leak.
594 request = self._request_factory(
595 message,
596 payload,
597 self,
598 writer,
599 self._task_handler or asyncio.current_task(loop), # type: ignore[arg-type]
600 )
601 try:
602 # a new task is used for copy context vars (#3406)
603 coro = self._handle_request(request, start, request_handler)
604 if sys.version_info >= (3, 12):
605 task = asyncio.Task(coro, loop=loop, eager_start=True)
606 else:
607 task = loop.create_task(coro)
608 try:
609 resp, reset = await task
610 except ConnectionError:
611 self.log_debug("Ignored premature client disconnection")
612 break
613
614 # Drop the processed task from asyncio.Task.all_tasks() early
615 del task
616 if reset:
617 self.log_debug("Ignored premature client disconnection 2")
618 break
619
620 # notify server about keep-alive
621 self._keepalive = bool(resp.keep_alive)
622
623 # check payload
624 if not payload.is_eof():
625 lingering_time = self._lingering_time
626 if not self._force_close and lingering_time:
627 self.log_debug(
628 "Start lingering close timer for %s sec.", lingering_time
629 )
630
631 now = loop.time()
632 end_t = now + lingering_time
633
634 try:
635 while not payload.is_eof() and now < end_t:
636 async with ceil_timeout(end_t - now):
637 # read and ignore
638 await payload.readany()
639 now = loop.time()
640 except (asyncio.CancelledError, asyncio.TimeoutError):
641 if (
642 sys.version_info >= (3, 11)
643 and (t := asyncio.current_task())
644 and t.cancelling()
645 ):
646 raise
647
648 # if payload still uncompleted
649 if not payload.is_eof() and not self._force_close:
650 self.log_debug("Uncompleted request.")
651 self.close()
652
653 payload.set_exception(_PAYLOAD_ACCESS_ERROR)
654
655 except asyncio.CancelledError:
656 self.log_debug("Ignored premature client disconnection")
657 self.force_close()
658 raise
659 except Exception as exc:
660 self.log_exception("Unhandled exception", exc_info=exc)
661 self.force_close()
662 except BaseException:
663 self.force_close()
664 raise
665 finally:
666 request._task = None # type: ignore[assignment] # Break reference cycle in case of exception
667 if self.transport is None and resp is not None:
668 self.log_debug("Ignored premature client disconnection.")
669
670 if self._keepalive and not self._close and not self._force_close:
671 # start keep-alive timer
672 close_time = loop.time() + keepalive_timeout
673 self._next_keepalive_close_time = close_time
674 if self._keepalive_handle is None:
675 self._keepalive_handle = loop.call_at(
676 close_time, self._process_keepalive
677 )
678 else:
679 break
680
681 # remove handler, close transport if no handlers left
682 if not self._force_close:
683 self._task_handler = None
684 if self.transport is not None:
685 self.transport.close()
686
687 async def finish_response(
688 self, request: BaseRequest, resp: StreamResponse, start_time: Optional[float]
689 ) -> Tuple[StreamResponse, bool]:
690 """Prepare the response and write_eof, then log access.
691
692 This has to
693 be called within the context of any exception so the access logger
694 can get exception information. Returns True if the client disconnects
695 prematurely.
696 """
697 request._finish()
698 if self._request_parser is not None:
699 self._request_parser.set_upgraded(False)
700 self._upgrade = False
701 if self._message_tail:
702 self._request_parser.feed_data(self._message_tail)
703 self._message_tail = b""
704 try:
705 prepare_meth = resp.prepare
706 except AttributeError:
707 if resp is None:
708 self.log_exception("Missing return statement on request handler")
709 else:
710 self.log_exception(
711 "Web-handler should return a response instance, "
712 "got {!r}".format(resp)
713 )
714 exc = HTTPInternalServerError()
715 resp = Response(
716 status=exc.status, reason=exc.reason, text=exc.text, headers=exc.headers
717 )
718 prepare_meth = resp.prepare
719 try:
720 await prepare_meth(request)
721 await resp.write_eof()
722 except ConnectionError:
723 self.log_access(request, resp, start_time)
724 return resp, True
725
726 self.log_access(request, resp, start_time)
727 return resp, False
728
729 def handle_error(
730 self,
731 request: BaseRequest,
732 status: int = 500,
733 exc: Optional[BaseException] = None,
734 message: Optional[str] = None,
735 ) -> StreamResponse:
736 """Handle errors.
737
738 Returns HTTP response with specific status code. Logs additional
739 information. It always closes current connection.
740 """
741 if self._request_count == 1 and isinstance(exc, BadHttpMethod):
742 # BadHttpMethod is common when a client sends non-HTTP
743 # or encrypted traffic to an HTTP port. This is expected
744 # to happen when connected to the public internet so we log
745 # it at the debug level as to not fill logs with noise.
746 self.logger.debug(
747 "Error handling request from %s", request.remote, exc_info=exc
748 )
749 else:
750 self.log_exception(
751 "Error handling request from %s", request.remote, exc_info=exc
752 )
753
754 # some data already got sent, connection is broken
755 if request.writer.output_size > 0:
756 raise ConnectionError(
757 "Response is sent already, cannot send another response "
758 "with the error message"
759 )
760
761 ct = "text/plain"
762 if status == HTTPStatus.INTERNAL_SERVER_ERROR:
763 title = "{0.value} {0.phrase}".format(HTTPStatus.INTERNAL_SERVER_ERROR)
764 msg = HTTPStatus.INTERNAL_SERVER_ERROR.description
765 tb = None
766 if self.debug:
767 with suppress(Exception):
768 tb = traceback.format_exc()
769
770 if "text/html" in request.headers.get("Accept", ""):
771 if tb:
772 tb = html_escape(tb)
773 msg = f"<h2>Traceback:</h2>\n<pre>{tb}</pre>"
774 message = (
775 "<html><head>"
776 "<title>{title}</title>"
777 "</head><body>\n<h1>{title}</h1>"
778 "\n{msg}\n</body></html>\n"
779 ).format(title=title, msg=msg)
780 ct = "text/html"
781 else:
782 if tb:
783 msg = tb
784 message = title + "\n\n" + msg
785
786 resp = Response(status=status, text=message, content_type=ct)
787 resp.force_close()
788
789 return resp
790
791 def _make_error_handler(
792 self, err_info: _ErrInfo
793 ) -> Callable[[BaseRequest], Awaitable[StreamResponse]]:
794 async def handler(request: BaseRequest) -> StreamResponse:
795 return self.handle_error(
796 request, err_info.status, err_info.exc, err_info.message
797 )
798
799 return handler
800 