From e9729b71d2c5bfaa31e5aae77121cba6977c47c9 Mon Sep 17 00:00:00 2001 From: RobLe3 Date: Fri, 28 Aug 2026 15:24:53 +0200 Subject: [PATCH] fix: bound native frame parsing before allocation --- src/iicp_client/iicp_tcp.py | 105 ++++++++++++--------- src/iicp_client/relay_session.py | 43 ++++----- src/iicp_client/relay_worker_client.py | 9 +- tests/fixtures/native-framing-v1.json | 114 +++++++++++++++++++---- tests/test_iicp_tcp.py | 123 +++++++++++++++++++++++-- tests/test_native_framing_fixture.py | 16 +++- tests/test_serve_multiplex.py | 1 + 7 files changed, 313 insertions(+), 98 deletions(-) diff --git a/src/iicp_client/iicp_tcp.py b/src/iicp_client/iicp_tcp.py index 1ed185d..5403abc 100644 --- a/src/iicp_client/iicp_tcp.py +++ b/src/iicp_client/iicp_tcp.py @@ -48,7 +48,7 @@ _HEADER_STRUCT = struct.Struct("!4sBBBBI") _READ_CHUNK = 4096 -_MAX_PAYLOAD = 16 * 1024 * 1024 # 16 MiB +MAX_FRAME_PAYLOAD = 16 * 1024 * 1024 # Length-field payload bytes; header excluded. class MsgType(IntEnum): @@ -83,6 +83,10 @@ class IicpFrame: payload: bytes def encode(self) -> bytes: + if self.version != FRAMING_VERSION: + raise ValueError(f"Unsupported IICP framing version: {self.version}; expected {FRAMING_VERSION}") + if len(self.payload) > MAX_FRAME_PAYLOAD: + raise ValueError(f"IICP frame payload too large: {len(self.payload)} > {MAX_FRAME_PAYLOAD}") header = _HEADER_STRUCT.pack( IICP_MAGIC, self.version, @@ -100,6 +104,10 @@ def decode(cls, data: bytes) -> tuple[IicpFrame, int]: magic, version, msg_type, flags, _res, payload_len = _HEADER_STRUCT.unpack_from(data) if magic != IICP_MAGIC: raise ValueError(f"Invalid IICP magic: {magic!r}") + if version != FRAMING_VERSION: + raise ValueError(f"Unsupported IICP framing version: {version}; expected {FRAMING_VERSION}") + if payload_len > MAX_FRAME_PAYLOAD: + raise ValueError(f"IICP frame payload too large: {payload_len} > {MAX_FRAME_PAYLOAD}") total = FRAME_HEADER_LEN + payload_len if len(data) < total: raise ValueError(f"IICP payload truncated: need {total}, have {len(data)}") @@ -120,8 +128,7 @@ def _cbor2() -> Any: import cbor2 # type: ignore[import-untyped] except ImportError as exc: raise ImportError( - "cbor2 is required for the native IICP transport. " - "Install with: pip install 'iicp-client[iicp-tcp]'" + "cbor2 is required for the native IICP transport. Install with: pip install 'iicp-client[iicp-tcp]'" ) from exc return cbor2 @@ -327,9 +334,7 @@ def __init__( async def start(self) -> None: # Validate cbor2 is importable before opening the socket so we fail fast. _cbor2() - self._server = await asyncio.start_server( - self._handle_connection, host=self.host, port=self.port - ) + self._server = await asyncio.start_server(self._handle_connection, host=self.host, port=self.port) logger.info("IICP TCP server listening on %s:%d", self.host, self.port) async def stop(self) -> None: @@ -348,9 +353,7 @@ async def serve_forever(self) -> None: # ── connection handling ────────────────────────────────────────────────── - async def _handle_connection( - self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter - ) -> None: + async def _handle_connection(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: peer = writer.get_extra_info("peername") logger.debug("IICP TCP connection from %s", peer) buf = bytearray() @@ -380,6 +383,7 @@ async def _session( return rest = await reader.readexactly(FRAME_HEADER_LEN - 4) buf += magic + rest + initialized = False while True: # Stage 1: ensure header is complete @@ -393,11 +397,14 @@ async def _session( # This was the iter-1410 adapter fix — pre-fix the session loop # closed on every frame with a non-empty CBOR payload because # decode requires header + payload and only the header had arrived. - magic_bytes, _ver, _mt, _flags, _res, payload_len = _HEADER_STRUCT.unpack_from(buf) + magic_bytes, version, msg_type, _flags, _res, payload_len = _HEADER_STRUCT.unpack_from(buf) if magic_bytes != IICP_MAGIC: logger.warning("Mid-stream magic drift — closing") return - if payload_len + FRAME_HEADER_LEN > _MAX_PAYLOAD: + if version != FRAMING_VERSION: + logger.warning("Unsupported IICP framing version — closing") + return + if payload_len > MAX_FRAME_PAYLOAD: logger.warning("IICP frame payload exceeds limit — closing") return total_len = FRAME_HEADER_LEN + payload_len @@ -414,7 +421,13 @@ async def _session( return del buf[:consumed] + if (not initialized and msg_type != MsgType.INIT) or (initialized and msg_type == MsgType.INIT): + logger.warning("Invalid IICP handshake state — closing") + return + keep_open = await self._dispatch(frame, writer) + if not initialized: + initialized = True if not keep_open: return @@ -426,7 +439,7 @@ async def _dispatch(self, frame: IicpFrame, writer: asyncio.StreamWriter) -> boo return True if mt == MsgType.INIT: - return await self._on_init(writer) + return await self._on_init(frame, writer) if mt == MsgType.PING: return await self._on_ping(frame, writer) if mt == MsgType.DISCOVER: @@ -444,10 +457,15 @@ async def _dispatch(self, frame: IicpFrame, writer: asyncio.StreamWriter) -> boo # ── handlers ───────────────────────────────────────────────────────────── - async def _on_init(self, writer: asyncio.StreamWriter) -> bool: - ack = IicpFrame.make( - MsgType.ACK, encode_ack(framing_version=FRAMING_VERSION, node_id=self.node_id) - ) + async def _on_init(self, frame: IicpFrame, writer: asyncio.StreamWriter) -> bool: + try: + body = decode_cbor(frame.payload) + except Exception: # noqa: BLE001 + return False + if not isinstance(body, dict) or body.get(1) != FRAMING_VERSION: + logger.warning("INIT requested an unsupported framing version — closing") + return False + ack = IicpFrame.make(MsgType.ACK, encode_ack(framing_version=FRAMING_VERSION, node_id=self.node_id)) writer.write(ack.encode()) await writer.drain() return True @@ -512,11 +530,7 @@ async def _on_call(self, frame: IicpFrame, writer: asyncio.StreamWriter) -> bool if isinstance(raw5, dict): payload_obj = raw5 else: - raw5_str = ( - raw5.decode("utf-8", errors="replace") - if isinstance(raw5, bytes) - else str(raw5) - ) + raw5_str = raw5.decode("utf-8", errors="replace") if isinstance(raw5, bytes) else str(raw5) if raw5_str: try: decoded = json.loads(raw5_str) @@ -557,11 +571,7 @@ async def _on_call(self, frame: IicpFrame, writer: asyncio.StreamWriter) -> bool # across HTTP and native IICP transports. from iicp_client.concurrency import CapacityExceededError, ConcurrencyGate - gate = ( - self.concurrency_gate - if isinstance(self.concurrency_gate, ConcurrencyGate) - else None - ) + gate = self.concurrency_gate if isinstance(self.concurrency_gate, ConcurrencyGate) else None async def _run_handler() -> None: nonlocal result, error_code, error_message @@ -570,9 +580,7 @@ async def _run_handler() -> None: if isinstance(handler_result, dict): if "error_code" in handler_result: error_code = int(handler_result["error_code"]) - error_message = str( - handler_result.get("error_message", "handler error") - ) + error_message = str(handler_result.get("error_message", "handler error")) else: result = encode_cbor(handler_result.get("result", handler_result)) else: @@ -661,11 +669,7 @@ async def run_handler() -> None: try: from iicp_client.concurrency import ConcurrencyGate - gate = ( - self.concurrency_gate - if isinstance(self.concurrency_gate, ConcurrencyGate) - else None - ) + gate = self.concurrency_gate if isinstance(self.concurrency_gate, ConcurrencyGate) else None if gate is None: await run_handler() else: @@ -772,14 +776,21 @@ async def handshake(self) -> None: if mt != MsgType.ACK: raise IicpTcpClientError(f"expected ACK (0x02), got 0x{mt:02x}") body = decode_cbor(payload) if payload else {} - if isinstance(body, dict): - self.framing_version = body.get(1) - v = body.get(2) - self.peer_node_id = v if isinstance(v, str) else None + if not isinstance(body, dict) or body.get(1) != FRAMING_VERSION: + negotiated = body.get(1) if isinstance(body, dict) else None + raise IicpTcpClientError(f"ACK negotiated unsupported framing version {negotiated!r}") + self.framing_version = FRAMING_VERSION + v = body.get(2) + self.peer_node_id = v if isinstance(v, str) else None + + def _require_handshake(self) -> None: + if self.framing_version != FRAMING_VERSION: + raise IicpTcpClientError("native session handshake is not complete") async def ping(self, echo: bytes | None = None) -> bytes | None: """Send PING; return the echoed bytes from the PONG (or None if not echoed).""" assert self._writer is not None + self._require_handshake() payload = encode_cbor({1: echo}) if echo else encode_cbor({}) self._writer.write(IicpFrame.make(MsgType.PING, payload).encode()) await self._writer.drain() @@ -792,6 +803,7 @@ async def ping(self, echo: bytes | None = None) -> bytes | None: async def discover(self, intent: str, *, session_id: str = "discover-1") -> list[dict]: """Send DISCOVER for `intent`; return the nodes list from the RESPONSE.""" assert self._writer is not None + self._require_handshake() payload = encode_cbor({2: session_id, 3: intent}) self._writer.write(IicpFrame.make(MsgType.DISCOVER, payload).encode()) await self._writer.drain() @@ -817,6 +829,7 @@ async def call( Raises IicpTcpClientError if the server replies with an error code. """ assert self._writer is not None + self._require_handshake() body: dict[int, object] = { 2: session_id, 3: intent, @@ -859,6 +872,7 @@ async def stream_call( contract and does not add lifecycle fields or wait for partial frames. """ assert self._writer is not None + self._require_handshake() if not task_id: raise ValueError("task_id is required for lifecycle streaming") attempt_id = call_id or str(uuid.uuid4()) @@ -901,6 +915,9 @@ async def close(self) -> None: """Send CLOSE (graceful teardown). Server hangs up; caller should disconnect.""" if self._writer is None or self._writer.is_closing(): return + if self.framing_version != FRAMING_VERSION: + await self.disconnect() + return self._writer.write(IicpFrame.make(MsgType.CLOSE, b"").encode()) try: await self._writer.drain() @@ -914,14 +931,14 @@ async def _read_frame(self, timeout_s: float | None = None) -> tuple[int, bytes] assert self._reader is not None t = timeout_s if timeout_s is not None else self.timeout_s head = await asyncio.wait_for(self._reader.readexactly(FRAME_HEADER_LEN), timeout=t) - magic, _ver, mt, _flags, _res, payload_len = _HEADER_STRUCT.unpack_from(head) + magic, version, mt, _flags, _res, payload_len = _HEADER_STRUCT.unpack_from(head) if magic != IICP_MAGIC: raise IicpTcpClientError(f"bad magic in response: {magic!r}") - payload = ( - await asyncio.wait_for(self._reader.readexactly(payload_len), timeout=t) - if payload_len - else b"" - ) + if version != FRAMING_VERSION: + raise IicpTcpClientError(f"unsupported framing version {version} in response") + if payload_len > MAX_FRAME_PAYLOAD: + raise IicpTcpClientError(f"response frame payload too large: {payload_len} > {MAX_FRAME_PAYLOAD}") + payload = await asyncio.wait_for(self._reader.readexactly(payload_len), timeout=t) if payload_len else b"" return mt, payload diff --git a/src/iicp_client/relay_session.py b/src/iicp_client/relay_session.py index 2744921..b6f567d 100644 --- a/src/iicp_client/relay_session.py +++ b/src/iicp_client/relay_session.py @@ -40,7 +40,7 @@ from iicp_client.relay_ticket import consume_relay_bind_ticket, verify_relay_bind_ticket -from .iicp_tcp import _decode_lifecycle_response +from .iicp_tcp import MAX_FRAME_PAYLOAD, _decode_lifecycle_response from .native_response_sequence import NativeResponseSequence, NativeResponseSequenceError logger = logging.getLogger(__name__) @@ -64,6 +64,8 @@ def _make_frame(msg_type: int, payload: bytes) -> bytes: + if len(payload) > MAX_FRAME_PAYLOAD: + raise ValueError(f"relay frame payload too large: {len(payload)} > {MAX_FRAME_PAYLOAD}") header = _HEADER_STRUCT.pack(_IICP_MAGIC, _FRAMING_VERSION, msg_type, 0, 0, len(payload)) return header + payload @@ -73,8 +75,7 @@ def _cbor2() -> Any: import cbor2 # type: ignore[import-untyped] except ImportError as exc: raise ImportError( - "cbor2 is required for relay sessions. " - "Install with: pip install 'iicp-client[iicp-tcp]'" + "cbor2 is required for relay sessions. Install with: pip install 'iicp-client[iicp-tcp]'" ) from exc return cbor2 @@ -131,9 +132,7 @@ def on_response(self, call_id: str, result: dict) -> None: if fut is not None and not fut.done(): fut.set_result(result) - async def forward_stream( - self, task: dict, timeout: float = 120.0 - ) -> AsyncIterator[dict[str, Any]]: + async def forward_stream(self, task: dict, timeout: float = 120.0) -> AsyncIterator[dict[str, Any]]: """Push a negotiated streaming CALL and yield validated lifecycle events.""" call_id = str(uuid.uuid4()) task_id = str(task.get("task_id") or call_id) @@ -253,9 +252,7 @@ def on_response(self, call_id: str, result: dict) -> None: if fut is not None and not fut.done(): fut.set_result(result) - async def forward_stream( - self, task: dict, timeout: float = 120.0 - ) -> AsyncIterator[dict[str, Any]]: + async def forward_stream(self, task: dict, timeout: float = 120.0) -> AsyncIterator[dict[str, Any]]: call_id = str(uuid.uuid4()) task_id = str(task.get("task_id") or call_id) session_id = str(task.get("session_id") or call_id) @@ -326,7 +323,9 @@ def __init__(self, max_sessions: int = MAX_RELAY_SESSIONS) -> None: self._lock = threading.Lock() self._max = max_sessions try: - self._bind_rate_limit = max(0, int(os.getenv("IICP_RELAY_BIND_RATE_LIMIT", str(DEFAULT_RELAY_BIND_RATE_LIMIT)))) + self._bind_rate_limit = max( + 0, int(os.getenv("IICP_RELAY_BIND_RATE_LIMIT", str(DEFAULT_RELAY_BIND_RATE_LIMIT))) + ) except ValueError: self._bind_rate_limit = DEFAULT_RELAY_BIND_RATE_LIMIT self._bind_rate_buckets: dict[str, tuple[float, int]] = {} @@ -422,9 +421,7 @@ def __init__( # so workers can register the correct {relay}/v1/relay-for/ endpoint. self.http_port = http_port self.require_bind_ticket = ( - os.getenv("IICP_RELAY_REQUIRE_BIND_TICKET") == "1" - if require_bind_ticket is None - else require_bind_ticket + os.getenv("IICP_RELAY_REQUIRE_BIND_TICKET") == "1" if require_bind_ticket is None else require_bind_ticket ) self.bind_ticket_public_key_hex = bind_ticket_public_key_hex or os.getenv("IICP_RELAY_BIND_TICKET_PUBLIC_KEY") self.relay_node_id = relay_node_id or os.getenv("IICP_NODE_ID", "*") @@ -432,9 +429,7 @@ def __init__( async def start(self) -> None: _cbor2() # validate import early - self._server = await asyncio.start_server( - self._handle_connection, host=self.host, port=self.port - ) + self._server = await asyncio.start_server(self._handle_connection, host=self.host, port=self.port) logger.info("Relay accept server listening on %s:%d", self.host, self.port) async def stop(self) -> None: @@ -449,9 +444,7 @@ async def serve_forever(self) -> None: async with self._server: await self._server.serve_forever() - async def _handle_connection( - self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter - ) -> None: + async def _handle_connection(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: peer = writer.get_extra_info("peername") logger.debug("Relay accept: connection from %s", peer) try: @@ -467,9 +460,7 @@ async def _handle_connection( except Exception: # noqa: BLE001 pass - async def _session( - self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter - ) -> None: + async def _session(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: """Handshake + relay-worker frame loop.""" # ── Step 1: INIT/ACK ────────────────────────────────────────────────── magic = await reader.readexactly(4) @@ -479,6 +470,9 @@ async def _session( rest = await reader.readexactly(_FRAME_HEADER_LEN - 4) header_bytes = magic + rest _, ver, msg_type, flags, _res, payload_len = _HEADER_STRUCT.unpack(header_bytes) + if ver != _FRAMING_VERSION or payload_len > MAX_FRAME_PAYLOAD: + logger.warning("Relay accept: unsupported version or oversized INIT") + return if msg_type != _MT_INIT: logger.warning("Relay accept: expected INIT, got 0x%02x", msg_type) return @@ -640,6 +634,7 @@ async def _relay_worker_loop( call_id = str(rb.get(15, "")) raw5 = rb.get(5, b"") import json as _json + if isinstance(raw5, (bytes, bytearray)): result = _json.loads(raw5) elif isinstance(raw5, str): @@ -662,8 +657,8 @@ async def _read_frame(self, reader: asyncio.StreamReader) -> bytes | None: header = await reader.readexactly(_FRAME_HEADER_LEN) except (asyncio.IncompleteReadError, EOFError, ConnectionResetError): return None - _, _, _, _, _, payload_len = _HEADER_STRUCT.unpack(header) - if payload_len > 16 * 1024 * 1024: + magic, version, _, _, _, payload_len = _HEADER_STRUCT.unpack(header) + if magic != _IICP_MAGIC or version != _FRAMING_VERSION or payload_len > MAX_FRAME_PAYLOAD: return None try: payload = await reader.readexactly(payload_len) if payload_len else b"" diff --git a/src/iicp_client/relay_worker_client.py b/src/iicp_client/relay_worker_client.py index d0f6fe9..a8d6169 100644 --- a/src/iicp_client/relay_worker_client.py +++ b/src/iicp_client/relay_worker_client.py @@ -35,6 +35,7 @@ from collections.abc import Awaitable, Callable from typing import Any +from iicp_client.iicp_tcp import MAX_FRAME_PAYLOAD from iicp_client.relay_ticket import fetch_relay_bind_ticket logger = logging.getLogger(__name__) @@ -83,6 +84,8 @@ def _dec(data: bytes) -> dict: def _make_frame(msg_type: int, payload: bytes) -> bytes: + if len(payload) > MAX_FRAME_PAYLOAD: + raise ValueError(f"relay frame payload too large: {len(payload)} > {MAX_FRAME_PAYLOAD}") header = _HEADER_STRUCT.pack(_IICP_MAGIC, _FRAMING_VERSION, msg_type, 0, 0, len(payload)) return header + payload @@ -94,13 +97,11 @@ async def _read_frame( header = await reader.readexactly(_HEADER_LEN) except (asyncio.IncompleteReadError, ConnectionResetError, EOFError): return None - magic = header[:4] + magic, version, msg_type, _flags, _reserved, payload_len = _HEADER_STRUCT.unpack(header) if magic != _IICP_MAGIC: logger.warning("Relay worker: bad magic %r", magic) return None - msg_type = header[5] - payload_len = _HEADER_STRUCT.unpack(header)[5] - if payload_len > 16 * 1024 * 1024: + if version != _FRAMING_VERSION or payload_len > MAX_FRAME_PAYLOAD: return None try: payload = await reader.readexactly(payload_len) if payload_len else b"" diff --git a/tests/fixtures/native-framing-v1.json b/tests/fixtures/native-framing-v1.json index e89517f..fb2993b 100644 --- a/tests/fixtures/native-framing-v1.json +++ b/tests/fixtures/native-framing-v1.json @@ -1,59 +1,141 @@ { "fixture_version": "1.0.0-draft", "status": "implementation-backed-pre-ratification", - "purpose": "Cross-implementation native framing vectors derived from the current Rust, Python, TypeScript, adapter, node and REACH implementations. This fixture covers frame decoding only; dispatch and lifecycle behavior are covered separately.", + "purpose": "Cross-implementation native framing vectors for the current ordered-stream binding. They cover bounded frame decoding only; dispatch, TLS, lifecycle, experimental relay opcodes, logical fragmentation and unsupported QUIC behavior are outside this fixture.", "frame": { "framing_version": 1, "header_bytes": 12, "layout": [ - {"name": "magic", "offset": 0, "bytes": 4, "value_hex": "49494350"}, - {"name": "version", "offset": 4, "bytes": 1}, - {"name": "type", "offset": 5, "bytes": 1}, - {"name": "flags", "offset": 6, "bytes": 1}, - {"name": "reserved", "offset": 7, "bytes": 1}, - {"name": "payload_length", "offset": 8, "bytes": 4, "encoding": "u32be"} + { + "name": "magic", + "offset": 0, + "bytes": 4, + "value_hex": "49494350" + }, + { + "name": "version", + "offset": 4, + "bytes": 1 + }, + { + "name": "type", + "offset": 5, + "bytes": 1 + }, + { + "name": "flags", + "offset": 6, + "bytes": 1 + }, + { + "name": "reserved", + "offset": 7, + "bytes": 1 + }, + { + "name": "payload_length", + "offset": 8, + "bytes": 4, + "encoding": "u32be" + } ], "decode_contract": { "reserved_byte": "ignored_on_receive", "unknown_flag_bits": "ignored_on_receive", - "message_type_validation": "outside_frame_decoder" - } + "message_type_validation": "outside_frame_decoder", + "framing_version": "exactly_1" + }, + "max_payload_bytes": 16777216, + "length_semantics": "payload_bytes_excluding_12_byte_header" }, "scenarios": [ { "name": "ping_empty", "wire_hex": "494943500109000000000000", - "expected": {"outcome": "accept", "version": 1, "message_type": 9, "flags": 0, "payload_hex": "", "consumed": 12} + "expected": { + "outcome": "accept", + "version": 1, + "message_type": 9, + "flags": 0, + "payload_hex": "", + "consumed": 12 + } }, { "name": "call_minimal_cbor", "wire_hex": "494943500105000000000003a10101", - "expected": {"outcome": "accept", "version": 1, "message_type": 5, "flags": 0, "payload_hex": "a10101", "consumed": 15} + "expected": { + "outcome": "accept", + "version": 1, + "message_type": 5, + "flags": 0, + "payload_hex": "a10101", + "consumed": 15 + } }, { "name": "reserved_byte_is_ignored_by_decoder", "wire_hex": "49494350010900ff00000000", - "expected": {"outcome": "accept", "version": 1, "message_type": 9, "flags": 0, "payload_hex": "", "consumed": 12} + "expected": { + "outcome": "accept", + "version": 1, + "message_type": 9, + "flags": 0, + "payload_hex": "", + "consumed": 12 + } }, { "name": "unknown_flag_bits_are_ignored_by_decoder", "wire_hex": "494943500109800000000000", - "expected": {"outcome": "accept", "version": 1, "message_type": 9, "flags": 128, "payload_hex": "", "consumed": 12} + "expected": { + "outcome": "accept", + "version": 1, + "message_type": 9, + "flags": 128, + "payload_hex": "", + "consumed": 12 + } }, { "name": "bad_magic", "wire_hex": "584943500109000000000000", - "expected": {"outcome": "reject", "reason": "invalid_magic"} + "expected": { + "outcome": "reject", + "reason": "invalid_magic" + } }, { "name": "truncated_header", "wire_hex": "4949435001090000000000", - "expected": {"outcome": "reject", "reason": "truncated_header"} + "expected": { + "outcome": "reject", + "reason": "truncated_header" + } }, { "name": "truncated_payload", "wire_hex": "494943500105000000000003a101", - "expected": {"outcome": "reject", "reason": "truncated_payload"} + "expected": { + "outcome": "reject", + "reason": "truncated_payload" + } + }, + { + "name": "unsupported_framing_version", + "wire_hex": "494943500209000000000000", + "expected": { + "outcome": "reject", + "reason": "unsupported_version" + } + }, + { + "name": "payload_length_exceeds_limit_before_body_read", + "wire_hex": "494943500109000001000001", + "expected": { + "outcome": "reject", + "reason": "payload_too_large" + } } ] } diff --git a/tests/test_iicp_tcp.py b/tests/test_iicp_tcp.py index 545aabb..f444138 100644 --- a/tests/test_iicp_tcp.py +++ b/tests/test_iicp_tcp.py @@ -22,6 +22,7 @@ FRAME_HEADER_LEN, FRAMING_VERSION, IICP_MAGIC, + MAX_FRAME_PAYLOAD, IicpTcpClient, IicpTcpClientError, IicpTcpServer, @@ -145,9 +146,7 @@ async def test_discover_invokes_lookup_returns_nodes(server_port): await _read_frame(reader) intent = "urn:iicp:intent:llm:chat:v1" - writer.write( - _frame(MsgType.DISCOVER, cbor2.dumps({2: "sess-1", 3: intent}, canonical=True)) - ) + writer.write(_frame(MsgType.DISCOVER, cbor2.dumps({2: "sess-1", 3: intent}, canonical=True))) await writer.drain() mt, payload = await _read_frame(reader) assert mt == MsgType.RESPONSE @@ -224,6 +223,63 @@ async def test_bad_magic_closes_connection(server_port): await writer.wait_closed() +async def test_server_rejects_application_frame_before_init(server_port): + reader, writer = await asyncio.open_connection("127.0.0.1", server_port) + try: + writer.write(_frame(MsgType.PING)) + await writer.drain() + assert await asyncio.wait_for(reader.read(1), timeout=TIMEOUT) == b"" + finally: + writer.close() + await writer.wait_closed() + + +async def test_server_rejects_duplicate_init(server_port): + reader, writer = await asyncio.open_connection("127.0.0.1", server_port) + init = _frame(MsgType.INIT, cbor2.dumps({1: FRAMING_VERSION}, canonical=True)) + try: + writer.write(init) + await writer.drain() + await _read_frame(reader) + writer.write(init) + await writer.drain() + assert await asyncio.wait_for(reader.read(1), timeout=TIMEOUT) == b"" + finally: + writer.close() + await writer.wait_closed() + + +async def test_server_rejects_init_with_unsupported_negotiated_version(server_port): + reader, writer = await asyncio.open_connection("127.0.0.1", server_port) + try: + writer.write(_frame(MsgType.INIT, cbor2.dumps({1: 2}, canonical=True))) + await writer.drain() + assert await asyncio.wait_for(reader.read(1), timeout=TIMEOUT) == b"" + finally: + writer.close() + await writer.wait_closed() + + +async def test_server_rejects_oversized_length_before_body_read(server_port): + reader, writer = await asyncio.open_connection("127.0.0.1", server_port) + try: + writer.write( + _HEADER.pack( + IICP_MAGIC, + FRAMING_VERSION, + MsgType.INIT, + 0, + 0, + MAX_FRAME_PAYLOAD + 1, + ) + ) + await writer.drain() + assert await asyncio.wait_for(reader.read(1), timeout=TIMEOUT) == b"" + finally: + writer.close() + await writer.wait_closed() + + async def test_payload_bearing_frame_does_not_close_session(server_port): """Regression guard for the iter-1410 adapter bug — pre-fix the session loop closed on every frame with a non-empty CBOR payload because IicpFrame.decode @@ -259,6 +315,57 @@ async def test_client_context_manager_handshake_and_close(server_port): assert client.peer_node_id == "test-node-id" +async def test_client_requires_handshake_before_application_frames(server_port): + async with IicpTcpClient("127.0.0.1", server_port) as client: + with pytest.raises(IicpTcpClientError, match="handshake is not complete"): + await client.ping() + + +async def test_client_rejects_oversized_response_header_before_body_read(): + async def oversized_ack(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + await _read_frame(reader) + writer.write( + _HEADER.pack( + IICP_MAGIC, + FRAMING_VERSION, + MsgType.ACK, + 0, + 0, + MAX_FRAME_PAYLOAD + 1, + ) + ) + await writer.drain() + writer.close() + + server = await asyncio.start_server(oversized_ack, "127.0.0.1", 0) + port = server.sockets[0].getsockname()[1] + try: + async with IicpTcpClient("127.0.0.1", port) as client: + with pytest.raises(IicpTcpClientError, match="response frame payload too large"): + await client.handshake() + finally: + server.close() + await server.wait_closed() + + +async def test_client_rejects_ack_with_wrong_negotiated_version(): + async def wrong_ack(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + await _read_frame(reader) + writer.write(_frame(MsgType.ACK, cbor2.dumps({1: 2}, canonical=True))) + await writer.drain() + writer.close() + + server = await asyncio.start_server(wrong_ack, "127.0.0.1", 0) + port = server.sockets[0].getsockname()[1] + try: + async with IicpTcpClient("127.0.0.1", port) as client: + with pytest.raises(IicpTcpClientError, match="ACK negotiated unsupported"): + await client.handshake() + finally: + server.close() + await server.wait_closed() + + async def test_client_ping_with_echo(server_port): async with IicpTcpClient("127.0.0.1", server_port) as client: await client.handshake() @@ -335,6 +442,7 @@ async def drain(self): return None client = IicpTcpClient("127.0.0.1") + client.framing_version = FRAMING_VERSION writer = Writer() client._writer = writer responses = iter( @@ -387,9 +495,7 @@ async def stream_handler(task): yield {"status": "partial", "result": b"hel", "tokens_used": 1} yield {"status": "success", "result": b"lo", "tokens_used": 2} - server = IicpTcpServer( - host="127.0.0.1", port=port, node_id="stream-node", streaming_handler=stream_handler - ) + server = IicpTcpServer(host="127.0.0.1", port=port, node_id="stream-node", streaming_handler=stream_handler) await server.start() try: async with IicpTcpClient("127.0.0.1", port) as client: @@ -419,9 +525,7 @@ async def stream_handler(_task): yield {"status": "partial", "result": b"some"} raise RuntimeError("sensitive backend detail") - server = IicpTcpServer( - host="127.0.0.1", port=port, node_id="stream-node", streaming_handler=stream_handler - ) + server = IicpTcpServer(host="127.0.0.1", port=port, node_id="stream-node", streaming_handler=stream_handler) await server.start() try: async with IicpTcpClient("127.0.0.1", port) as client: @@ -455,6 +559,7 @@ async def drain(self): return None client = IicpTcpClient("127.0.0.1") + client.framing_version = FRAMING_VERSION client._writer = Writer() async def read_frame(timeout_s=None): diff --git a/tests/test_native_framing_fixture.py b/tests/test_native_framing_fixture.py index 44ec7ad..c22f345 100644 --- a/tests/test_native_framing_fixture.py +++ b/tests/test_native_framing_fixture.py @@ -1,10 +1,11 @@ """Implementation-backed vectors for the established 12-byte native frame.""" + from __future__ import annotations import json from pathlib import Path -from iicp_client.iicp_tcp import FRAME_HEADER_LEN, IicpFrame, MsgType +from iicp_client.iicp_tcp import FRAME_HEADER_LEN, MAX_FRAME_PAYLOAD, IicpFrame, MsgType FIXTURE = Path(__file__).parent / "fixtures" / "native-framing-v1.json" @@ -12,11 +13,14 @@ def test_native_frame_decoder_matches_canonical_implementation_backed_vectors() -> None: data = json.loads(FIXTURE.read_text()) assert data["frame"]["header_bytes"] == FRAME_HEADER_LEN == 12 + assert data["frame"]["max_payload_bytes"] == MAX_FRAME_PAYLOAD == 16 * 1024 * 1024 expected_errors = { "invalid_magic": "Invalid IICP magic", "truncated_header": "frame too short", "truncated_payload": "payload truncated", + "unsupported_version": "Unsupported IICP framing version", + "payload_too_large": "frame payload too large", } for scenario in data["scenarios"]: name = scenario["name"] @@ -42,3 +46,13 @@ def test_native_frame_encoder_emits_the_canonical_empty_ping_vector() -> None: data = json.loads(FIXTURE.read_text()) ping = next(scenario for scenario in data["scenarios"] if scenario["name"] == "ping_empty") assert IicpFrame.make(MsgType.PING, b"").encode() == bytes.fromhex(ping["wire_hex"]) + + +def test_native_frame_encoder_rejects_payload_above_the_declared_limit() -> None: + payload = b"x" * (MAX_FRAME_PAYLOAD + 1) + try: + IicpFrame.make(MsgType.CALL, payload).encode() + except ValueError as error: + assert "frame payload too large" in str(error) + else: + raise AssertionError("oversized payload must be rejected") diff --git a/tests/test_serve_multiplex.py b/tests/test_serve_multiplex.py index cbb2e97..9d3d689 100644 --- a/tests/test_serve_multiplex.py +++ b/tests/test_serve_multiplex.py @@ -90,6 +90,7 @@ async def test_http_and_native_call_share_one_port() -> None: # Native IICP CALL answers on the SAME port (pre-#457 this hit the HTTP parser). async with IicpTcpClient("127.0.0.1", port) as client: + await client.handshake() result = await client.call(CHAT, {"messages": [{"role": "user", "content": "hi"}]}) assert isinstance(result, dict), "native CALL returned a RESPONSE result over the shared port" finally: