exoskeleton/code/core/network_emulator.py

253 lines
8.1 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,
2026-07-27 18:00:41 +08:00
common_random_numbers: bool = False,
) -> tuple[NetworkTraceEntry, ...]:
2026-07-27 18:00:41 +08:00
"""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")
2026-07-27 18:00:41 +08:00
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):
2026-07-27 18:00:41 +08:00
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),
2026-07-27 18:00:41 +08:00
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)