exoskeleton/code/core/network_emulator.py

234 lines
7.3 KiB
Python
Raw Normal View History

"""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)