"""Deterministic packet channels and timeout semantics for G0c experiments.""" from __future__ import annotations from dataclasses import dataclass import heapq from typing import Generic, Iterable, TypeVar import numpy as np from core.feedback_protocol import ( PacketRejectReason, PacketState, ) PacketT = TypeVar("PacketT") @dataclass(frozen=True) class NetworkTraceEntry: """Transport decision for one sequence number.""" seq: int delay_s: float lost: bool = False duplicate: bool = False corrupt: bool = False def __post_init__(self) -> None: if self.seq < 0: raise ValueError("seq must be non-negative") if not np.isfinite(self.delay_s) or self.delay_s < 0.0: raise ValueError("delay_s must be finite and non-negative") def generate_network_trace( count: int, *, base_delay_s: float, jitter_s: float = 0.0, loss_probability: float = 0.0, duplicate_probability: float = 0.0, corrupt_probability: float = 0.0, seed: int = 0, ) -> tuple[NetworkTraceEntry, ...]: """Generate a reusable trace without coupling other random streams.""" if count < 0: raise ValueError("count must be non-negative") probabilities = ( loss_probability, duplicate_probability, corrupt_probability, ) if any(not 0.0 <= value <= 1.0 for value in probabilities): raise ValueError("network probabilities must lie in [0, 1]") if base_delay_s < 0.0 or jitter_s < 0.0: raise ValueError("delay and jitter must be non-negative") rng = np.random.default_rng(seed) entries = [] for seq in range(count): jitter = rng.uniform(-jitter_s, jitter_s) if jitter_s else 0.0 entries.append( NetworkTraceEntry( seq=seq, delay_s=max(0.0, base_delay_s + jitter), lost=bool(rng.random() < loss_probability), duplicate=bool(rng.random() < duplicate_probability), corrupt=bool(rng.random() < corrupt_probability), ) ) return tuple(entries) @dataclass(frozen=True) class Delivery(Generic[PacketT]): packet: PacketT source_time: float arrival_time: float corrupt: bool duplicate_copy: bool class DeterministicChannel(Generic[PacketT]): """One scheduled channel driven entirely by a frozen trace.""" def __init__(self, trace: Iterable[NetworkTraceEntry]): self._trace = {entry.seq: entry for entry in trace} self._pending: list[tuple[float, int, Delivery[PacketT]]] = [] self._insertion_order = 0 def send(self, packet: PacketT, now: float) -> None: seq = int(getattr(packet, "seq")) entry = self._trace.get(seq, NetworkTraceEntry(seq=seq, delay_s=0.0)) if entry.lost: return arrival = float(now) + entry.delay_s delivery = Delivery( packet=packet, source_time=float(getattr(packet, "source_time")), arrival_time=arrival, corrupt=entry.corrupt, duplicate_copy=False, ) heapq.heappush( self._pending, (arrival, self._insertion_order, delivery) ) self._insertion_order += 1 if entry.duplicate: duplicate = Delivery( packet=packet, source_time=delivery.source_time, arrival_time=arrival + 1e-12, corrupt=entry.corrupt, duplicate_copy=True, ) heapq.heappush( self._pending, (duplicate.arrival_time, self._insertion_order, duplicate), ) self._insertion_order += 1 def poll(self, now: float) -> list[Delivery[PacketT]]: ready: list[Delivery[PacketT]] = [] while self._pending and self._pending[0][0] <= float(now) + 1e-15: _, _, delivery = heapq.heappop(self._pending) ready.append(delivery) return ready @property def pending_count(self) -> int: return len(self._pending) @dataclass(frozen=True) class Reception: seq: int arrival_time: float age_s: float accepted: bool reason: PacketRejectReason state: PacketState @dataclass(frozen=True) class HeldPacket(Generic[PacketT]): packet: PacketT | None state: PacketState age_s: float | None fresh: bool class PacketReceiver(Generic[PacketT]): """Newest-sequence hold, timeout, and recovery state machine.""" def __init__(self, timeout_s: float): if not np.isfinite(timeout_s) or timeout_s <= 0.0: raise ValueError("timeout_s must be finite and positive") self.timeout_s = float(timeout_s) self.last_seq = -1 self.last_packet: PacketT | None = None self.last_arrival_time: float | None = None self.state = PacketState.EMPTY self._fresh = False self._recovering_sample_pending = False def accept(self, delivery: Delivery[PacketT]) -> Reception: packet = delivery.packet seq = int(getattr(packet, "seq")) valid = bool(getattr(packet, "valid", True)) source_time = float(getattr(packet, "source_time")) age = max(0.0, delivery.arrival_time - source_time) if delivery.corrupt: reason = PacketRejectReason.CORRUPT elif not valid: reason = PacketRejectReason.INVALID elif seq <= self.last_seq: reason = PacketRejectReason.DUPLICATE_OR_STALE else: previous = self.state self.last_seq = seq self.last_packet = packet self.last_arrival_time = delivery.arrival_time self._fresh = True self._recovering_sample_pending = previous in ( PacketState.EMPTY, PacketState.TIMED_OUT, ) self.state = ( PacketState.RECOVERING if self._recovering_sample_pending else PacketState.ACTIVE ) return Reception( seq=seq, arrival_time=delivery.arrival_time, age_s=age, accepted=True, reason=PacketRejectReason.ACCEPTED, state=self.state, ) return Reception( seq=seq, arrival_time=delivery.arrival_time, age_s=age, accepted=False, reason=reason, state=self.state, ) def sample(self, now: float) -> HeldPacket[PacketT]: if self.last_packet is None or self.last_arrival_time is None: self.state = PacketState.EMPTY return HeldPacket(None, self.state, None, False) age = max(0.0, float(now) - self.last_arrival_time) if age > self.timeout_s: self.state = PacketState.TIMED_OUT self._fresh = False return HeldPacket(None, self.state, age, False) if self._recovering_sample_pending: state = PacketState.RECOVERING self._recovering_sample_pending = False elif self._fresh: state = PacketState.ACTIVE else: state = PacketState.HELD fresh = self._fresh self._fresh = False self.state = state return HeldPacket(self.last_packet, state, age, fresh)