Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions controls/sae_2025_ws/src/uav/test/test_comms_diagnostics.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
from __future__ import annotations

from types import SimpleNamespace

from uav.runtime.comms_diagnostics import (
log_bootstrap_diagnostics,
log_comms_diag,
topic_graph_counts,
)


class _FakeLogger:
def __init__(self) -> None:
self.messages: list[str] = []

def info(self, message: str) -> None:
self.messages.append(message)


def test_log_comms_diag_prefixes_and_serializes_fields() -> None:
logger = _FakeLogger()

log_comms_diag(
logger,
"MODE",
"runtime comm snapshot",
connection_status={"payload_1": False},
peer_vehicle_names=("payload_1",),
ready=False,
)

assert len(logger.messages) == 1
assert logger.messages[0].startswith("COMMS DIAG | MODE | runtime comm snapshot")
assert 'connection_status={"payload_1": false}' in logger.messages[0]
assert 'peer_vehicle_names=["payload_1"]' in logger.messages[0]
assert "ready=false" in logger.messages[0]


def test_log_bootstrap_diagnostics_emits_boot_and_network_snapshots(
monkeypatch,
) -> None:
logger = _FakeLogger()

monkeypatch.setattr("uav.runtime.comms_diagnostics.socket.gethostname", lambda: "test-host")
monkeypatch.setattr("uav.runtime.comms_diagnostics.os.getpid", lambda: 4242)
monkeypatch.setattr(
"uav.runtime.comms_diagnostics._run_snapshot_command",
lambda command, timeout_sec=1.5: (
0,
"snapshot for " + " ".join(command),
),
)

log_bootstrap_diagnostics(
logger,
runtime_kind="payload_mission",
vehicle_name="payload_0",
mission_path="/tmp/mission.yaml",
mission_target="payload",
auto_launch=True,
peer_heartbeat_hz=10.0,
peer_stale_timeout_s=0.5,
peer_vehicle_names=("payload_1",),
vision_nodes=(),
)

assert len(logger.messages) == 5
assert logger.messages[0].startswith("COMMS DIAG | BOOT |")
assert "vehicle_name=payload_0" in logger.messages[0]
assert "hostname=test-host" in logger.messages[0]
assert any("COMMS DIAG | NET | snapshot=ip_addr" in message for message in logger.messages)
assert any(
"snapshot for nmcli -t -f NAME,TYPE,DEVICE,STATE con show --active" in message
for message in logger.messages
)


def test_topic_graph_counts_handles_missing_methods_and_errors() -> None:
node = SimpleNamespace(
count_publishers=lambda _topic: 2,
count_subscribers=lambda _topic: (_ for _ in ()).throw(RuntimeError("boom")),
)

counts = topic_graph_counts(node, "/payload_1/mode_manager/heartbeat")

assert counts["publishers"] == 2
assert counts["subscribers"] == "error: boom"
41 changes: 41 additions & 0 deletions controls/sae_2025_ws/src/uav/test/test_runtime_behavior.py
Original file line number Diff line number Diff line change
Expand Up @@ -694,6 +694,47 @@ def test_run_active_mode_handles_update_error():
assert handled == ["error"]


def test_run_active_mode_logs_connectivity_transition_once():
_require_runtime_support()

start_mode = _TrackingMode("start")
start_mode.peer_vehicle_names = ("payload_1",)
start_mode.connection_ready = lambda _connection_status: False
manager = _make_mode_manager()
manager.modes = {"start": start_mode}
manager.active_mode = "start"
manager._peer_connections = SimpleNamespace(
status=lambda: {"payload_1": False},
peer_names=("payload_1",),
debug_snapshot=lambda: {
"local_heartbeat_topic": "/payload_0/mode_manager/heartbeat",
"remote_heartbeat_topics": {
"payload_1": "/payload_1/mode_manager/heartbeat"
},
"peer_vehicle_names": ("payload_1",),
"connection_status": {"payload_1": False},
"last_seen_age_s": {"payload_1": None},
},
close=lambda: None,
)
manager.count_publishers = lambda _topic: 0
manager.count_subscribers = lambda _topic: 1

ModeManager._run_active_mode(manager, 1.0)
ModeManager._run_active_mode(manager, 2.0)

diag_messages = [
msg
for level, msg in manager._logger.messages
if level == "info"
and "COMMS DIAG | MODE | runtime comm snapshot" in msg
and "reason=mode_connectivity_transition" in msg
]
assert len(diag_messages) == 1
assert "connection_ready=false" in diag_messages[0]
assert start_mode.disconnects == [(1.0, {"payload_1": False}), (1.0, {"payload_1": False})]


def test_handle_mode_state_requires_exact_transition_label():
_require_runtime_support()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,13 @@ def __init__(
peer_name: 0 for peer_name in self._peer_names
}
self._last_disconnect_signature: tuple[str, ...] = ()
self._peer_state_topics = {
peer_name: vehicle.namespaced_path(local_topic, namespace=f"/{peer_name}")
for peer_name in self._peer_names
}
self._last_operating_state: str | None = None
self._logged_first_peer_receipt: set[str] = set()
self._logged_first_shared_receipt: set[str] = set()

def _now(self) -> float:
return self.node.get_clock().now().nanoseconds * 1e-9
Expand All @@ -108,6 +115,12 @@ def _on_peer_message(self, peer_name: str, message: String) -> None:
self._peer_message_counts[peer_name] = (
self._peer_message_counts.get(peer_name, 0) + 1
)
if peer_name not in self._logged_first_peer_receipt:
self._logged_first_peer_receipt.add(peer_name)
self.log(
f"COMMS DIAG | PEER_TEST | first peer-local message received "
f"peer_name={peer_name} topic={self._peer_state_topics.get(peer_name, '')}"
)

def _on_shared_message(self, message: String) -> None:
self._shared_message_count += 1
Expand All @@ -119,6 +132,25 @@ def _on_shared_message(self, message: String) -> None:
self._shared_remote_message_counts[sender] = (
self._shared_remote_message_counts.get(sender, 0) + 1
)
if sender not in self._logged_first_shared_receipt:
self._logged_first_shared_receipt.add(sender)
self.log(
f"COMMS DIAG | PEER_TEST | first shared remote message received "
f"peer_name={sender} topic={self._shared_topic}"
)

def _log_state_transition(
self, *, state: str, disconnected_peers: tuple[str, ...] = ()
) -> None:
if state == self._last_operating_state:
return
self._last_operating_state = state
self.log(
"COMMS DIAG | PEER_TEST | state transition "
f"state={state} disconnected_peers={list(disconnected_peers)} "
f"peer_total={sum(self._peer_message_counts.values())} "
f"shared_remote_total={sum(self._shared_remote_message_counts.values())}"
)

def _status_payload(
self, *, state: str, disconnected_peers: tuple[str, ...]
Expand Down Expand Up @@ -183,13 +215,22 @@ def on_enter(self) -> None:
self._peer_message_counts[peer_name] = 0
self._shared_remote_message_counts[peer_name] = 0
self._last_disconnect_signature = ()
self._last_operating_state = None
self._logged_first_peer_receipt.clear()
self._logged_first_shared_receipt.clear()
self.log(
f"starting peer fleet test for {self.vehicle.name}; expecting peers {self._peer_names}"
)
self.log(
"COMMS DIAG | PEER_TEST | topics "
f"local_topic={self._local_topic} shared_topic={self._shared_topic} "
f"peer_topics={self._peer_state_topics}"
)

def on_update(self, time_delta: float) -> None:
del time_delta
now = self._now()
self._log_state_transition(state="connected")
self._publish_status(now=now, state="connected")
self._log_status(now=now, state="connected")

Expand All @@ -207,6 +248,9 @@ def on_disconnect(
if disconnect_signature != self._last_disconnect_signature:
self.log("waiting for peers: " + ", ".join(disconnect_signature))
self._last_disconnect_signature = disconnect_signature
self._log_state_transition(
state="waiting", disconnected_peers=disconnect_signature
)
self._publish_status(
now=now,
state="waiting",
Expand Down
92 changes: 90 additions & 2 deletions controls/sae_2025_ws/src/uav/uav/runtime/ModeManager.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@

from uav.vehicles.Vehicle import Vehicle
from uav.modes.Mode import Mode
from .comms_diagnostics import log_comms_diag, topic_graph_counts
from .comm_naming import ManagedCommKind, resolve_managed_target
from .comm_policy import ManagedCommLifetime, managed_creation_phase_error
from .managed_comms import ModeCommBuilder
Expand Down Expand Up @@ -76,6 +77,7 @@ def __init__(
self._shared_mode_state = {}
self._current_comm_builder = None
self._runtime_closed = False
self._last_connectivity_diag_signature = None
self._initialize_comms_runtime(vehicle_name=vehicle_name)
self.start_mission_service = self.create_service(
Trigger, "mode_manager/start_mission", self._start_mission_callback
Expand All @@ -94,12 +96,71 @@ def shared_state_for(self, mode_or_class: object) -> dict[str, Any]:
return self._shared_mode_state.setdefault(key, {})

def _comm_debug_snapshot(self) -> dict[str, Any]:
peer_snapshot = None
if hasattr(self._peer_connections, "debug_snapshot"):
peer_snapshot = self._peer_connections.debug_snapshot()
return {
"peer_vehicle_names": tuple(self._peer_connections.peer_names),
"connection_status": self._peer_connections.status(),
"descriptors": self._managed_registry.debug_descriptors(),
"peer_snapshot": peer_snapshot,
}

def _diag_log(self, category: str, message: str | None = None, **fields) -> None:
log_comms_diag(self.get_logger(), category, message=message, **fields)

def _heartbeat_graph_snapshot(self) -> dict[str, dict[str, object]]:
peer_snapshot = getattr(self._peer_connections, "debug_snapshot", lambda: {})()
topics: dict[str, str] = {}
local_topic = peer_snapshot.get("local_heartbeat_topic")
if local_topic:
topics["local"] = str(local_topic)
for peer_name, topic in sorted(
dict(peer_snapshot.get("remote_heartbeat_topics", {})).items()
):
topics[peer_name] = str(topic)
return {
label: {"topic": topic, **topic_graph_counts(self, topic)}
for label, topic in topics.items()
}

def _log_comm_snapshot(
self,
*,
reason: str,
include_descriptors: bool = False,
**fields,
) -> None:
snapshot = self._comm_debug_snapshot()
descriptor_payload = None
if include_descriptors:
descriptor_payload = [
vars(descriptor).copy() for descriptor in snapshot["descriptors"]
]
payload = {
"reason": reason,
"peer_vehicle_names": snapshot["peer_vehicle_names"],
"connection_status": snapshot["connection_status"],
"peer_snapshot": snapshot["peer_snapshot"],
"heartbeat_graph": self._heartbeat_graph_snapshot(),
"descriptor_count": len(snapshot["descriptors"]),
"bound_descriptor_count": sum(
1 for descriptor in snapshot["descriptors"] if descriptor.bound
),
"descriptors": descriptor_payload,
}
payload.update(fields)
self._diag_log("MODE", "runtime comm snapshot", **payload)

def _handle_peer_connection_change(self, peer_name: str, *, connected: bool) -> None:
self._managed_registry.on_peer_connection_change(peer_name, connected=connected)
self._log_comm_snapshot(
reason="peer_connection_change",
peer_name=peer_name,
connected=connected,
active_mode=self.active_mode,
)

def _initialize_comms_runtime(self, *, vehicle_name: str) -> None:
self._runtime_vehicle_name = normalize_vehicle_name(vehicle_name)
raw_node = super(ModeManager, self)
Expand All @@ -122,13 +183,15 @@ def _initialize_comms_runtime(self, *, vehicle_name: str) -> None:
raw_destroy_publisher=self._raw_node_api.destroy_publisher,
raw_destroy_subscription=self._raw_node_api.destroy_subscription,
raw_destroy_client=self._raw_node_api.destroy_client,
diagnostic_logger=self._diag_log,
)
self._peer_connections = PeerConnectionTracker(
runtime_vehicle_name=self._runtime_vehicle_name,
peer_heartbeat_hz=self.peer_heartbeat_hz,
peer_stale_timeout_s=self.peer_stale_timeout_s,
now_seconds=self._now_seconds,
on_connection_change=self._managed_registry.on_peer_connection_change,
on_connection_change=self._handle_peer_connection_change,
diagnostic_logger=self._diag_log,
raw_create_publisher=self._raw_node_api.create_publisher,
raw_create_subscription=self._raw_node_api.create_subscription,
raw_destroy_publisher=self._raw_node_api.destroy_publisher,
Expand All @@ -141,6 +204,7 @@ def configure_peer_vehicle_names(
self, peer_vehicle_names: tuple[str, ...] | list[str] | set[str]
) -> None:
self._peer_connections.configure(peer_vehicle_names)
self._log_comm_snapshot(reason="after_peer_config", include_descriptors=True)

def setup_vision(self, vision_nodes: list[str]) -> None:
nodes_to_setup = [node for node in vision_nodes if node]
Expand Down Expand Up @@ -376,6 +440,7 @@ def switch_mode(self, mode_name: str) -> None:

if mode_name in self.modes:
self.active_mode = mode_name
self._last_connectivity_diag_signature = None
mode = self.get_active_mode()
builder = self._make_comm_builder(
owner=mode,
Expand All @@ -389,6 +454,11 @@ def switch_mode(self, mode_name: str) -> None:
except Exception:
self._managed_registry.destroy_for_owner(mode, lifetime="active")
raise
self._log_comm_snapshot(
reason="mode_activated",
include_descriptors=True,
active_mode=mode_name,
)
else:
self.get_logger().error(f"Mode {mode_name} not found.")

Expand All @@ -400,8 +470,25 @@ def _run_active_mode(self, current_time: float) -> None:
self.last_update_time = current_time
mode = self.get_active_mode()
connection_status = self._mode_connection_status(mode)
connection_ready = mode.connection_ready(connection_status)
connectivity_signature = (
self.active_mode,
bool(connection_ready),
tuple(sorted(connection_status.items())),
)
last_connectivity_diag_signature = getattr(
self, "_last_connectivity_diag_signature", None
)
if connectivity_signature != last_connectivity_diag_signature:
self._last_connectivity_diag_signature = connectivity_signature
self._log_comm_snapshot(
reason="mode_connectivity_transition",
active_mode=self.active_mode,
connection_ready=connection_ready,
connection_status=connection_status,
)
try:
if mode.connection_ready(connection_status):
if connection_ready:
mode.update(time_delta)
else:
mode.disconnect(time_delta, connection_status)
Expand All @@ -420,6 +507,7 @@ def _deactivate_active_mode(self) -> None:
mode_name = self.active_mode
mode = self.modes.get(mode_name)
self.active_mode = None
self._last_connectivity_diag_signature = None
if mode is None or not getattr(mode, "active", False):
return

Expand Down
Loading
Loading