diff --git a/RNS/Channel.py b/RNS/Channel.py index 3072b852..799f002e 100644 --- a/RNS/Channel.py +++ b/RNS/Channel.py @@ -45,63 +45,6 @@ TPacket = TypeVar("TPacket") class SystemMessageTypes(enum.IntEnum): SMT_STREAM_DATA = 0xff00 -class ChannelOutletBase(ABC, Generic[TPacket]): - """ - An abstract transport layer interface used by Channel. - - DEPRECATED: This was created for testing; eventually - Channel will use Link or a LinkBase interface - directly. - """ - @abstractmethod - def send(self, raw: bytes) -> TPacket: - raise NotImplemented() - - @abstractmethod - def resend(self, packet: TPacket) -> TPacket: - raise NotImplemented() - - @property - @abstractmethod - def mdu(self): - raise NotImplemented() - - @property - @abstractmethod - def rtt(self): - raise NotImplemented() - - @property - @abstractmethod - def is_usable(self): - raise NotImplemented() - - @abstractmethod - def get_packet_state(self, packet: TPacket) -> MessageState: - raise NotImplemented() - - @abstractmethod - def timed_out(self): - raise NotImplemented() - - @abstractmethod - def __str__(self): - raise NotImplemented() - - @abstractmethod - def set_packet_timeout_callback(self, packet: TPacket, callback: Callable[[TPacket], None] | None, - timeout: float | None = None): - raise NotImplemented() - - @abstractmethod - def set_packet_delivered_callback(self, packet: TPacket, callback: Callable[[TPacket], None] | None): - raise NotImplemented() - - @abstractmethod - def get_packet_id(self, packet: TPacket) -> any: - raise NotImplemented() - - class CEType(enum.IntEnum): """ ChannelException type codes @@ -113,7 +56,6 @@ class CEType(enum.IntEnum): ME_ALREADY_SENT = 4 ME_TOO_BIG = 5 - class ChannelException(Exception): """ An exception thrown by Channel, with a type code. @@ -122,7 +64,6 @@ class ChannelException(Exception): super().__init__(args) self.type = ce_type - class MessageState(enum.IntEnum): """ Set of possible states for a Message @@ -132,7 +73,6 @@ class MessageState(enum.IntEnum): MSGSTATE_DELIVERED = 2 MSGSTATE_FAILED = 3 - class MessageBase(abc.ABC): """ Base type for any messages sent or received on a Channel. @@ -170,7 +110,6 @@ class MessageBase(abc.ABC): MessageCallbackType = NewType("MessageCallbackType", Callable[[MessageBase], bool]) - class Envelope: """ Internal wrapper used to transport messages over a channel and @@ -180,8 +119,7 @@ class Envelope: msgtype, self.sequence, length = struct.unpack(">HHH", self.raw[:6]) raw = self.raw[6:] ctor = message_factories.get(msgtype, None) - if ctor is None: - raise ChannelException(CEType.ME_NOT_REGISTERED, f"Unable to find constructor for Channel MSGTYPE {hex(msgtype)}") + if ctor is None: raise ChannelException(CEType.ME_NOT_REGISTERED, f"Unable to find constructor for Channel MSGTYPE {hex(msgtype)}") message = ctor() message.unpack(raw) self.unpacked = True @@ -190,14 +128,13 @@ class Envelope: return message def pack(self) -> bytes: - if self.message.__class__.MSGTYPE is None: - raise ChannelException(CEType.ME_NO_MSG_TYPE, f"{self.message.__class__} lacks MSGTYPE") + if self.message.__class__.MSGTYPE is None: raise ChannelException(CEType.ME_NO_MSG_TYPE, f"{self.message.__class__} lacks MSGTYPE") data = self.message.pack() self.raw = struct.pack(">HHH", self.message.MSGTYPE, self.sequence, len(data)) + data self.packed = True return self.raw - def __init__(self, outlet: ChannelOutletBase, message: MessageBase = None, raw: bytes = None, sequence: int = None): + def __init__(self, outlet: LinkChannelOutlet, message: MessageBase = None, raw: bytes = None, sequence: int = None): self.ts = time.time() self.id = id(self) self.message = message @@ -278,7 +215,7 @@ class Channel(contextlib.AbstractContextManager): SEQ_MAX = 0xFFFF SEQ_MODULUS = SEQ_MAX+1 - def __init__(self, outlet: ChannelOutletBase): + def __init__(self, outlet: LinkChannelOutlet): """ @param outlet: @@ -307,11 +244,9 @@ class Channel(contextlib.AbstractContextManager): self.window_min = Channel.WINDOW_MIN self.window_flexibility = Channel.WINDOW_FLEXIBILITY - def __enter__(self) -> Channel: - return self + def __enter__(self) -> Channel: return self - def __exit__(self, __exc_type: Type[BaseException] | None, __exc_value: BaseException | None, - __traceback: TracebackType | None) -> bool | None: + def __exit__(self, __exc_type: Type[BaseException] | None, __exc_value: BaseException | None, __traceback: TracebackType | None) -> bool | None: self._shutdown() return False @@ -327,20 +262,13 @@ class Channel(contextlib.AbstractContextManager): def _register_message_type(self, message_class: Type[MessageBase], *, is_system_type: bool = False): with self._lock: - if not issubclass(message_class, MessageBase): - raise ChannelException(CEType.ME_INVALID_MSG_TYPE, - f"{message_class} is not a subclass of {MessageBase}.") - if message_class.MSGTYPE is None: - raise ChannelException(CEType.ME_INVALID_MSG_TYPE, - f"{message_class} has invalid MSGTYPE class attribute.") - if message_class.MSGTYPE >= 0xf000 and not is_system_type: - raise ChannelException(CEType.ME_INVALID_MSG_TYPE, - f"{message_class} has system-reserved message type.") - try: - message_class() + if not issubclass(message_class, MessageBase): raise ChannelException(CEType.ME_INVALID_MSG_TYPE, f"{message_class} is not a subclass of {MessageBase}.") + if message_class.MSGTYPE is None: raise ChannelException(CEType.ME_INVALID_MSG_TYPE, f"{message_class} has invalid MSGTYPE class attribute.") + if message_class.MSGTYPE >= 0xf000 and not is_system_type: raise ChannelException(CEType.ME_INVALID_MSG_TYPE, f"{message_class} has system-reserved message type.") + + try: message_class() except Exception as ex: - raise ChannelException(CEType.ME_INVALID_MSG_TYPE, - f"{message_class} raised an exception when constructed with no arguments: {ex}") + raise ChannelException(CEType.ME_INVALID_MSG_TYPE, f"{message_class} raised an exception when constructed with no arguments: {ex}") self._message_factories[message_class.MSGTYPE] = message_class @@ -384,8 +312,10 @@ class Channel(contextlib.AbstractContextManager): self._outlet.set_packet_timeout_callback(envelope.packet, None) self._outlet.set_packet_delivered_callback(envelope.packet, None) envelope.tracked = False + for envelope in self._rx_ring: envelope.tracked = False + self._tx_ring.clear() self._rx_ring.clear() @@ -394,14 +324,12 @@ class Channel(contextlib.AbstractContextManager): i = 0 for existing in ring: - if envelope.sequence == existing.sequence: - RNS.log(f"Envelope: Emplacement of duplicate envelope with sequence "+str(envelope.sequence), RNS.LOG_EXTREME) + RNS.log(f"Envelope: Emplacement of duplicate envelope with sequence "+str(envelope.sequence), RNS.LOG_EXTREME) if RNS.sl(RNS.LOG_EXTREME) else None return False if envelope.sequence < existing.sequence and not (self._next_rx_sequence - envelope.sequence) > (Channel.SEQ_MAX//2): ring.insert(i, envelope) - envelope.tracked = True return True @@ -417,10 +345,8 @@ class Channel(contextlib.AbstractContextManager): for cb in cbs: try: - if cb(message): - return - except Exception as e: - RNS.log("Channel "+str(self)+" experienced an error while running a message callback. The contained exception was: "+str(e), RNS.LOG_ERROR) + if cb(message): return + except Exception as e: RNS.log(f"Channel {str(self)} experienced an error while running a message callback. The contained exception was: {e}", RNS.LOG_ERROR) def _receive(self, raw: bytes): try: @@ -432,16 +358,20 @@ class Channel(contextlib.AbstractContextManager): window_overflow = (self._next_rx_sequence+Channel.WINDOW_MAX) % Channel.SEQ_MODULUS if window_overflow < self._next_rx_sequence: if envelope.sequence > window_overflow: - RNS.log("Invalid packet sequence ("+str(envelope.sequence)+") received on channel "+str(self), RNS.LOG_EXTREME) + RNS.log(f"Invalid packet sequence ({envelope.sequence}) received on channel {self}", RNS.LOG_DEBUG) if RNS.sl(RNS.LOG_DEBUG) else None return else: - RNS.log("Invalid packet sequence ("+str(envelope.sequence)+") received on channel "+str(self), RNS.LOG_EXTREME) + RNS.log(f"Invalid packet sequence ({envelope.sequence}) received on channel {self}", RNS.LOG_DEBUG) if RNS.sl(RNS.LOG_DEBUG) else None return + elif envelope.sequence > self._next_rx_sequence + self.WINDOW_MAX: + RNS.log(f"Invalid packet sequence ({envelope.sequence}) received on channel {self}", RNS.LOG_DEBUG) if RNS.sl(RNS.LOG_DEBUG) else None + return + is_new = self._emplace_envelope(envelope, self._rx_ring) if not is_new: - RNS.log("Duplicate message received on channel "+str(self), RNS.LOG_EXTREME) + RNS.log("Duplicate message received on channel "+str(self), RNS.LOG_EXTREME) if RNS.sl(RNS.LOG_EXTREME) else None return else: with self._lock: @@ -457,10 +387,8 @@ class Channel(contextlib.AbstractContextManager): self._next_rx_sequence = (self._next_rx_sequence + 1) % Channel.SEQ_MODULUS for e in contigous: - if not e.unpacked: - m = e.unpack(self._message_factories) - else: - m = e.message + if not e.unpacked: m = e.unpack(self._message_factories) + else: m = e.message self._rx_ring.remove(e) self._run_callbacks(m) @@ -474,9 +402,7 @@ class Channel(contextlib.AbstractContextManager): :return: True if ready """ - if not self._outlet.is_usable: - return False - + if not self._outlet.is_usable: return False with self._lock: outstanding = 0 for envelope in self._tx_ring: @@ -484,59 +410,38 @@ class Channel(contextlib.AbstractContextManager): if not envelope.packet or not self._outlet.get_packet_state(envelope.packet) == MessageState.MSGSTATE_DELIVERED: outstanding += 1 - if outstanding >= self.window: - return False + if outstanding >= self.window: return False return True def _packet_tx_op(self, packet: TPacket, op: Callable[[TPacket], bool]): target_id = self._outlet.get_packet_id(packet) with self._lock: - envelope = next(filter(lambda e: e.packet is not None - and self._outlet.get_packet_id(e.packet) == target_id, - self._tx_ring), None) - + envelope = next(filter(lambda e: e.packet is not None and self._outlet.get_packet_id(e.packet) == target_id, self._tx_ring), None) if envelope and op(envelope): envelope.tracked = False if envelope in self._tx_ring: self._tx_ring.remove(envelope) - if self.window < self.window_max: - self.window += 1 - - # TODO: Remove at some point - # RNS.log("Increased "+str(self)+" window to "+str(self.window), RNS.LOG_DEBUG) + if self.window < self.window_max: self.window += 1 if self._outlet.rtt != 0: if self._outlet.rtt > Channel.RTT_FAST: self.fast_rate_rounds = 0 - - if self._outlet.rtt > Channel.RTT_MEDIUM: - self.medium_rate_rounds = 0 - + if self._outlet.rtt > Channel.RTT_MEDIUM: self.medium_rate_rounds = 0 else: self.medium_rate_rounds += 1 if self.window_max < Channel.WINDOW_MAX_MEDIUM and self.medium_rate_rounds == Channel.FAST_RATE_THRESHOLD: self.window_max = Channel.WINDOW_MAX_MEDIUM self.window_min = Channel.WINDOW_MIN_LIMIT_MEDIUM - # TODO: Remove at some point - # RNS.log("Increased "+str(self)+" max window to "+str(self.window_max), RNS.LOG_DEBUG) - # RNS.log("Increased "+str(self)+" min window to "+str(self.window_min), RNS.LOG_DEBUG) - else: self.fast_rate_rounds += 1 if self.window_max < Channel.WINDOW_MAX_FAST and self.fast_rate_rounds == Channel.FAST_RATE_THRESHOLD: self.window_max = Channel.WINDOW_MAX_FAST self.window_min = Channel.WINDOW_MIN_LIMIT_FAST - # TODO: Remove at some point - # RNS.log("Increased "+str(self)+" max window to "+str(self.window_max), RNS.LOG_DEBUG) - # RNS.log("Increased "+str(self)+" min window to "+str(self.window_min), RNS.LOG_DEBUG) - - else: - RNS.log("Envelope not found in TX ring for "+str(self), RNS.LOG_EXTREME) - if not envelope: - RNS.log("Spurious message received on "+str(self), RNS.LOG_EXTREME) + else: RNS.log("Envelope not found in TX ring for "+str(self), RNS.LOG_EXTREME) if RNS.sl(RNS.LOG_EXTREME) else None + if not envelope: RNS.log("Spurious message received on "+str(self), RNS.LOG_EXTREME) if RNS.sl(RNS.LOG_EXTREME) else None def _packet_delivered(self, packet: TPacket): self._packet_tx_op(packet, lambda env: True) @@ -545,29 +450,23 @@ class Channel(contextlib.AbstractContextManager): for envelope in self._tx_ring: updated_timeout = self._get_packet_timeout_time(envelope.tries) if envelope.packet and hasattr(envelope.packet, "receipt") and envelope.packet.receipt and envelope.packet.receipt.timeout: - if updated_timeout > envelope.packet.receipt.timeout: - envelope.packet.receipt.set_timeout(updated_timeout) + if updated_timeout > envelope.packet.receipt.timeout: envelope.packet.receipt.set_timeout(updated_timeout) def _get_packet_timeout_time(self, tries: int) -> float: to = pow(1.5, tries - 1) * max(self._outlet.rtt*2.5, 0.025) * (len(self._tx_ring)+1.5) return to def _packet_timeout(self, packet: TPacket): - if self._outlet.get_packet_state(packet) == MessageState.MSGSTATE_DELIVERED: - return + if self._outlet.get_packet_state(packet) == MessageState.MSGSTATE_DELIVERED: return target_id = self._outlet.get_packet_id(packet) envelope_to_resend: Envelope | None = None should_teardown = False with self._lock: - envelope = next(filter( - lambda e: e.packet is not None and self._outlet.get_packet_id(e.packet) == target_id, - self._tx_ring), None) - if envelope is None: - return + envelope = next(filter(lambda e: e.packet is not None and self._outlet.get_packet_id(e.packet) == target_id, self._tx_ring), None) + if envelope is None: return - if envelope.tries >= self._max_tries: - should_teardown = True + if envelope.tries >= self._max_tries: should_teardown = True else: envelope.tries += 1 envelope_to_resend = envelope @@ -587,14 +486,11 @@ class Channel(contextlib.AbstractContextManager): self._outlet.resend(envelope_to_resend.packet) with self._lock: self._outlet.set_packet_delivered_callback(envelope_to_resend.packet, self._packet_delivered) - self._outlet.set_packet_timeout_callback( - envelope_to_resend.packet, self._packet_timeout, - self._get_packet_timeout_time(envelope_to_resend.tries)) + self._outlet.set_packet_timeout_callback(envelope_to_resend.packet, self._packet_timeout, self._get_packet_timeout_time(envelope_to_resend.tries)) self._update_packet_timeouts() already_delivered = (self._outlet.get_packet_state(envelope_to_resend.packet) == MessageState.MSGSTATE_DELIVERED) - if already_delivered: - self._packet_delivered(envelope_to_resend.packet) + if already_delivered: self._packet_delivered(envelope_to_resend.packet) def send(self, message: MessageBase) -> Envelope: """ @@ -605,24 +501,19 @@ class Channel(contextlib.AbstractContextManager): """ with self._send_lock: with self._lock: - if not self.is_ready_to_send(): - raise ChannelException(CEType.ME_LINK_NOT_READY, f"Link is not ready") + if not self.is_ready_to_send(): raise ChannelException(CEType.ME_LINK_NOT_READY, f"Link is not ready") reserved_sequence = self._next_sequence envelope = Envelope(self._outlet, message=message, sequence=reserved_sequence) envelope.pack() if len(envelope.raw) > self._outlet.mdu: - raise ChannelException(CEType.ME_TOO_BIG, - f"Packed message too big for packet: {len(envelope.raw)} > {self._outlet.mdu}") + raise ChannelException(CEType.ME_TOO_BIG, f"Packed message too big for packet: {len(envelope.raw)} > {self._outlet.mdu}") self._next_sequence = (reserved_sequence + 1) % Channel.SEQ_MODULUS envelope.packet = self._outlet.send(envelope.raw) - if (envelope.packet is None - or getattr(envelope.packet, "raw", None) is None - or (hasattr(envelope.packet, "receipt") and envelope.packet.receipt is None)): - with self._lock: - self._next_sequence = reserved_sequence + if envelope.packet is None or getattr(envelope.packet, "raw", None) is None or (hasattr(envelope.packet, "receipt") and envelope.packet.receipt is None): + with self._lock: self._next_sequence = reserved_sequence raise ChannelException(CEType.ME_LINK_NOT_READY, "Outlet did not transmit packet") with self._lock: @@ -633,9 +524,8 @@ class Channel(contextlib.AbstractContextManager): self._update_packet_timeouts() already_delivered = (self._outlet.get_packet_state(envelope.packet) == MessageState.MSGSTATE_DELIVERED) - # prevent _tx_ring envelope leak - if already_delivered: - self._packet_delivered(envelope.packet) + # Prevent _tx_ring envelope leak + if already_delivered: self._packet_delivered(envelope.packet) return envelope @@ -650,14 +540,13 @@ class Channel(contextlib.AbstractContextManager): :return: number of bytes available """ mdu = self._outlet.mdu - 6 # sizeof(msgtype) + sizeof(length) + sizeof(sequence) - if mdu > 0xFFFF: - mdu = 0xFFFF + if mdu > 0xFFFF: mdu = 0xFFFF return mdu -class LinkChannelOutlet(ChannelOutletBase): +class LinkChannelOutlet(ABC, Generic[TPacket]): """ - An implementation of ChannelOutletBase for RNS.Link. + A channel outlet implementation for RNS.Link. Allows Channel to send packets over an RNS Link with Packets. @@ -668,71 +557,51 @@ class LinkChannelOutlet(ChannelOutletBase): def send(self, raw: bytes) -> RNS.Packet: packet = RNS.Packet(self.link, raw, context=RNS.Packet.CHANNEL) - if self.link.status == RNS.Link.ACTIVE: - packet.send() + if self.link.status == RNS.Link.ACTIVE: packet.send() return packet def resend(self, packet: RNS.Packet) -> RNS.Packet: receipt = packet.resend() - if not receipt: - RNS.log("Failed to resend packet", RNS.LOG_ERROR) + if not receipt: RNS.log("Failed to resend packet", RNS.LOG_ERROR) return packet @property - def mdu(self): - return self.link.mdu + def mdu(self): return self.link.mdu @property - def rtt(self): - return self.link.rtt + def rtt(self): return self.link.rtt + # TODO: Initial implementation had issues + # proxying this to link status. Should + # return False when link is not ACTIVE. + # Investigate what happened here. @property - def is_usable(self): - return True # had issues looking at Link.status + def is_usable(self): return True def get_packet_state(self, packet: TPacket) -> MessageState: if packet.receipt == None: return MessageState.MSGSTATE_FAILED status = packet.receipt.get_status() - if status == RNS.PacketReceipt.SENT: - return MessageState.MSGSTATE_SENT - if status == RNS.PacketReceipt.DELIVERED: - return MessageState.MSGSTATE_DELIVERED - if status == RNS.PacketReceipt.FAILED: - return MessageState.MSGSTATE_FAILED - else: - raise Exception(f"Unexpected receipt state: {status}") + if status == RNS.PacketReceipt.SENT: return MessageState.MSGSTATE_SENT + if status == RNS.PacketReceipt.DELIVERED: return MessageState.MSGSTATE_DELIVERED + if status == RNS.PacketReceipt.FAILED: return MessageState.MSGSTATE_FAILED + else: raise Exception(f"Unexpected receipt state: {status}") - def timed_out(self): - self.link.teardown() - - def __str__(self): - return f"{self.__class__.__name__}({self.link})" - - def set_packet_timeout_callback(self, packet: RNS.Packet, callback: Callable[[RNS.Packet], None] | None, - timeout: float | None = None): - if timeout and packet.receipt: - packet.receipt.set_timeout(timeout) - - def inner(receipt: RNS.PacketReceipt): - callback(packet) - - if packet and packet.receipt: - packet.receipt.set_timeout_callback(inner if callback else None) + def set_packet_timeout_callback(self, packet: RNS.Packet, callback: Callable[[RNS.Packet], None] | None, timeout: float | None = None): + if timeout and packet.receipt: packet.receipt.set_timeout(timeout) + def inner(receipt: RNS.PacketReceipt): callback(packet) + if packet and packet.receipt: packet.receipt.set_timeout_callback(inner if callback else None) def set_packet_delivered_callback(self, packet: RNS.Packet, callback: Callable[[RNS.Packet], None] | None): - def inner(receipt: RNS.PacketReceipt): - callback(packet) - - if packet and packet.receipt: - packet.receipt.set_delivery_callback(inner if callback else None) + def inner(receipt: RNS.PacketReceipt): callback(packet) + if packet and packet.receipt: packet.receipt.set_delivery_callback(inner if callback else None) def get_packet_id(self, packet: RNS.Packet) -> any: - if (packet - and getattr(packet, "raw", None) is not None - and hasattr(packet, "get_hash") - and callable(packet.get_hash)): + if packet and getattr(packet, "raw", None) is not None and hasattr(packet, "get_hash") and callable(packet.get_hash): return packet.get_hash() - else: - return None + else: return None + + def timed_out(self): self.link.teardown() + + def __str__(self): return f"{self.__class__.__name__}({self.link})" \ No newline at end of file