codekingpro/portable-devtools
114k
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 