Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
105 changes: 61 additions & 44 deletions src/iicp_client/iicp_tcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand All @@ -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)}")
Expand All @@ -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

Expand Down Expand Up @@ -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:
Expand All @@ -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()
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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

Expand All @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand All @@ -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()
Expand All @@ -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,
Expand Down Expand Up @@ -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())
Expand Down Expand Up @@ -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()
Expand All @@ -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


Expand Down
43 changes: 19 additions & 24 deletions src/iicp_client/relay_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand All @@ -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

Expand All @@ -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

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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]] = {}
Expand Down Expand Up @@ -422,19 +421,15 @@ def __init__(
# so workers can register the correct {relay}/v1/relay-for/<wid> 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", "*")
self._server: asyncio.AbstractServer | None = None

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:
Expand All @@ -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:
Expand All @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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):
Expand All @@ -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""
Expand Down
Loading
Loading