Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
163 changes: 94 additions & 69 deletions dimos/manipulation/execution_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,8 @@

from __future__ import annotations

from collections.abc import Sequence
from collections.abc import Callable, Sequence
from dataclasses import dataclass
import math
import threading
import time
Expand All @@ -40,6 +41,14 @@ class _PlanRejectedError(Exception):
"""Expected rejection while mapping a generated plan."""


@dataclass
class _Execution:
"""Keep a waiter's result attached to the execution it started waiting for."""

result: ExecutionResult | None = None
active: bool = False


class PlanExecutionManager:
"""Own mapping, dispatch, and polling for one trajectory execution."""

Expand All @@ -50,6 +59,7 @@ def __init__(
coordinator: ControlCoordinator,
default_timeout: float,
poll_interval: float = 0.1,
on_result: Callable[[ExecutionResult], None] | None = None,
) -> None:
self._joint_names = frozenset(joint_names)
if not self._joint_names or len(self._joint_names) != len(joint_names):
Expand All @@ -59,16 +69,16 @@ def __init__(
self._state_lock = threading.Lock()
self._default_timeout = default_timeout
self._poll_interval = poll_interval
self._active = False
self._latest_result: ExecutionResult | None = None
self._execution = _Execution()
self._on_result = on_result

@property
def status(self) -> ExecutionStatus:
"""Return the latest known execution status."""
with self._state_lock:
if self._latest_result is None:
if self._execution.result is None:
return ExecutionStatus.IDLE
return self._latest_result.status
return self._execution.result.status

def execute(
self,
Expand All @@ -80,12 +90,15 @@ def execute(
"""Dispatch a plan and optionally poll until it reaches a terminal state."""
with self._operation_lock:
with self._state_lock:
if self._active:
if self._execution.active:
return ExecutionResult(ExecutionStatus.REJECTED, "Another trajectory is active")
execution = self._execution = _Execution()
try:
trajectory = self._prepare_trajectory(plan)
except _PlanRejectedError as exc:
return ExecutionResult(ExecutionStatus.REJECTED, str(exc))
rejected = ExecutionResult(ExecutionStatus.REJECTED, str(exc))
self._store(rejected, active=False)
return rejected

try:
result = self._coordinator.execute_trajectory(trajectory)
Expand Down Expand Up @@ -116,35 +129,41 @@ def execute(

if not blocking:
return accepted
return self.wait(timeout)
return self._wait(execution, timeout)

def wait(self, timeout: float | None = None) -> ExecutionResult:
"""Poll JTT status until terminal, preserving the active execution on timeout."""
with self._state_lock:
execution = self._execution
return self._wait(execution, timeout)

def _wait(self, execution: _Execution, timeout: float | None) -> ExecutionResult:
wait_timeout = self._default_timeout if timeout is None else timeout
if not math.isfinite(wait_timeout) or wait_timeout < 0.0:
return ExecutionResult(ExecutionStatus.REJECTED, "timeout must be finite and >= 0")
with self._state_lock:
latest = self._latest_result
active = self._active
if latest is None:
return ExecutionResult(ExecutionStatus.NO_EXECUTION, "No execution exists")
if not active:
return latest

deadline = time.monotonic() + wait_timeout
while True:
status = self._get_status()
if isinstance(status, ExecutionResult):
return status
mapped = self._result_from_status(status)
if mapped.status in {
ExecutionStatus.COMPLETED,
ExecutionStatus.ABORTED,
ExecutionStatus.FAULT,
}:
self._store(mapped, active=False)
return mapped
self._store(mapped, active=True)
# Serialize each RPC and its result, but let cancel/execute run between polls.
with self._operation_lock:
with self._state_lock:
latest = execution.result
active = execution.active
if latest is None:
return ExecutionResult(ExecutionStatus.NO_EXECUTION, "No execution exists")
if not active:
return latest
status = self._get_status()
if isinstance(status, ExecutionResult):
return status
mapped = self._result_from_status(status)
if mapped.status in {
ExecutionStatus.COMPLETED,
ExecutionStatus.ABORTED,
ExecutionStatus.FAULT,
}:
self._store(mapped, active=False)
return mapped
self._store(mapped, active=True)
remaining = deadline - time.monotonic()
if remaining <= 0.0:
return ExecutionResult(
Expand All @@ -167,46 +186,49 @@ def cancel(self, timeout: float = 1.0) -> ExecutionResult:
)
self._store(result, active=False)
return result
if cancellation.status is TrajectoryCancellationStatus.UNCERTAIN:
result = ExecutionResult(
ExecutionStatus.UNCERTAIN,
cancellation.message or "Coordinator cancellation outcome is uncertain",
)
self._store(result, active=False)
return result
if cancellation.status is TrajectoryCancellationStatus.UNCERTAIN:
result = ExecutionResult(
ExecutionStatus.UNCERTAIN,
cancellation.message or "Coordinator cancellation outcome is uncertain",
)
self._store(result, active=False)
return result

with self._state_lock:
latest = self._latest_result
status = self._get_status()
if isinstance(status, ExecutionResult):
return status
mapped = self._result_from_status(status)
if mapped.status in {
ExecutionStatus.COMPLETED,
ExecutionStatus.ABORTED,
ExecutionStatus.FAULT,
}:
self._store(mapped, active=False)
return mapped
if cancellation.status is TrajectoryCancellationStatus.CANCELLED:
return self.wait(timeout)
if latest is not None and latest.status in {
ExecutionStatus.COMPLETED,
ExecutionStatus.ABORTED,
ExecutionStatus.FAULT,
}:
return latest
if status.state is TrajectoryState.IDLE:
result = ExecutionResult(ExecutionStatus.NO_EXECUTION, cancellation.message)
self._store(result, active=False)
return result
result = ExecutionResult(
ExecutionStatus.UNCERTAIN,
"Coordinator reported no active trajectory while JTT is still executing",
trajectory_status=status,
)
self._store(result, active=False)
return result
with self._state_lock:
latest = self._execution.result
status = self._get_status()
if isinstance(status, ExecutionResult):
return status
mapped = self._result_from_status(status)
if mapped.status in {
ExecutionStatus.COMPLETED,
ExecutionStatus.ABORTED,
ExecutionStatus.FAULT,
}:
self._store(mapped, active=False)
return mapped
if cancellation.status is TrajectoryCancellationStatus.CANCELLED:
self._store(mapped, active=True)
execution = self._execution
else:
if latest is not None and latest.status in {
ExecutionStatus.COMPLETED,
ExecutionStatus.ABORTED,
ExecutionStatus.FAULT,
}:
return latest
if status.state is TrajectoryState.IDLE:
result = ExecutionResult(ExecutionStatus.NO_EXECUTION, cancellation.message)
self._store(result, active=False)
return result
result = ExecutionResult(
ExecutionStatus.UNCERTAIN,
"Coordinator reported no active trajectory while JTT is still executing",
trajectory_status=status,
)
self._store(result, active=False)
return result
return self._wait(execution, timeout)

def _get_status(self) -> TrajectoryStatus | ExecutionResult:
try:
Expand Down Expand Up @@ -244,9 +266,12 @@ def _result_from_status(status: TrajectoryStatus) -> ExecutionResult:
return ExecutionResult(mapped, status.error, trajectory_status=status)

def _store(self, result: ExecutionResult, *, active: bool) -> None:
"""Store and project a result while the caller holds the operation lock."""
with self._state_lock:
self._latest_result = result
self._active = active
self._execution.result = result
self._execution.active = active
if self._on_result is not None:
self._on_result(result)

def _prepare_trajectory(self, plan: GeneratedPlan) -> JointTrajectory:
if not isinstance(plan, GeneratedPlan):
Expand Down
29 changes: 19 additions & 10 deletions dimos/manipulation/manipulation_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -374,6 +374,7 @@ def _tf_publish_loop(self) -> None:
@rpc
def get_state(self) -> ManipulationSnapshot:
"""Return one snapshot containing every planning group."""
self._refresh_execution_status()
groups: dict[PlanningGroupID, PlanningGroupState] = {}
if self._world_monitor is not None:
for group in self._world_monitor.planning_groups.list():
Expand All @@ -398,11 +399,7 @@ def get_state(self) -> ManipulationSnapshot:
operation_status = OperationStatus[self._state.name]
error = self._error_message or None
has_pending_plan = self._last_plan is not None
execution_status = (
self._execution_manager.status
if hasattr(self, "_execution_manager")
else ExecutionStatus.IDLE
)
execution_status = self._execution_manager.status
Comment thread
TomCC7 marked this conversation as resolved.
Outdated
return ManipulationSnapshot(
timestamp=time.time(),
operation_status=operation_status,
Expand All @@ -414,9 +411,18 @@ def get_state(self) -> ManipulationSnapshot:

def get_operation_status(self) -> OperationStatus:
"""Return the current operation status without collecting telemetry."""
self._refresh_execution_status()
with self._lock:
return OperationStatus[self._state.name]

def _refresh_execution_status(self) -> None:
"""Make one synchronous status RPC, without waiting for motion to finish."""
if self._execution_manager.status in {
ExecutionStatus.ACCEPTED,
ExecutionStatus.EXECUTING,
}:
self.wait_for_execution(timeout=0.0)
Comment thread
TomCC7 marked this conversation as resolved.

@rpc
def get_error(self) -> str:
"""Get last error message.
Expand All @@ -441,7 +447,6 @@ def cancel(self) -> ExecutionResult:
self._dismiss_preview(plan.group_ids)
if is_planning and result.status is ExecutionStatus.NO_EXECUTION:
result = ExecutionResult(ExecutionStatus.ABORTED, "Planning cancelled")
self._apply_execution_result(result)
return result

@rpc
Expand Down Expand Up @@ -1204,6 +1209,7 @@ def _initialize_execution(self) -> None:
joint_names=self.config.model.joint_names,
coordinator=self._control_coordinator,
default_timeout=self.config.execution_timeout,
on_result=self._apply_execution_result,
)

@rpc
Expand All @@ -1224,19 +1230,22 @@ def execute(self, blocking: bool = True, timeout: float | None = None) -> Execut
except Exception as exc:
logger.exception("Failed to dispatch generated plan")
result = ExecutionResult(ExecutionStatus.UNCERTAIN, str(exc))
self._apply_execution_result(result)
self._apply_execution_result(result)
return result

@rpc
def wait_for_execution(self, timeout: float | None = None) -> ExecutionResult:
"""Wait for the active trajectory or return its cached terminal result."""
result = self._execution_manager.wait(timeout)
self._apply_execution_result(result)
return result
return self._execution_manager.wait(timeout)

def _apply_execution_result(self, result: ExecutionResult) -> None:
"""Mirror execution ownership into the broader manipulation snapshot."""
with self._lock:
if (
self._state is ManipulationState.PLANNING
and result.status is ExecutionStatus.NO_EXECUTION
):
result = ExecutionResult(ExecutionStatus.ABORTED, "Planning cancelled")
if result.status in {ExecutionStatus.ACCEPTED, ExecutionStatus.EXECUTING}:
self._state = ManipulationState.EXECUTING
self._error_message = ""
Expand Down
Loading
Loading