Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
messages.py349 linesDownload Raw Back to sync
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 
codekingpro/portable-devtools · Team Ai