@neoline/hostbridge 2.0.4 → 2.0.5
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.
- package/README.md +11 -9
- package/docs/ARCHITECTURE.md +62 -0
- package/package.json +2 -1
- package/pyproject.toml +1 -1
- package/src/server_control_mcp/__init__.py +4 -3
- package/src/server_control_mcp/async_lifecycle.py +41 -0
- package/src/server_control_mcp/cli.py +1 -1
- package/src/server_control_mcp/client.py +302 -287
- package/src/server_control_mcp/config.py +2 -1
- package/src/server_control_mcp/daemon.py +233 -526
- package/src/server_control_mcp/doctor.py +34 -2
- package/src/server_control_mcp/hosts.py +136 -2
- package/src/server_control_mcp/mock_ssh.py +52 -14
- package/src/server_control_mcp/mock_ssh_exec.py +147 -58
- package/src/server_control_mcp/mock_ssh_session.py +3 -2
- package/src/server_control_mcp/mock_ssh_sftp.py +126 -28
- package/src/server_control_mcp/mock_ssh_tunnel.py +141 -444
- package/src/server_control_mcp/mux_connection.py +326 -0
- package/src/server_control_mcp/mux_daemon.py +186 -0
- package/src/server_control_mcp/mux_protocol.py +279 -0
- package/src/server_control_mcp/mux_records.py +71 -0
- package/src/server_control_mcp/mux_rpc.py +1779 -0
- package/src/server_control_mcp/mux_service.py +506 -0
- package/src/server_control_mcp/mux_stream.py +283 -0
- package/src/server_control_mcp/mux_sync_client.py +224 -0
- package/src/server_control_mcp/remote_agent_bundle.py +56 -0
- package/src/server_control_mcp/remote_mux_agent.py +769 -0
- package/src/server_control_mcp/runtime_services.py +862 -0
- package/src/server_control_mcp/server.py +33 -18
- package/src/server_control_mcp/services.py +2 -666
- package/src/server_control_mcp/transports/__init__.py +4 -8
- package/src/server_control_mcp/transports/base.py +60 -38
- package/src/server_control_mcp/transports/ssh.py +206 -259
- package/src/server_control_mcp/tunnel_manager.py +1245 -0
- package/src/server_control_mcp/tunnel_native_ssh.py +454 -0
- package/src/server_control_mcp/tunnel_providers.py +20 -0
- package/src/server_control_mcp/tunnel_pty_agent.py +909 -0
- package/src/server_control_mcp/protocol.py +0 -363
- package/src/server_control_mcp/remote_tunnel_agent.py +0 -143
- package/src/server_control_mcp/transports/pty.py +0 -515
|
@@ -0,0 +1,1779 @@
|
|
|
1
|
+
"""Persistent, concurrent RPC over the HostBridge multiplexed wire protocol."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import math
|
|
7
|
+
import time
|
|
8
|
+
from collections import deque
|
|
9
|
+
from collections.abc import Awaitable, Callable, Mapping
|
|
10
|
+
from contextlib import suppress
|
|
11
|
+
from dataclasses import dataclass
|
|
12
|
+
from pathlib import Path
|
|
13
|
+
|
|
14
|
+
from .mux_connection import ByteCreditWindow, ConnectionLifecycle, ConnectionState, MuxStateError
|
|
15
|
+
from .mux_protocol import (
|
|
16
|
+
MAX_FRAME_PAYLOAD_BYTES,
|
|
17
|
+
FrameType,
|
|
18
|
+
MuxCodec,
|
|
19
|
+
MuxFrame,
|
|
20
|
+
MuxProtocolError,
|
|
21
|
+
SequenceTracker,
|
|
22
|
+
StreamIdAllocator,
|
|
23
|
+
StreamParity,
|
|
24
|
+
decode_control_payload,
|
|
25
|
+
decode_window_update,
|
|
26
|
+
encode_control_payload,
|
|
27
|
+
encode_frame,
|
|
28
|
+
encode_window_update,
|
|
29
|
+
)
|
|
30
|
+
from .mux_stream import AsyncByteCreditWindow, MuxRpcStream
|
|
31
|
+
|
|
32
|
+
MuxRpcHandler = Callable[[str, dict[str, object]], Awaitable[dict[str, object]]]
|
|
33
|
+
MuxRpcStreamHandler = Callable[[str, dict[str, object], MuxRpcStream], Awaitable[dict[str, object]]]
|
|
34
|
+
MuxRpcRequestClosedHandler = Callable[[str], Awaitable[None]]
|
|
35
|
+
MuxFrameWriteFailureHandler = Callable[[BaseException], Awaitable[None]]
|
|
36
|
+
|
|
37
|
+
DEFAULT_STREAM_WINDOW_BYTES = 1024 * 1024
|
|
38
|
+
DEFAULT_CONNECTION_WINDOW_BYTES = 32 * 1024 * 1024
|
|
39
|
+
DEFAULT_DATA_CHUNK_BYTES = 64 * 1024
|
|
40
|
+
DEFAULT_TOMBSTONE_RETENTION_SECONDS = 5.0
|
|
41
|
+
DEFAULT_MAX_TOMBSTONES = 65_536
|
|
42
|
+
DEFAULT_MAX_STREAM_QUEUE_BYTES = 2 * 1024 * 1024
|
|
43
|
+
DEFAULT_MAX_DATA_QUEUE_BYTES = 32 * 1024 * 1024
|
|
44
|
+
DEFAULT_MAX_CONTROL_QUEUE_BYTES = 64 * 1024
|
|
45
|
+
DEFAULT_HEARTBEAT_INTERVAL_SECONDS = 30.0
|
|
46
|
+
DEFAULT_HEARTBEAT_TIMEOUT_SECONDS = 10.0
|
|
47
|
+
DEFAULT_HANDSHAKE_TIMEOUT_SECONDS = 10.0
|
|
48
|
+
DEFAULT_CLEANUP_TIMEOUT_SECONDS = 5.0
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class MuxRpcError(RuntimeError):
|
|
52
|
+
pass
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class MuxCapacityError(MuxRpcError):
|
|
56
|
+
pass
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
class MuxConnectionClosed(MuxRpcError):
|
|
60
|
+
pass
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
class MuxRemoteError(MuxRpcError):
|
|
64
|
+
def __init__(
|
|
65
|
+
self,
|
|
66
|
+
code: str,
|
|
67
|
+
message: str,
|
|
68
|
+
*,
|
|
69
|
+
retryable: bool = False,
|
|
70
|
+
details: dict[str, object] | None = None,
|
|
71
|
+
) -> None:
|
|
72
|
+
super().__init__(message)
|
|
73
|
+
self.code = code
|
|
74
|
+
self.retryable = retryable
|
|
75
|
+
self.details = details or {}
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
class MuxRpcFailure(MuxRemoteError):
|
|
79
|
+
"""Structured failure raised by a server-side RPC handler."""
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
@dataclass(frozen=True, slots=True)
|
|
83
|
+
class _TerminalStream:
|
|
84
|
+
sequences: SequenceTracker
|
|
85
|
+
expires_at: float
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
class _TerminalStreams:
|
|
89
|
+
def __init__(
|
|
90
|
+
self,
|
|
91
|
+
*,
|
|
92
|
+
retention_seconds: float = DEFAULT_TOMBSTONE_RETENTION_SECONDS,
|
|
93
|
+
max_entries: int = DEFAULT_MAX_TOMBSTONES,
|
|
94
|
+
clock: Callable[[], float] = time.monotonic,
|
|
95
|
+
) -> None:
|
|
96
|
+
if (
|
|
97
|
+
not isinstance(retention_seconds, (int, float))
|
|
98
|
+
or isinstance(retention_seconds, bool)
|
|
99
|
+
or not math.isfinite(retention_seconds)
|
|
100
|
+
or retention_seconds <= 0
|
|
101
|
+
):
|
|
102
|
+
raise ValueError("retention_seconds must be a finite positive number")
|
|
103
|
+
_positive(max_entries, "max_entries")
|
|
104
|
+
self._retention_seconds = float(retention_seconds)
|
|
105
|
+
self._max_entries = max_entries
|
|
106
|
+
self._clock = clock
|
|
107
|
+
self._entries: dict[int, _TerminalStream] = {}
|
|
108
|
+
|
|
109
|
+
def add(self, stream_id: int, sequences: SequenceTracker) -> None:
|
|
110
|
+
self._expire()
|
|
111
|
+
if stream_id not in self._entries and len(self._entries) >= self._max_entries:
|
|
112
|
+
raise MuxCapacityError("terminal stream capacity is exhausted")
|
|
113
|
+
self._entries[stream_id] = _TerminalStream(
|
|
114
|
+
sequences,
|
|
115
|
+
self._clock() + self._retention_seconds,
|
|
116
|
+
)
|
|
117
|
+
|
|
118
|
+
def get(self, stream_id: int) -> _TerminalStream | None:
|
|
119
|
+
self._expire()
|
|
120
|
+
return self._entries.get(stream_id)
|
|
121
|
+
|
|
122
|
+
def clear(self) -> None:
|
|
123
|
+
self._entries.clear()
|
|
124
|
+
|
|
125
|
+
def _expire(self) -> None:
|
|
126
|
+
now = self._clock()
|
|
127
|
+
expired = [stream_id for stream_id, entry in self._entries.items() if entry.expires_at <= now]
|
|
128
|
+
for stream_id in expired:
|
|
129
|
+
del self._entries[stream_id]
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
class _FrameReader:
|
|
133
|
+
def __init__(self, reader: asyncio.StreamReader) -> None:
|
|
134
|
+
self._reader = reader
|
|
135
|
+
self._codec = MuxCodec()
|
|
136
|
+
self._ready: deque[MuxFrame] = deque()
|
|
137
|
+
|
|
138
|
+
async def read(self) -> MuxFrame:
|
|
139
|
+
while not self._ready:
|
|
140
|
+
data = await self._reader.read(64 * 1024)
|
|
141
|
+
if not data:
|
|
142
|
+
if self._codec.buffered_bytes:
|
|
143
|
+
raise MuxConnectionClosed("multiplexed connection ended with an incomplete frame")
|
|
144
|
+
raise MuxConnectionClosed("multiplexed connection closed")
|
|
145
|
+
self._ready.extend(self._codec.feed(data))
|
|
146
|
+
return self._ready.popleft()
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
@dataclass(slots=True)
|
|
150
|
+
class _QueuedFrame:
|
|
151
|
+
frame_type: FrameType
|
|
152
|
+
stream_id: int
|
|
153
|
+
sequences: SequenceTracker
|
|
154
|
+
flags: int
|
|
155
|
+
payload: bytes
|
|
156
|
+
completed: asyncio.Future[None]
|
|
157
|
+
queued_bytes: int
|
|
158
|
+
is_data: bool
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
class _FrameWriter:
|
|
162
|
+
def __init__(
|
|
163
|
+
self,
|
|
164
|
+
writer: asyncio.StreamWriter,
|
|
165
|
+
*,
|
|
166
|
+
max_stream_queue_bytes: int = DEFAULT_MAX_STREAM_QUEUE_BYTES,
|
|
167
|
+
max_data_queue_bytes: int = DEFAULT_MAX_DATA_QUEUE_BYTES,
|
|
168
|
+
max_control_queue_bytes: int = DEFAULT_MAX_CONTROL_QUEUE_BYTES,
|
|
169
|
+
failure_handler: MuxFrameWriteFailureHandler | None = None,
|
|
170
|
+
) -> None:
|
|
171
|
+
_positive(max_stream_queue_bytes, "max_stream_queue_bytes")
|
|
172
|
+
_positive(max_data_queue_bytes, "max_data_queue_bytes")
|
|
173
|
+
_positive(max_control_queue_bytes, "max_control_queue_bytes")
|
|
174
|
+
self._writer = writer
|
|
175
|
+
self._lock = asyncio.Lock()
|
|
176
|
+
self._max_stream_queue_bytes = max_stream_queue_bytes
|
|
177
|
+
self._max_data_queue_bytes = max_data_queue_bytes
|
|
178
|
+
self._max_control_queue_bytes = max_control_queue_bytes
|
|
179
|
+
self._stream_queue_bytes: dict[int, int] = {}
|
|
180
|
+
self._data_queue_bytes = 0
|
|
181
|
+
self._control_queue_bytes = 0
|
|
182
|
+
self._control_queue: deque[_QueuedFrame] = deque()
|
|
183
|
+
self._stream_queues: dict[int, deque[_QueuedFrame]] = {}
|
|
184
|
+
self._ready_streams: deque[int] = deque()
|
|
185
|
+
self._writer_task: asyncio.Task[None] | None = None
|
|
186
|
+
self._closed_error: BaseException | None = None
|
|
187
|
+
self._failure_handler = failure_handler
|
|
188
|
+
|
|
189
|
+
async def write_sequenced(
|
|
190
|
+
self,
|
|
191
|
+
frame_type: FrameType,
|
|
192
|
+
stream_id: int,
|
|
193
|
+
sequences: SequenceTracker,
|
|
194
|
+
*,
|
|
195
|
+
flags: int = 0,
|
|
196
|
+
payload: bytes = b"",
|
|
197
|
+
) -> None:
|
|
198
|
+
completed: asyncio.Future[None] = asyncio.get_running_loop().create_future()
|
|
199
|
+
queued_bytes = max(1, len(payload))
|
|
200
|
+
is_data = frame_type is FrameType.DATA
|
|
201
|
+
queued = _QueuedFrame(frame_type, stream_id, sequences, flags, payload, completed, queued_bytes, is_data)
|
|
202
|
+
async with self._lock:
|
|
203
|
+
if self._closed_error is not None:
|
|
204
|
+
raise self._closed_error
|
|
205
|
+
self._reserve_budget(queued)
|
|
206
|
+
if stream_id == 0:
|
|
207
|
+
self._control_queue.append(queued)
|
|
208
|
+
else:
|
|
209
|
+
stream_queue = self._stream_queues.setdefault(stream_id, deque())
|
|
210
|
+
stream_queue.append(queued)
|
|
211
|
+
if len(stream_queue) == 1:
|
|
212
|
+
self._ready_streams.append(stream_id)
|
|
213
|
+
if self._writer_task is None:
|
|
214
|
+
self._writer_task = asyncio.create_task(
|
|
215
|
+
self._run(),
|
|
216
|
+
name="hostbridge-mux-frame-writer",
|
|
217
|
+
)
|
|
218
|
+
await asyncio.shield(completed)
|
|
219
|
+
|
|
220
|
+
async def _run(self) -> None:
|
|
221
|
+
control_burst = 0
|
|
222
|
+
while True:
|
|
223
|
+
async with self._lock:
|
|
224
|
+
queued, control_burst = self._next_frame(control_burst)
|
|
225
|
+
if queued is None:
|
|
226
|
+
self._writer_task = None
|
|
227
|
+
return
|
|
228
|
+
try:
|
|
229
|
+
frame = MuxFrame(
|
|
230
|
+
queued.frame_type,
|
|
231
|
+
queued.stream_id,
|
|
232
|
+
queued.sequences.next_outbound(),
|
|
233
|
+
flags=queued.flags,
|
|
234
|
+
payload=queued.payload,
|
|
235
|
+
)
|
|
236
|
+
self._writer.write(encode_frame(frame))
|
|
237
|
+
await self._writer.drain()
|
|
238
|
+
except (ConnectionError, OSError, RuntimeError) as exc:
|
|
239
|
+
error = MuxConnectionClosed("multiplexed connection write failed")
|
|
240
|
+
error.__cause__ = exc
|
|
241
|
+
await self._fail(error, queued)
|
|
242
|
+
return
|
|
243
|
+
except BaseException as exc:
|
|
244
|
+
await self._fail(exc, queued)
|
|
245
|
+
return
|
|
246
|
+
async with self._lock:
|
|
247
|
+
self._release_budget(queued)
|
|
248
|
+
if not queued.completed.done():
|
|
249
|
+
queued.completed.set_result(None)
|
|
250
|
+
|
|
251
|
+
def _reserve_budget(self, queued: _QueuedFrame) -> None:
|
|
252
|
+
if queued.is_data:
|
|
253
|
+
stream_bytes = self._stream_queue_bytes.get(queued.stream_id, 0)
|
|
254
|
+
if stream_bytes + queued.queued_bytes > self._max_stream_queue_bytes:
|
|
255
|
+
raise MuxCapacityError(
|
|
256
|
+
f"stream {queued.stream_id} data queue limit is {self._max_stream_queue_bytes} bytes"
|
|
257
|
+
)
|
|
258
|
+
if self._data_queue_bytes + queued.queued_bytes > self._max_data_queue_bytes:
|
|
259
|
+
raise MuxCapacityError(f"connection data queue limit is {self._max_data_queue_bytes} bytes")
|
|
260
|
+
self._stream_queue_bytes[queued.stream_id] = stream_bytes + queued.queued_bytes
|
|
261
|
+
self._data_queue_bytes += queued.queued_bytes
|
|
262
|
+
return
|
|
263
|
+
if self._control_queue_bytes + queued.queued_bytes > self._max_control_queue_bytes:
|
|
264
|
+
raise MuxCapacityError(f"connection control queue limit is {self._max_control_queue_bytes} bytes")
|
|
265
|
+
self._control_queue_bytes += queued.queued_bytes
|
|
266
|
+
|
|
267
|
+
def _release_budget(self, queued: _QueuedFrame) -> None:
|
|
268
|
+
if queued.is_data:
|
|
269
|
+
stream_bytes = self._stream_queue_bytes[queued.stream_id] - queued.queued_bytes
|
|
270
|
+
if stream_bytes:
|
|
271
|
+
self._stream_queue_bytes[queued.stream_id] = stream_bytes
|
|
272
|
+
else:
|
|
273
|
+
del self._stream_queue_bytes[queued.stream_id]
|
|
274
|
+
self._data_queue_bytes -= queued.queued_bytes
|
|
275
|
+
return
|
|
276
|
+
self._control_queue_bytes -= queued.queued_bytes
|
|
277
|
+
|
|
278
|
+
def _next_frame(self, control_burst: int) -> tuple[_QueuedFrame | None, int]:
|
|
279
|
+
if self._control_queue and (control_burst < 8 or not self._ready_streams):
|
|
280
|
+
return self._control_queue.popleft(), control_burst + 1
|
|
281
|
+
if self._ready_streams:
|
|
282
|
+
stream_id = self._ready_streams.popleft()
|
|
283
|
+
stream_queue = self._stream_queues[stream_id]
|
|
284
|
+
queued = stream_queue.popleft()
|
|
285
|
+
if stream_queue:
|
|
286
|
+
self._ready_streams.append(stream_id)
|
|
287
|
+
else:
|
|
288
|
+
del self._stream_queues[stream_id]
|
|
289
|
+
return queued, 0
|
|
290
|
+
if self._control_queue:
|
|
291
|
+
return self._control_queue.popleft(), control_burst + 1
|
|
292
|
+
return None, control_burst
|
|
293
|
+
|
|
294
|
+
async def _fail(self, error: BaseException, active: _QueuedFrame) -> None:
|
|
295
|
+
async with self._lock:
|
|
296
|
+
self._closed_error = error
|
|
297
|
+
queued = [active, *self._control_queue]
|
|
298
|
+
for stream_queue in self._stream_queues.values():
|
|
299
|
+
queued.extend(stream_queue)
|
|
300
|
+
self._control_queue.clear()
|
|
301
|
+
self._stream_queues.clear()
|
|
302
|
+
self._ready_streams.clear()
|
|
303
|
+
self._stream_queue_bytes.clear()
|
|
304
|
+
self._data_queue_bytes = 0
|
|
305
|
+
self._control_queue_bytes = 0
|
|
306
|
+
self._writer_task = None
|
|
307
|
+
failure_handler = self._failure_handler
|
|
308
|
+
self._failure_handler = None
|
|
309
|
+
if failure_handler is not None:
|
|
310
|
+
try:
|
|
311
|
+
await failure_handler(error)
|
|
312
|
+
except BaseException as exc:
|
|
313
|
+
error.add_note(f"mux write failure cleanup also failed: {exc!r}")
|
|
314
|
+
for item in queued:
|
|
315
|
+
if not item.completed.done():
|
|
316
|
+
item.completed.set_exception(error)
|
|
317
|
+
item.completed.exception()
|
|
318
|
+
|
|
319
|
+
|
|
320
|
+
@dataclass(slots=True)
|
|
321
|
+
class _PendingRequest:
|
|
322
|
+
future: asyncio.Future[dict[str, object]]
|
|
323
|
+
sequences: SequenceTracker
|
|
324
|
+
result: dict[str, object] | None = None
|
|
325
|
+
|
|
326
|
+
|
|
327
|
+
@dataclass(slots=True)
|
|
328
|
+
class _ClientStreamRequest:
|
|
329
|
+
sequences: SequenceTracker
|
|
330
|
+
receive_window: int
|
|
331
|
+
opened: asyncio.Future[MuxRpcStream]
|
|
332
|
+
stream: MuxRpcStream | None = None
|
|
333
|
+
local_eof_sent: bool = False
|
|
334
|
+
|
|
335
|
+
|
|
336
|
+
class MuxRpcClient:
|
|
337
|
+
def __init__(
|
|
338
|
+
self,
|
|
339
|
+
reader: asyncio.StreamReader,
|
|
340
|
+
writer: asyncio.StreamWriter,
|
|
341
|
+
*,
|
|
342
|
+
max_pending: int,
|
|
343
|
+
connection_window: int,
|
|
344
|
+
stream_window: int,
|
|
345
|
+
heartbeat_interval: float = DEFAULT_HEARTBEAT_INTERVAL_SECONDS,
|
|
346
|
+
heartbeat_timeout: float = DEFAULT_HEARTBEAT_TIMEOUT_SECONDS,
|
|
347
|
+
handshake_timeout: float = DEFAULT_HANDSHAKE_TIMEOUT_SECONDS,
|
|
348
|
+
cleanup_timeout: float = DEFAULT_CLEANUP_TIMEOUT_SECONDS,
|
|
349
|
+
) -> None:
|
|
350
|
+
_positive(max_pending, "max_pending")
|
|
351
|
+
_positive(connection_window, "connection_window")
|
|
352
|
+
_positive(stream_window, "stream_window")
|
|
353
|
+
_finite_positive(heartbeat_interval, "heartbeat_interval")
|
|
354
|
+
_finite_positive(heartbeat_timeout, "heartbeat_timeout")
|
|
355
|
+
_finite_positive(handshake_timeout, "handshake_timeout")
|
|
356
|
+
_finite_positive(cleanup_timeout, "cleanup_timeout")
|
|
357
|
+
self._reader = _FrameReader(reader)
|
|
358
|
+
self._writer = _FrameWriter(writer, failure_handler=self._connection_failed)
|
|
359
|
+
self._stream_writer = writer
|
|
360
|
+
self._max_pending = max_pending
|
|
361
|
+
self._default_stream_window = stream_window
|
|
362
|
+
self._allocator = StreamIdAllocator(StreamParity.ODD)
|
|
363
|
+
self._pending: dict[int, _PendingRequest] = {}
|
|
364
|
+
self._streams: dict[int, _ClientStreamRequest] = {}
|
|
365
|
+
self._terminal_streams = _TerminalStreams()
|
|
366
|
+
self._reader_task: asyncio.Task[None] | None = None
|
|
367
|
+
self._lifecycle = ConnectionLifecycle()
|
|
368
|
+
self._lifecycle.start_connecting()
|
|
369
|
+
self._local_close_started = False
|
|
370
|
+
self._close_error: MuxConnectionClosed | None = None
|
|
371
|
+
self._close_lock = asyncio.Lock()
|
|
372
|
+
self._close_task: asyncio.Task[None] | None = None
|
|
373
|
+
self._connection_sequences = SequenceTracker()
|
|
374
|
+
self._connection_receive_window = ByteCreditWindow(
|
|
375
|
+
initial_credit=connection_window,
|
|
376
|
+
maximum_credit=connection_window,
|
|
377
|
+
)
|
|
378
|
+
self._connection_send_window: AsyncByteCreditWindow | None = None
|
|
379
|
+
self._window_stalls = 0
|
|
380
|
+
self._protocol_failures = 0
|
|
381
|
+
self._heartbeat_interval = float(heartbeat_interval)
|
|
382
|
+
self._heartbeat_timeout = float(heartbeat_timeout)
|
|
383
|
+
self._handshake_timeout = float(handshake_timeout)
|
|
384
|
+
self._cleanup_timeout = float(cleanup_timeout)
|
|
385
|
+
self._heartbeat_task: asyncio.Task[None] | None = None
|
|
386
|
+
self._heartbeat_waiter: tuple[bytes, asyncio.Future[None]] | None = None
|
|
387
|
+
self._heartbeat_nonce = 0
|
|
388
|
+
|
|
389
|
+
@classmethod
|
|
390
|
+
async def connect_unix(
|
|
391
|
+
cls,
|
|
392
|
+
path: Path | str,
|
|
393
|
+
*,
|
|
394
|
+
max_pending: int = 256,
|
|
395
|
+
connection_window: int = DEFAULT_CONNECTION_WINDOW_BYTES,
|
|
396
|
+
stream_window: int = DEFAULT_STREAM_WINDOW_BYTES,
|
|
397
|
+
heartbeat_interval: float = DEFAULT_HEARTBEAT_INTERVAL_SECONDS,
|
|
398
|
+
heartbeat_timeout: float = DEFAULT_HEARTBEAT_TIMEOUT_SECONDS,
|
|
399
|
+
handshake_timeout: float = DEFAULT_HANDSHAKE_TIMEOUT_SECONDS,
|
|
400
|
+
cleanup_timeout: float = DEFAULT_CLEANUP_TIMEOUT_SECONDS,
|
|
401
|
+
) -> MuxRpcClient:
|
|
402
|
+
_positive(max_pending, "max_pending")
|
|
403
|
+
_positive(connection_window, "connection_window")
|
|
404
|
+
_positive(stream_window, "stream_window")
|
|
405
|
+
_finite_positive(heartbeat_interval, "heartbeat_interval")
|
|
406
|
+
_finite_positive(heartbeat_timeout, "heartbeat_timeout")
|
|
407
|
+
_finite_positive(handshake_timeout, "handshake_timeout")
|
|
408
|
+
_finite_positive(cleanup_timeout, "cleanup_timeout")
|
|
409
|
+
try:
|
|
410
|
+
reader, writer = await asyncio.open_unix_connection(str(path))
|
|
411
|
+
except (ConnectionError, OSError) as exc:
|
|
412
|
+
raise MuxConnectionClosed(f"failed to connect to multiplexed daemon at {path}") from exc
|
|
413
|
+
client = cls(
|
|
414
|
+
reader,
|
|
415
|
+
writer,
|
|
416
|
+
max_pending=max_pending,
|
|
417
|
+
connection_window=connection_window,
|
|
418
|
+
stream_window=stream_window,
|
|
419
|
+
heartbeat_interval=heartbeat_interval,
|
|
420
|
+
heartbeat_timeout=heartbeat_timeout,
|
|
421
|
+
handshake_timeout=handshake_timeout,
|
|
422
|
+
cleanup_timeout=cleanup_timeout,
|
|
423
|
+
)
|
|
424
|
+
try:
|
|
425
|
+
await client._start()
|
|
426
|
+
except BaseException:
|
|
427
|
+
writer.close()
|
|
428
|
+
await _wait_closed_bounded(writer, cleanup_timeout)
|
|
429
|
+
raise
|
|
430
|
+
return client
|
|
431
|
+
|
|
432
|
+
@property
|
|
433
|
+
def pending_count(self) -> int:
|
|
434
|
+
return len(self._pending) + len(self._streams)
|
|
435
|
+
|
|
436
|
+
@property
|
|
437
|
+
def closed(self) -> bool:
|
|
438
|
+
return self._lifecycle.terminal
|
|
439
|
+
|
|
440
|
+
@property
|
|
441
|
+
def draining(self) -> bool:
|
|
442
|
+
return self._lifecycle.draining
|
|
443
|
+
|
|
444
|
+
@property
|
|
445
|
+
def state(self) -> ConnectionState:
|
|
446
|
+
return self._lifecycle.state
|
|
447
|
+
|
|
448
|
+
@property
|
|
449
|
+
def window_stalls(self) -> int:
|
|
450
|
+
return self._window_stalls
|
|
451
|
+
|
|
452
|
+
@property
|
|
453
|
+
def protocol_failures(self) -> int:
|
|
454
|
+
return self._protocol_failures
|
|
455
|
+
|
|
456
|
+
def _record_window_stall(self) -> None:
|
|
457
|
+
self._window_stalls += 1
|
|
458
|
+
|
|
459
|
+
async def request(self, method: str, params: Mapping[str, object]) -> dict[str, object]:
|
|
460
|
+
if self.closed:
|
|
461
|
+
raise self._close_error or MuxConnectionClosed("multiplexed client is closed")
|
|
462
|
+
if self.draining:
|
|
463
|
+
raise MuxConnectionClosed("multiplexed server connection is draining")
|
|
464
|
+
if not isinstance(method, str) or not method.strip():
|
|
465
|
+
raise ValueError("method must be a non-empty string")
|
|
466
|
+
if not isinstance(params, Mapping):
|
|
467
|
+
raise TypeError("params must be a mapping")
|
|
468
|
+
if self.pending_count >= self._max_pending:
|
|
469
|
+
raise MuxCapacityError(f"pending request limit is {self._max_pending}")
|
|
470
|
+
|
|
471
|
+
stream_id = self._allocator.allocate()
|
|
472
|
+
sequences = SequenceTracker()
|
|
473
|
+
future: asyncio.Future[dict[str, object]] = asyncio.get_running_loop().create_future()
|
|
474
|
+
pending = _PendingRequest(future, sequences)
|
|
475
|
+
self._pending[stream_id] = pending
|
|
476
|
+
control = encode_control_payload({"method": method, "params": dict(params)})
|
|
477
|
+
try:
|
|
478
|
+
await _write_sequenced_before_honouring_cancellation(
|
|
479
|
+
self._writer,
|
|
480
|
+
FrameType.OPEN,
|
|
481
|
+
stream_id,
|
|
482
|
+
sequences,
|
|
483
|
+
payload=control,
|
|
484
|
+
)
|
|
485
|
+
return await future
|
|
486
|
+
except asyncio.CancelledError:
|
|
487
|
+
if self._pending.pop(stream_id, None) is not None:
|
|
488
|
+
self._terminal_streams.add(stream_id, sequences)
|
|
489
|
+
await _complete_sequenced_write_after_cancellation(
|
|
490
|
+
self._writer,
|
|
491
|
+
FrameType.RESET,
|
|
492
|
+
stream_id,
|
|
493
|
+
sequences,
|
|
494
|
+
payload=_error_payload("cancelled", "request cancelled"),
|
|
495
|
+
)
|
|
496
|
+
raise
|
|
497
|
+
except MuxConnectionClosed:
|
|
498
|
+
self._pending.pop(stream_id, None)
|
|
499
|
+
if future.done() and not future.cancelled():
|
|
500
|
+
future.exception()
|
|
501
|
+
raise
|
|
502
|
+
|
|
503
|
+
async def open_stream(
|
|
504
|
+
self,
|
|
505
|
+
method: str,
|
|
506
|
+
params: Mapping[str, object],
|
|
507
|
+
*,
|
|
508
|
+
receive_window: int | None = None,
|
|
509
|
+
) -> MuxRpcStream:
|
|
510
|
+
if self.closed:
|
|
511
|
+
raise self._close_error or MuxConnectionClosed("multiplexed client is closed")
|
|
512
|
+
if self.draining:
|
|
513
|
+
raise MuxConnectionClosed("multiplexed server connection is draining")
|
|
514
|
+
if not isinstance(method, str) or not method.strip():
|
|
515
|
+
raise ValueError("method must be a non-empty string")
|
|
516
|
+
if not isinstance(params, Mapping):
|
|
517
|
+
raise TypeError("params must be a mapping")
|
|
518
|
+
if self.pending_count >= self._max_pending:
|
|
519
|
+
raise MuxCapacityError(f"pending request limit is {self._max_pending}")
|
|
520
|
+
window = self._default_stream_window if receive_window is None else receive_window
|
|
521
|
+
_positive(window, "receive_window")
|
|
522
|
+
|
|
523
|
+
stream_id = self._allocator.allocate()
|
|
524
|
+
sequences = SequenceTracker()
|
|
525
|
+
opened: asyncio.Future[MuxRpcStream] = asyncio.get_running_loop().create_future()
|
|
526
|
+
request = _ClientStreamRequest(sequences, window, opened)
|
|
527
|
+
self._streams[stream_id] = request
|
|
528
|
+
control = encode_control_payload(
|
|
529
|
+
{
|
|
530
|
+
"mode": "stream",
|
|
531
|
+
"method": method,
|
|
532
|
+
"params": dict(params),
|
|
533
|
+
"receive_window": window,
|
|
534
|
+
}
|
|
535
|
+
)
|
|
536
|
+
try:
|
|
537
|
+
await _write_sequenced_before_honouring_cancellation(
|
|
538
|
+
self._writer,
|
|
539
|
+
FrameType.OPEN,
|
|
540
|
+
stream_id,
|
|
541
|
+
sequences,
|
|
542
|
+
payload=control,
|
|
543
|
+
)
|
|
544
|
+
return await opened
|
|
545
|
+
except asyncio.CancelledError:
|
|
546
|
+
if self._streams.pop(stream_id, None) is not None:
|
|
547
|
+
self._terminal_streams.add(stream_id, sequences)
|
|
548
|
+
await _complete_sequenced_write_after_cancellation(
|
|
549
|
+
self._writer,
|
|
550
|
+
FrameType.RESET,
|
|
551
|
+
stream_id,
|
|
552
|
+
sequences,
|
|
553
|
+
payload=_error_payload("cancelled", "stream open cancelled"),
|
|
554
|
+
)
|
|
555
|
+
raise
|
|
556
|
+
except MuxConnectionClosed:
|
|
557
|
+
self._streams.pop(stream_id, None)
|
|
558
|
+
if opened.done() and not opened.cancelled():
|
|
559
|
+
opened.exception()
|
|
560
|
+
raise
|
|
561
|
+
|
|
562
|
+
async def _stream_send_data(self, stream: MuxRpcStream, data: bytes) -> None:
|
|
563
|
+
await self._send_stream_data(stream, data)
|
|
564
|
+
|
|
565
|
+
async def _stream_accept(self, stream: MuxRpcStream) -> None:
|
|
566
|
+
raise RuntimeError("client-owned streams are already accepted")
|
|
567
|
+
|
|
568
|
+
async def _stream_send_eof(self, stream: MuxRpcStream) -> None:
|
|
569
|
+
await _write_sequenced_before_honouring_cancellation(
|
|
570
|
+
self._writer,
|
|
571
|
+
FrameType.EOF,
|
|
572
|
+
stream.stream_id,
|
|
573
|
+
stream._sequences,
|
|
574
|
+
)
|
|
575
|
+
request = self._streams.get(stream.stream_id)
|
|
576
|
+
if request is not None and request.stream is stream:
|
|
577
|
+
request.local_eof_sent = True
|
|
578
|
+
await self._maybe_close_stream(stream.stream_id)
|
|
579
|
+
|
|
580
|
+
async def _stream_release_credit(self, stream: MuxRpcStream, amount: int) -> None:
|
|
581
|
+
await stream._release_receive_credit(amount)
|
|
582
|
+
self._connection_receive_window.grant(amount)
|
|
583
|
+
await self._writer.write_sequenced(
|
|
584
|
+
FrameType.WINDOW_UPDATE,
|
|
585
|
+
stream.stream_id,
|
|
586
|
+
stream._sequences,
|
|
587
|
+
payload=encode_window_update(amount),
|
|
588
|
+
)
|
|
589
|
+
await self._writer.write_sequenced(
|
|
590
|
+
FrameType.WINDOW_UPDATE,
|
|
591
|
+
0,
|
|
592
|
+
self._connection_sequences,
|
|
593
|
+
payload=encode_window_update(amount),
|
|
594
|
+
)
|
|
595
|
+
await self._maybe_close_stream(stream.stream_id)
|
|
596
|
+
|
|
597
|
+
async def _stream_reset(self, stream: MuxRpcStream, message: str) -> None:
|
|
598
|
+
request = self._streams.pop(stream.stream_id, None)
|
|
599
|
+
if request is None:
|
|
600
|
+
return
|
|
601
|
+
self._terminal_streams.add(stream.stream_id, request.sequences)
|
|
602
|
+
error = MuxConnectionClosed(f"stream reset: {message}")
|
|
603
|
+
await stream._fail(error)
|
|
604
|
+
await self._writer.write_sequenced(
|
|
605
|
+
FrameType.RESET,
|
|
606
|
+
stream.stream_id,
|
|
607
|
+
request.sequences,
|
|
608
|
+
payload=_error_payload("cancelled", message),
|
|
609
|
+
)
|
|
610
|
+
|
|
611
|
+
async def close(self) -> None:
|
|
612
|
+
task = self._close_task
|
|
613
|
+
if task is None:
|
|
614
|
+
task = asyncio.create_task(
|
|
615
|
+
self._close_owned(),
|
|
616
|
+
name="hostbridge-mux-rpc-client-close",
|
|
617
|
+
)
|
|
618
|
+
task.add_done_callback(_consume_task_result)
|
|
619
|
+
self._close_task = task
|
|
620
|
+
await asyncio.shield(task)
|
|
621
|
+
|
|
622
|
+
async def _close_owned(self) -> None:
|
|
623
|
+
async with self._close_lock:
|
|
624
|
+
if self.closed:
|
|
625
|
+
await self._stop_heartbeat()
|
|
626
|
+
if self._reader_task is not None and self._reader_task is not asyncio.current_task():
|
|
627
|
+
with suppress(MuxConnectionClosed, asyncio.CancelledError):
|
|
628
|
+
await self._reader_task
|
|
629
|
+
return
|
|
630
|
+
self._local_close_started = True
|
|
631
|
+
if self.state is ConnectionState.READY:
|
|
632
|
+
self._lifecycle.start_draining()
|
|
633
|
+
elif self.state in {ConnectionState.CONNECTING, ConnectionState.HANDSHAKING}:
|
|
634
|
+
self._lifecycle.mark_closed()
|
|
635
|
+
await self._stop_heartbeat()
|
|
636
|
+
self._close_error = MuxConnectionClosed("multiplexed client closed")
|
|
637
|
+
pending_items = list(self._pending.items())
|
|
638
|
+
self._pending.clear()
|
|
639
|
+
stream_items = list(self._streams.items())
|
|
640
|
+
self._streams.clear()
|
|
641
|
+
for _stream_id, pending in pending_items:
|
|
642
|
+
if not pending.future.done():
|
|
643
|
+
pending.future.set_exception(MuxConnectionClosed("multiplexed client closed"))
|
|
644
|
+
for _stream_id, request in stream_items:
|
|
645
|
+
error = MuxConnectionClosed("multiplexed client closed")
|
|
646
|
+
if request.stream is not None:
|
|
647
|
+
await request.stream._fail(error)
|
|
648
|
+
elif not request.opened.done():
|
|
649
|
+
request.opened.set_exception(error)
|
|
650
|
+
try:
|
|
651
|
+
async with asyncio.timeout(self._cleanup_timeout):
|
|
652
|
+
for stream_id, pending in pending_items:
|
|
653
|
+
with suppress(MuxRpcError, MuxProtocolError):
|
|
654
|
+
await self._writer.write_sequenced(
|
|
655
|
+
FrameType.RESET,
|
|
656
|
+
stream_id,
|
|
657
|
+
pending.sequences,
|
|
658
|
+
payload=_error_payload("cancelled", "client closed"),
|
|
659
|
+
)
|
|
660
|
+
for stream_id, request in stream_items:
|
|
661
|
+
with suppress(MuxRpcError, MuxProtocolError):
|
|
662
|
+
await self._writer.write_sequenced(
|
|
663
|
+
FrameType.RESET,
|
|
664
|
+
stream_id,
|
|
665
|
+
request.sequences,
|
|
666
|
+
payload=_error_payload("cancelled", "client closed"),
|
|
667
|
+
)
|
|
668
|
+
await self._writer.write_sequenced(
|
|
669
|
+
FrameType.GOAWAY,
|
|
670
|
+
0,
|
|
671
|
+
self._connection_sequences,
|
|
672
|
+
payload=encode_control_payload({"reason": "client closed"}),
|
|
673
|
+
)
|
|
674
|
+
except (TimeoutError, MuxRpcError, MuxProtocolError):
|
|
675
|
+
pass
|
|
676
|
+
self._stream_writer.close()
|
|
677
|
+
await _wait_closed_bounded(self._stream_writer, self._cleanup_timeout)
|
|
678
|
+
if self._reader_task is not None and self._reader_task is not asyncio.current_task():
|
|
679
|
+
with suppress(MuxConnectionClosed, asyncio.CancelledError):
|
|
680
|
+
await self._reader_task
|
|
681
|
+
if self.state is ConnectionState.DRAINING:
|
|
682
|
+
self._lifecycle.mark_closed()
|
|
683
|
+
|
|
684
|
+
async def _start(self) -> None:
|
|
685
|
+
self._lifecycle.start_handshake()
|
|
686
|
+
try:
|
|
687
|
+
async with asyncio.timeout(self._handshake_timeout):
|
|
688
|
+
await self._handshake()
|
|
689
|
+
except TimeoutError as exc:
|
|
690
|
+
self._lifecycle.mark_failed()
|
|
691
|
+
raise MuxConnectionClosed("multiplexed client handshake timed out") from exc
|
|
692
|
+
except BaseException:
|
|
693
|
+
self._lifecycle.mark_failed()
|
|
694
|
+
raise
|
|
695
|
+
self._lifecycle.mark_ready()
|
|
696
|
+
self._reader_task = asyncio.create_task(self._read_loop(), name="hostbridge-mux-rpc-client")
|
|
697
|
+
self._heartbeat_task = asyncio.create_task(
|
|
698
|
+
self._heartbeat_loop(),
|
|
699
|
+
name="hostbridge-mux-rpc-heartbeat",
|
|
700
|
+
)
|
|
701
|
+
|
|
702
|
+
async def _handshake(self) -> None:
|
|
703
|
+
await self._writer.write_sequenced(
|
|
704
|
+
FrameType.HELLO,
|
|
705
|
+
0,
|
|
706
|
+
self._connection_sequences,
|
|
707
|
+
payload=encode_control_payload(
|
|
708
|
+
{
|
|
709
|
+
"role": "client",
|
|
710
|
+
"rpc": True,
|
|
711
|
+
"connection_window": self._connection_receive_window.available,
|
|
712
|
+
}
|
|
713
|
+
),
|
|
714
|
+
)
|
|
715
|
+
frame = await self._reader.read()
|
|
716
|
+
self._connection_sequences.accept_inbound(frame.sequence)
|
|
717
|
+
if frame.frame_type is not FrameType.HELLO_OK or frame.stream_id != 0:
|
|
718
|
+
raise MuxProtocolError("server did not complete the mux handshake")
|
|
719
|
+
hello = decode_control_payload(frame.payload)
|
|
720
|
+
if hello.get("role") != "server" or hello.get("rpc") is not True:
|
|
721
|
+
raise MuxProtocolError("server returned incompatible mux capabilities")
|
|
722
|
+
connection_window = hello.get("connection_window")
|
|
723
|
+
if not isinstance(connection_window, int) or isinstance(connection_window, bool) or connection_window <= 0:
|
|
724
|
+
raise MuxProtocolError("server returned an invalid connection window")
|
|
725
|
+
self._connection_send_window = AsyncByteCreditWindow(
|
|
726
|
+
connection_window,
|
|
727
|
+
on_stall=self._record_window_stall,
|
|
728
|
+
)
|
|
729
|
+
|
|
730
|
+
async def _read_loop(self) -> None:
|
|
731
|
+
try:
|
|
732
|
+
while True:
|
|
733
|
+
frame = await self._reader.read()
|
|
734
|
+
if frame.stream_id == 0:
|
|
735
|
+
await self._handle_connection_frame(frame)
|
|
736
|
+
continue
|
|
737
|
+
await self._handle_stream_frame(frame)
|
|
738
|
+
except asyncio.CancelledError:
|
|
739
|
+
raise
|
|
740
|
+
except Exception as exc:
|
|
741
|
+
await self._connection_failed(exc)
|
|
742
|
+
|
|
743
|
+
async def _handle_connection_frame(self, frame: MuxFrame) -> None:
|
|
744
|
+
self._connection_sequences.accept_inbound(frame.sequence)
|
|
745
|
+
if frame.frame_type is FrameType.WINDOW_UPDATE:
|
|
746
|
+
if self._connection_send_window is None:
|
|
747
|
+
raise MuxProtocolError("connection window update arrived before handshake")
|
|
748
|
+
try:
|
|
749
|
+
await self._connection_send_window.grant(decode_window_update(frame.payload))
|
|
750
|
+
except ValueError as exc:
|
|
751
|
+
raise MuxProtocolError("connection window credit exceeds the negotiated maximum") from exc
|
|
752
|
+
return
|
|
753
|
+
if frame.frame_type is FrameType.GOAWAY:
|
|
754
|
+
decode_control_payload(frame.payload)
|
|
755
|
+
self._lifecycle.start_draining()
|
|
756
|
+
return
|
|
757
|
+
if frame.frame_type is FrameType.PING:
|
|
758
|
+
await self._writer.write_sequenced(
|
|
759
|
+
FrameType.PONG,
|
|
760
|
+
0,
|
|
761
|
+
self._connection_sequences,
|
|
762
|
+
payload=frame.payload,
|
|
763
|
+
)
|
|
764
|
+
return
|
|
765
|
+
if frame.frame_type is FrameType.PONG:
|
|
766
|
+
waiter = self._heartbeat_waiter
|
|
767
|
+
if waiter is not None and frame.payload == waiter[0] and not waiter[1].done():
|
|
768
|
+
waiter[1].set_result(None)
|
|
769
|
+
return
|
|
770
|
+
raise MuxProtocolError(f"unexpected connection frame {frame.frame_type.name}")
|
|
771
|
+
|
|
772
|
+
async def _handle_stream_frame(self, frame: MuxFrame) -> None:
|
|
773
|
+
stream_request = self._streams.get(frame.stream_id)
|
|
774
|
+
if stream_request is not None:
|
|
775
|
+
await self._handle_client_stream_frame(frame, stream_request)
|
|
776
|
+
return
|
|
777
|
+
pending = self._pending.get(frame.stream_id)
|
|
778
|
+
if pending is None:
|
|
779
|
+
terminal = self._terminal_streams.get(frame.stream_id)
|
|
780
|
+
if terminal is None:
|
|
781
|
+
raise MuxProtocolError(f"frame for unknown stream {frame.stream_id}")
|
|
782
|
+
await self._handle_terminal_frame(frame, terminal)
|
|
783
|
+
return
|
|
784
|
+
|
|
785
|
+
pending.sequences.accept_inbound(frame.sequence)
|
|
786
|
+
if frame.frame_type is FrameType.OPEN_OK:
|
|
787
|
+
control = decode_control_payload(frame.payload)
|
|
788
|
+
result = control.get("result")
|
|
789
|
+
if not isinstance(result, dict):
|
|
790
|
+
raise MuxProtocolError("OPEN_OK response result must be an object")
|
|
791
|
+
pending.result = result
|
|
792
|
+
await self._writer.write_sequenced(
|
|
793
|
+
FrameType.CLOSE,
|
|
794
|
+
frame.stream_id,
|
|
795
|
+
pending.sequences,
|
|
796
|
+
payload=encode_control_payload({"reason": "complete"}),
|
|
797
|
+
)
|
|
798
|
+
return
|
|
799
|
+
if frame.frame_type is FrameType.CLOSE_ACK:
|
|
800
|
+
if pending.result is None:
|
|
801
|
+
raise MuxProtocolError("CLOSE_ACK arrived before OPEN_OK")
|
|
802
|
+
self._pending.pop(frame.stream_id, None)
|
|
803
|
+
self._terminal_streams.add(frame.stream_id, pending.sequences)
|
|
804
|
+
if not pending.future.done():
|
|
805
|
+
pending.future.set_result(pending.result)
|
|
806
|
+
return
|
|
807
|
+
if frame.frame_type is FrameType.RESET:
|
|
808
|
+
self._pending.pop(frame.stream_id, None)
|
|
809
|
+
self._terminal_streams.add(frame.stream_id, pending.sequences)
|
|
810
|
+
error = _remote_error(frame.payload)
|
|
811
|
+
with suppress(MuxRpcError, MuxProtocolError):
|
|
812
|
+
await self._writer.write_sequenced(
|
|
813
|
+
FrameType.RESET,
|
|
814
|
+
frame.stream_id,
|
|
815
|
+
pending.sequences,
|
|
816
|
+
payload=_error_payload("acknowledged", "remote reset acknowledged"),
|
|
817
|
+
)
|
|
818
|
+
if not pending.future.done():
|
|
819
|
+
pending.future.set_exception(error)
|
|
820
|
+
return
|
|
821
|
+
raise MuxProtocolError(f"unexpected response frame {frame.frame_type.name}")
|
|
822
|
+
|
|
823
|
+
async def _handle_client_stream_frame(self, frame: MuxFrame, request: _ClientStreamRequest) -> None:
|
|
824
|
+
request.sequences.accept_inbound(frame.sequence)
|
|
825
|
+
if frame.frame_type is FrameType.OPEN_OK:
|
|
826
|
+
if request.stream is not None:
|
|
827
|
+
raise MuxProtocolError("stream received duplicate OPEN_OK")
|
|
828
|
+
control = decode_control_payload(frame.payload)
|
|
829
|
+
send_credit = control.get("receive_window")
|
|
830
|
+
if not isinstance(send_credit, int) or isinstance(send_credit, bool) or send_credit <= 0:
|
|
831
|
+
raise MuxProtocolError("OPEN_OK contains an invalid receive window")
|
|
832
|
+
request.stream = MuxRpcStream(
|
|
833
|
+
self,
|
|
834
|
+
frame.stream_id,
|
|
835
|
+
request.sequences,
|
|
836
|
+
send_credit=send_credit,
|
|
837
|
+
receive_credit=request.receive_window,
|
|
838
|
+
open_initiator=True,
|
|
839
|
+
on_send_stall=self._record_window_stall,
|
|
840
|
+
)
|
|
841
|
+
if not request.opened.done():
|
|
842
|
+
request.opened.set_result(request.stream)
|
|
843
|
+
return
|
|
844
|
+
if frame.frame_type is FrameType.RESET:
|
|
845
|
+
self._streams.pop(frame.stream_id, None)
|
|
846
|
+
self._terminal_streams.add(frame.stream_id, request.sequences)
|
|
847
|
+
error = _remote_error(frame.payload)
|
|
848
|
+
if request.stream is not None:
|
|
849
|
+
await request.stream._fail(error)
|
|
850
|
+
elif not request.opened.done():
|
|
851
|
+
request.opened.set_exception(error)
|
|
852
|
+
with suppress(MuxRpcError, MuxProtocolError):
|
|
853
|
+
await self._writer.write_sequenced(
|
|
854
|
+
FrameType.RESET,
|
|
855
|
+
frame.stream_id,
|
|
856
|
+
request.sequences,
|
|
857
|
+
payload=_error_payload("acknowledged", "remote reset acknowledged"),
|
|
858
|
+
)
|
|
859
|
+
return
|
|
860
|
+
stream = request.stream
|
|
861
|
+
if stream is None:
|
|
862
|
+
raise MuxProtocolError(f"{frame.frame_type.name} arrived before OPEN_OK")
|
|
863
|
+
if frame.frame_type is FrameType.DATA:
|
|
864
|
+
try:
|
|
865
|
+
self._connection_receive_window.consume(len(frame.payload))
|
|
866
|
+
await stream._accept_data(frame.payload)
|
|
867
|
+
except MuxStateError as exc:
|
|
868
|
+
raise MuxProtocolError("stream DATA exceeds advertised receive credit") from exc
|
|
869
|
+
return
|
|
870
|
+
if frame.frame_type is FrameType.WINDOW_UPDATE:
|
|
871
|
+
try:
|
|
872
|
+
await stream._grant_send_credit(decode_window_update(frame.payload))
|
|
873
|
+
except ValueError as exc:
|
|
874
|
+
raise MuxProtocolError("stream window credit exceeds the negotiated maximum") from exc
|
|
875
|
+
return
|
|
876
|
+
if frame.frame_type is FrameType.EOF:
|
|
877
|
+
await stream._accept_eof()
|
|
878
|
+
await self._maybe_close_stream(frame.stream_id)
|
|
879
|
+
return
|
|
880
|
+
if frame.frame_type is FrameType.CLOSE_ACK:
|
|
881
|
+
if not stream.close_sent:
|
|
882
|
+
raise MuxProtocolError("CLOSE_ACK arrived before CLOSE")
|
|
883
|
+
control = decode_control_payload(frame.payload)
|
|
884
|
+
result = control.get("result")
|
|
885
|
+
if not isinstance(result, dict):
|
|
886
|
+
raise MuxProtocolError("CLOSE_ACK result must be an object")
|
|
887
|
+
self._streams.pop(frame.stream_id, None)
|
|
888
|
+
self._terminal_streams.add(frame.stream_id, request.sequences)
|
|
889
|
+
await stream._finish(result)
|
|
890
|
+
return
|
|
891
|
+
raise MuxProtocolError(f"unexpected stream frame {frame.frame_type.name}")
|
|
892
|
+
|
|
893
|
+
async def _handle_terminal_frame(self, frame: MuxFrame, terminal: _TerminalStream) -> None:
|
|
894
|
+
terminal.sequences.accept_inbound(frame.sequence)
|
|
895
|
+
if frame.frame_type is FrameType.DATA:
|
|
896
|
+
await self._discard_late_data(frame)
|
|
897
|
+
return
|
|
898
|
+
if frame.frame_type is FrameType.WINDOW_UPDATE:
|
|
899
|
+
decode_window_update(frame.payload)
|
|
900
|
+
return
|
|
901
|
+
if frame.frame_type in {
|
|
902
|
+
FrameType.OPEN_OK,
|
|
903
|
+
FrameType.EOF,
|
|
904
|
+
FrameType.CLOSE,
|
|
905
|
+
FrameType.CLOSE_ACK,
|
|
906
|
+
FrameType.RESET,
|
|
907
|
+
}:
|
|
908
|
+
return
|
|
909
|
+
raise MuxProtocolError(f"unexpected late frame {frame.frame_type.name}")
|
|
910
|
+
|
|
911
|
+
async def _discard_late_data(self, frame: MuxFrame) -> None:
|
|
912
|
+
try:
|
|
913
|
+
self._connection_receive_window.consume(len(frame.payload))
|
|
914
|
+
self._connection_receive_window.grant(len(frame.payload))
|
|
915
|
+
except MuxStateError as exc:
|
|
916
|
+
raise MuxProtocolError("late stream DATA exceeds connection credit") from exc
|
|
917
|
+
await self._writer.write_sequenced(
|
|
918
|
+
FrameType.WINDOW_UPDATE,
|
|
919
|
+
0,
|
|
920
|
+
self._connection_sequences,
|
|
921
|
+
payload=encode_window_update(len(frame.payload)),
|
|
922
|
+
)
|
|
923
|
+
|
|
924
|
+
async def _send_stream_data(self, stream: MuxRpcStream, data: bytes) -> None:
|
|
925
|
+
connection_window = self._connection_send_window
|
|
926
|
+
if connection_window is None:
|
|
927
|
+
raise MuxConnectionClosed("connection send window is unavailable")
|
|
928
|
+
chunk_size = min(
|
|
929
|
+
DEFAULT_DATA_CHUNK_BYTES,
|
|
930
|
+
MAX_FRAME_PAYLOAD_BYTES,
|
|
931
|
+
stream._send_window.maximum,
|
|
932
|
+
connection_window.maximum,
|
|
933
|
+
)
|
|
934
|
+
for offset in range(0, len(data), chunk_size):
|
|
935
|
+
chunk = data[offset : offset + chunk_size]
|
|
936
|
+
await stream._send_window.acquire(len(chunk))
|
|
937
|
+
try:
|
|
938
|
+
await connection_window.acquire(len(chunk))
|
|
939
|
+
except BaseException:
|
|
940
|
+
await stream._send_window.grant(len(chunk))
|
|
941
|
+
raise
|
|
942
|
+
try:
|
|
943
|
+
await _write_sequenced_before_honouring_cancellation(
|
|
944
|
+
self._writer,
|
|
945
|
+
FrameType.DATA,
|
|
946
|
+
stream.stream_id,
|
|
947
|
+
stream._sequences,
|
|
948
|
+
payload=chunk,
|
|
949
|
+
)
|
|
950
|
+
except MuxCapacityError as exc:
|
|
951
|
+
await connection_window.grant(len(chunk))
|
|
952
|
+
await stream._send_window.grant(len(chunk))
|
|
953
|
+
await self._resource_limit_stream(stream, exc)
|
|
954
|
+
raise
|
|
955
|
+
|
|
956
|
+
async def _resource_limit_stream(self, stream: MuxRpcStream, error: MuxCapacityError) -> None:
|
|
957
|
+
request = self._streams.pop(stream.stream_id, None)
|
|
958
|
+
if request is None:
|
|
959
|
+
return
|
|
960
|
+
self._terminal_streams.add(stream.stream_id, request.sequences)
|
|
961
|
+
await stream._fail(error)
|
|
962
|
+
await self._writer.write_sequenced(
|
|
963
|
+
FrameType.RESET,
|
|
964
|
+
stream.stream_id,
|
|
965
|
+
request.sequences,
|
|
966
|
+
payload=_error_payload("resource_limit", str(error), retryable=True),
|
|
967
|
+
)
|
|
968
|
+
|
|
969
|
+
async def _maybe_close_stream(self, stream_id: int) -> None:
|
|
970
|
+
request = self._streams.get(stream_id)
|
|
971
|
+
if request is None or request.stream is None or request.stream.close_sent:
|
|
972
|
+
return
|
|
973
|
+
if (
|
|
974
|
+
not request.local_eof_sent
|
|
975
|
+
or not request.stream.local_eof
|
|
976
|
+
or not request.stream.remote_eof
|
|
977
|
+
or request.stream.buffered_bytes != 0
|
|
978
|
+
):
|
|
979
|
+
return
|
|
980
|
+
request.stream._send_close()
|
|
981
|
+
await self._writer.write_sequenced(
|
|
982
|
+
FrameType.CLOSE,
|
|
983
|
+
stream_id,
|
|
984
|
+
request.sequences,
|
|
985
|
+
payload=encode_control_payload({"reason": "complete"}),
|
|
986
|
+
)
|
|
987
|
+
|
|
988
|
+
async def _connection_failed(self, cause: BaseException) -> None:
|
|
989
|
+
if self.closed:
|
|
990
|
+
return
|
|
991
|
+
if self._local_close_started:
|
|
992
|
+
if self.state is ConnectionState.DRAINING:
|
|
993
|
+
self._lifecycle.mark_closed()
|
|
994
|
+
return
|
|
995
|
+
if isinstance(cause, MuxProtocolError):
|
|
996
|
+
self._protocol_failures += 1
|
|
997
|
+
clean_drain = self.draining and not self._pending and not self._streams
|
|
998
|
+
if clean_drain:
|
|
999
|
+
self._lifecycle.mark_closed()
|
|
1000
|
+
else:
|
|
1001
|
+
self._lifecycle.mark_failed()
|
|
1002
|
+
await self._stop_heartbeat()
|
|
1003
|
+
error = MuxConnectionClosed(str(cause) or "multiplexed connection failed")
|
|
1004
|
+
self._close_error = error
|
|
1005
|
+
pending = list(self._pending.values())
|
|
1006
|
+
self._pending.clear()
|
|
1007
|
+
streams = list(self._streams.values())
|
|
1008
|
+
self._streams.clear()
|
|
1009
|
+
self._terminal_streams.clear()
|
|
1010
|
+
for pending_request in pending:
|
|
1011
|
+
if not pending_request.future.done():
|
|
1012
|
+
pending_request.future.set_exception(error)
|
|
1013
|
+
for stream_request in streams:
|
|
1014
|
+
if stream_request.stream is not None:
|
|
1015
|
+
await stream_request.stream._fail(error)
|
|
1016
|
+
elif not stream_request.opened.done():
|
|
1017
|
+
stream_request.opened.set_exception(error)
|
|
1018
|
+
if self._connection_send_window is not None:
|
|
1019
|
+
await self._connection_send_window.close(error)
|
|
1020
|
+
self._stream_writer.close()
|
|
1021
|
+
await _wait_closed_bounded(self._stream_writer, self._cleanup_timeout)
|
|
1022
|
+
|
|
1023
|
+
async def _heartbeat_loop(self) -> None:
|
|
1024
|
+
try:
|
|
1025
|
+
while True:
|
|
1026
|
+
await asyncio.sleep(self._heartbeat_interval)
|
|
1027
|
+
if self.closed:
|
|
1028
|
+
return
|
|
1029
|
+
self._heartbeat_nonce = (self._heartbeat_nonce + 1) & 0xFFFFFFFFFFFFFFFF
|
|
1030
|
+
payload = self._heartbeat_nonce.to_bytes(8, "big")
|
|
1031
|
+
acknowledged = asyncio.get_running_loop().create_future()
|
|
1032
|
+
self._heartbeat_waiter = (payload, acknowledged)
|
|
1033
|
+
try:
|
|
1034
|
+
await self._writer.write_sequenced(
|
|
1035
|
+
FrameType.PING,
|
|
1036
|
+
0,
|
|
1037
|
+
self._connection_sequences,
|
|
1038
|
+
payload=payload,
|
|
1039
|
+
)
|
|
1040
|
+
async with asyncio.timeout(self._heartbeat_timeout):
|
|
1041
|
+
await acknowledged
|
|
1042
|
+
except TimeoutError:
|
|
1043
|
+
await self._connection_failed(MuxConnectionClosed("multiplexed heartbeat timed out"))
|
|
1044
|
+
return
|
|
1045
|
+
finally:
|
|
1046
|
+
if self._heartbeat_waiter is not None and self._heartbeat_waiter[1] is acknowledged:
|
|
1047
|
+
self._heartbeat_waiter = None
|
|
1048
|
+
except asyncio.CancelledError:
|
|
1049
|
+
raise
|
|
1050
|
+
except (MuxConnectionClosed, MuxProtocolError, OSError) as exc:
|
|
1051
|
+
await self._connection_failed(exc)
|
|
1052
|
+
|
|
1053
|
+
async def _stop_heartbeat(self) -> None:
|
|
1054
|
+
heartbeat = self._heartbeat_task
|
|
1055
|
+
if heartbeat is None or heartbeat is asyncio.current_task():
|
|
1056
|
+
return
|
|
1057
|
+
self._heartbeat_task = None
|
|
1058
|
+
heartbeat.cancel()
|
|
1059
|
+
await asyncio.gather(heartbeat, return_exceptions=True)
|
|
1060
|
+
|
|
1061
|
+
|
|
1062
|
+
@dataclass(slots=True)
|
|
1063
|
+
class _ServerRequest:
|
|
1064
|
+
sequences: SequenceTracker
|
|
1065
|
+
method: str | None = None
|
|
1066
|
+
task: asyncio.Task[None] | None = None
|
|
1067
|
+
response_sent: bool = False
|
|
1068
|
+
reset_sent: bool = False
|
|
1069
|
+
stream: MuxRpcStream | None = None
|
|
1070
|
+
stream_result: dict[str, object] | None = None
|
|
1071
|
+
|
|
1072
|
+
|
|
1073
|
+
class MuxRpcServerConnection:
|
|
1074
|
+
def __init__(
|
|
1075
|
+
self,
|
|
1076
|
+
reader: asyncio.StreamReader,
|
|
1077
|
+
writer: asyncio.StreamWriter,
|
|
1078
|
+
handler: MuxRpcHandler,
|
|
1079
|
+
*,
|
|
1080
|
+
max_handlers: int,
|
|
1081
|
+
stream_handler: MuxRpcStreamHandler | None,
|
|
1082
|
+
request_closed_handler: MuxRpcRequestClosedHandler | None,
|
|
1083
|
+
connection_window: int,
|
|
1084
|
+
stream_window: int,
|
|
1085
|
+
handshake_timeout: float = DEFAULT_HANDSHAKE_TIMEOUT_SECONDS,
|
|
1086
|
+
cleanup_timeout: float = DEFAULT_CLEANUP_TIMEOUT_SECONDS,
|
|
1087
|
+
) -> None:
|
|
1088
|
+
_positive(max_handlers, "max_handlers")
|
|
1089
|
+
_positive(connection_window, "connection_window")
|
|
1090
|
+
_positive(stream_window, "stream_window")
|
|
1091
|
+
_finite_positive(handshake_timeout, "handshake_timeout")
|
|
1092
|
+
_finite_positive(cleanup_timeout, "cleanup_timeout")
|
|
1093
|
+
self._reader = _FrameReader(reader)
|
|
1094
|
+
self._writer = _FrameWriter(writer, failure_handler=self._writer_failed)
|
|
1095
|
+
self._stream_writer = writer
|
|
1096
|
+
self._handler = handler
|
|
1097
|
+
self._stream_handler = stream_handler
|
|
1098
|
+
self._request_closed_handler = request_closed_handler
|
|
1099
|
+
self._max_handlers = max_handlers
|
|
1100
|
+
self._default_stream_window = stream_window
|
|
1101
|
+
self._requests: dict[int, _ServerRequest] = {}
|
|
1102
|
+
self._terminal_streams = _TerminalStreams()
|
|
1103
|
+
self._connection_sequences = SequenceTracker()
|
|
1104
|
+
self._connection_receive_window = ByteCreditWindow(
|
|
1105
|
+
initial_credit=connection_window,
|
|
1106
|
+
maximum_credit=connection_window,
|
|
1107
|
+
)
|
|
1108
|
+
self._connection_send_window: AsyncByteCreditWindow | None = None
|
|
1109
|
+
self._window_stalls = 0
|
|
1110
|
+
self._protocol_failures = 0
|
|
1111
|
+
self._lifecycle = ConnectionLifecycle()
|
|
1112
|
+
self._lifecycle.start_connecting()
|
|
1113
|
+
self._drain_requested = False
|
|
1114
|
+
self._goaway_sent = False
|
|
1115
|
+
self._ready = asyncio.Event()
|
|
1116
|
+
self._handshake_timeout = float(handshake_timeout)
|
|
1117
|
+
self._cleanup_timeout = float(cleanup_timeout)
|
|
1118
|
+
|
|
1119
|
+
@property
|
|
1120
|
+
def active_requests(self) -> int:
|
|
1121
|
+
return len(self._requests)
|
|
1122
|
+
|
|
1123
|
+
@property
|
|
1124
|
+
def window_stalls(self) -> int:
|
|
1125
|
+
return self._window_stalls
|
|
1126
|
+
|
|
1127
|
+
@property
|
|
1128
|
+
def protocol_failures(self) -> int:
|
|
1129
|
+
return self._protocol_failures
|
|
1130
|
+
|
|
1131
|
+
@property
|
|
1132
|
+
def state(self) -> ConnectionState:
|
|
1133
|
+
return self._lifecycle.state
|
|
1134
|
+
|
|
1135
|
+
def _record_window_stall(self) -> None:
|
|
1136
|
+
self._window_stalls += 1
|
|
1137
|
+
|
|
1138
|
+
async def _writer_failed(self, _cause: BaseException) -> None:
|
|
1139
|
+
self._stream_writer.close()
|
|
1140
|
+
|
|
1141
|
+
async def start_draining(self, reason: str = "server shutting down") -> None:
|
|
1142
|
+
if self._lifecycle.terminal:
|
|
1143
|
+
return
|
|
1144
|
+
self._drain_requested = True
|
|
1145
|
+
await self._ready.wait()
|
|
1146
|
+
if self._lifecycle.terminal or self._goaway_sent:
|
|
1147
|
+
return
|
|
1148
|
+
self._lifecycle.start_draining()
|
|
1149
|
+
self._goaway_sent = True
|
|
1150
|
+
await self._writer.write_sequenced(
|
|
1151
|
+
FrameType.GOAWAY,
|
|
1152
|
+
0,
|
|
1153
|
+
self._connection_sequences,
|
|
1154
|
+
payload=encode_control_payload({"reason": reason}),
|
|
1155
|
+
)
|
|
1156
|
+
if not self._requests:
|
|
1157
|
+
self._stream_writer.close()
|
|
1158
|
+
|
|
1159
|
+
async def serve(self) -> None:
|
|
1160
|
+
try:
|
|
1161
|
+
self._lifecycle.start_handshake()
|
|
1162
|
+
try:
|
|
1163
|
+
async with asyncio.timeout(self._handshake_timeout):
|
|
1164
|
+
await self._handshake()
|
|
1165
|
+
except TimeoutError:
|
|
1166
|
+
self._lifecycle.mark_failed()
|
|
1167
|
+
return
|
|
1168
|
+
self._lifecycle.mark_ready()
|
|
1169
|
+
if self._drain_requested:
|
|
1170
|
+
self._lifecycle.start_draining()
|
|
1171
|
+
self._ready.set()
|
|
1172
|
+
while True:
|
|
1173
|
+
frame = await self._reader.read()
|
|
1174
|
+
if frame.stream_id == 0:
|
|
1175
|
+
if await self._handle_connection_frame(frame):
|
|
1176
|
+
return
|
|
1177
|
+
else:
|
|
1178
|
+
await self._handle_stream_frame(frame)
|
|
1179
|
+
if self._lifecycle.draining and not self._requests:
|
|
1180
|
+
return
|
|
1181
|
+
except (MuxConnectionClosed, MuxProtocolError, OSError) as exc:
|
|
1182
|
+
if isinstance(exc, MuxProtocolError):
|
|
1183
|
+
self._protocol_failures += 1
|
|
1184
|
+
self._lifecycle.mark_failed()
|
|
1185
|
+
return
|
|
1186
|
+
finally:
|
|
1187
|
+
self._ready.set()
|
|
1188
|
+
tasks = [request.task for request in self._requests.values() if request.task is not None]
|
|
1189
|
+
for task in tasks:
|
|
1190
|
+
task.cancel()
|
|
1191
|
+
if tasks:
|
|
1192
|
+
await asyncio.gather(*tasks, return_exceptions=True)
|
|
1193
|
+
error = MuxConnectionClosed("multiplexed server connection closed")
|
|
1194
|
+
for request in self._requests.values():
|
|
1195
|
+
if request.stream is not None:
|
|
1196
|
+
await request.stream._fail(error)
|
|
1197
|
+
if self._connection_send_window is not None:
|
|
1198
|
+
await self._connection_send_window.close(error)
|
|
1199
|
+
self._requests.clear()
|
|
1200
|
+
self._terminal_streams.clear()
|
|
1201
|
+
self._stream_writer.close()
|
|
1202
|
+
await _wait_closed_bounded(self._stream_writer, self._cleanup_timeout)
|
|
1203
|
+
if self.state is ConnectionState.DRAINING:
|
|
1204
|
+
self._lifecycle.mark_closed()
|
|
1205
|
+
elif not self._lifecycle.terminal:
|
|
1206
|
+
self._lifecycle.mark_failed()
|
|
1207
|
+
|
|
1208
|
+
async def _handshake(self) -> None:
|
|
1209
|
+
frame = await self._reader.read()
|
|
1210
|
+
self._connection_sequences.accept_inbound(frame.sequence)
|
|
1211
|
+
if frame.frame_type is not FrameType.HELLO or frame.stream_id != 0:
|
|
1212
|
+
raise MuxProtocolError("client did not start with HELLO")
|
|
1213
|
+
hello = decode_control_payload(frame.payload)
|
|
1214
|
+
if hello.get("role") != "client" or hello.get("rpc") is not True:
|
|
1215
|
+
raise MuxProtocolError("client offered incompatible mux capabilities")
|
|
1216
|
+
connection_window = hello.get("connection_window")
|
|
1217
|
+
if not isinstance(connection_window, int) or isinstance(connection_window, bool) or connection_window <= 0:
|
|
1218
|
+
raise MuxProtocolError("client offered an invalid connection window")
|
|
1219
|
+
self._connection_send_window = AsyncByteCreditWindow(
|
|
1220
|
+
connection_window,
|
|
1221
|
+
on_stall=self._record_window_stall,
|
|
1222
|
+
)
|
|
1223
|
+
await self._writer.write_sequenced(
|
|
1224
|
+
FrameType.HELLO_OK,
|
|
1225
|
+
0,
|
|
1226
|
+
self._connection_sequences,
|
|
1227
|
+
payload=encode_control_payload(
|
|
1228
|
+
{
|
|
1229
|
+
"role": "server",
|
|
1230
|
+
"rpc": True,
|
|
1231
|
+
"connection_window": self._connection_receive_window.available,
|
|
1232
|
+
}
|
|
1233
|
+
),
|
|
1234
|
+
)
|
|
1235
|
+
|
|
1236
|
+
async def _handle_connection_frame(self, frame: MuxFrame) -> bool:
|
|
1237
|
+
self._connection_sequences.accept_inbound(frame.sequence)
|
|
1238
|
+
if frame.frame_type is FrameType.WINDOW_UPDATE:
|
|
1239
|
+
if self._connection_send_window is None:
|
|
1240
|
+
raise MuxProtocolError("connection window update arrived before handshake")
|
|
1241
|
+
try:
|
|
1242
|
+
await self._connection_send_window.grant(decode_window_update(frame.payload))
|
|
1243
|
+
except ValueError as exc:
|
|
1244
|
+
raise MuxProtocolError("connection window credit exceeds the negotiated maximum") from exc
|
|
1245
|
+
return False
|
|
1246
|
+
if frame.frame_type is FrameType.GOAWAY:
|
|
1247
|
+
self._lifecycle.start_draining()
|
|
1248
|
+
return True
|
|
1249
|
+
if frame.frame_type is FrameType.PING:
|
|
1250
|
+
await self._writer.write_sequenced(
|
|
1251
|
+
FrameType.PONG,
|
|
1252
|
+
0,
|
|
1253
|
+
self._connection_sequences,
|
|
1254
|
+
payload=frame.payload,
|
|
1255
|
+
)
|
|
1256
|
+
return False
|
|
1257
|
+
if frame.frame_type is FrameType.PONG:
|
|
1258
|
+
return False
|
|
1259
|
+
raise MuxProtocolError(f"unexpected connection frame {frame.frame_type.name}")
|
|
1260
|
+
|
|
1261
|
+
async def _handle_stream_frame(self, frame: MuxFrame) -> None:
|
|
1262
|
+
request = self._requests.get(frame.stream_id)
|
|
1263
|
+
if frame.frame_type is FrameType.OPEN:
|
|
1264
|
+
if (
|
|
1265
|
+
request is not None
|
|
1266
|
+
or self._terminal_streams.get(frame.stream_id) is not None
|
|
1267
|
+
or frame.stream_id % 2 != int(StreamParity.ODD)
|
|
1268
|
+
):
|
|
1269
|
+
raise MuxProtocolError(f"invalid or reused client stream {frame.stream_id}")
|
|
1270
|
+
request = _ServerRequest(SequenceTracker())
|
|
1271
|
+
request.sequences.accept_inbound(frame.sequence)
|
|
1272
|
+
self._requests[frame.stream_id] = request
|
|
1273
|
+
if self._drain_requested or self._lifecycle.draining:
|
|
1274
|
+
await self._send_failure(
|
|
1275
|
+
frame.stream_id,
|
|
1276
|
+
request,
|
|
1277
|
+
MuxRpcFailure("connection_lost", "multiplexed server connection is draining", retryable=True),
|
|
1278
|
+
)
|
|
1279
|
+
return
|
|
1280
|
+
control = decode_control_payload(frame.payload)
|
|
1281
|
+
mode = control.get("mode", "unary")
|
|
1282
|
+
method = control.get("method")
|
|
1283
|
+
params = control.get("params")
|
|
1284
|
+
if not isinstance(method, str) or not method.strip() or not isinstance(params, dict):
|
|
1285
|
+
await self._send_failure(frame.stream_id, request, MuxRpcFailure("invalid_request", "invalid RPC"))
|
|
1286
|
+
return
|
|
1287
|
+
request.method = method
|
|
1288
|
+
active_handlers = sum(item.task is not None and not item.task.done() for item in self._requests.values())
|
|
1289
|
+
if active_handlers >= self._max_handlers:
|
|
1290
|
+
await self._send_failure(
|
|
1291
|
+
frame.stream_id,
|
|
1292
|
+
request,
|
|
1293
|
+
MuxRpcFailure("busy", "server handler capacity is exhausted", retryable=True),
|
|
1294
|
+
)
|
|
1295
|
+
return
|
|
1296
|
+
if mode == "stream":
|
|
1297
|
+
receive_window = control.get("receive_window")
|
|
1298
|
+
if (
|
|
1299
|
+
self._stream_handler is None
|
|
1300
|
+
or not isinstance(receive_window, int)
|
|
1301
|
+
or isinstance(receive_window, bool)
|
|
1302
|
+
or receive_window <= 0
|
|
1303
|
+
):
|
|
1304
|
+
await self._send_failure(
|
|
1305
|
+
frame.stream_id,
|
|
1306
|
+
request,
|
|
1307
|
+
MuxRpcFailure("invalid_request", "streaming RPC is unavailable or invalid"),
|
|
1308
|
+
)
|
|
1309
|
+
return
|
|
1310
|
+
request.stream = MuxRpcStream(
|
|
1311
|
+
self,
|
|
1312
|
+
frame.stream_id,
|
|
1313
|
+
request.sequences,
|
|
1314
|
+
send_credit=receive_window,
|
|
1315
|
+
receive_credit=self._default_stream_window,
|
|
1316
|
+
open_initiator=False,
|
|
1317
|
+
result_expected=False,
|
|
1318
|
+
accepted=False,
|
|
1319
|
+
on_send_stall=self._record_window_stall,
|
|
1320
|
+
)
|
|
1321
|
+
request.task = asyncio.create_task(
|
|
1322
|
+
self._run_stream_handler(frame.stream_id, request, method, params),
|
|
1323
|
+
name=f"hostbridge-mux-stream-{frame.stream_id}",
|
|
1324
|
+
)
|
|
1325
|
+
return
|
|
1326
|
+
if mode != "unary":
|
|
1327
|
+
await self._send_failure(
|
|
1328
|
+
frame.stream_id,
|
|
1329
|
+
request,
|
|
1330
|
+
MuxRpcFailure("invalid_request", "unsupported RPC mode"),
|
|
1331
|
+
)
|
|
1332
|
+
return
|
|
1333
|
+
request.task = asyncio.create_task(
|
|
1334
|
+
self._run_handler(frame.stream_id, request, method, params),
|
|
1335
|
+
name=f"hostbridge-mux-rpc-{frame.stream_id}",
|
|
1336
|
+
)
|
|
1337
|
+
return
|
|
1338
|
+
if request is None:
|
|
1339
|
+
terminal = self._terminal_streams.get(frame.stream_id)
|
|
1340
|
+
if terminal is not None:
|
|
1341
|
+
await self._handle_terminal_frame(frame, terminal)
|
|
1342
|
+
return
|
|
1343
|
+
raise MuxProtocolError(f"frame for unknown stream {frame.stream_id}")
|
|
1344
|
+
|
|
1345
|
+
request.sequences.accept_inbound(frame.sequence)
|
|
1346
|
+
if request.stream is not None:
|
|
1347
|
+
if not request.response_sent and frame.frame_type is not FrameType.RESET:
|
|
1348
|
+
raise MuxProtocolError(f"{frame.frame_type.name} arrived before stream acceptance")
|
|
1349
|
+
if frame.frame_type is FrameType.DATA:
|
|
1350
|
+
try:
|
|
1351
|
+
self._connection_receive_window.consume(len(frame.payload))
|
|
1352
|
+
await request.stream._accept_data(frame.payload)
|
|
1353
|
+
except MuxStateError as exc:
|
|
1354
|
+
raise MuxProtocolError("stream DATA exceeds advertised receive credit") from exc
|
|
1355
|
+
return
|
|
1356
|
+
if frame.frame_type is FrameType.WINDOW_UPDATE:
|
|
1357
|
+
try:
|
|
1358
|
+
await request.stream._grant_send_credit(decode_window_update(frame.payload))
|
|
1359
|
+
except ValueError as exc:
|
|
1360
|
+
raise MuxProtocolError("stream window credit exceeds the negotiated maximum") from exc
|
|
1361
|
+
return
|
|
1362
|
+
if frame.frame_type is FrameType.EOF:
|
|
1363
|
+
await request.stream._accept_eof()
|
|
1364
|
+
return
|
|
1365
|
+
if frame.frame_type is FrameType.CLOSE:
|
|
1366
|
+
if not request.response_sent:
|
|
1367
|
+
raise MuxProtocolError("CLOSE arrived before OPEN_OK")
|
|
1368
|
+
if request.stream is not None:
|
|
1369
|
+
try:
|
|
1370
|
+
request.stream._receive_close()
|
|
1371
|
+
except MuxStateError as exc:
|
|
1372
|
+
raise MuxProtocolError("stream CLOSE arrived before both EOF") from exc
|
|
1373
|
+
if request.stream_result is None:
|
|
1374
|
+
return
|
|
1375
|
+
await self._acknowledge_stream_close(frame.stream_id, request)
|
|
1376
|
+
return
|
|
1377
|
+
await self._writer.write_sequenced(
|
|
1378
|
+
FrameType.CLOSE_ACK,
|
|
1379
|
+
frame.stream_id,
|
|
1380
|
+
request.sequences,
|
|
1381
|
+
payload=b"",
|
|
1382
|
+
)
|
|
1383
|
+
self._retire_request(frame.stream_id, request)
|
|
1384
|
+
if self._request_closed_handler is not None and request.method is not None:
|
|
1385
|
+
await self._request_closed_handler(request.method)
|
|
1386
|
+
return
|
|
1387
|
+
if frame.frame_type is FrameType.RESET:
|
|
1388
|
+
if request.task is not None and not request.task.done():
|
|
1389
|
+
request.task.cancel()
|
|
1390
|
+
if not request.reset_sent:
|
|
1391
|
+
request.reset_sent = True
|
|
1392
|
+
await self._writer.write_sequenced(
|
|
1393
|
+
FrameType.RESET,
|
|
1394
|
+
frame.stream_id,
|
|
1395
|
+
request.sequences,
|
|
1396
|
+
payload=_error_payload("cancelled", "request cancelled"),
|
|
1397
|
+
)
|
|
1398
|
+
if request.stream is not None:
|
|
1399
|
+
await request.stream._fail(MuxConnectionClosed("stream reset by peer"))
|
|
1400
|
+
self._retire_request(frame.stream_id, request)
|
|
1401
|
+
return
|
|
1402
|
+
raise MuxProtocolError(f"unexpected request frame {frame.frame_type.name}")
|
|
1403
|
+
|
|
1404
|
+
async def _handle_terminal_frame(self, frame: MuxFrame, terminal: _TerminalStream) -> None:
|
|
1405
|
+
terminal.sequences.accept_inbound(frame.sequence)
|
|
1406
|
+
if frame.frame_type is FrameType.DATA:
|
|
1407
|
+
await self._discard_late_data(frame)
|
|
1408
|
+
return
|
|
1409
|
+
if frame.frame_type is FrameType.WINDOW_UPDATE:
|
|
1410
|
+
decode_window_update(frame.payload)
|
|
1411
|
+
return
|
|
1412
|
+
if frame.frame_type in {
|
|
1413
|
+
FrameType.EOF,
|
|
1414
|
+
FrameType.CLOSE,
|
|
1415
|
+
FrameType.CLOSE_ACK,
|
|
1416
|
+
FrameType.RESET,
|
|
1417
|
+
}:
|
|
1418
|
+
return
|
|
1419
|
+
raise MuxProtocolError(f"unexpected late frame {frame.frame_type.name}")
|
|
1420
|
+
|
|
1421
|
+
async def _discard_late_data(self, frame: MuxFrame) -> None:
|
|
1422
|
+
try:
|
|
1423
|
+
self._connection_receive_window.consume(len(frame.payload))
|
|
1424
|
+
self._connection_receive_window.grant(len(frame.payload))
|
|
1425
|
+
except MuxStateError as exc:
|
|
1426
|
+
raise MuxProtocolError("late stream DATA exceeds connection credit") from exc
|
|
1427
|
+
await self._writer.write_sequenced(
|
|
1428
|
+
FrameType.WINDOW_UPDATE,
|
|
1429
|
+
0,
|
|
1430
|
+
self._connection_sequences,
|
|
1431
|
+
payload=encode_window_update(len(frame.payload)),
|
|
1432
|
+
)
|
|
1433
|
+
|
|
1434
|
+
def _retire_request(self, stream_id: int, request: _ServerRequest) -> None:
|
|
1435
|
+
self._terminal_streams.add(stream_id, request.sequences)
|
|
1436
|
+
self._requests.pop(stream_id, None)
|
|
1437
|
+
|
|
1438
|
+
async def _run_handler(
|
|
1439
|
+
self,
|
|
1440
|
+
stream_id: int,
|
|
1441
|
+
request: _ServerRequest,
|
|
1442
|
+
method: str,
|
|
1443
|
+
params: dict[str, object],
|
|
1444
|
+
) -> None:
|
|
1445
|
+
try:
|
|
1446
|
+
result = await self._handler(method, params)
|
|
1447
|
+
if not isinstance(result, dict):
|
|
1448
|
+
raise MuxRpcFailure("internal", "RPC handler returned a non-object result")
|
|
1449
|
+
if self._requests.get(stream_id) is not request:
|
|
1450
|
+
return
|
|
1451
|
+
await self._writer.write_sequenced(
|
|
1452
|
+
FrameType.OPEN_OK,
|
|
1453
|
+
stream_id,
|
|
1454
|
+
request.sequences,
|
|
1455
|
+
payload=encode_control_payload({"result": result}),
|
|
1456
|
+
)
|
|
1457
|
+
request.response_sent = True
|
|
1458
|
+
except asyncio.CancelledError:
|
|
1459
|
+
raise
|
|
1460
|
+
except MuxRpcFailure as exc:
|
|
1461
|
+
if self._requests.get(stream_id) is request:
|
|
1462
|
+
await self._send_failure(stream_id, request, exc)
|
|
1463
|
+
except Exception:
|
|
1464
|
+
if self._requests.get(stream_id) is request:
|
|
1465
|
+
await self._send_failure(
|
|
1466
|
+
stream_id,
|
|
1467
|
+
request,
|
|
1468
|
+
MuxRpcFailure("internal", "internal HostBridge error", retryable=True),
|
|
1469
|
+
)
|
|
1470
|
+
|
|
1471
|
+
async def _run_stream_handler(
|
|
1472
|
+
self,
|
|
1473
|
+
stream_id: int,
|
|
1474
|
+
request: _ServerRequest,
|
|
1475
|
+
method: str,
|
|
1476
|
+
params: dict[str, object],
|
|
1477
|
+
) -> None:
|
|
1478
|
+
stream = request.stream
|
|
1479
|
+
handler = self._stream_handler
|
|
1480
|
+
try:
|
|
1481
|
+
if stream is None or handler is None:
|
|
1482
|
+
raise MuxRpcFailure("internal", "stream handler state is unavailable")
|
|
1483
|
+
result = await handler(method, params, stream)
|
|
1484
|
+
if not isinstance(result, dict):
|
|
1485
|
+
raise MuxRpcFailure("internal", "stream handler returned a non-object result")
|
|
1486
|
+
if not stream.accepted:
|
|
1487
|
+
raise MuxRpcFailure("internal", "stream handler returned before accepting the stream")
|
|
1488
|
+
if self._requests.get(stream_id) is not request:
|
|
1489
|
+
return
|
|
1490
|
+
request.stream_result = result
|
|
1491
|
+
await stream.send_eof()
|
|
1492
|
+
if stream.close_received and self._requests.get(stream_id) is request:
|
|
1493
|
+
await self._acknowledge_stream_close(stream_id, request)
|
|
1494
|
+
except asyncio.CancelledError:
|
|
1495
|
+
raise
|
|
1496
|
+
except MuxRpcFailure as exc:
|
|
1497
|
+
if self._requests.get(stream_id) is request:
|
|
1498
|
+
await self._send_failure(stream_id, request, exc)
|
|
1499
|
+
if stream is not None:
|
|
1500
|
+
await stream._fail(MuxConnectionClosed(f"stream reset: {exc}"))
|
|
1501
|
+
except Exception:
|
|
1502
|
+
if self._requests.get(stream_id) is request:
|
|
1503
|
+
await self._send_failure(
|
|
1504
|
+
stream_id,
|
|
1505
|
+
request,
|
|
1506
|
+
MuxRpcFailure("internal", "internal HostBridge error", retryable=True),
|
|
1507
|
+
)
|
|
1508
|
+
if stream is not None:
|
|
1509
|
+
await stream._fail(MuxConnectionClosed("stream handler failed"))
|
|
1510
|
+
|
|
1511
|
+
async def _stream_send_data(self, stream: MuxRpcStream, data: bytes) -> None:
|
|
1512
|
+
await self._send_stream_data(stream, data)
|
|
1513
|
+
|
|
1514
|
+
async def _acknowledge_stream_close(self, stream_id: int, request: _ServerRequest) -> None:
|
|
1515
|
+
stream = request.stream
|
|
1516
|
+
result = request.stream_result
|
|
1517
|
+
if stream is None or result is None:
|
|
1518
|
+
raise MuxProtocolError("stream CLOSE acknowledgement is missing terminal state")
|
|
1519
|
+
await self._writer.write_sequenced(
|
|
1520
|
+
FrameType.CLOSE_ACK,
|
|
1521
|
+
stream_id,
|
|
1522
|
+
request.sequences,
|
|
1523
|
+
payload=encode_control_payload({"result": result}),
|
|
1524
|
+
)
|
|
1525
|
+
await stream._finish(result)
|
|
1526
|
+
self._retire_request(stream_id, request)
|
|
1527
|
+
if self._request_closed_handler is not None and request.method is not None:
|
|
1528
|
+
await self._request_closed_handler(request.method)
|
|
1529
|
+
|
|
1530
|
+
async def _stream_accept(self, stream: MuxRpcStream) -> None:
|
|
1531
|
+
request = self._requests.get(stream.stream_id)
|
|
1532
|
+
if request is None or request.stream is not stream:
|
|
1533
|
+
raise MuxConnectionClosed("stream request is no longer active")
|
|
1534
|
+
if request.response_sent:
|
|
1535
|
+
return
|
|
1536
|
+
await self._writer.write_sequenced(
|
|
1537
|
+
FrameType.OPEN_OK,
|
|
1538
|
+
stream.stream_id,
|
|
1539
|
+
request.sequences,
|
|
1540
|
+
payload=encode_control_payload({"stream": True, "receive_window": self._default_stream_window}),
|
|
1541
|
+
)
|
|
1542
|
+
request.response_sent = True
|
|
1543
|
+
|
|
1544
|
+
async def _stream_send_eof(self, stream: MuxRpcStream) -> None:
|
|
1545
|
+
await _write_sequenced_before_honouring_cancellation(
|
|
1546
|
+
self._writer,
|
|
1547
|
+
FrameType.EOF,
|
|
1548
|
+
stream.stream_id,
|
|
1549
|
+
stream._sequences,
|
|
1550
|
+
)
|
|
1551
|
+
|
|
1552
|
+
async def _stream_release_credit(self, stream: MuxRpcStream, amount: int) -> None:
|
|
1553
|
+
await stream._release_receive_credit(amount)
|
|
1554
|
+
self._connection_receive_window.grant(amount)
|
|
1555
|
+
await self._writer.write_sequenced(
|
|
1556
|
+
FrameType.WINDOW_UPDATE,
|
|
1557
|
+
stream.stream_id,
|
|
1558
|
+
stream._sequences,
|
|
1559
|
+
payload=encode_window_update(amount),
|
|
1560
|
+
)
|
|
1561
|
+
await self._writer.write_sequenced(
|
|
1562
|
+
FrameType.WINDOW_UPDATE,
|
|
1563
|
+
0,
|
|
1564
|
+
self._connection_sequences,
|
|
1565
|
+
payload=encode_window_update(amount),
|
|
1566
|
+
)
|
|
1567
|
+
|
|
1568
|
+
async def _stream_reset(self, stream: MuxRpcStream, message: str) -> None:
|
|
1569
|
+
request = self._requests.get(stream.stream_id)
|
|
1570
|
+
if request is None:
|
|
1571
|
+
return
|
|
1572
|
+
if request.task is not None and request.task is not asyncio.current_task():
|
|
1573
|
+
request.task.cancel()
|
|
1574
|
+
request.reset_sent = True
|
|
1575
|
+
await stream._fail(MuxConnectionClosed(f"stream reset: {message}"))
|
|
1576
|
+
await self._writer.write_sequenced(
|
|
1577
|
+
FrameType.RESET,
|
|
1578
|
+
stream.stream_id,
|
|
1579
|
+
request.sequences,
|
|
1580
|
+
payload=_error_payload("cancelled", message),
|
|
1581
|
+
)
|
|
1582
|
+
|
|
1583
|
+
async def _send_stream_data(self, stream: MuxRpcStream, data: bytes) -> None:
|
|
1584
|
+
connection_window = self._connection_send_window
|
|
1585
|
+
if connection_window is None:
|
|
1586
|
+
raise MuxConnectionClosed("connection send window is unavailable")
|
|
1587
|
+
chunk_size = min(
|
|
1588
|
+
DEFAULT_DATA_CHUNK_BYTES,
|
|
1589
|
+
MAX_FRAME_PAYLOAD_BYTES,
|
|
1590
|
+
stream._send_window.maximum,
|
|
1591
|
+
connection_window.maximum,
|
|
1592
|
+
)
|
|
1593
|
+
for offset in range(0, len(data), chunk_size):
|
|
1594
|
+
chunk = data[offset : offset + chunk_size]
|
|
1595
|
+
await stream._send_window.acquire(len(chunk))
|
|
1596
|
+
try:
|
|
1597
|
+
await connection_window.acquire(len(chunk))
|
|
1598
|
+
except BaseException:
|
|
1599
|
+
await stream._send_window.grant(len(chunk))
|
|
1600
|
+
raise
|
|
1601
|
+
try:
|
|
1602
|
+
await _write_sequenced_before_honouring_cancellation(
|
|
1603
|
+
self._writer,
|
|
1604
|
+
FrameType.DATA,
|
|
1605
|
+
stream.stream_id,
|
|
1606
|
+
stream._sequences,
|
|
1607
|
+
payload=chunk,
|
|
1608
|
+
)
|
|
1609
|
+
except MuxCapacityError as exc:
|
|
1610
|
+
await connection_window.grant(len(chunk))
|
|
1611
|
+
await stream._send_window.grant(len(chunk))
|
|
1612
|
+
await self._resource_limit_stream(stream, exc)
|
|
1613
|
+
raise
|
|
1614
|
+
|
|
1615
|
+
async def _resource_limit_stream(self, stream: MuxRpcStream, error: MuxCapacityError) -> None:
|
|
1616
|
+
request = self._requests.get(stream.stream_id)
|
|
1617
|
+
if request is None:
|
|
1618
|
+
return
|
|
1619
|
+
if request.task is not None and request.task is not asyncio.current_task():
|
|
1620
|
+
request.task.cancel()
|
|
1621
|
+
request.reset_sent = True
|
|
1622
|
+
await stream._fail(error)
|
|
1623
|
+
self._retire_request(stream.stream_id, request)
|
|
1624
|
+
await self._writer.write_sequenced(
|
|
1625
|
+
FrameType.RESET,
|
|
1626
|
+
stream.stream_id,
|
|
1627
|
+
request.sequences,
|
|
1628
|
+
payload=_error_payload("resource_limit", str(error), retryable=True),
|
|
1629
|
+
)
|
|
1630
|
+
|
|
1631
|
+
async def _send_failure(self, stream_id: int, request: _ServerRequest, error: MuxRpcFailure) -> None:
|
|
1632
|
+
request.reset_sent = True
|
|
1633
|
+
await self._writer.write_sequenced(
|
|
1634
|
+
FrameType.RESET,
|
|
1635
|
+
stream_id,
|
|
1636
|
+
request.sequences,
|
|
1637
|
+
payload=_error_payload(error.code, str(error), retryable=error.retryable, details=error.details),
|
|
1638
|
+
)
|
|
1639
|
+
|
|
1640
|
+
|
|
1641
|
+
async def serve_mux_rpc(
|
|
1642
|
+
reader: asyncio.StreamReader,
|
|
1643
|
+
writer: asyncio.StreamWriter,
|
|
1644
|
+
handler: MuxRpcHandler,
|
|
1645
|
+
*,
|
|
1646
|
+
max_handlers: int = 256,
|
|
1647
|
+
stream_handler: MuxRpcStreamHandler | None = None,
|
|
1648
|
+
request_closed_handler: MuxRpcRequestClosedHandler | None = None,
|
|
1649
|
+
connection_window: int = DEFAULT_CONNECTION_WINDOW_BYTES,
|
|
1650
|
+
stream_window: int = DEFAULT_STREAM_WINDOW_BYTES,
|
|
1651
|
+
handshake_timeout: float = DEFAULT_HANDSHAKE_TIMEOUT_SECONDS,
|
|
1652
|
+
cleanup_timeout: float = DEFAULT_CLEANUP_TIMEOUT_SECONDS,
|
|
1653
|
+
) -> None:
|
|
1654
|
+
connection = MuxRpcServerConnection(
|
|
1655
|
+
reader,
|
|
1656
|
+
writer,
|
|
1657
|
+
handler,
|
|
1658
|
+
max_handlers=max_handlers,
|
|
1659
|
+
stream_handler=stream_handler,
|
|
1660
|
+
request_closed_handler=request_closed_handler,
|
|
1661
|
+
connection_window=connection_window,
|
|
1662
|
+
stream_window=stream_window,
|
|
1663
|
+
handshake_timeout=handshake_timeout,
|
|
1664
|
+
cleanup_timeout=cleanup_timeout,
|
|
1665
|
+
)
|
|
1666
|
+
await connection.serve()
|
|
1667
|
+
|
|
1668
|
+
|
|
1669
|
+
async def _wait_closed_bounded(writer: asyncio.StreamWriter, timeout: float) -> None:
|
|
1670
|
+
try:
|
|
1671
|
+
async with asyncio.timeout(timeout):
|
|
1672
|
+
await writer.wait_closed()
|
|
1673
|
+
except Exception:
|
|
1674
|
+
return
|
|
1675
|
+
|
|
1676
|
+
|
|
1677
|
+
def _consume_task_result(task: asyncio.Task[None]) -> None:
|
|
1678
|
+
with suppress(BaseException):
|
|
1679
|
+
task.result()
|
|
1680
|
+
|
|
1681
|
+
|
|
1682
|
+
async def _complete_sequenced_write_after_cancellation(
|
|
1683
|
+
writer: _FrameWriter,
|
|
1684
|
+
frame_type: FrameType,
|
|
1685
|
+
stream_id: int,
|
|
1686
|
+
sequences: SequenceTracker,
|
|
1687
|
+
*,
|
|
1688
|
+
flags: int = 0,
|
|
1689
|
+
payload: bytes = b"",
|
|
1690
|
+
) -> None:
|
|
1691
|
+
task = asyncio.create_task(
|
|
1692
|
+
writer.write_sequenced(
|
|
1693
|
+
frame_type,
|
|
1694
|
+
stream_id,
|
|
1695
|
+
sequences,
|
|
1696
|
+
flags=flags,
|
|
1697
|
+
payload=payload,
|
|
1698
|
+
)
|
|
1699
|
+
)
|
|
1700
|
+
while not task.done():
|
|
1701
|
+
try:
|
|
1702
|
+
await asyncio.shield(task)
|
|
1703
|
+
except asyncio.CancelledError:
|
|
1704
|
+
continue
|
|
1705
|
+
with suppress(MuxRpcError):
|
|
1706
|
+
task.result()
|
|
1707
|
+
|
|
1708
|
+
|
|
1709
|
+
async def _write_sequenced_before_honouring_cancellation(
|
|
1710
|
+
writer: _FrameWriter,
|
|
1711
|
+
frame_type: FrameType,
|
|
1712
|
+
stream_id: int,
|
|
1713
|
+
sequences: SequenceTracker,
|
|
1714
|
+
*,
|
|
1715
|
+
flags: int = 0,
|
|
1716
|
+
payload: bytes = b"",
|
|
1717
|
+
) -> None:
|
|
1718
|
+
task = asyncio.create_task(
|
|
1719
|
+
writer.write_sequenced(
|
|
1720
|
+
frame_type,
|
|
1721
|
+
stream_id,
|
|
1722
|
+
sequences,
|
|
1723
|
+
flags=flags,
|
|
1724
|
+
payload=payload,
|
|
1725
|
+
)
|
|
1726
|
+
)
|
|
1727
|
+
cancellation: asyncio.CancelledError | None = None
|
|
1728
|
+
while not task.done():
|
|
1729
|
+
try:
|
|
1730
|
+
await asyncio.shield(task)
|
|
1731
|
+
except asyncio.CancelledError as exc:
|
|
1732
|
+
cancellation = exc
|
|
1733
|
+
task.result()
|
|
1734
|
+
if cancellation is not None:
|
|
1735
|
+
raise cancellation
|
|
1736
|
+
|
|
1737
|
+
|
|
1738
|
+
def _error_payload(
|
|
1739
|
+
code: str,
|
|
1740
|
+
message: str,
|
|
1741
|
+
*,
|
|
1742
|
+
retryable: bool = False,
|
|
1743
|
+
details: dict[str, object] | None = None,
|
|
1744
|
+
) -> bytes:
|
|
1745
|
+
return encode_control_payload(
|
|
1746
|
+
{
|
|
1747
|
+
"code": code,
|
|
1748
|
+
"message": message,
|
|
1749
|
+
"retryable": retryable,
|
|
1750
|
+
"details": details or {},
|
|
1751
|
+
}
|
|
1752
|
+
)
|
|
1753
|
+
|
|
1754
|
+
|
|
1755
|
+
def _remote_error(payload: bytes) -> MuxRemoteError:
|
|
1756
|
+
control = decode_control_payload(payload)
|
|
1757
|
+
code = control.get("code")
|
|
1758
|
+
message = control.get("message")
|
|
1759
|
+
retryable = control.get("retryable", False)
|
|
1760
|
+
details = control.get("details", {})
|
|
1761
|
+
if (
|
|
1762
|
+
not isinstance(code, str)
|
|
1763
|
+
or not code
|
|
1764
|
+
or not isinstance(message, str)
|
|
1765
|
+
or not isinstance(retryable, bool)
|
|
1766
|
+
or not isinstance(details, dict)
|
|
1767
|
+
):
|
|
1768
|
+
raise MuxProtocolError("RESET contains an invalid structured error")
|
|
1769
|
+
return MuxRemoteError(code, message, retryable=retryable, details=details)
|
|
1770
|
+
|
|
1771
|
+
|
|
1772
|
+
def _positive(value: int, name: str) -> None:
|
|
1773
|
+
if not isinstance(value, int) or isinstance(value, bool) or value <= 0:
|
|
1774
|
+
raise ValueError(f"{name} must be a positive integer")
|
|
1775
|
+
|
|
1776
|
+
|
|
1777
|
+
def _finite_positive(value: float, name: str) -> None:
|
|
1778
|
+
if not isinstance(value, int | float) or isinstance(value, bool) or not math.isfinite(value) or value <= 0:
|
|
1779
|
+
raise ValueError(f"{name} must be a finite positive number")
|