Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/zmq/eventloop/zmqstream.py: 27%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

278 statements  

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)