"""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, common_random_numbers: bool = False, ) -> tuple[NetworkTraceEntry, ...]: """Generate a reusable trace without coupling other random streams. The default branch retains the historical conditional draw order exactly. With common random numbers enabled, every packet consumes four uniforms in the fixed order jitter/loss/duplicate/corrupt, even when a corresponding magnitude or probability is zero. Network profiles can then threshold and scale the same latent draws without shifting later packet decisions. """ 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") if not isinstance(common_random_numbers, (bool, np.bool_)): raise ValueError("common_random_numbers must be boolean") rng = np.random.default_rng(seed) entries = [] for seq in range(count): if common_random_numbers: jitter_draw, loss_draw, duplicate_draw, corrupt_draw = ( rng.random(4) ) jitter = jitter_s * (2.0 * jitter_draw - 1.0) else: jitter = rng.uniform(-jitter_s, jitter_s) if jitter_s else 0.0 loss_draw = rng.random() duplicate_draw = rng.random() corrupt_draw = rng.random() entries.append( NetworkTraceEntry( seq=seq, delay_s=max(0.0, base_delay_s + jitter), lost=bool(loss_draw < loss_probability), duplicate=bool(duplicate_draw < duplicate_probability), corrupt=bool(corrupt_draw < 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)