diff --git a/solstone/convey/secure_listener/framing.py b/solstone/convey/secure_listener/framing.py index 40cc309fe..c2d9d6fb5 100644 --- a/solstone/convey/secure_listener/framing.py +++ b/solstone/convey/secure_listener/framing.py @@ -25,10 +25,10 @@ from __future__ import annotations from dataclasses import dataclass, field from typing import Final -# Flag bits — each frame must carry exactly one of OPEN / DATA / CLOSE / -# RESET / WINDOW / PING / PONG, except OPEN|DATA (open with initial bytes) -# and DATA|CLOSE (last data + half-close). PING and PONG ride on stream_id -# 0 only (control channel) and carry an 8-byte nonce. +# Flag bits — each received frame must carry one of the canonical SPL flag +# sets: OPEN, DATA, OPEN|DATA, CLOSE, OPEN|CLOSE, DATA|CLOSE, +# OPEN|DATA|CLOSE, RESET, WINDOW, PING, or PONG. PING and PONG ride on +# stream_id 0 only (control channel) and carry an 8-byte nonce. FLAG_OPEN: Final[int] = 0x01 FLAG_DATA: Final[int] = 0x02 FLAG_CLOSE: Final[int] = 0x04 @@ -50,6 +50,7 @@ RESET_UNSPECIFIED: Final[int] = 0xFF HEADER_LEN: Final[int] = 8 MAX_PAYLOAD: Final[int] = (1 << 24) - 1 # 16 MiB - 1 INITIAL_WINDOW: Final[int] = 1 << 20 # 1 MiB +MAX_SEND_CREDIT: Final[int] = (1 << 31) - 1 MAX_CONCURRENT_STREAMS: Final[int] = 256 RECOMMENDED_CHUNK: Final[int] = 64 * 1024 CONTROL_NONCE_LEN: Final[int] = 8 @@ -195,7 +196,9 @@ def validate_flags(flags: int) -> None: FLAG_PING, FLAG_PONG, FLAG_OPEN | FLAG_DATA, + FLAG_OPEN | FLAG_CLOSE, FLAG_DATA | FLAG_CLOSE, + FLAG_OPEN | FLAG_DATA | FLAG_CLOSE, } if exclusive not in allowed: raise ProtocolError(f"illegal flag combination: {flags:#x}") diff --git a/solstone/convey/secure_listener/mux.py b/solstone/convey/secure_listener/mux.py index 46c3454dc..248c742b8 100644 --- a/solstone/convey/secure_listener/mux.py +++ b/solstone/convey/secure_listener/mux.py @@ -33,6 +33,7 @@ from .framing import ( FLAG_WINDOW, INITIAL_WINDOW, MAX_CONCURRENT_STREAMS, + MAX_SEND_CREDIT, RECOMMENDED_CHUNK, RESET_CANCEL, RESET_FLOW_CONTROL_ERROR, @@ -52,6 +53,7 @@ from .framing import ( parse_control_nonce, parse_reset_reason, parse_window_credit, + validate_flags, ) if TYPE_CHECKING: @@ -67,6 +69,9 @@ RESET_CTX_UNKNOWN_STREAM: Final[str] = "unknown_stream" RESET_CTX_OVER_CREDIT_DATA: Final[str] = "over_credit_data" RESET_CTX_OVER_CREDIT_OPEN: Final[str] = "over_credit_open" RESET_CTX_BAD_WINDOW_FRAME: Final[str] = "bad_window_frame" +RESET_CTX_WINDOW_OVERFLOW: Final[str] = "window_overflow" +RESET_CTX_INVALID_FLAGS: Final[str] = "invalid_flag_combination" +RESET_CTX_MISPLACED_CONTROL: Final[str] = "misplaced_control_frame" RESET_CTX_HANDLER_EXCEPTION: Final[str] = "handler_exception" RESET_CTX_NO_IDENTITY: Final[str] = "no_identity" RESET_CTX_APP_CANCELLATION: Final[str] = "app_cancellation" @@ -93,6 +98,10 @@ class ResetDiagnostic: ResetDiag = Callable[[ResetDiagnostic], None] +class TunnelFatalError(Exception): + """Raised when a framing violation terminates the whole tunnel.""" + + @dataclass class _StreamState: stream_id: int @@ -173,19 +182,17 @@ class Multiplexer: self._on_reset = on_reset self._streams: dict[int, _StreamState] = {} self._closed = False + self._next_local_id = 2 if is_listener else 1 async def feed(self, plaintext: bytes) -> None: - if not plaintext: + if self._closed or not plaintext: return self._decoder.feed(plaintext) while True: try: frame = self._decoder.next() except ProtocolError: - await self._reset_all( - RESET_PROTOCOL_ERROR, - RESET_CTX_MALFORMED_FRAME, - ) + await self._tunnel_fatal(RESET_CTX_MALFORMED_FRAME) return if frame is None: return @@ -207,7 +214,12 @@ class Multiplexer: await self._dispatch_control(frame) return if frame.flags & (FLAG_PING | FLAG_PONG): - await self._reset_all(RESET_PROTOCOL_ERROR, RESET_CTX_MALFORMED_FRAME) + await self._reject_stream(frame, RESET_CTX_MISPLACED_CONTROL) + return + try: + validate_flags(frame.flags) + except ProtocolError: + await self._reject_stream(frame, RESET_CTX_INVALID_FLAGS) return if frame.flags & FLAG_OPEN: @@ -320,6 +332,14 @@ class Multiplexer: ) self._terminate(state) return + if state.send_credit + credit > MAX_SEND_CREDIT: + await self._emit_reset( + frame.stream_id, + RESET_FLOW_CONTROL_ERROR, + RESET_CTX_WINDOW_OVERFLOW, + ) + self._terminate(state) + return state.send_credit += credit state.credit_event.set() if frame.flags & FLAG_RESET: @@ -334,15 +354,15 @@ class Multiplexer: is_ping = bool(frame.flags & FLAG_PING) is_pong = bool(frame.flags & FLAG_PONG) if is_ping == is_pong: - await self._reset_all(RESET_PROTOCOL_ERROR, RESET_CTX_MALFORMED_FRAME) + await self._tunnel_fatal(RESET_CTX_MALFORMED_FRAME) return if frame.flags & ~(FLAG_PING | FLAG_PONG): - await self._reset_all(RESET_PROTOCOL_ERROR, RESET_CTX_MALFORMED_FRAME) + await self._tunnel_fatal(RESET_CTX_MALFORMED_FRAME) return try: nonce = parse_control_nonce(frame) except ProtocolError: - await self._reset_all(RESET_PROTOCOL_ERROR, RESET_CTX_MALFORMED_FRAME) + await self._tunnel_fatal(RESET_CTX_MALFORMED_FRAME) return if is_ping: await self._emit(build_pong(nonce)) @@ -418,14 +438,18 @@ class Multiplexer: await self._emit(build_reset(stream_id, reason)) self._fire_diag(stream_id, reason, context) - async def _reset_all(self, reason: int, context: str) -> None: - count = 0 - for state in list(self._streams.values()): - count += 1 - await self._emit_reset(state.stream_id, reason, context) + async def _reject_stream(self, frame: Frame, context: str) -> None: + state = self._streams.get(frame.stream_id) + await self._emit_reset(frame.stream_id, RESET_PROTOCOL_ERROR, context) + if state is not None: self._terminate(state) - if count == 0: - self._fire_diag(0, reason, context) + + async def _tunnel_fatal(self, context: str) -> None: + if self._closed: + return + self._fire_diag(0, RESET_PROTOCOL_ERROR, context) + await self.close() + raise TunnelFatalError(context) async def open_stream( self, @@ -444,10 +468,8 @@ class Multiplexer: return reader, writer def _next_local_stream_id(self) -> int: - start = 2 if self._is_listener else 1 - cur = start - while cur in self._streams: - cur += 2 - if cur > 0xFFFFFFFF: - raise RuntimeError("stream_id space exhausted") - return cur + if self._next_local_id > 0xFFFFFFFF: + raise RuntimeError("stream_id space exhausted") + stream_id = self._next_local_id + self._next_local_id += 2 + return stream_id diff --git a/solstone/think/link/client.py b/solstone/think/link/client.py index 6dc11afcc..d7f2eb1fc 100644 --- a/solstone/think/link/client.py +++ b/solstone/think/link/client.py @@ -38,6 +38,7 @@ from solstone.convey.secure_listener.framing import ( INITIAL_WINDOW, MAX_CONCURRENT_STREAMS, MAX_PAYLOAD, + MAX_SEND_CREDIT, RECOMMENDED_CHUNK, RESET_CANCEL, RESET_FLOW_CONTROL_ERROR, @@ -56,6 +57,7 @@ from solstone.convey.secure_listener.framing import ( parse_control_nonce, parse_reset_reason, parse_window_credit, + validate_flags, ) from solstone.think.link.tls import TlsError as _TlsError @@ -296,6 +298,8 @@ class _DialerMultiplexer: if frame is None: return await self._dispatch(frame) + if self._closed: + return def close(self) -> None: if self._closed: @@ -311,7 +315,12 @@ class _DialerMultiplexer: await self._dispatch_control(frame) return if frame.flags & (FLAG_PING | FLAG_PONG): - self.close() + await self._reject_stream(frame) + return + try: + validate_flags(frame.flags) + except ProtocolError: + await self._reject_stream(frame) return if frame.flags & FLAG_OPEN: @@ -320,7 +329,8 @@ class _DialerMultiplexer: state = self._streams.get(frame.stream_id) if state is None: - await self._emit(build_reset(frame.stream_id, RESET_PROTOCOL_ERROR)) + if frame.flags & (FLAG_DATA | FLAG_WINDOW): + await self._emit(build_reset(frame.stream_id, RESET_PROTOCOL_ERROR)) return if frame.flags & FLAG_DATA: @@ -355,6 +365,11 @@ class _DialerMultiplexer: state.reset_reason = RESET_PROTOCOL_ERROR self._close_stream(state, forget=True) return + if state.send_credit + credit > MAX_SEND_CREDIT: + await self._emit(build_reset(frame.stream_id, RESET_FLOW_CONTROL_ERROR)) + state.reset_reason = RESET_FLOW_CONTROL_ERROR + self._close_stream(state, forget=True) + return state.send_credit += credit state.credit_event.set() @@ -382,6 +397,13 @@ class _DialerMultiplexer: if is_ping: await self._emit(build_pong(nonce)) + async def _reject_stream(self, frame: Frame) -> None: + state = self._streams.get(frame.stream_id) + await self._emit(build_reset(frame.stream_id, RESET_PROTOCOL_ERROR)) + if state is not None: + state.reset_reason = RESET_PROTOCOL_ERROR + self._close_stream(state, forget=True) + async def _emit(self, frame: Frame) -> None: if self._closed: return @@ -636,6 +658,8 @@ class TunnelSession: await self._transport.send(outbound) if plaintext: await self._mux.feed(plaintext) + if self._mux._closed: + return finally: await self._cancel_keepalive() self._mux.close() diff --git a/tests/link/test_client_session.py b/tests/link/test_client_session.py index d631b0958..b97cb21bf 100644 --- a/tests/link/test_client_session.py +++ b/tests/link/test_client_session.py @@ -18,10 +18,12 @@ from solstone.convey.secure_listener.framing import ( FLAG_WINDOW, INITIAL_WINDOW, MAX_PAYLOAD, + MAX_SEND_CREDIT, RECOMMENDED_CHUNK, RESET_CANCEL, RESET_FLOW_CONTROL_ERROR, RESET_INTERNAL_ERROR, + RESET_PROTOCOL_ERROR, Frame, FrameDecoder, build_close, @@ -29,6 +31,7 @@ from solstone.convey.secure_listener.framing import ( build_ping, build_pong, build_reset, + build_window, parse_reset_reason, parse_window_credit, ) @@ -212,15 +215,259 @@ async def test_dialer_mux_pong_records_liveness_without_emit() -> None: @pytest.mark.asyncio async def test_dialer_mux_malformed_control_frame_closes_without_raising() -> None: sent: list[bytes] = [] + dispatched: list[Frame] = [] async def send(data: bytes) -> None: sent.append(data) mux = client._DialerMultiplexer(send) - await mux.feed(Frame(0, FLAG_PING | FLAG_PONG, b"abcdefgh").encode()) + original_dispatch = mux._dispatch + + async def dispatch_spy(frame: Frame) -> None: + dispatched.append(frame) + await original_dispatch(frame) + + mux._dispatch = dispatch_spy + await mux.feed( + Frame(0, FLAG_PING | FLAG_PONG, b"abcdefgh").encode() + + build_ping(b"12345678").encode() + ) assert mux._closed is True assert sent == [] + assert [frame.flags for frame in dispatched] == [FLAG_PING | FLAG_PONG] + + +@pytest.mark.asyncio +async def test_dialer_window_credit_exact_cap_is_accepted() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + stream = await mux.open_stream() + + await mux.feed(build_window(stream.id, MAX_SEND_CREDIT - INITIAL_WINDOW).encode()) + + assert stream._state.send_credit == MAX_SEND_CREDIT + assert _reset_frames(sent, stream.id) == [] + + +@pytest.mark.asyncio +async def test_dialer_window_credit_overflow_resets_and_forgets() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + stream = await mux.open_stream() + + await mux.feed( + build_window(stream.id, MAX_SEND_CREDIT - INITIAL_WINDOW + 1).encode() + ) + + resets = _reset_frames(sent, stream.id) + assert len(resets) == 1 + assert parse_reset_reason(resets[0]) == RESET_FLOW_CONTROL_ERROR + assert stream._state.reset_reason == RESET_FLOW_CONTROL_ERROR + assert stream.id not in mux._streams + + +@pytest.mark.asyncio +async def test_dialer_window_credit_uses_remaining_credit_accounting() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + stream = await mux.open_stream() + await stream.write(b"x") + assert stream._state.send_credit == INITIAL_WINDOW - 1 + + await mux.feed( + build_window(stream.id, MAX_SEND_CREDIT - INITIAL_WINDOW + 1).encode() + ) + + assert stream._state.send_credit == MAX_SEND_CREDIT + assert _reset_frames(sent, stream.id) == [] + + no_consume_sent: list[bytes] = [] + + async def no_consume_send(data: bytes) -> None: + no_consume_sent.append(data) + + no_consume_mux = client._DialerMultiplexer(no_consume_send) + no_consume_stream = await no_consume_mux.open_stream() + await no_consume_mux.feed( + build_window( + no_consume_stream.id, + MAX_SEND_CREDIT - INITIAL_WINDOW + 1, + ).encode() + ) + no_consume_resets = _reset_frames(no_consume_sent, no_consume_stream.id) + assert len(no_consume_resets) == 1 + assert parse_reset_reason(no_consume_resets[0]) == RESET_FLOW_CONTROL_ERROR + + accum_sent: list[bytes] = [] + + async def accum_send(data: bytes) -> None: + accum_sent.append(data) + + accum_mux = client._DialerMultiplexer(accum_send) + accum_stream = await accum_mux.open_stream() + await accum_mux.feed(build_window(accum_stream.id, 0x40000000).encode()) + assert accum_stream._state.send_credit == 1_074_790_400 + assert _reset_frames(accum_sent, accum_stream.id) == [] + await accum_mux.feed(build_window(accum_stream.id, 0x40000000).encode()) + accum_resets = _reset_frames(accum_sent, accum_stream.id) + assert len(accum_resets) == 1 + assert parse_reset_reason(accum_resets[0]) == RESET_FLOW_CONTROL_ERROR + + single_sent: list[bytes] = [] + + async def single_send(data: bytes) -> None: + single_sent.append(data) + + single_mux = client._DialerMultiplexer(single_send) + single_stream = await single_mux.open_stream() + await single_mux.feed(build_window(single_stream.id, 0xFFFFFFFF).encode()) + single_resets = _reset_frames(single_sent, single_stream.id) + assert len(single_resets) == 1 + assert parse_reset_reason(single_resets[0]) == RESET_FLOW_CONTROL_ERROR + + +@pytest.mark.asyncio +async def test_dialer_invalid_flags_on_known_stream_reset_and_forget() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + stream = await mux.open_stream() + + await mux.feed(Frame(stream.id, FLAG_DATA | FLAG_WINDOW, b"x").encode()) + + resets = _reset_frames(sent, stream.id) + assert len(resets) == 1 + assert parse_reset_reason(resets[0]) == RESET_PROTOCOL_ERROR + assert stream._state.buffered == [] + assert stream._state.reset_reason == RESET_PROTOCOL_ERROR + assert stream.id not in mux._streams + + +@pytest.mark.asyncio +async def test_dialer_invalid_open_flags_close_existing_stream_before_open_policy() -> ( + None +): + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + stream = await mux.open_stream() + + await mux.feed(Frame(stream.id, FLAG_OPEN | FLAG_WINDOW, b"").encode()) + + resets = _reset_frames(sent, stream.id) + assert len(resets) == 1 + assert parse_reset_reason(resets[0]) == RESET_PROTOCOL_ERROR + assert stream._state.reset_reason == RESET_PROTOCOL_ERROR + assert stream.id not in mux._streams + + +@pytest.mark.asyncio +async def test_dialer_misplaced_control_on_unknown_stream_resets() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + + await mux.feed(Frame(99, FLAG_PING, b"\x00" * 8).encode()) + + resets = _reset_frames(sent, 99) + assert len(resets) == 1 + assert parse_reset_reason(resets[0]) == RESET_PROTOCOL_ERROR + assert 99 not in mux._streams + assert mux._closed is False + + +@pytest.mark.asyncio +async def test_dialer_misplaced_control_on_known_stream_beats_invalid_data() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + stream = await mux.open_stream() + + await mux.feed(Frame(stream.id, FLAG_PING | FLAG_DATA, b"x").encode()) + + resets = _reset_frames(sent, stream.id) + assert len(resets) == 1 + assert parse_reset_reason(resets[0]) == RESET_PROTOCOL_ERROR + assert stream._state.buffered == [] + assert stream._state.reset_reason == RESET_PROTOCOL_ERROR + assert stream.id not in mux._streams + + +@pytest.mark.asyncio +async def test_dialer_unknown_stream_close_is_ignored() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + + await mux.feed(build_close(99).encode()) + + assert sent == [] + assert mux._closed is False + + +@pytest.mark.asyncio +async def test_dialer_unknown_stream_reset_is_ignored() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + + await mux.feed(build_reset(99, RESET_CANCEL).encode()) + + assert sent == [] + assert mux._closed is False + + +@pytest.mark.asyncio +async def test_dialer_unknown_stream_data_and_window_get_reset() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + mux = client._DialerMultiplexer(send) + + await mux.feed(build_data(99, b"x").encode()) + await mux.feed(build_window(101, 1).encode()) + + data_resets = _reset_frames(sent, 99) + window_resets = _reset_frames(sent, 101) + assert len(data_resets) == 1 + assert len(window_resets) == 1 + assert parse_reset_reason(data_resets[0]) == RESET_PROTOCOL_ERROR + assert parse_reset_reason(window_resets[0]) == RESET_PROTOCOL_ERROR + assert 99 not in mux._streams + assert 101 not in mux._streams @pytest.mark.asyncio @@ -272,6 +519,33 @@ async def test_tunnel_session_pongs_keep_session_alive_and_requests_work( await session.close() +@pytest.mark.asyncio +async def test_tunnel_session_mux_fatal_closes_session_without_keepalive_wait( + pass_through_tls: None, +) -> None: + transport = FakeTransport() + session = _session( + transport, + keepalive_interval=60.0, + keepalive_timeout=60.0, + ) + try: + transport.inbound.put_nowait( + Frame(0, FLAG_PING | FLAG_PONG, b"abcdefgh").encode() + ) + + await _wait_for( + lambda: session._closed.is_set() and session._reader_task.done(), + timeout=1.0, + ) + + assert session._mux._closed is True + assert transport.closed is True + assert session._reader_task.exception() is None + finally: + await session.close() + + @pytest.mark.asyncio async def test_tunnel_request_uses_head_only_open_for_body_over_max_payload( pass_through_tls: None, diff --git a/tests/link/test_framing.py b/tests/link/test_framing.py index bb34e3b18..9e967d15f 100644 --- a/tests/link/test_framing.py +++ b/tests/link/test_framing.py @@ -18,6 +18,7 @@ from solstone.convey.secure_listener.framing import ( FLAG_RESET, FLAG_WINDOW, HEADER_LEN, + MAX_SEND_CREDIT, RESET_PROTOCOL_ERROR, Frame, FrameDecoder, @@ -127,12 +128,18 @@ def test_reset_reason_parse_roundtrip() -> None: def test_validate_flags_allows_only_legal_combos() -> None: - for flag in (FLAG_OPEN, FLAG_DATA, FLAG_CLOSE, FLAG_RESET, FLAG_WINDOW): + for flag in ( + FLAG_OPEN, + FLAG_DATA, + FLAG_CLOSE, + FLAG_RESET, + FLAG_WINDOW, + FLAG_OPEN | FLAG_DATA, + FLAG_OPEN | FLAG_CLOSE, + FLAG_DATA | FLAG_CLOSE, + FLAG_OPEN | FLAG_DATA | FLAG_CLOSE, + ): validate_flags(flag) - validate_flags(FLAG_OPEN | FLAG_DATA) - validate_flags(FLAG_DATA | FLAG_CLOSE) - with pytest.raises(ProtocolError): - validate_flags(FLAG_OPEN | FLAG_CLOSE) with pytest.raises(ProtocolError): validate_flags(FLAG_DATA | FLAG_WINDOW) @@ -225,6 +232,10 @@ def test_control_nonce_len_is_eight() -> None: assert CONTROL_NONCE_LEN == 8 +def test_max_send_credit_is_signed_31_bit_cap() -> None: + assert MAX_SEND_CREDIT == (1 << 31) - 1 + + def test_reserved_mask_collapsed_to_bit_seven() -> None: # bits 5 and 6 are PING/PONG; only bit 7 (0x80) remains reserved. assert FLAG_RESERVED_MASK == 0x80 diff --git a/tests/link/test_mux.py b/tests/link/test_mux.py index 7d0bd251d..9548adbe1 100644 --- a/tests/link/test_mux.py +++ b/tests/link/test_mux.py @@ -13,12 +13,14 @@ import pytest from solstone.convey.secure_listener.framing import ( FLAG_CLOSE, FLAG_DATA, + FLAG_OPEN, FLAG_PING, FLAG_PONG, FLAG_RESET, FLAG_WINDOW, INITIAL_WINDOW, MAX_CONCURRENT_STREAMS, + MAX_SEND_CREDIT, RESET_CANCEL, RESET_FLOW_CONTROL_ERROR, RESET_INTERNAL_ERROR, @@ -41,15 +43,19 @@ from solstone.convey.secure_listener.mux import ( RESET_CTX_BODY_DISCARD_CANCELLATION, RESET_CTX_DUPLICATE_OPEN, RESET_CTX_HANDLER_EXCEPTION, + RESET_CTX_INVALID_FLAGS, RESET_CTX_MALFORMED_FRAME, + RESET_CTX_MISPLACED_CONTROL, RESET_CTX_NO_IDENTITY, RESET_CTX_OVER_CREDIT_DATA, RESET_CTX_OVER_CREDIT_OPEN, RESET_CTX_PARITY_VIOLATION, RESET_CTX_STREAM_CAP_OVERFLOW, RESET_CTX_UNKNOWN_STREAM, + RESET_CTX_WINDOW_OVERFLOW, Multiplexer, ResetDiagnostic, + TunnelFatalError, ) from solstone.convey.secure_listener.wsgi import dispatch_stream from solstone.think.link.client import _http_head_bytes @@ -105,6 +111,24 @@ def _assert_single_reset( assert _reset_reasons(frames, stream_id) == [reason] +def _assert_single_diag( + diags: list[ResetDiagnostic], + stream_id: int, + reason: int, + context: str, +) -> None: + assert diags == [ + ResetDiagnostic( + stream_id=stream_id, + reason_code=reason, + reason_name="protocol_error" + if reason == RESET_PROTOCOL_ERROR + else "flow_control_error", + context=context, + ) + ] + + @pytest.mark.asyncio async def test_open_with_initial_payload_hits_handler() -> None: handler_seen: dict[int, bytes] = {} @@ -324,6 +348,187 @@ async def test_unknown_stream_window_gets_reset() -> None: await mux.close() +@pytest.mark.asyncio +async def test_listener_window_credit_exact_cap_is_accepted() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + await asyncio.Event().wait() + + mux = Multiplexer(send, handler, is_listener=True) + try: + await mux.feed(build_open(1).encode()) + state = mux._streams[1] + + await mux.feed(build_window(1, MAX_SEND_CREDIT - state.send_credit).encode()) + + assert state.send_credit == MAX_SEND_CREDIT + assert _reset_reasons(_decode_frames(sent), 1) == [] + assert 1 in mux._streams + finally: + await mux.close() + + +@pytest.mark.asyncio +async def test_listener_window_credit_overflow_resets_and_terminates() -> None: + sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + await asyncio.Event().wait() + + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) + try: + await mux.feed(build_open(1).encode()) + state = mux._streams[1] + + await mux.feed( + build_window(1, MAX_SEND_CREDIT - state.send_credit + 1).encode() + ) + + frames = _decode_frames(sent) + _assert_single_reset(frames, 1, RESET_FLOW_CONTROL_ERROR) + _assert_single_diag( + diags, + 1, + RESET_FLOW_CONTROL_ERROR, + RESET_CTX_WINDOW_OVERFLOW, + ) + assert 1 not in mux._streams + finally: + await mux.close() + + +@pytest.mark.asyncio +async def test_listener_invalid_flags_on_unknown_stream_reset_without_state() -> None: + sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + return + + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) + try: + await mux.feed(Frame(99, FLAG_DATA | FLAG_WINDOW, b"x").encode()) + + frames = _decode_frames(sent) + _assert_single_reset(frames, 99, RESET_PROTOCOL_ERROR) + _assert_single_diag( + diags, + 99, + RESET_PROTOCOL_ERROR, + RESET_CTX_INVALID_FLAGS, + ) + assert 99 not in mux._streams + finally: + await mux.close() + + +@pytest.mark.asyncio +async def test_listener_invalid_open_flags_reject_before_opening() -> None: + sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] + handler_invoked = False + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + nonlocal handler_invoked + handler_invoked = True + + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) + try: + await mux.feed(Frame(1, FLAG_OPEN | FLAG_WINDOW, b"").encode()) + + frames = _decode_frames(sent) + _assert_single_reset(frames, 1, RESET_PROTOCOL_ERROR) + _assert_single_diag( + diags, + 1, + RESET_PROTOCOL_ERROR, + RESET_CTX_INVALID_FLAGS, + ) + assert handler_invoked is False + assert 1 not in mux._streams + finally: + await mux.close() + + +@pytest.mark.asyncio +async def test_listener_invalid_flags_on_known_stream_terminate_without_payload() -> ( + None +): + sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + await asyncio.Event().wait() + + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) + try: + await mux.feed(build_open(1).encode()) + state = mux._streams[1] + + await mux.feed(Frame(1, FLAG_DATA | FLAG_WINDOW, b"x").encode()) + + frames = _decode_frames(sent) + _assert_single_reset(frames, 1, RESET_PROTOCOL_ERROR) + assert bytes(state.reader._buffer) == b"" + assert diags[-1] == ResetDiagnostic( + stream_id=1, + reason_code=RESET_PROTOCOL_ERROR, + reason_name="protocol_error", + context=RESET_CTX_INVALID_FLAGS, + ) + assert 1 not in mux._streams + finally: + await mux.close() + + +@pytest.mark.asyncio +async def test_listener_misplaced_pong_on_known_stream_resets_and_terminates() -> None: + sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + await asyncio.Event().wait() + + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) + try: + await mux.feed(build_open(1).encode()) + + await mux.feed(Frame(1, FLAG_PONG, b"\x00" * 8).encode()) + + frames = _decode_frames(sent) + _assert_single_reset(frames, 1, RESET_PROTOCOL_ERROR) + assert diags[-1] == ResetDiagnostic( + stream_id=1, + reason_code=RESET_PROTOCOL_ERROR, + reason_name="protocol_error", + context=RESET_CTX_MISPLACED_CONTROL, + ) + assert 1 not in mux._streams + assert mux._closed is False + finally: + await mux.close() + + @pytest.mark.asyncio async def test_concurrent_streams_do_not_interfere() -> None: responses: dict[int, bytes] = {} @@ -390,6 +595,49 @@ async def test_validates_open_reopen_is_protocol_error() -> None: await mux.close() +@pytest.mark.asyncio +async def test_listener_local_stream_ids_do_not_recycle_after_forget() -> None: + sent: list[bytes] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + await asyncio.Event().wait() + + mux = Multiplexer(send, handler, is_listener=True) + try: + _reader1, writer1 = await mux.open_stream() + assert writer1.stream_id == 2 + await writer1.reset(RESET_CANCEL, RESET_CTX_APP_CANCELLATION) + + _reader2, writer2 = await mux.open_stream() + + assert writer2.stream_id == 4 + assert [ + frame.stream_id for frame in _decode_frames(sent) if frame.flags & FLAG_OPEN + ] == [ + 2, + 4, + ] + finally: + await mux.close() + + +def test_listener_local_stream_id_exhaustion_still_raises() -> None: + async def send(_: bytes) -> None: + return + + async def handler(*_: object) -> None: + return + + mux = Multiplexer(send, handler, is_listener=True) + mux._next_local_id = 0x1_0000_0000 + + with pytest.raises(RuntimeError, match="stream_id space exhausted"): + mux._next_local_stream_id() + + # ---- streamID==0 PING/PONG keepalive responder ------------------------------ @@ -452,9 +700,109 @@ async def test_unsolicited_pong_is_silently_dropped() -> None: await mux.close() +@pytest.mark.parametrize( + "frame", + [ + Frame(0, FLAG_PING | FLAG_PONG, b"\x00" * 8), + Frame(0, FLAG_PING | FLAG_DATA, b"\x00" * 8), + ], +) +@pytest.mark.asyncio +async def test_stream_zero_malformed_control_raises_tunnel_fatal_once( + frame: Frame, +) -> None: + sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + pytest.fail("handler should not be invoked for control frames") + + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) + with pytest.raises(TunnelFatalError, match=RESET_CTX_MALFORMED_FRAME): + await mux.feed(frame.encode()) + + assert sent == [] + _assert_single_diag( + diags, + 0, + RESET_PROTOCOL_ERROR, + RESET_CTX_MALFORMED_FRAME, + ) + assert mux._closed is True + + +@pytest.mark.asyncio +async def test_decoder_corrupt_frame_fatal_tears_down_streams_without_diag_storm() -> ( + None +): + sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + await asyncio.Event().wait() + + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) + await mux.feed(build_open(1).encode() + build_open(3).encode()) + tasks = [state.task for state in mux._streams.values()] + assert all(task is not None for task in tasks) + corrupt = bytearray(build_data(1, b"x").encode()) + corrupt[4] |= 0x80 + + with pytest.raises(TunnelFatalError, match=RESET_CTX_MALFORMED_FRAME): + await mux.feed(bytes(corrupt)) + + for _ in range(20): + await asyncio.sleep(0) + if all(task.done() for task in tasks if task is not None): + break + + assert mux._closed is True + assert mux._streams == {} + assert all(task.done() for task in tasks if task is not None) + _assert_single_diag( + diags, + 0, + RESET_PROTOCOL_ERROR, + RESET_CTX_MALFORMED_FRAME, + ) + + +@pytest.mark.asyncio +async def test_second_feed_after_tunnel_fatal_is_inert() -> None: + sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] + + async def send(data: bytes) -> None: + sent.append(data) + + async def handler(*_: object) -> None: + return + + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) + # accept.py awaits reader_task at line 379 and logs this raised fatal in the + # generic Exception handler at line 187; direct mux callers should see it. + with pytest.raises(TunnelFatalError, match=RESET_CTX_MALFORMED_FRAME): + await mux.feed(Frame(0, FLAG_PING, b"short").encode()) + sent_count = len(sent) + diag_count = len(diags) + + await mux.feed(build_ping(b"12345678").encode()) + + assert len(sent) == sent_count + assert len(diags) == diag_count + assert mux._closed is True + + @pytest.mark.asyncio async def test_ping_on_nonzero_stream_is_protocol_error() -> None: sent: list[bytes] = [] + diags: list[ResetDiagnostic] = [] async def send(data: bytes) -> None: sent.append(data) @@ -462,16 +810,22 @@ async def test_ping_on_nonzero_stream_is_protocol_error() -> None: async def handler(*_: object) -> None: return - mux = Multiplexer(send, handler, is_listener=True) + mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) # PING on stream 5 (illegal — control frames are streamID==0 only). illegal = Frame(stream_id=5, flags=FLAG_PING, payload=b"\x00" * 8).encode() await mux.feed(illegal) - # Behavior parity with other top-level protocol errors: a RESET stamps the - # tunnel as broken; we don't have streams to reset here, so the side effect - # is internal teardown. The wire effect is no PONG emission. frames = _decode_frames(sent) + _assert_single_reset(frames, 5, RESET_PROTOCOL_ERROR) + _assert_single_diag( + diags, + 5, + RESET_PROTOCOL_ERROR, + RESET_CTX_MISPLACED_CONTROL, + ) assert not any(f.flags & FLAG_PONG for f in frames) + assert 5 not in mux._streams + assert mux._closed is False await mux.close() @@ -886,16 +1240,32 @@ async def test_reset_diagnostics_distinguish_contexts_and_are_privacy_clean( await collect_with_handler(raising_handler, build_open(1).encode()) - malformed_mux = Multiplexer( + misplaced_mux = Multiplexer( send, waiting_handler, is_listener=True, on_reset=diags.append, ) try: - await malformed_mux.feed( + await misplaced_mux.feed( Frame(stream_id=5, flags=FLAG_PING, payload=b"\x00" * 8).encode() ) + finally: + await misplaced_mux.close() + + malformed_mux = Multiplexer( + send, + waiting_handler, + is_listener=True, + on_reset=diags.append, + ) + try: + with pytest.raises(TunnelFatalError): + await malformed_mux.feed( + Frame( + stream_id=0, flags=FLAG_PING | FLAG_PONG, payload=b"\x00" * 8 + ).encode() + ) finally: await malformed_mux.close() @@ -973,6 +1343,7 @@ async def test_reset_diagnostics_distinguish_contexts_and_are_privacy_clean( RESET_CTX_NO_IDENTITY, RESET_CTX_HANDLER_EXCEPTION, RESET_CTX_MALFORMED_FRAME, + RESET_CTX_MISPLACED_CONTROL, RESET_CTX_BODY_DISCARD_CANCELLATION, RESET_CTX_APP_CANCELLATION, } <= contexts