202 lines
6.6 KiB
Python
202 lines
6.6 KiB
Python
#!/usr/bin/env python3
|
|
"""Pure numerical contracts for matched H3 feedback and packet semantics."""
|
|
|
|
from pathlib import Path
|
|
import sys
|
|
import unittest
|
|
|
|
import numpy as np
|
|
|
|
|
|
CODE_ROOT = Path(__file__).resolve().parents[1]
|
|
sys.path.insert(0, str(CODE_ROOT))
|
|
|
|
from core.feedback_protocol import ( # noqa: E402
|
|
ForwardPacket,
|
|
MapPolicy,
|
|
MapRegistry,
|
|
MapSnapshot,
|
|
MappingKind,
|
|
PacketRejectReason,
|
|
PacketState,
|
|
ReturnPacket,
|
|
map_return_feedback,
|
|
normalized_actual_power_mismatch,
|
|
)
|
|
from core.network_emulator import ( # noqa: E402
|
|
DeterministicChannel,
|
|
NetworkTraceEntry,
|
|
PacketReceiver,
|
|
generate_network_trace,
|
|
)
|
|
|
|
|
|
class FeedbackProtocolTest(unittest.TestCase):
|
|
def make_packet(self, seq=3, map_id=1):
|
|
wrench = np.array([1.0, -2.0, 0.5, 0.2, -0.3, 0.1])
|
|
Js = np.arange(42, dtype=float).reshape(6, 7) / 50.0
|
|
return ReturnPacket(
|
|
seq=seq,
|
|
source_index=10 + seq,
|
|
source_time=0.02 * seq,
|
|
echoed_map_id=map_id,
|
|
residual=np.linspace(-1.0, 1.0, 7),
|
|
wrench=wrench,
|
|
js_t_wrench=Js.T @ wrench,
|
|
qd_slave_actual=np.linspace(0.1, 0.7, 7),
|
|
)
|
|
|
|
def test_matched_conditions_share_one_input_hash(self):
|
|
registry = MapRegistry()
|
|
A = np.eye(7)
|
|
registry.add(MapSnapshot(1, 2, 0.04, A))
|
|
packet = self.make_packet()
|
|
Jm = np.arange(42, dtype=float).reshape(6, 7) / 70.0
|
|
|
|
differential = map_return_feedback(
|
|
kind=MappingKind.MATCHED_DIFFERENTIAL_WRENCH,
|
|
packet=packet,
|
|
master_jacobian=Jm,
|
|
maps=registry,
|
|
)
|
|
direct = map_return_feedback(
|
|
kind=MappingKind.DIRECT_MASTER_JACOBIAN,
|
|
packet=packet,
|
|
master_jacobian=Jm,
|
|
maps=registry,
|
|
)
|
|
|
|
self.assertEqual(
|
|
differential.matched_input_hash, direct.matched_input_hash
|
|
)
|
|
np.testing.assert_allclose(
|
|
differential.tau_master_raw, A.T @ packet.js_t_wrench
|
|
)
|
|
np.testing.assert_allclose(
|
|
direct.tau_master_raw, Jm.T @ packet.wrench
|
|
)
|
|
|
|
def test_source_stamped_policy_uses_echoed_map(self):
|
|
registry = MapRegistry()
|
|
registry.add(MapSnapshot(1, 0, 0.0, np.eye(7)))
|
|
registry.add(MapSnapshot(2, 1, 0.02, 2.0 * np.eye(7)))
|
|
packet = self.make_packet(map_id=1)
|
|
Jm = np.zeros((6, 7))
|
|
source = map_return_feedback(
|
|
kind=MappingKind.MATCHED_DIFFERENTIAL_WRENCH,
|
|
packet=packet,
|
|
master_jacobian=Jm,
|
|
maps=registry,
|
|
map_policy=MapPolicy.SOURCE_STAMPED,
|
|
)
|
|
current = map_return_feedback(
|
|
kind=MappingKind.MATCHED_DIFFERENTIAL_WRENCH,
|
|
packet=packet,
|
|
master_jacobian=Jm,
|
|
maps=registry,
|
|
map_policy=MapPolicy.CURRENT,
|
|
)
|
|
self.assertEqual(source.selected_map_id, 1)
|
|
self.assertEqual(current.selected_map_id, 2)
|
|
np.testing.assert_allclose(
|
|
current.tau_master_raw, 2.0 * source.tau_master_raw
|
|
)
|
|
|
|
def test_power_endpoint_is_zero_for_exact_source_alignment(self):
|
|
tau_s = np.array([[1.0, 2.0], [-1.0, 0.5]])
|
|
qd_s = np.array([[0.3, -0.1], [0.2, 0.4]])
|
|
tau_m = tau_s.copy()
|
|
qd_m = qd_s.copy()
|
|
mismatch = normalized_actual_power_mismatch(
|
|
tau_m, qd_m, tau_s, qd_s, np.array([0.01, 0.02])
|
|
)
|
|
self.assertAlmostEqual(mismatch, 0.0)
|
|
|
|
def test_duplicate_out_of_order_and_timeout_are_explicit(self):
|
|
packets = [
|
|
ForwardPacket(i, i, i * 0.01, i, np.zeros(2), np.zeros(2))
|
|
for i in range(3)
|
|
]
|
|
# seq 0 arrives after seq 1, and seq 1 is duplicated.
|
|
channel = DeterministicChannel(
|
|
[
|
|
NetworkTraceEntry(0, 0.04),
|
|
NetworkTraceEntry(1, 0.01, duplicate=True),
|
|
NetworkTraceEntry(2, 0.02),
|
|
]
|
|
)
|
|
for packet in packets:
|
|
channel.send(packet, packet.source_time)
|
|
receiver = PacketReceiver(timeout_s=0.05)
|
|
receptions = []
|
|
for delivery in channel.poll(0.05):
|
|
receptions.append(receiver.accept(delivery))
|
|
|
|
self.assertTrue(receptions[0].accepted)
|
|
self.assertIn(
|
|
PacketRejectReason.DUPLICATE_OR_STALE,
|
|
[item.reason for item in receptions if not item.accepted],
|
|
)
|
|
held = receiver.sample(0.051)
|
|
self.assertIn(held.state, (PacketState.ACTIVE, PacketState.RECOVERING))
|
|
self.assertEqual(receiver.sample(0.2).state, PacketState.TIMED_OUT)
|
|
|
|
def test_default_network_trace_retains_historical_draw_sequence(self):
|
|
trace = generate_network_trace(
|
|
4,
|
|
base_delay_s=0.05,
|
|
jitter_s=0.01,
|
|
loss_probability=0.25,
|
|
duplicate_probability=0.4,
|
|
corrupt_probability=0.1,
|
|
seed=123,
|
|
)
|
|
expected = (
|
|
(0.05364703726496287, True, True, False),
|
|
(0.04351811802170061, False, False, False),
|
|
(0.056395091231860046, False, False, False),
|
|
(0.05648483192194823, True, False, False),
|
|
)
|
|
for entry, values in zip(trace, expected, strict=True):
|
|
self.assertEqual(
|
|
(
|
|
entry.delay_s,
|
|
entry.lost,
|
|
entry.duplicate,
|
|
entry.corrupt,
|
|
),
|
|
values,
|
|
)
|
|
|
|
def test_common_random_numbers_fix_packet_draw_alignment(self):
|
|
common = {
|
|
"count": 32,
|
|
"base_delay_s": 0.1,
|
|
"loss_probability": 0.35,
|
|
"duplicate_probability": 0.25,
|
|
"corrupt_probability": 0.15,
|
|
"seed": 91,
|
|
"common_random_numbers": True,
|
|
}
|
|
zero_jitter = generate_network_trace(jitter_s=0.0, **common)
|
|
small_jitter = generate_network_trace(jitter_s=0.01, **common)
|
|
large_jitter = generate_network_trace(jitter_s=0.03, **common)
|
|
|
|
decisions = lambda trace: [
|
|
(entry.lost, entry.duplicate, entry.corrupt)
|
|
for entry in trace
|
|
]
|
|
self.assertEqual(decisions(zero_jitter), decisions(small_jitter))
|
|
self.assertEqual(decisions(small_jitter), decisions(large_jitter))
|
|
small_latent = np.array(
|
|
[(entry.delay_s - 0.1) / 0.01 for entry in small_jitter]
|
|
)
|
|
large_latent = np.array(
|
|
[(entry.delay_s - 0.1) / 0.03 for entry in large_jitter]
|
|
)
|
|
np.testing.assert_allclose(small_latent, large_latent)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|