234 lines
7.3 KiB
Python
234 lines
7.3 KiB
Python
"""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)
|