File size: 3,563 Bytes
31dc8dc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
from __future__ import annotations

from typing import Callable, Generic, TypeVar

T = TypeVar("T")


def _load_msgpack():
    try:
        import msgpack
    except ImportError as exc:
        raise RuntimeError("Diffulex ZMQ serving requires msgpack. Install the project dependencies first.") from exc
    return msgpack


def _load_zmq():
    try:
        import zmq
    except ImportError as exc:
        raise RuntimeError("Diffulex ZMQ serving requires pyzmq. Install the project dependencies first.") from exc
    return zmq


def _load_zmq_asyncio():
    try:
        import zmq.asyncio
    except ImportError as exc:
        raise RuntimeError("Diffulex ZMQ serving requires pyzmq. Install the project dependencies first.") from exc
    return zmq.asyncio


class ZmqPushQueue(Generic[T]):
    def __init__(self, addr: str, *, create: bool, encoder: Callable[[T], dict]):
        zmq = _load_zmq()
        self._msgpack = _load_msgpack()
        self.context = zmq.Context()
        self.socket = self.context.socket(zmq.PUSH)
        self.socket.setsockopt(zmq.LINGER, 0)
        self.socket.bind(addr) if create else self.socket.connect(addr)
        self.encoder = encoder

    def put(self, obj: T) -> None:
        event = self._msgpack.packb(self.encoder(obj), use_bin_type=True)
        self.socket.send(event, copy=False)

    def stop(self) -> None:
        self.socket.close()
        self.context.term()


class ZmqPullQueue(Generic[T]):
    def __init__(self, addr: str, *, create: bool, decoder: Callable[[dict], T]):
        zmq = _load_zmq()
        self._msgpack = _load_msgpack()
        self.context = zmq.Context()
        self.socket = self.context.socket(zmq.PULL)
        self.socket.setsockopt(zmq.LINGER, 0)
        self.socket.bind(addr) if create else self.socket.connect(addr)
        self.decoder = decoder

    def get(self) -> T:
        event = self.socket.recv()
        return self.decoder(self._msgpack.unpackb(event, raw=False))

    def empty(self) -> bool:
        return self.socket.poll(timeout=0) == 0

    def stop(self) -> None:
        self.socket.close()
        self.context.term()


class ZmqAsyncPushQueue(Generic[T]):
    def __init__(self, addr: str, *, create: bool, encoder: Callable[[T], dict]):
        zmq_asyncio = _load_zmq_asyncio()
        self._msgpack = _load_msgpack()
        self.context = zmq_asyncio.Context()
        zmq = _load_zmq()
        self.socket = self.context.socket(zmq.PUSH)
        self.socket.setsockopt(zmq.LINGER, 0)
        self.socket.bind(addr) if create else self.socket.connect(addr)
        self.encoder = encoder

    async def put(self, obj: T) -> None:
        event = self._msgpack.packb(self.encoder(obj), use_bin_type=True)
        await self.socket.send(event, copy=False)

    def stop(self) -> None:
        self.socket.close()
        self.context.term()


class ZmqAsyncPullQueue(Generic[T]):
    def __init__(self, addr: str, *, create: bool, decoder: Callable[[dict], T]):
        zmq_asyncio = _load_zmq_asyncio()
        self._msgpack = _load_msgpack()
        self.context = zmq_asyncio.Context()
        zmq = _load_zmq()
        self.socket = self.context.socket(zmq.PULL)
        self.socket.setsockopt(zmq.LINGER, 0)
        self.socket.bind(addr) if create else self.socket.connect(addr)
        self.decoder = decoder

    async def get(self) -> T:
        event = await self.socket.recv()
        return self.decoder(self._msgpack.unpackb(event, raw=False))

    def stop(self) -> None:
        self.socket.close()
        self.context.term()