Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
protocol.py273 linesDownload Raw Back to asyncio
1import asyncio
2from typing import Any, Callable, Dict, Optional, Text, Tuple, Union, cast
3
4from ..quic import events
5from ..quic.connection import NetworkAddress, QuicConnection
6from ..quic.packet import QuicErrorCode
7
8QuicConnectionIdHandler = Callable[[bytes], None]
9QuicStreamHandler = Callable[[asyncio.StreamReader, asyncio.StreamWriter], None]
10
11
12class QuicConnectionProtocol(asyncio.DatagramProtocol):
13    def __init__(
14        self, quic: QuicConnection, stream_handler: Optional[QuicStreamHandler] = None
15    ):
16        loop = asyncio.get_event_loop()
17
18        self._closed = asyncio.Event()
19        self._connected = False
20        self._connected_waiter: Optional[asyncio.Future[None]] = None
21        self._loop = loop
22        self._ping_waiters: Dict[int, asyncio.Future[None]] = {}
23        self._quic = quic
24        self._stream_readers: Dict[int, asyncio.StreamReader] = {}
25        self._timer: Optional[asyncio.TimerHandle] = None
26        self._timer_at: Optional[float] = None
27        self._transmit_task: Optional[asyncio.Handle] = None
28        self._transport: Optional[asyncio.DatagramTransport] = None
29
30        # callbacks
31        self._connection_id_issued_handler: QuicConnectionIdHandler = lambda c: None
32        self._connection_id_retired_handler: QuicConnectionIdHandler = lambda c: None
33        self._connection_terminated_handler: Callable[[], None] = lambda: None
34        if stream_handler is not None:
35            self._stream_handler = stream_handler
36        else:
37            self._stream_handler = lambda r, w: None
38
39    def change_connection_id(self) -> None:
40        """
41        Change the connection ID used to communicate with the peer.
42
43        The previous connection ID will be retired.
44        """
45        self._quic.change_connection_id()
46        self.transmit()
47
48    def close(
49        self,
50        error_code: int = QuicErrorCode.NO_ERROR,
51        reason_phrase: str = "",
52    ) -> None:
53        """
54        Close the connection.
55
56        :param error_code: An error code indicating why the connection is
57                           being closed.
58        :param reason_phrase: A human-readable explanation of why the
59                              connection is being closed.
60        """
61        self._quic.close(
62            error_code=error_code,
63            reason_phrase=reason_phrase,
64        )
65        self.transmit()
66
67    def connect(self, addr: NetworkAddress, transmit=True) -> None:
68        """
69        Initiate the TLS handshake.
70
71        This method can only be called for clients and a single time.
72        """
73        self._quic.connect(addr, now=self._loop.time())
74        if transmit:
75            self.transmit()
76
77    async def create_stream(
78        self, is_unidirectional: bool = False
79    ) -> Tuple[asyncio.StreamReader, asyncio.StreamWriter]:
80        """
81        Create a QUIC stream and return a pair of (reader, writer) objects.
82
83        The returned reader and writer objects are instances of
84        :class:`asyncio.StreamReader` and :class:`asyncio.StreamWriter` classes.
85        """
86        stream_id = self._quic.get_next_available_stream_id(
87            is_unidirectional=is_unidirectional
88        )
89        return self._create_stream(stream_id)
90
91    def request_key_update(self) -> None:
92        """
93        Request an update of the encryption keys.
94        """
95        self._quic.request_key_update()
96        self.transmit()
97
98    async def ping(self) -> None:
99        """
100        Ping the peer and wait for the response.
101        """
102        waiter = self._loop.create_future()
103        uid = id(waiter)
104        self._ping_waiters[uid] = waiter
105        self._quic.send_ping(uid)
106        self.transmit()
107        await asyncio.shield(waiter)
108
109    def transmit(self) -> None:
110        """
111        Send pending datagrams to the peer and arm the timer if needed.
112
113        This method is called automatically when data is received from the peer
114        or when a timer goes off. If you interact directly with the underlying
115        :class:`~aioquic.quic.connection.QuicConnection`, make sure you call this
116        method whenever data needs to be sent out to the network.
117        """
118        self._transmit_task = None
119
120        # send datagrams
121        for data, addr in self._quic.datagrams_to_send(now=self._loop.time()):
122            self._transport.sendto(data, addr)
123
124        # re-arm timer
125        timer_at = self._quic.get_timer()
126        if self._timer is not None and self._timer_at != timer_at:
127            self._timer.cancel()
128            self._timer = None
129        if self._timer is None and timer_at is not None:
130            self._timer = self._loop.call_at(timer_at, self._handle_timer)
131        self._timer_at = timer_at
132
133    async def wait_closed(self) -> None:
134        """
135        Wait for the connection to be closed.
136        """
137        await self._closed.wait()
138
139    async def wait_connected(self) -> None:
140        """
141        Wait for the TLS handshake to complete.
142        """
143        assert self._connected_waiter is None, "already awaiting connected"
144        if not self._connected:
145            self._connected_waiter = self._loop.create_future()
146            await asyncio.shield(self._connected_waiter)
147
148    # asyncio.Transport
149
150    def connection_made(self, transport: asyncio.BaseTransport) -> None:
151        """:meta private:"""
152        self._transport = cast(asyncio.DatagramTransport, transport)
153
154    def datagram_received(self, data: Union[bytes, Text], addr: NetworkAddress) -> None:
155        """:meta private:"""
156        self._quic.receive_datagram(cast(bytes, data), addr, now=self._loop.time())
157        self._process_events()
158        self.transmit()
159
160    # overridable
161
162    def quic_event_received(self, event: events.QuicEvent) -> None:
163        """
164        Called when a QUIC event is received.
165
166        Reimplement this in your subclass to handle the events.
167        """
168        # FIXME: move this to a subclass
169        if isinstance(event, events.ConnectionTerminated):
170            for reader in self._stream_readers.values():
171                reader.feed_eof()
172        elif isinstance(event, events.StreamDataReceived):
173            reader = self._stream_readers.get(event.stream_id, None)
174            if reader is None:
175                reader, writer = self._create_stream(event.stream_id)
176                self._stream_handler(reader, writer)
177            reader.feed_data(event.data)
178            if event.end_stream:
179                reader.feed_eof()
180
181    # private
182
183    def _create_stream(
184        self, stream_id: int
185    ) -> Tuple[asyncio.StreamReader, asyncio.StreamWriter]:
186        adapter = QuicStreamAdapter(self, stream_id)
187        reader = asyncio.StreamReader()
188        protocol = asyncio.streams.StreamReaderProtocol(reader)
189        writer = asyncio.StreamWriter(adapter, protocol, reader, self._loop)
190        self._stream_readers[stream_id] = reader
191        return reader, writer
192
193    def _handle_timer(self) -> None:
194        now = max(self._timer_at, self._loop.time())
195        self._timer = None
196        self._timer_at = None
197        self._quic.handle_timer(now=now)
198        self._process_events()
199        self.transmit()
200
201    def _process_events(self) -> None:
202        event = self._quic.next_event()
203        while event is not None:
204            if isinstance(event, events.ConnectionIdIssued):
205                self._connection_id_issued_handler(event.connection_id)
206            elif isinstance(event, events.ConnectionIdRetired):
207                self._connection_id_retired_handler(event.connection_id)
208            elif isinstance(event, events.ConnectionTerminated):
209                self._connection_terminated_handler()
210
211                # abort connection waiter
212                if self._connected_waiter is not None:
213                    waiter = self._connected_waiter
214                    self._connected_waiter = None
215                    waiter.set_exception(ConnectionError)
216
217                # abort ping waiters
218                for waiter in self._ping_waiters.values():
219                    waiter.set_exception(ConnectionError)
220                self._ping_waiters.clear()
221
222                self._closed.set()
223            elif isinstance(event, events.HandshakeCompleted):
224                if self._connected_waiter is not None:
225                    waiter = self._connected_waiter
226                    self._connected = True
227                    self._connected_waiter = None
228                    waiter.set_result(None)
229            elif isinstance(event, events.PingAcknowledged):
230                waiter = self._ping_waiters.pop(event.uid, None)
231                if waiter is not None:
232                    waiter.set_result(None)
233            self.quic_event_received(event)
234            event = self._quic.next_event()
235
236    def _transmit_soon(self) -> None:
237        if self._transmit_task is None:
238            self._transmit_task = self._loop.call_soon(self.transmit)
239
240
241class QuicStreamAdapter(asyncio.Transport):
242    def __init__(self, protocol: QuicConnectionProtocol, stream_id: int):
243        self.protocol = protocol
244        self.stream_id = stream_id
245        self._closing = False
246
247    def can_write_eof(self) -> bool:
248        return True
249
250    def get_extra_info(self, name: str, default: Any = None) -> Any:
251        """
252        Get information about the underlying QUIC stream.
253        """
254        if name == "stream_id":
255            return self.stream_id
256
257    def write(self, data):
258        self.protocol._quic.send_stream_data(self.stream_id, data)
259        self.protocol._transmit_soon()
260
261    def write_eof(self):
262        if self._closing:
263            return
264        self._closing = True
265        self.protocol._quic.send_stream_data(self.stream_id, b"", end_stream=True)
266        self.protocol._transmit_soon()
267
268    def close(self):
269        self.write_eof()
270
271    def is_closing(self) -> bool:
272        return self._closing
273 
codekingpro/portable-devtools · Team Ai