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
92 changes: 56 additions & 36 deletions controls/sae_2025_ws/src/uav/uav/runtime/ModeManager.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import os
from pathlib import Path
from time import time
from typing import Any, Callable, Iterator, get_type_hints
from typing import Any, Callable, Iterator, cast, get_type_hints

from rclpy.node import Node
from std_srvs.srv import Trigger
Expand All @@ -30,15 +30,15 @@

@dataclass(frozen=True)
class _RawNodeApi:
create_publisher: Callable[..., object]
create_subscription: Callable[..., object]
create_client: Callable[..., object]
create_service: Callable[..., object]
create_timer: Callable[..., object]
destroy_publisher: Callable[[object], bool | None]
destroy_subscription: Callable[[object], bool | None]
destroy_client: Callable[[object], bool | None]
destroy_timer: Callable[[object], bool | None] | None
create_publisher: Callable[..., Any]
create_subscription: Callable[..., Any]
create_client: Callable[..., Any]
create_service: Callable[..., Any]
create_timer: Callable[..., Any]
destroy_publisher: Callable[..., bool | None]
destroy_subscription: Callable[..., bool | None]
destroy_client: Callable[..., bool | None]
destroy_timer: Callable[..., bool | None] | None


MISSION_STARTED_MARKER_ENV = "PENNAIR_MISSION_STARTED_MARKER_PATH"
Expand All @@ -57,15 +57,15 @@ def __init__(
peer_stale_timeout_s: float = 0.5,
) -> None:
super().__init__(node_name)
self.vehicle = None
self.modes = {}
self.transitions = {}
self.active_mode = None
self.vehicle: Vehicle | None = None
self.modes: dict[str, Mode] = {}
self.transitions: dict[str, dict[str, str]] = {}
self.active_mode: str | None = None
self.last_update_time = time()
self._vision_clients = {}
self.timer = None
self._vision_clients: dict[str, Any] = {}
self.timer: Any | None = None
self.auto_launch = bool(auto_launch)
self._auto_launch_timer = None
self._auto_launch_timer: Any | None = None
self._runtime_vehicle_name = normalize_vehicle_name(vehicle_name)
self.peer_heartbeat_hz = float(peer_heartbeat_hz)
self.peer_stale_timeout_s = float(peer_stale_timeout_s)
Expand All @@ -78,8 +78,8 @@ def __init__(
"peer_stale_timeout_s must be positive, "
f"got {self.peer_stale_timeout_s!r}."
)
self._shared_mode_state = {}
self._current_comm_builder = None
self._shared_mode_state: dict[str, dict[str, Any]] = {}
self._current_comm_builder: ModeCommBuilder | None = None
self._runtime_closed = False
self._initialize_comms_runtime(vehicle_name=vehicle_name)
self.start_mission_service = self.create_service(
Expand All @@ -90,7 +90,10 @@ def __init__(
self._auto_launch_timer = self.create_timer(0.1, self._maybe_auto_launch)

def get_active_mode(self) -> Mode:
return self.modes[self.active_mode]
active_mode = self.active_mode
if active_mode is None:
raise RuntimeError("No active mode is set.")
return self.modes[active_mode]

def _now_seconds(self) -> float:
return self.get_clock().now().nanoseconds * 1e-9
Expand Down Expand Up @@ -204,11 +207,14 @@ def setup_vision(self, vision_nodes: list[str]) -> None:
)

@property
def vision_clients(self) -> dict:
def vision_clients(self) -> dict[str, Any]:
return self._vision_clients

def _connect_vision_client(self, vision_class):
service_name = self.vehicle.vision_service_name(vision_class)
def _connect_vision_client(self, vision_class: type[Any]) -> tuple[Any, str]:
vehicle = self.vehicle
if vehicle is None:
raise ValueError("Vision nodes require an active vehicle camera contract.")
service_name = vehicle.vision_service_name(vision_class)
while True:
client = super().create_client(vision_class.srv, service_name)
if client.wait_for_service(timeout_sec=1.0):
Expand Down Expand Up @@ -258,14 +264,14 @@ def _use_comm_builder(self, builder: ModeCommBuilder) -> Iterator[None]:
finally:
self._current_comm_builder = previous_builder

def _raw_method(self, name: str):
def _raw_method(self, name: str) -> Callable[..., Any]:
raw_node_api = getattr(self, "_raw_node_api", None)
if raw_node_api is not None:
return getattr(raw_node_api, name)
raw_method = getattr(super(ModeManager, self), name, None)
if raw_method is not None:
return raw_method
return getattr(Node, name)
return cast(Callable[..., Any], getattr(Node, name))

def _instantiate_mode(
self,
Expand Down Expand Up @@ -304,7 +310,7 @@ def _create_entity(
name: str,
args: tuple[Any, ...],
kwargs: dict[str, Any],
):
) -> Any:
if not hasattr(self, "_runtime_vehicle_name"):
return self._raw_method(f"create_{kind}")(
interface_type,
Expand Down Expand Up @@ -335,7 +341,7 @@ def _create_entity(
kwargs=kwargs,
)

def _destroy_entity(self, *, kind: str, entity) -> bool:
def _destroy_entity(self, *, kind: str, entity: object) -> bool:
registry = getattr(self, "_managed_registry", None)
if registry is None:
return bool(self._raw_method(f"destroy_{kind}")(entity))
Expand Down Expand Up @@ -409,7 +415,10 @@ def transition(self, state: str) -> str:
self.get_logger().info(
f"Transitioning from {self.active_mode} based on state {state}."
)
return self.transitions[self.active_mode][state]
active_mode = self.active_mode
if active_mode is None:
raise RuntimeError("No active mode is set.")
return self.transitions[active_mode][state]

def switch_mode(self, mode_name: str) -> None:
if self.active_mode:
Expand Down Expand Up @@ -502,6 +511,9 @@ def start_mission(self) -> bool:
self._write_mission_started_marker()
return True

def spin_once(self) -> None:
raise NotImplementedError("ModeManager subclasses must implement spin_once().")

def _cancel_auto_launch_timer(self) -> None:
if self._auto_launch_timer is None:
return
Expand Down Expand Up @@ -544,7 +556,9 @@ def handle_mode_state(self, state: str) -> None:
else:
self.switch_mode(next_mode)

def create_publisher(self, msg_type, topic: str, *args, **kwargs):
def create_publisher(
self, msg_type: object, topic: str, *args: Any, **kwargs: Any
) -> Any:
return self._create_entity(
kind="publisher",
interface_type=msg_type,
Expand All @@ -553,7 +567,9 @@ def create_publisher(self, msg_type, topic: str, *args, **kwargs):
kwargs=kwargs,
)

def create_subscription(self, msg_type, topic: str, *args, **kwargs):
def create_subscription(
self, msg_type: object, topic: str, *args: Any, **kwargs: Any
) -> Any:
return self._create_entity(
kind="subscription",
interface_type=msg_type,
Expand All @@ -562,7 +578,9 @@ def create_subscription(self, msg_type, topic: str, *args, **kwargs):
kwargs=kwargs,
)

def create_client(self, srv_type, srv_name: str, *args, **kwargs):
def create_client(
self, srv_type: object, srv_name: str, *args: Any, **kwargs: Any
) -> Any:
return self._create_entity(
kind="client",
interface_type=srv_type,
Expand All @@ -571,7 +589,9 @@ def create_client(self, srv_type, srv_name: str, *args, **kwargs):
kwargs=kwargs,
)

def create_service(self, srv_type, srv_name: str, *args, **kwargs):
def create_service(
self, srv_type: object, srv_name: str, *args: Any, **kwargs: Any
) -> Any:
return self._create_entity(
kind="service",
interface_type=srv_type,
Expand All @@ -580,13 +600,13 @@ def create_service(self, srv_type, srv_name: str, *args, **kwargs):
kwargs=kwargs,
)

def destroy_publisher(self, publisher) -> bool:
def destroy_publisher(self, publisher: object) -> bool:
return self._destroy_entity(kind="publisher", entity=publisher)

def destroy_subscription(self, subscription) -> bool:
def destroy_subscription(self, subscription: object) -> bool:
return self._destroy_entity(kind="subscription", entity=subscription)

def destroy_client(self, client) -> bool:
def destroy_client(self, client: object) -> bool:
return self._destroy_entity(kind="client", entity=client)

def _close_runtime_helpers(self) -> None:
Expand All @@ -600,7 +620,7 @@ def _close_runtime_helpers(self) -> None:
if peer_connections is not None:
peer_connections.close()

def destroy_node(self) -> bool:
def destroy_node(self) -> Any:
self._clear_mission_started_marker()
if getattr(self, "active_mode", None):
self._deactivate_active_mode()
Expand Down
74 changes: 42 additions & 32 deletions controls/sae_2025_ws/src/uav/uav/runtime/UAVModeManager.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,12 @@
from time import time
from typing import cast

from px4_msgs.msg import VehicleStatus
from std_srvs.srv import Trigger

from uav.vehicles.AirframeClass import AirframeClass
from uav.vehicles.Multicopter import Multicopter
from uav.vehicles.UAV import UAV
from uav.vehicles.VTOL import VTOL
from uav.modes.uav.LandingMode import LandingMode
from .ModeManager import ModeManager
Expand Down Expand Up @@ -71,72 +73,79 @@ def __init__(
)
self.setup_modes(mission_spec)

def _uav_vehicle(self) -> UAV | None:
vehicle = self.vehicle
if vehicle is None:
return None
return cast(UAV, vehicle)

def _auto_launch_ready(self) -> bool:
if self.vehicle is None:
vehicle = self._uav_vehicle()
if vehicle is None:
return False
return (
self.vehicle.vehicle_status is not None
and self.vehicle.vehicle_attitude is not None
and self.vehicle.yaw is not None
and self.vehicle.local_position is not None
and self.vehicle.global_position is not None
and bool(self.vehicle.flight_check)
vehicle.vehicle_status is not None
and vehicle.vehicle_attitude is not None
and vehicle.yaw is not None
and vehicle.local_position is not None
and vehicle.global_position is not None
and bool(vehicle.flight_check)
)

def trigger_failsafe(self, request, response):
self.get_logger().info("Failsafe triggered via service call")
if self.vehicle is not None and hasattr(self.vehicle, "failsafe_trigger"):
self.vehicle.failsafe_trigger = True
self.vehicle.failsafe = (
self.vehicle.failsafe_px4 or self.vehicle.failsafe_trigger
)
vehicle = self._uav_vehicle()
if vehicle is not None:
vehicle.failsafe_trigger = True
vehicle.failsafe = vehicle.failsafe_px4 or vehicle.failsafe_trigger
response.success = True
response.message = "Failsafe triggered."
return response

def spin_once(self) -> None:
current_time = time()
if self.vehicle is None:
vehicle = self._uav_vehicle()
if vehicle is None:
return

if self.active_mode is None:
self.switch_mode("start")

if self.vehicle.failsafe:
if not self.vehicle.emergency_landing:
self.vehicle.hover()
if vehicle.failsafe:
if not vehicle.emergency_landing:
vehicle.hover()
self.get_logger().warn("Failsafe: Switching to AUTO_LOITER mode.")
self.vehicle.emergency_landing = True
vehicle.emergency_landing = True
if (
self.vehicle.nav_state == VehicleStatus.NAVIGATION_STATE_AUTO_LOITER
or self.vehicle.arm_state != VehicleStatus.ARMING_STATE_ARMED
vehicle.nav_state == VehicleStatus.NAVIGATION_STATE_AUTO_LOITER
or vehicle.arm_state != VehicleStatus.ARMING_STATE_ARMED
):
self.vehicle.land()
vehicle.land()
self.get_logger().warn("Failsafe: Initiating landing.")
return

if self.servo_only:
self._run_active_mode(current_time)
return

if not self.vehicle.origin_set:
self.vehicle.set_origin()
if not vehicle.origin_set:
vehicle.set_origin()

if self.vehicle.arm_state != VehicleStatus.ARMING_STATE_ARMED:
if vehicle.arm_state != VehicleStatus.ARMING_STATE_ARMED:
self.get_logger().info(
f"UAV is not armed. Current arm state: {self.vehicle.arm_state}"
f"UAV is not armed. Current arm state: {vehicle.arm_state}"
)
if (
self.active_mode is not None
and isinstance(self.get_active_mode(), LandingMode)
and self.vehicle.nav_state != VehicleStatus.NAVIGATION_STATE_AUTO_LAND
and vehicle.nav_state != VehicleStatus.NAVIGATION_STATE_AUTO_LAND
):
self.get_logger().info("Successfully Landed UAV")
self.get_logger().info("Finishing Mission")
self.destroy_node()
return

if self.vehicle.attempted_takeoff and self.active_mode is not None:
if vehicle.attempted_takeoff and self.active_mode is not None:
self.get_logger().error(
"UAV disarmed unexpectedly after takeoff attempt. Terminating to prevent infinite cycle."
)
Expand All @@ -146,27 +155,28 @@ def spin_once(self) -> None:
self.destroy_node()
return

self.vehicle.arm()
vehicle.arm()
self.get_logger().info("Arming UAV")
self.start_time = current_time
return

if self.vehicle.local_position is None or self.vehicle.global_position is None:
if vehicle.local_position is None or vehicle.global_position is None:
return

self.vehicle.publish_offboard_control_heartbeat_signal()
vehicle.publish_offboard_control_heartbeat_signal()
self._run_active_mode(current_time)

if self.vehicle.nav_state == VehicleStatus.NAVIGATION_STATE_AUTO_LAND:
if vehicle.nav_state == VehicleStatus.NAVIGATION_STATE_AUTO_LAND:
self.get_logger().info("Landing")

def handle_mode_state(self, state: str) -> None:
if state == "error":
self.get_logger().error(
f"Error in mode {self.active_mode}. Switching to failsafe."
)
if self.vehicle is not None:
self.vehicle.failsafe = True
vehicle = self._uav_vehicle()
if vehicle is not None:
vehicle.failsafe = True
return
if state == "terminate":
self.get_logger().info("Mission has completed.")
Expand Down
Loading
Loading