codekingpro/portable-devtools
114k
1from __future__ import annotations
2
3import codecs
4import queue
5import threading
6from typing import Any, Callable, Iterable, Iterator, Literal, overload
7
8from ..exceptions import ConcurrencyError
9from ..frames import OP_BINARY, OP_CONT, OP_TEXT, Frame
10from ..typing import Data
11from .utils import Deadline
12
13
14__all__ = ["Assembler"]
15
16UTF8Decoder = codecs.getincrementaldecoder("utf-8")
17
18
19class Assembler:
20 """
21 Assemble messages from frames.
22
23 :class:`Assembler` expects only data frames. The stream of frames must
24 respect the protocol; if it doesn't, the behavior is undefined.
25
26 Args:
27 pause: Called when the buffer of frames goes above the high water mark;
28 should pause reading from the network.
29 resume: Called when the buffer of frames goes below the low water mark;
30 should resume reading from the network.
31
32 """
33
34 def __init__(
35 self,
36 high: int | None = None,
37 low: int | None = None,
38 pause: Callable[[], Any] = lambda: None,
39 resume: Callable[[], Any] = lambda: None,
40 ) -> None:
41 # Serialize reads and writes -- except for reads via synchronization
42 # primitives provided by the threading and queue modules.
43 self.mutex = threading.Lock()
44
45 # Queue of incoming frames.
46 self.frames: queue.SimpleQueue[Frame | None] = queue.SimpleQueue()
47
48 # We cannot put a hard limit on the size of the queue because a single
49 # call to Protocol.data_received() could produce thousands of frames,
50 # which must be buffered. Instead, we pause reading when the buffer goes
51 # above the high limit and we resume when it goes under the low limit.
52 if high is not None and low is None:
53 low = high // 4
54 if high is None and low is not None:
55 high = low * 4
56 if high is not None and low is not None:
57 if low < 0:
58 raise ValueError("low must be positive or equal to zero")
59 if high < low:
60 raise ValueError("high must be greater than or equal to low")
61 self.high, self.low = high, low
62 self.pause = pause
63 self.resume = resume
64 self.paused = False
65
66 # This flag prevents concurrent calls to get() by user code.
67 self.get_in_progress = False
68
69 # This flag marks the end of the connection.
70 self.closed = False
71
72 def get_next_frame(self, timeout: float | None = None) -> Frame:
73 # Helper to factor out the logic for getting the next frame from the
74 # queue, while handling timeouts and reaching the end of the stream.
75 if self.closed:
76 try:
77 frame = self.frames.get(block=False)
78 except queue.Empty:
79 raise EOFError("stream of frames ended") from None
80 else:
81 try:
82 # Check for a frame that's already received if timeout <= 0.
83 # SimpleQueue.get() doesn't support negative timeout values.
84 if timeout is not None and timeout <= 0:
85 frame = self.frames.get(block=False)
86 else:
87 frame = self.frames.get(block=True, timeout=timeout)
88 except queue.Empty:
89 raise TimeoutError(f"timed out in {timeout:.1f}s") from None
90 if frame is None:
91 raise EOFError("stream of frames ended")
92 return frame
93
94 def reset_queue(self, frames: Iterable[Frame]) -> None:
95 # Helper to put frames back into the queue after they were fetched.
96 # This happens only when the queue is empty. However, by the time
97 # we acquire self.mutex, put() may have added items in the queue.
98 # Therefore, we must handle the case where the queue is not empty.
99 frame: Frame | None
100 with self.mutex:
101 queued = []
102 try:
103 while True:
104 queued.append(self.frames.get(block=False))
105 except queue.Empty:
106 pass
107 for frame in frames:
108 self.frames.put(frame)
109 # This loop runs only when a race condition occurs.
110 for frame in queued: # pragma: no cover
111 self.frames.put(frame)
112
113 # This overload structure is required to avoid the error:
114 # "parameter without a default follows parameter with a default"
115
116 @overload
117 def get(self, timeout: float | None, decode: Literal[True]) -> str: ...
118
119 @overload
120 def get(self, timeout: float | None, decode: Literal[False]) -> bytes: ...
121
122 @overload
123 def get(self, timeout: float | None = None, *, decode: Literal[True]) -> str: ...
124
125 @overload
126 def get(self, timeout: float | None = None, *, decode: Literal[False]) -> bytes: ...
127
128 @overload
129 def get(self, timeout: float | None = None, decode: bool | None = None) -> Data: ...
130
131 def get(self, timeout: float | None = None, decode: bool | None = None) -> Data:
132 """
133 Read the next message.
134
135 :meth:`get` returns a single :class:`str` or :class:`bytes`.
136
137 If the message is fragmented, :meth:`get` waits until the last frame is
138 received, then it reassembles the message and returns it. To receive
139 messages frame by frame, use :meth:`get_iter` instead.
140
141 Args:
142 timeout: If a timeout is provided and elapses before a complete
143 message is received, :meth:`get` raises :exc:`TimeoutError`.
144 decode: :obj:`False` disables UTF-8 decoding of text frames and
145 returns :class:`bytes`. :obj:`True` forces UTF-8 decoding of
146 binary frames and returns :class:`str`.
147
148 Raises:
149 EOFError: If the stream of frames has ended.
150 UnicodeDecodeError: If a text frame contains invalid UTF-8.
151 ConcurrencyError: If two coroutines run :meth:`get` or
152 :meth:`get_iter` concurrently.
153 TimeoutError: If a timeout is provided and elapses before a
154 complete message is received.
155
156 """
157 with self.mutex:
158 if self.get_in_progress:
159 raise ConcurrencyError("get() or get_iter() is already running")
160 self.get_in_progress = True
161
162 # Locking with get_in_progress prevents concurrent execution
163 # until get() fetches a complete message or times out.
164
165 try:
166 deadline = Deadline(timeout)
167
168 # Fetch the first frame.
169 frame = self.get_next_frame(deadline.timeout(raise_if_elapsed=False))
170 with self.mutex:
171 self.maybe_resume()
172 assert frame.opcode is OP_TEXT or frame.opcode is OP_BINARY
173 if decode is None:
174 decode = frame.opcode is OP_TEXT
175 frames = [frame]
176
177 # Fetch subsequent frames for fragmented messages.
178 while not frame.fin:
179 try:
180 frame = self.get_next_frame(
181 deadline.timeout(raise_if_elapsed=False)
182 )
183 except TimeoutError:
184 # Put frames already received back into the queue
185 # so that future calls to get() can return them.
186 self.reset_queue(frames)
187 raise
188 with self.mutex:
189 self.maybe_resume()
190 assert frame.opcode is OP_CONT
191 frames.append(frame)
192
193 finally:
194 self.get_in_progress = False
195
196 # This converts frame.data to bytes when it's a bytearray.
197 data = b"".join(frame.data for frame in frames)
198 if decode:
199 return data.decode()
200 else:
201 return data
202
203 @overload
204 def get_iter(self, decode: Literal[True]) -> Iterator[str]: ...
205
206 @overload
207 def get_iter(self, decode: Literal[False]) -> Iterator[bytes]: ...
208
209 @overload
210 def get_iter(self, decode: bool | None = None) -> Iterator[Data]: ...
211
212 def get_iter(self, decode: bool | None = None) -> Iterator[Data]:
213 """
214 Stream the next message.
215
216 Iterating the return value of :meth:`get_iter` yields a :class:`str` or
217 :class:`bytes` for each frame in the message.
218
219 The iterator must be fully consumed before calling :meth:`get_iter` or
220 :meth:`get` again. Else, :exc:`ConcurrencyError` is raised.
221
222 This method only makes sense for fragmented messages. If messages aren't
223 fragmented, use :meth:`get` instead.
224
225 Args:
226 decode: :obj:`False` disables UTF-8 decoding of text frames and
227 returns :class:`bytes`. :obj:`True` forces UTF-8 decoding of
228 binary frames and returns :class:`str`.
229
230 Raises:
231 EOFError: If the stream of frames has ended.
232 UnicodeDecodeError: If a text frame contains invalid UTF-8.
233 ConcurrencyError: If two coroutines run :meth:`get` or
234 :meth:`get_iter` concurrently.
235
236 """
237 with self.mutex:
238 if self.get_in_progress:
239 raise ConcurrencyError("get() or get_iter() is already running")
240 self.get_in_progress = True
241
242 # Locking with get_in_progress prevents concurrent execution
243 # until get_iter() fetches a complete message or times out.
244
245 # If get_iter() raises an exception e.g. in decoder.decode(),
246 # get_in_progress remains set and the connection becomes unusable.
247
248 # Yield the first frame.
249 frame = self.get_next_frame()
250 with self.mutex:
251 self.maybe_resume()
252 assert frame.opcode is OP_TEXT or frame.opcode is OP_BINARY
253 if decode is None:
254 decode = frame.opcode is OP_TEXT
255 if decode:
256 decoder = UTF8Decoder()
257 yield decoder.decode(frame.data, frame.fin)
258 else:
259 # Convert to bytes when frame.data is a bytearray.
260 yield bytes(frame.data)
261
262 # Yield subsequent frames for fragmented messages.
263 while not frame.fin:
264 frame = self.get_next_frame()
265 with self.mutex:
266 self.maybe_resume()
267 assert frame.opcode is OP_CONT
268 if decode:
269 yield decoder.decode(frame.data, frame.fin)
270 else:
271 # Convert to bytes when frame.data is a bytearray.
272 yield bytes(frame.data)
273
274 self.get_in_progress = False
275
276 def put(self, frame: Frame) -> None:
277 """
278 Add ``frame`` to the next message.
279
280 Raises:
281 EOFError: If the stream of frames has ended.
282
283 """
284 with self.mutex:
285 if self.closed:
286 raise EOFError("stream of frames ended")
287
288 self.frames.put(frame)
289 self.maybe_pause()
290
291 # put() and get/get_iter() call maybe_pause() and maybe_resume() while
292 # holding self.mutex. This guarantees that the calls interleave properly.
293 # Specifically, it prevents a race condition where maybe_resume() would
294 # run before maybe_pause(), leaving the connection incorrectly paused.
295
296 # A race condition is possible when get/get_iter() call self.frames.get()
297 # without holding self.mutex. However, it's harmless — and even beneficial!
298 # It can only result in popping an item from the queue before maybe_resume()
299 # runs and skipping a pause() - resume() cycle that would otherwise occur.
300
301 def maybe_pause(self) -> None:
302 """Pause the writer if queue is above the high water mark."""
303 # Skip if flow control is disabled.
304 if self.high is None:
305 return
306
307 assert self.mutex.locked()
308
309 # Check for "> high" to support high = 0.
310 if self.frames.qsize() > self.high and not self.paused:
311 self.paused = True
312 self.pause()
313
314 def maybe_resume(self) -> None:
315 """Resume the writer if queue is below the low water mark."""
316 # Skip if flow control is disabled.
317 if self.low is None:
318 return
319
320 assert self.mutex.locked()
321
322 # Check for "<= low" to support low = 0.
323 if self.frames.qsize() <= self.low and self.paused:
324 self.paused = False
325 self.resume()
326
327 def close(self) -> None:
328 """
329 End the stream of frames.
330
331 Calling :meth:`close` concurrently with :meth:`get`, :meth:`get_iter`,
332 or :meth:`put` is safe. They will raise :exc:`EOFError`.
333
334 """
335 with self.mutex:
336 if self.closed:
337 return
338
339 self.closed = True
340
341 if self.get_in_progress:
342 # Unblock get() or get_iter().
343 self.frames.put(None)
344
345 if self.paused:
346 # Unblock recv_events().
347 self.paused = False
348 self.resume()
349 