1# Derived from iostream.py from tornado 1.0, Copyright 2009 Facebook
2# Used under Apache License Version 2.0
3#
4# Modifications are Copyright (C) PyZMQ Developers
5# Distributed under the terms of the Modified BSD License.
6"""A utility class for event-based messaging on a zmq socket using tornado.
7
8.. seealso::
9
10 - :mod:`zmq.asyncio`
11 - :mod:`zmq.eventloop.future`
12"""
13
14from __future__ import annotations
15
16import asyncio
17import pickle
18import warnings
19from collections.abc import Awaitable, Sequence
20from queue import Queue
21from typing import Any, Callable, Literal, overload
22
23from tornado.ioloop import IOLoop
24from tornado.log import gen_log
25
26import zmq
27import zmq._future
28from zmq import POLLIN, POLLOUT
29from zmq.utils import jsonapi
30
31
32class ZMQStream:
33 """A utility class to register callbacks when a zmq socket sends and receives
34
35 For use with tornado IOLoop.
36
37 There are three main methods
38
39 Methods:
40
41 * **on_recv(callback, copy=True):**
42 register a callback to be run every time the socket has something to receive
43 * **on_send(callback):**
44 register a callback to be run every time you call send
45 * **send_multipart(self, msg, flags=0, copy=False, callback=None):**
46 perform a send that will trigger the callback
47 if callback is passed, on_send is also called.
48
49 There are also send_multipart(), send_json(), send_pyobj()
50
51 Three other methods for deactivating the callbacks:
52
53 * **stop_on_recv():**
54 turn off the recv callback
55 * **stop_on_send():**
56 turn off the send callback
57
58 which simply call ``on_<evt>(None)``.
59
60 The entire socket interface, excluding direct recv methods, is also
61 provided, primarily through direct-linking the methods.
62 e.g.
63
64 >>> stream.bind is stream.socket.bind
65 True
66
67
68 .. versionadded:: 25
69
70 send/recv callbacks can be coroutines.
71
72 .. versionchanged:: 25
73
74 ZMQStreams only support base zmq.Socket classes (this has always been true, but not enforced).
75 If ZMQStreams are created with e.g. async Socket subclasses,
76 a RuntimeWarning will be shown,
77 and the socket cast back to the default zmq.Socket
78 before connecting events.
79
80 Previously, using async sockets (or any zmq.Socket subclass) would result in undefined behavior for the
81 arguments passed to callback functions.
82 Now, the callback functions reliably get the return value of the base `zmq.Socket` send/recv_multipart methods
83 (the list of message frames).
84 """
85
86 socket: zmq.Socket
87 io_loop: IOLoop
88 poller: zmq.Poller
89 _send_queue: Queue
90 _recv_callback: Callable | None
91 _send_callback: Callable | None
92 _close_callback: Callable | None
93 _state: int = 0
94 _flushed: bool = False
95 _recv_copy: bool = False
96 _fd: int
97
98 def __init__(self, socket: zmq.Socket, io_loop: IOLoop | None = None):
99 if isinstance(socket, zmq._future._AsyncSocket):
100 warnings.warn(
101 f"""ZMQStream only supports the base zmq.Socket class.
102
103 Use zmq.Socket(shadow=other_socket)
104 or `ctx.socket(zmq.{socket._type_name}, socket_class=zmq.Socket)`
105 to create a base zmq.Socket object,
106 no matter what other kind of socket your Context creates.
107 """,
108 RuntimeWarning,
109 stacklevel=2,
110 )
111 # shadow back to base zmq.Socket,
112 # otherwise callbacks like `on_recv` will get the wrong types.
113 socket = zmq.Socket(shadow=socket)
114 self.socket = socket
115
116 # IOLoop.current() is deprecated if called outside the event loop
117 # that means
118 self.io_loop = io_loop or IOLoop.current()
119 self.poller = zmq.Poller()
120 self._fd = self.socket.FD
121
122 self._send_queue = Queue()
123 self._recv_callback = None
124 self._send_callback = None
125 self._close_callback = None
126 self._recv_copy = False
127 self._flushed = False
128
129 self._state = 0
130 self._init_io_state()
131
132 # shortcircuit some socket methods
133 self.bind = self.socket.bind
134 self.bind_to_random_port = self.socket.bind_to_random_port
135 self.connect = self.socket.connect
136 self.setsockopt = self.socket.setsockopt
137 self.getsockopt = self.socket.getsockopt
138 self.setsockopt_string = self.socket.setsockopt_string
139 self.getsockopt_string = self.socket.getsockopt_string
140 self.setsockopt_unicode = self.socket.setsockopt_unicode
141 self.getsockopt_unicode = self.socket.getsockopt_unicode
142
143 def stop_on_recv(self):
144 """Disable callback and automatic receiving."""
145 return self.on_recv(None)
146
147 def stop_on_send(self):
148 """Disable callback on sending."""
149 return self.on_send(None)
150
151 def stop_on_err(self):
152 """DEPRECATED, does nothing"""
153 gen_log.warn("on_err does nothing, and will be removed")
154
155 def on_err(self, callback: Callable):
156 """DEPRECATED, does nothing"""
157 gen_log.warn("on_err does nothing, and will be removed")
158
159 @overload
160 def on_recv(
161 self,
162 callback: Callable[[list[bytes]], Any],
163 ) -> None: ...
164
165 @overload
166 def on_recv(
167 self,
168 callback: Callable[[list[bytes]], Any],
169 copy: Literal[True],
170 ) -> None: ...
171
172 @overload
173 def on_recv(
174 self,
175 callback: Callable[[list[zmq.Frame]], Any],
176 copy: Literal[False],
177 ) -> None: ...
178
179 @overload
180 def on_recv(
181 self,
182 callback: Callable[[list[zmq.Frame]], Any] | Callable[[list[bytes]], Any],
183 copy: bool = ...,
184 ): ...
185
186 def on_recv(
187 self,
188 callback: Callable[[list[zmq.Frame]], Any] | Callable[[list[bytes]], Any],
189 copy: bool = True,
190 ) -> None:
191 """Register a callback for when a message is ready to recv.
192
193 There can be only one callback registered at a time, so each
194 call to `on_recv` replaces previously registered callbacks.
195
196 on_recv(None) disables recv event polling.
197
198 Use on_recv_stream(callback) instead, to register a callback that will receive
199 both this ZMQStream and the message, instead of just the message.
200
201 Parameters
202 ----------
203
204 callback : callable
205 callback must take exactly one argument, which will be a
206 list, as returned by socket.recv_multipart()
207 if callback is None, recv callbacks are disabled.
208 copy : bool
209 copy is passed directly to recv, so if copy is False,
210 callback will receive Message objects. If copy is True,
211 then callback will receive bytes/str objects.
212
213 Returns : None
214 """
215
216 self._check_closed()
217 assert callback is None or callable(callback)
218 self._recv_callback = callback
219 self._recv_copy = copy
220 if callback is None:
221 self._drop_io_state(zmq.POLLIN)
222 else:
223 self._add_io_state(zmq.POLLIN)
224
225 @overload
226 def on_recv_stream(
227 self,
228 callback: Callable[[ZMQStream, list[bytes]], Any],
229 ) -> None: ...
230
231 @overload
232 def on_recv_stream(
233 self,
234 callback: Callable[[ZMQStream, list[bytes]], Any],
235 copy: Literal[True],
236 ) -> None: ...
237
238 @overload
239 def on_recv_stream(
240 self,
241 callback: Callable[[ZMQStream, list[zmq.Frame]], Any],
242 copy: Literal[False],
243 ) -> None: ...
244
245 @overload
246 def on_recv_stream(
247 self,
248 callback: (
249 Callable[[ZMQStream, list[zmq.Frame]], Any]
250 | Callable[[ZMQStream, list[bytes]], Any]
251 ),
252 copy: bool = ...,
253 ): ...
254
255 def on_recv_stream(
256 self,
257 callback: (
258 Callable[[ZMQStream, list[zmq.Frame]], Any]
259 | Callable[[ZMQStream, list[bytes]], Any]
260 ),
261 copy: bool = True,
262 ):
263 """Same as on_recv, but callback will get this stream as first argument
264
265 callback must take exactly two arguments, as it will be called as::
266
267 callback(stream, msg)
268
269 Useful when a single callback should be used with multiple streams.
270 """
271 if callback is None:
272 self.stop_on_recv()
273 else:
274
275 def stream_callback(msg):
276 return callback(self, msg)
277
278 self.on_recv(stream_callback, copy=copy)
279
280 def on_send(
281 self, callback: Callable[[Sequence[Any], zmq.MessageTracker | None], Any]
282 ):
283 """Register a callback to be called on each send
284
285 There will be two arguments::
286
287 callback(msg, status)
288
289 * `msg` will be the list of sendable objects that was just sent
290 * `status` will be the return result of socket.send_multipart(msg) -
291 MessageTracker or None.
292
293 Non-copying sends return a MessageTracker object whose
294 `done` attribute will be True when the send is complete.
295 This allows users to track when an object is safe to write to
296 again.
297
298 The second argument will always be None if copy=True
299 on the send.
300
301 Use on_send_stream(callback) to register a callback that will be passed
302 this ZMQStream as the first argument, in addition to the other two.
303
304 on_send(None) disables recv event polling.
305
306 Parameters
307 ----------
308
309 callback : callable
310 callback must take exactly two arguments, which will be
311 the message being sent (always a list),
312 and the return result of socket.send_multipart(msg) -
313 MessageTracker or None.
314
315 if callback is None, send callbacks are disabled.
316 """
317
318 self._check_closed()
319 assert callback is None or callable(callback)
320 self._send_callback = callback
321
322 def on_send_stream(
323 self,
324 callback: Callable[[ZMQStream, Sequence[Any], zmq.MessageTracker | None], Any],
325 ):
326 """Same as on_send, but callback will get this stream as first argument
327
328 Callback will be passed three arguments::
329
330 callback(stream, msg, status)
331
332 Useful when a single callback should be used with multiple streams.
333 """
334 if callback is None:
335 self.stop_on_send()
336 else:
337 self.on_send(lambda msg, status: callback(self, msg, status))
338
339 def send(self, msg, flags=0, copy=True, track=False, callback=None, **kwargs):
340 """Send a message, optionally also register a new callback for sends.
341 See zmq.socket.send for details.
342 """
343 return self.send_multipart(
344 [msg], flags=flags, copy=copy, track=track, callback=callback, **kwargs
345 )
346
347 def send_multipart(
348 self,
349 msg: Sequence[Any],
350 flags: int = 0,
351 copy: bool = True,
352 track: bool = False,
353 callback: Callable | None = None,
354 **kwargs: Any,
355 ) -> None:
356 """Send a multipart message, optionally also register a new callback for sends.
357 See zmq.socket.send_multipart for details.
358 """
359 kwargs.update(dict(flags=flags, copy=copy, track=track))
360 self._send_queue.put((msg, kwargs))
361 callback = callback or self._send_callback
362 if callback is not None:
363 self.on_send(callback)
364 else:
365 # noop callback
366 self.on_send(lambda *args: None)
367 self._add_io_state(zmq.POLLOUT)
368
369 def send_string(
370 self,
371 u: str,
372 flags: int = 0,
373 encoding: str = 'utf-8',
374 callback: Callable | None = None,
375 **kwargs: Any,
376 ):
377 """Send a unicode message with an encoding.
378 See zmq.socket.send_unicode for details.
379 """
380 if not isinstance(u, str):
381 raise TypeError("unicode/str objects only")
382 return self.send(u.encode(encoding), flags=flags, callback=callback, **kwargs)
383
384 send_unicode = send_string
385
386 def send_json(
387 self,
388 obj: Any,
389 flags: int = 0,
390 callback: Callable | None = None,
391 **kwargs: Any,
392 ):
393 """Send json-serialized version of an object.
394 See zmq.socket.send_json for details.
395 """
396 msg = jsonapi.dumps(obj)
397 return self.send(msg, flags=flags, callback=callback, **kwargs)
398
399 def send_pyobj(
400 self,
401 obj: Any,
402 flags: int = 0,
403 protocol: int = -1,
404 callback: Callable | None = None,
405 **kwargs: Any,
406 ):
407 """Send a Python object as a message using pickle to serialize.
408
409 See zmq.socket.send_json for details.
410 """
411 msg = pickle.dumps(obj, protocol)
412 return self.send(msg, flags, callback=callback, **kwargs)
413
414 def _finish_flush(self):
415 """callback for unsetting _flushed flag."""
416 self._flushed = False
417
418 def flush(self, flag: int = zmq.POLLIN | zmq.POLLOUT, limit: int | None = None):
419 """Flush pending messages.
420
421 This method safely handles all pending incoming and/or outgoing messages,
422 bypassing the inner loop, passing them to the registered callbacks.
423
424 A limit can be specified, to prevent blocking under high load.
425
426 flush will return the first time ANY of these conditions are met:
427 * No more events matching the flag are pending.
428 * the total number of events handled reaches the limit.
429
430 Note that if ``flag|POLLIN != 0``, recv events will be flushed even if no callback
431 is registered, unlike normal IOLoop operation. This allows flush to be
432 used to remove *and ignore* incoming messages.
433
434 Parameters
435 ----------
436 flag : int
437 default=POLLIN|POLLOUT
438 0MQ poll flags.
439 If flag|POLLIN, recv events will be flushed.
440 If flag|POLLOUT, send events will be flushed.
441 Both flags can be set at once, which is the default.
442 limit : None or int, optional
443 The maximum number of messages to send or receive.
444 Both send and recv count against this limit.
445
446 Returns
447 -------
448 int :
449 count of events handled (both send and recv)
450 """
451 self._check_closed()
452 # unset self._flushed, so callbacks will execute, in case flush has
453 # already been called this iteration
454 already_flushed = self._flushed
455 self._flushed = False
456 # initialize counters
457 count = 0
458
459 def update_flag():
460 """Update the poll flag, to prevent registering POLLOUT events
461 if we don't have pending sends."""
462 return flag & zmq.POLLIN | (self.sending() and flag & zmq.POLLOUT)
463
464 flag = update_flag()
465 if not flag:
466 # nothing to do
467 return 0
468 self.poller.register(self.socket, flag)
469 events = self.poller.poll(0)
470 while events and (not limit or count < limit):
471 s, event = events[0]
472 if event & POLLIN: # receiving
473 self._handle_recv()
474 count += 1
475 if self.socket is None:
476 # break if socket was closed during callback
477 break
478 if event & POLLOUT and self.sending():
479 self._handle_send()
480 count += 1
481 if self.socket is None:
482 # break if socket was closed during callback
483 break
484
485 flag = update_flag()
486 if flag:
487 self.poller.register(self.socket, flag)
488 events = self.poller.poll(0)
489 else:
490 events = []
491 if count: # only bypass loop if we actually flushed something
492 # skip send/recv callbacks this iteration
493 self._flushed = True
494 # reregister them at the end of the loop
495 if not already_flushed: # don't need to do it again
496 self.io_loop.add_callback(self._finish_flush)
497 elif already_flushed:
498 self._flushed = True
499
500 # update ioloop poll state, which may have changed
501 self._rebuild_io_state()
502 return count
503
504 def set_close_callback(self, callback: Callable | None):
505 """Call the given callback when the stream is closed."""
506 self._close_callback = callback
507
508 def close(self, linger: int | None = None) -> None:
509 """Close this stream."""
510 if self.socket is not None:
511 if self.socket.closed:
512 # fallback on raw fd for closed sockets
513 # hopefully this happened promptly after close,
514 # otherwise somebody else may have the FD
515 warnings.warn(
516 f"Unregistering FD {self._fd} after closing socket. "
517 "This could result in unregistering handlers for the wrong socket. "
518 "Please use stream.close() instead of closing the socket directly.",
519 stacklevel=2,
520 )
521 self.io_loop.remove_handler(self._fd)
522 else:
523 self.io_loop.remove_handler(self.socket)
524 self.socket.close(linger)
525 self.socket = None # type: ignore
526 if self._close_callback:
527 self._run_callback(self._close_callback)
528
529 def receiving(self) -> bool:
530 """Returns True if we are currently receiving from the stream."""
531 return self._recv_callback is not None
532
533 def sending(self) -> bool:
534 """Returns True if we are currently sending to the stream."""
535 return not self._send_queue.empty()
536
537 def closed(self) -> bool:
538 if self.socket is None:
539 return True
540 if self.socket.closed:
541 # underlying socket has been closed, but not by us!
542 # trigger our cleanup
543 self.close()
544 return True
545 return False
546
547 def _run_callback(self, callback, *args, **kwargs):
548 """Wrap running callbacks in try/except to allow us to
549 close our socket."""
550 try:
551 f = callback(*args, **kwargs)
552 if isinstance(f, Awaitable):
553 f = asyncio.ensure_future(f)
554 else:
555 f = None
556 except Exception:
557 gen_log.error("Uncaught exception in ZMQStream callback", exc_info=True)
558 # Re-raise the exception so that IOLoop.handle_callback_exception
559 # can see it and log the error
560 raise
561
562 if f is not None:
563 # handle async callbacks
564 def _log_error(f):
565 try:
566 f.result()
567 except Exception:
568 gen_log.error(
569 "Uncaught exception in ZMQStream callback", exc_info=True
570 )
571
572 f.add_done_callback(_log_error)
573
574 def _handle_events(self, fd, events):
575 """This method is the actual handler for IOLoop, that gets called whenever
576 an event on my socket is posted. It dispatches to _handle_recv, etc."""
577 if not self.socket:
578 gen_log.warning("Got events for closed stream %s", self)
579 return
580 try:
581 zmq_events = self.socket.EVENTS
582 except zmq.ContextTerminated:
583 gen_log.warning("Got events for stream %s after terminating context", self)
584 # trigger close check, this will unregister callbacks
585 self.closed()
586 return
587 except zmq.ZMQError as e:
588 # run close check
589 # shadow sockets may have been closed elsewhere,
590 # which should show up as ENOTSOCK here
591 if self.closed():
592 gen_log.warning(
593 "Got events for stream %s attached to closed socket: %s", self, e
594 )
595 else:
596 gen_log.error("Error getting events for %s: %s", self, e)
597 return
598 try:
599 # dispatch events:
600 if zmq_events & zmq.POLLIN and self.receiving():
601 self._handle_recv()
602 if not self.socket:
603 return
604 if zmq_events & zmq.POLLOUT and self.sending():
605 self._handle_send()
606 if not self.socket:
607 return
608
609 # rebuild the poll state
610 self._rebuild_io_state()
611 except Exception:
612 gen_log.error("Uncaught exception in zmqstream callback", exc_info=True)
613 raise
614
615 def _handle_recv(self):
616 """Handle a recv event."""
617 if self._flushed:
618 return
619 try:
620 msg = self.socket.recv_multipart(zmq.NOBLOCK, copy=self._recv_copy)
621 except zmq.ZMQError as e:
622 if e.errno == zmq.EAGAIN:
623 # state changed since poll event
624 pass
625 else:
626 raise
627 else:
628 if self._recv_callback:
629 callback = self._recv_callback
630 self._run_callback(callback, msg)
631
632 def _handle_send(self):
633 """Handle a send event."""
634 if self._flushed:
635 return
636 if not self.sending():
637 gen_log.error("Shouldn't have handled a send event")
638 return
639
640 msg, kwargs = self._send_queue.get()
641 try:
642 status = self.socket.send_multipart(msg, **kwargs)
643 except zmq.ZMQError as e:
644 gen_log.error("SEND Error: %s", e)
645 status = e
646 if self._send_callback:
647 callback = self._send_callback
648 self._run_callback(callback, msg, status)
649
650 def _check_closed(self):
651 if not self.socket:
652 raise OSError("Stream is closed")
653
654 def _rebuild_io_state(self):
655 """rebuild io state based on self.sending() and receiving()"""
656 if self.socket is None:
657 return
658 state = 0
659 if self.receiving():
660 state |= zmq.POLLIN
661 if self.sending():
662 state |= zmq.POLLOUT
663
664 self._state = state
665 self._update_handler(state)
666
667 def _add_io_state(self, state):
668 """Add io_state to poller."""
669 self._state = self._state | state
670 self._update_handler(self._state)
671
672 def _drop_io_state(self, state):
673 """Stop poller from watching an io_state."""
674 self._state = self._state & (~state)
675 self._update_handler(self._state)
676
677 def _update_handler(self, state):
678 """Update IOLoop handler with state."""
679 if self.socket is None:
680 return
681
682 if state & self.socket.events:
683 # events still exist that haven't been processed
684 # explicitly schedule handling to avoid missing events due to edge-triggered FDs
685 self.io_loop.add_callback(lambda: self._handle_events(self.socket, 0))
686
687 def _init_io_state(self):
688 """initialize the ioloop event handler"""
689 self.io_loop.add_handler(self.socket, self._handle_events, self.io_loop.READ)