Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
web_protocol.py800 linesDownload Raw Back to aiohttp
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 
codekingpro/portable-devtools · Team Ai