socket-netty 0.3.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pynetty/__init__.py +137 -0
- pynetty/bootstrap/__init__.py +3 -0
- pynetty/bootstrap/bootstrap.py +162 -0
- pynetty/buffer/__init__.py +18 -0
- pynetty/buffer/allocator.py +113 -0
- pynetty/buffer/bytebuf.py +256 -0
- pynetty/channel/__init__.py +24 -0
- pynetty/channel/channel.py +130 -0
- pynetty/channel/channel_future.py +143 -0
- pynetty/channel/channel_option.py +59 -0
- pynetty/channel/channel_pipeline.py +173 -0
- pynetty/channel/datagram_channel.py +164 -0
- pynetty/channel/event_loop.py +157 -0
- pynetty/channel/flow_control.py +99 -0
- pynetty/channel/protocol_adapter.py +59 -0
- pynetty/exceptions.py +94 -0
- pynetty/handler/__init__.py +47 -0
- pynetty/handler/channel_handler.py +83 -0
- pynetty/handler/channel_handler_context.py +104 -0
- pynetty/handler/codec.py +152 -0
- pynetty/handler/protolib_codec.py +152 -0
- pynetty/handler/ssl_context.py +62 -0
- pynetty/handler/timeout.py +159 -0
- socket_netty-0.3.0.dist-info/METADATA +402 -0
- socket_netty-0.3.0.dist-info/RECORD +27 -0
- socket_netty-0.3.0.dist-info/WHEEL +5 -0
- socket_netty-0.3.0.dist-info/top_level.txt +1 -0
pynetty/handler/codec.py
ADDED
|
@@ -0,0 +1,152 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Common codecs, equivalent to io.netty.handler.codec.*
|
|
3
|
+
|
|
4
|
+
- LengthFieldBasedFrameDecoder: reassembles TCP frames using a length
|
|
5
|
+
field at the start of each message (the classic "TCP doesn't respect
|
|
6
|
+
message boundaries" problem).
|
|
7
|
+
- LengthFieldPrepender: prepends the length field before writing.
|
|
8
|
+
|
|
9
|
+
Very useful for game protocols (Minecraft, Free Fire, etc.) where each
|
|
10
|
+
packet is preceded by its size.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import struct
|
|
16
|
+
from typing import Any
|
|
17
|
+
|
|
18
|
+
from pynetty.buffer.bytebuf import ByteBuf
|
|
19
|
+
from pynetty.exceptions import (
|
|
20
|
+
CorruptedFrameException,
|
|
21
|
+
DecoderException,
|
|
22
|
+
EncoderException,
|
|
23
|
+
TooLongFrameException,
|
|
24
|
+
)
|
|
25
|
+
from pynetty.handler.channel_handler import ChannelInboundHandler, ChannelOutboundHandler
|
|
26
|
+
from pynetty.handler.channel_handler_context import ChannelHandlerContext
|
|
27
|
+
|
|
28
|
+
# Netty's own default cap for LengthFieldBasedFrameDecoder-style decoders
|
|
29
|
+
# is unbounded unless you specify max_frame_length; we mirror that but
|
|
30
|
+
# make it configurable, matching real Netty usage patterns.
|
|
31
|
+
_DEFAULT_MAX_FRAME_LENGTH = 8 * 1024 * 1024 # 8 MiB, a sane game-protocol default
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class LengthFieldBasedFrameDecoder(ChannelInboundHandler):
|
|
35
|
+
"""
|
|
36
|
+
Accumulates bytes until a full frame is available, using a
|
|
37
|
+
fixed-size length field (1, 2, or 4 bytes) at the start of each
|
|
38
|
+
message.
|
|
39
|
+
|
|
40
|
+
length_field_length: how many bytes the length field occupies (1/2/4).
|
|
41
|
+
strip_length_field: if True, the message delivered to channel_read
|
|
42
|
+
does not include the length field, only the payload.
|
|
43
|
+
max_frame_length: frames whose declared length exceeds this raise
|
|
44
|
+
TooLongFrameException instead of buffering unbounded memory.
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
_FORMATS = {1: ">B", 2: ">H", 4: ">I"}
|
|
48
|
+
|
|
49
|
+
def __init__(
|
|
50
|
+
self,
|
|
51
|
+
length_field_length: int = 4,
|
|
52
|
+
strip_length_field: bool = True,
|
|
53
|
+
max_frame_length: int = _DEFAULT_MAX_FRAME_LENGTH,
|
|
54
|
+
) -> None:
|
|
55
|
+
if length_field_length not in self._FORMATS:
|
|
56
|
+
raise ValueError("length_field_length must be 1, 2, or 4")
|
|
57
|
+
self._len_fmt = self._FORMATS[length_field_length]
|
|
58
|
+
self._len_size = length_field_length
|
|
59
|
+
self._strip = strip_length_field
|
|
60
|
+
self._max_frame_length = max_frame_length
|
|
61
|
+
self._buffer = bytearray()
|
|
62
|
+
|
|
63
|
+
async def channel_read(self, ctx: ChannelHandlerContext, msg: Any) -> None:
|
|
64
|
+
data = msg if isinstance(msg, (bytes, bytearray)) else bytes(msg)
|
|
65
|
+
self._buffer.extend(data)
|
|
66
|
+
|
|
67
|
+
try:
|
|
68
|
+
while True:
|
|
69
|
+
if len(self._buffer) < self._len_size:
|
|
70
|
+
break
|
|
71
|
+
(frame_len,) = struct.unpack(self._len_fmt, self._buffer[: self._len_size])
|
|
72
|
+
|
|
73
|
+
if frame_len < 0:
|
|
74
|
+
raise CorruptedFrameException(
|
|
75
|
+
f"negative pre-adjustment length field: {frame_len}"
|
|
76
|
+
)
|
|
77
|
+
if frame_len > self._max_frame_length:
|
|
78
|
+
# Drop what we have; there's no reliable way to resync
|
|
79
|
+
# mid-stream once a bogus length is trusted.
|
|
80
|
+
self._buffer.clear()
|
|
81
|
+
raise TooLongFrameException(
|
|
82
|
+
f"frame length ({frame_len}) exceeds the configured maximum "
|
|
83
|
+
f"({self._max_frame_length})"
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
total = self._len_size + frame_len
|
|
87
|
+
if len(self._buffer) < total:
|
|
88
|
+
break
|
|
89
|
+
|
|
90
|
+
if self._strip:
|
|
91
|
+
frame = bytes(self._buffer[self._len_size : total])
|
|
92
|
+
else:
|
|
93
|
+
frame = bytes(self._buffer[:total])
|
|
94
|
+
|
|
95
|
+
del self._buffer[:total]
|
|
96
|
+
await ctx.fire_channel_read(frame)
|
|
97
|
+
except (CorruptedFrameException, TooLongFrameException):
|
|
98
|
+
raise
|
|
99
|
+
except Exception as exc:
|
|
100
|
+
raise DecoderException("Failed to decode packet") from exc
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class LengthFieldPrepender(ChannelOutboundHandler):
|
|
104
|
+
"""Prepends the payload size before writing it to the socket."""
|
|
105
|
+
|
|
106
|
+
_FORMATS = {1: ">B", 2: ">H", 4: ">I"}
|
|
107
|
+
|
|
108
|
+
def __init__(self, length_field_length: int = 4) -> None:
|
|
109
|
+
if length_field_length not in self._FORMATS:
|
|
110
|
+
raise ValueError("length_field_length must be 1, 2, or 4")
|
|
111
|
+
self._len_fmt = self._FORMATS[length_field_length]
|
|
112
|
+
|
|
113
|
+
async def write(self, ctx: ChannelHandlerContext, msg: Any) -> None:
|
|
114
|
+
try:
|
|
115
|
+
if isinstance(msg, ByteBuf):
|
|
116
|
+
payload = msg.to_bytes()
|
|
117
|
+
elif isinstance(msg, (bytes, bytearray)):
|
|
118
|
+
payload = bytes(msg)
|
|
119
|
+
else:
|
|
120
|
+
raise TypeError(
|
|
121
|
+
f"LengthFieldPrepender expects bytes/ByteBuf, got {type(msg).__name__}"
|
|
122
|
+
)
|
|
123
|
+
header = struct.pack(self._len_fmt, len(payload))
|
|
124
|
+
except TypeError:
|
|
125
|
+
raise
|
|
126
|
+
except Exception as exc:
|
|
127
|
+
raise EncoderException("Failed to encode packet") from exc
|
|
128
|
+
await ctx.write(header + payload)
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
class ByteToMessageCodec(ChannelInboundHandler):
|
|
132
|
+
"""
|
|
133
|
+
Simple base for decoders that operate on an accumulated ByteBuf
|
|
134
|
+
instead of raw `bytes` fragments. Subclass and override `decode`.
|
|
135
|
+
"""
|
|
136
|
+
|
|
137
|
+
def __init__(self) -> None:
|
|
138
|
+
self._cumulation = ByteBuf()
|
|
139
|
+
|
|
140
|
+
async def channel_read(self, ctx: ChannelHandlerContext, msg: Any) -> None:
|
|
141
|
+
data = msg if isinstance(msg, (bytes, bytearray)) else bytes(msg)
|
|
142
|
+
self._cumulation.write_bytes(data)
|
|
143
|
+
try:
|
|
144
|
+
await self.decode(ctx, self._cumulation)
|
|
145
|
+
except (CorruptedFrameException, TooLongFrameException):
|
|
146
|
+
raise
|
|
147
|
+
except Exception as exc:
|
|
148
|
+
raise DecoderException("Failed to decode packet") from exc
|
|
149
|
+
|
|
150
|
+
async def decode(self, ctx: ChannelHandlerContext, buf: ByteBuf) -> None:
|
|
151
|
+
"""Override: read from `buf` while complete frames are available."""
|
|
152
|
+
raise NotImplementedError
|
|
@@ -0,0 +1,152 @@
|
|
|
1
|
+
"""
|
|
2
|
+
ProtolibCodec: bridges pynetty's ChannelPipeline with protolib's
|
|
3
|
+
declarative binary-protocol engine (https://pypi.org/project/protolib/).
|
|
4
|
+
|
|
5
|
+
protolib turns bytes into a dict (`{name, params}`) and back, based on
|
|
6
|
+
a `.yml`/`.json` protocol description — instead of you hand-writing a
|
|
7
|
+
parser for every packet type. This codec plugs that engine directly
|
|
8
|
+
into the pipeline: `channel_read` delivers a dict, and `write` accepts
|
|
9
|
+
a `(name, params)` tuple (or a dict with those two keys) and turns it
|
|
10
|
+
into bytes.
|
|
11
|
+
|
|
12
|
+
Two framing modes, controlled by `framed`:
|
|
13
|
+
|
|
14
|
+
- `framed=True` (default): the codec owns framing itself via
|
|
15
|
+
protolib's own `PacketFramer` (varint length-prefix, Minecraft-style
|
|
16
|
+
wire format). Do NOT also add a LengthFieldBasedFrameDecoder in the
|
|
17
|
+
pipeline in this mode — ProtolibCodec already consumes raw,
|
|
18
|
+
unframed bytes straight from the socket.
|
|
19
|
+
- `framed=False`: the codec assumes it's receiving an already-complete
|
|
20
|
+
frame (typically because a LengthFieldBasedFrameDecoder — or any
|
|
21
|
+
other decoder — sits earlier in the pipeline). Use this for
|
|
22
|
+
protocols with fixed-size or otherwise non-varint framing, such as
|
|
23
|
+
Minecraft Classic/ClassiCube. In this mode the codec calls
|
|
24
|
+
`Protocol.parse_packet`/`serialize_packet` directly and does no
|
|
25
|
+
buffering of its own.
|
|
26
|
+
|
|
27
|
+
Any error protolib raises while parsing/serializing (ProtolibError and
|
|
28
|
+
subclasses, including BufferUnderrun escaping in framed=False mode) is
|
|
29
|
+
wrapped into pynetty's own DecoderException/EncoderException, so a
|
|
30
|
+
broken packet looks the same everywhere in the pipeline:
|
|
31
|
+
|
|
32
|
+
io.netty.handler.codec.DecoderException: Failed to decode packet
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
from __future__ import annotations
|
|
36
|
+
|
|
37
|
+
from typing import Any, Optional, Tuple, Union
|
|
38
|
+
|
|
39
|
+
from pynetty.exceptions import DecoderException, EncoderException
|
|
40
|
+
from pynetty.handler.channel_handler import ChannelInboundHandler, ChannelOutboundHandler
|
|
41
|
+
from pynetty.handler.channel_handler_context import ChannelHandlerContext
|
|
42
|
+
|
|
43
|
+
try:
|
|
44
|
+
from protolib import Protocol, PacketFramer
|
|
45
|
+
from protolib.errors import ProtolibError
|
|
46
|
+
from protolib.io import BufferUnderrun
|
|
47
|
+
except ImportError as _exc: # pragma: no cover - exercised only when protolib is missing
|
|
48
|
+
Protocol = None # type: ignore[assignment]
|
|
49
|
+
PacketFramer = None # type: ignore[assignment]
|
|
50
|
+
ProtolibError = Exception # type: ignore[assignment,misc]
|
|
51
|
+
BufferUnderrun = Exception # type: ignore[assignment,misc]
|
|
52
|
+
_IMPORT_ERROR = _exc
|
|
53
|
+
else:
|
|
54
|
+
_IMPORT_ERROR = None
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _require_protolib() -> None:
|
|
58
|
+
if Protocol is None:
|
|
59
|
+
raise ImportError(
|
|
60
|
+
"ProtolibCodec requires the 'protolib' package. Install it with "
|
|
61
|
+
"'pip install protolib', or install pynetty with its declared "
|
|
62
|
+
"dependency (pip install pynetty) which pulls it in automatically."
|
|
63
|
+
) from _IMPORT_ERROR
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class ProtolibCodec(ChannelInboundHandler, ChannelOutboundHandler):
|
|
67
|
+
"""
|
|
68
|
+
Combined inbound+outbound handler: decodes incoming bytes into
|
|
69
|
+
`{name, params}` dicts using a protolib `Protocol`, and encodes
|
|
70
|
+
outgoing `(name, params)` tuples (or `{"name": ..., "params": ...}`
|
|
71
|
+
dicts) back into bytes.
|
|
72
|
+
|
|
73
|
+
protocol: a protolib `Protocol` instance, OR anything accepted by
|
|
74
|
+
`Protocol(...)` directly (a path to .yml/.json, an in-memory
|
|
75
|
+
string, or an already-parsed dict) — in which case this codec
|
|
76
|
+
builds the `Protocol` itself.
|
|
77
|
+
state / direction_in / direction_out: the state/direction keys
|
|
78
|
+
used to call `parse_packet`/`serialize_packet`. For stateless
|
|
79
|
+
protocols (e.g. Minecraft Bedrock), pass state=None — see
|
|
80
|
+
protolib's own docs on `packet_type_name`.
|
|
81
|
+
framed: True to have this codec do its own varint-length framing
|
|
82
|
+
via protolib's PacketFramer (raw socket bytes in); False to
|
|
83
|
+
assume a complete frame is already being delivered by an
|
|
84
|
+
earlier pipeline handler (e.g. LengthFieldBasedFrameDecoder).
|
|
85
|
+
"""
|
|
86
|
+
|
|
87
|
+
def __init__(
|
|
88
|
+
self,
|
|
89
|
+
protocol: Any,
|
|
90
|
+
state: Optional[str] = "play",
|
|
91
|
+
direction_in: str = "toServer",
|
|
92
|
+
direction_out: str = "toClient",
|
|
93
|
+
framed: bool = True,
|
|
94
|
+
) -> None:
|
|
95
|
+
_require_protolib()
|
|
96
|
+
self._protocol = protocol if isinstance(protocol, Protocol) else Protocol(protocol)
|
|
97
|
+
self._state = state
|
|
98
|
+
self._direction_in = direction_in
|
|
99
|
+
self._direction_out = direction_out
|
|
100
|
+
self._framed = framed
|
|
101
|
+
self._framer = PacketFramer() if framed else None
|
|
102
|
+
|
|
103
|
+
# ------------------------------------------------------------------
|
|
104
|
+
# Inbound: bytes -> {"name": ..., "params": ...}
|
|
105
|
+
# ------------------------------------------------------------------
|
|
106
|
+
async def channel_read(self, ctx: ChannelHandlerContext, msg: Any) -> None:
|
|
107
|
+
data = msg if isinstance(msg, (bytes, bytearray)) else bytes(msg)
|
|
108
|
+
|
|
109
|
+
try:
|
|
110
|
+
if self._framed:
|
|
111
|
+
frames = self._framer.feed(data)
|
|
112
|
+
else:
|
|
113
|
+
frames = [bytes(data)]
|
|
114
|
+
|
|
115
|
+
for frame in frames:
|
|
116
|
+
parsed = self._protocol.parse_packet(self._state, self._direction_in, frame)
|
|
117
|
+
await ctx.fire_channel_read({"name": parsed.name, "params": parsed.params})
|
|
118
|
+
except BufferUnderrun:
|
|
119
|
+
# Only reachable in framed=False mode: an earlier handler
|
|
120
|
+
# handed us an incomplete frame. In framed=True mode
|
|
121
|
+
# PacketFramer already guarantees complete frames.
|
|
122
|
+
raise DecoderException(
|
|
123
|
+
"Incomplete packet: an earlier handler must deliver a "
|
|
124
|
+
"complete frame when framed=False"
|
|
125
|
+
) from None
|
|
126
|
+
except ProtolibError as exc:
|
|
127
|
+
raise DecoderException("Failed to decode packet") from exc
|
|
128
|
+
|
|
129
|
+
# ------------------------------------------------------------------
|
|
130
|
+
# Outbound: (name, params) or {"name": ..., "params": ...} -> bytes
|
|
131
|
+
# ------------------------------------------------------------------
|
|
132
|
+
async def write(self, ctx: ChannelHandlerContext, msg: Any) -> None:
|
|
133
|
+
name, params = self._unpack_outbound(msg)
|
|
134
|
+
|
|
135
|
+
try:
|
|
136
|
+
raw = self._protocol.serialize_packet(self._state, self._direction_out, name, params)
|
|
137
|
+
payload = PacketFramer.wrap(raw) if self._framed else raw
|
|
138
|
+
except ProtolibError as exc:
|
|
139
|
+
raise EncoderException("Failed to encode packet") from exc
|
|
140
|
+
|
|
141
|
+
await ctx.write(payload)
|
|
142
|
+
|
|
143
|
+
@staticmethod
|
|
144
|
+
def _unpack_outbound(msg: Any) -> Tuple[str, dict]:
|
|
145
|
+
if isinstance(msg, tuple) and len(msg) == 2:
|
|
146
|
+
return msg[0], msg[1]
|
|
147
|
+
if isinstance(msg, dict) and "name" in msg and "params" in msg:
|
|
148
|
+
return msg["name"], msg["params"]
|
|
149
|
+
raise EncoderException(
|
|
150
|
+
f"ProtolibCodec expects a (name, params) tuple or a "
|
|
151
|
+
f"{{'name': ..., 'params': ...}} dict, got {type(msg).__name__}"
|
|
152
|
+
)
|
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
"""
|
|
2
|
+
SslContext: equivalent to io.netty.handler.ssl.SslContextBuilder.
|
|
3
|
+
|
|
4
|
+
asyncio already supports TLS natively via ssl.SSLContext passed to
|
|
5
|
+
create_server()/create_connection(). Here we expose a builder with the
|
|
6
|
+
same feel as Netty instead of making the user deal with the stdlib
|
|
7
|
+
`ssl` module directly.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import ssl
|
|
13
|
+
from typing import Optional
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class SslContextBuilder:
|
|
17
|
+
"""Simplified equivalent to io.netty.handler.ssl.SslContextBuilder."""
|
|
18
|
+
|
|
19
|
+
def __init__(self) -> None:
|
|
20
|
+
self._cert_file: Optional[str] = None
|
|
21
|
+
self._key_file: Optional[str] = None
|
|
22
|
+
self._ca_file: Optional[str] = None
|
|
23
|
+
self._for_client = False
|
|
24
|
+
self._verify_peer = True
|
|
25
|
+
|
|
26
|
+
@classmethod
|
|
27
|
+
def for_server(cls, cert_file: str, key_file: str) -> "SslContextBuilder":
|
|
28
|
+
builder = cls()
|
|
29
|
+
builder._cert_file = cert_file
|
|
30
|
+
builder._key_file = key_file
|
|
31
|
+
builder._for_client = False
|
|
32
|
+
return builder
|
|
33
|
+
|
|
34
|
+
@classmethod
|
|
35
|
+
def for_client(cls) -> "SslContextBuilder":
|
|
36
|
+
builder = cls()
|
|
37
|
+
builder._for_client = True
|
|
38
|
+
return builder
|
|
39
|
+
|
|
40
|
+
def trust_manager(self, ca_file: str) -> "SslContextBuilder":
|
|
41
|
+
self._ca_file = ca_file
|
|
42
|
+
return self
|
|
43
|
+
|
|
44
|
+
def insecure_skip_verify(self) -> "SslContextBuilder":
|
|
45
|
+
"""Equivalent to Netty's InsecureTrustManagerFactory — for development/testing ONLY."""
|
|
46
|
+
self._verify_peer = False
|
|
47
|
+
return self
|
|
48
|
+
|
|
49
|
+
def build(self) -> ssl.SSLContext:
|
|
50
|
+
if self._for_client:
|
|
51
|
+
ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
|
|
52
|
+
if self._ca_file:
|
|
53
|
+
ctx.load_verify_locations(cafile=self._ca_file)
|
|
54
|
+
if not self._verify_peer:
|
|
55
|
+
ctx.check_hostname = False
|
|
56
|
+
ctx.verify_mode = ssl.CERT_NONE
|
|
57
|
+
else:
|
|
58
|
+
ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
|
59
|
+
if not self._cert_file or not self._key_file:
|
|
60
|
+
raise ValueError("TLS server requires cert_file and key_file (use for_server(...))")
|
|
61
|
+
ctx.load_cert_chain(certfile=self._cert_file, keyfile=self._key_file)
|
|
62
|
+
return ctx
|
|
@@ -0,0 +1,159 @@
|
|
|
1
|
+
"""
|
|
2
|
+
IdleStateHandler / ReadTimeoutHandler / WriteTimeoutHandler: equivalent
|
|
3
|
+
to io.netty.handler.timeout.*
|
|
4
|
+
|
|
5
|
+
- IdleStateHandler: fires an event once the channel has gone N seconds
|
|
6
|
+
without read activity, write activity, or both. It does not close
|
|
7
|
+
the connection by itself — typically combined with a handler that
|
|
8
|
+
reacts to the event (e.g. sending a ping, or closing if there's no
|
|
9
|
+
response).
|
|
10
|
+
- ReadTimeoutHandler: closes the channel if too much time passes
|
|
11
|
+
without receiving data.
|
|
12
|
+
- WriteTimeoutHandler: closes the channel if a write takes too long to
|
|
13
|
+
complete (approximated here, since asyncio doesn't easily expose
|
|
14
|
+
"pending write" at the socket level).
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from __future__ import annotations
|
|
18
|
+
|
|
19
|
+
import asyncio
|
|
20
|
+
import time
|
|
21
|
+
from enum import Enum, auto
|
|
22
|
+
from typing import Optional
|
|
23
|
+
|
|
24
|
+
from pynetty.exceptions import ReadTimeoutError, WriteTimeoutError
|
|
25
|
+
from pynetty.handler.channel_handler import ChannelInboundHandler, ChannelOutboundHandler
|
|
26
|
+
from pynetty.handler.channel_handler_context import ChannelHandlerContext
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class IdleState(Enum):
|
|
30
|
+
READER_IDLE = auto()
|
|
31
|
+
WRITER_IDLE = auto()
|
|
32
|
+
ALL_IDLE = auto()
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class IdleStateEvent:
|
|
36
|
+
def __init__(self, state: IdleState) -> None:
|
|
37
|
+
self.state = state
|
|
38
|
+
|
|
39
|
+
def __repr__(self) -> str:
|
|
40
|
+
return f"IdleStateEvent({self.state.name})"
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
class IdleStateHandler(ChannelInboundHandler):
|
|
44
|
+
"""
|
|
45
|
+
Fires a user_event_triggered-equivalent instead of overloading
|
|
46
|
+
channel_read: since pynetty doesn't (yet) have a separate channel
|
|
47
|
+
for user events, this handler calls a user-provided `on_idle`
|
|
48
|
+
callback directly.
|
|
49
|
+
"""
|
|
50
|
+
|
|
51
|
+
def __init__(
|
|
52
|
+
self,
|
|
53
|
+
reader_idle_seconds: float = 0,
|
|
54
|
+
writer_idle_seconds: float = 0,
|
|
55
|
+
all_idle_seconds: float = 0,
|
|
56
|
+
) -> None:
|
|
57
|
+
self._reader_idle = reader_idle_seconds
|
|
58
|
+
self._writer_idle = writer_idle_seconds
|
|
59
|
+
self._all_idle = all_idle_seconds
|
|
60
|
+
self._last_read = time.monotonic()
|
|
61
|
+
self._last_write = time.monotonic()
|
|
62
|
+
self._task: Optional[asyncio.Task] = None
|
|
63
|
+
self.on_idle = None # optional async callback: async def on_idle(ctx, IdleStateEvent)
|
|
64
|
+
|
|
65
|
+
async def channel_active(self, ctx: ChannelHandlerContext) -> None:
|
|
66
|
+
self._last_read = time.monotonic()
|
|
67
|
+
self._last_write = time.monotonic()
|
|
68
|
+
self._task = asyncio.get_event_loop().create_task(self._watch(ctx))
|
|
69
|
+
await ctx.fire_channel_active()
|
|
70
|
+
|
|
71
|
+
async def channel_inactive(self, ctx: ChannelHandlerContext) -> None:
|
|
72
|
+
if self._task is not None:
|
|
73
|
+
self._task.cancel()
|
|
74
|
+
await ctx.fire_channel_inactive()
|
|
75
|
+
|
|
76
|
+
async def channel_read(self, ctx: ChannelHandlerContext, msg) -> None:
|
|
77
|
+
self._last_read = time.monotonic()
|
|
78
|
+
await ctx.fire_channel_read(msg)
|
|
79
|
+
|
|
80
|
+
def notify_write(self) -> None:
|
|
81
|
+
"""Call manually after each write (there's no automatic outbound hook here)."""
|
|
82
|
+
self._last_write = time.monotonic()
|
|
83
|
+
|
|
84
|
+
async def _watch(self, ctx: ChannelHandlerContext) -> None:
|
|
85
|
+
try:
|
|
86
|
+
while True:
|
|
87
|
+
await asyncio.sleep(0.5)
|
|
88
|
+
now = time.monotonic()
|
|
89
|
+
if self._reader_idle and (now - self._last_read) >= self._reader_idle:
|
|
90
|
+
await self._fire(ctx, IdleState.READER_IDLE)
|
|
91
|
+
self._last_read = now
|
|
92
|
+
if self._writer_idle and (now - self._last_write) >= self._writer_idle:
|
|
93
|
+
await self._fire(ctx, IdleState.WRITER_IDLE)
|
|
94
|
+
self._last_write = now
|
|
95
|
+
if self._all_idle and (now - min(self._last_read, self._last_write)) >= self._all_idle:
|
|
96
|
+
await self._fire(ctx, IdleState.ALL_IDLE)
|
|
97
|
+
except asyncio.CancelledError:
|
|
98
|
+
pass
|
|
99
|
+
|
|
100
|
+
async def _fire(self, ctx: ChannelHandlerContext, state: IdleState) -> None:
|
|
101
|
+
if self.on_idle is not None:
|
|
102
|
+
await self.on_idle(ctx, IdleStateEvent(state))
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
class ReadTimeoutHandler(ChannelInboundHandler):
|
|
106
|
+
"""Closes the channel if `timeout_seconds` pass without receiving any data."""
|
|
107
|
+
|
|
108
|
+
def __init__(self, timeout_seconds: float) -> None:
|
|
109
|
+
self._timeout = timeout_seconds
|
|
110
|
+
self._task: Optional[asyncio.Task] = None
|
|
111
|
+
self._last_read = time.monotonic()
|
|
112
|
+
|
|
113
|
+
async def channel_active(self, ctx: ChannelHandlerContext) -> None:
|
|
114
|
+
self._last_read = time.monotonic()
|
|
115
|
+
self._task = asyncio.get_event_loop().create_task(self._watch(ctx))
|
|
116
|
+
await ctx.fire_channel_active()
|
|
117
|
+
|
|
118
|
+
async def channel_inactive(self, ctx: ChannelHandlerContext) -> None:
|
|
119
|
+
if self._task is not None:
|
|
120
|
+
self._task.cancel()
|
|
121
|
+
await ctx.fire_channel_inactive()
|
|
122
|
+
|
|
123
|
+
async def channel_read(self, ctx: ChannelHandlerContext, msg) -> None:
|
|
124
|
+
self._last_read = time.monotonic()
|
|
125
|
+
await ctx.fire_channel_read(msg)
|
|
126
|
+
|
|
127
|
+
async def _watch(self, ctx: ChannelHandlerContext) -> None:
|
|
128
|
+
try:
|
|
129
|
+
while True:
|
|
130
|
+
await asyncio.sleep(0.25)
|
|
131
|
+
if (time.monotonic() - self._last_read) >= self._timeout:
|
|
132
|
+
await ctx.fire_exception_caught(
|
|
133
|
+
ReadTimeoutError(f"No data received in {self._timeout}s")
|
|
134
|
+
)
|
|
135
|
+
await ctx.close()
|
|
136
|
+
return
|
|
137
|
+
except asyncio.CancelledError:
|
|
138
|
+
pass
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
class WriteTimeoutHandler(ChannelOutboundHandler):
|
|
142
|
+
"""
|
|
143
|
+
Closes the channel if a single write doesn't complete within the
|
|
144
|
+
time limit. Since asyncio.Transport.write() isn't awaitable (it
|
|
145
|
+
doesn't block, it queues internally), this watches that the
|
|
146
|
+
transport's drain() doesn't take too long when the outbound buffer
|
|
147
|
+
is full.
|
|
148
|
+
"""
|
|
149
|
+
|
|
150
|
+
def __init__(self, timeout_seconds: float) -> None:
|
|
151
|
+
self._timeout = timeout_seconds
|
|
152
|
+
|
|
153
|
+
async def write(self, ctx: ChannelHandlerContext, msg) -> None:
|
|
154
|
+
try:
|
|
155
|
+
await asyncio.wait_for(ctx.write(msg), timeout=self._timeout)
|
|
156
|
+
except asyncio.TimeoutError:
|
|
157
|
+
exc = WriteTimeoutError(f"Write did not complete in {self._timeout}s")
|
|
158
|
+
await ctx.fire_exception_caught(exc)
|
|
159
|
+
await ctx.close()
|