Skip to content
Merged
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
20 changes: 19 additions & 1 deletion docs/examples/my_cool_arm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

import asyncio
import json
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
from typing import Any, AsyncIterator, Dict, List, Mapping, Optional, Tuple, Union

from viam.components.arm import Arm, JointPositions, KinematicsFileFormat, Pose
from viam.operations import run_with_operation
Expand Down Expand Up @@ -118,6 +118,24 @@ async def move_through_joint_positions(

self.is_stopped = True

async def move_through_joint_positions_streamed( # type: ignore
self,
batches: AsyncIterator[List[Arm.TrajectoryPoint]],
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
**kwargs,
) -> AsyncIterator[Arm.TrajectoryUpdate]:
# Each list the framework hands us is one wire batch. Walk it point by point to drive the
# arm, then send one update back per batch so the caller can watch progress and so a fault
# reaches it while the trajectory is still running.
self.is_stopped = False
async for batch in batches:
for point in batch:
self.joint_positions = JointPositions(values=point.positions)
yield Arm.TrajectoryUpdate()
self.is_stopped = True

async def get_3d_models(self, extra: Optional[Dict[str, Any]] = None, **kwargs) -> Mapping[str, Mesh]:
# Return the 3D meshes for this arm, keyed by name. This arm has none.
return {}
Expand Down
20 changes: 19 additions & 1 deletion examples/complex_module/src/arm/my_arm.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import asyncio
import os
from typing import Any, ClassVar, Dict, List, Mapping, Optional, Tuple
from typing import Any, AsyncIterator, ClassVar, Dict, List, Mapping, Optional, Tuple

from typing_extensions import Self

Expand Down Expand Up @@ -122,6 +122,24 @@ async def move_through_joint_positions(

self.is_stopped = True

async def move_through_joint_positions_streamed( # type: ignore
self,
batches: AsyncIterator[List[Arm.TrajectoryPoint]],
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
**kwargs,
) -> AsyncIterator[Arm.TrajectoryUpdate]:
# Each list the framework hands us is one wire batch. Walk it point by point to drive the
# arm, then send one update back per batch so the caller can watch progress and so a fault
# reaches it while the trajectory is still running.
self.is_stopped = False
async for batch in batches:
for point in batch:
self.joint_positions = JointPositions(values=point.positions)
yield Arm.TrajectoryUpdate()
self.is_stopped = True

async def get_3d_models(self, extra: Optional[Dict[str, Any]] = None, **kwargs) -> Mapping[str, Mesh]:
# This arm has no meshes to report.
return {}
Expand Down
15 changes: 15 additions & 0 deletions examples/server/v1/components.py
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,21 @@ async def move_through_joint_positions(
for position in positions:
self.joint_positions = position

async def move_through_joint_positions_streamed( # type: ignore
self,
batches: AsyncIterator[List[Arm.TrajectoryPoint]],
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
**kwargs,
) -> AsyncIterator[Arm.TrajectoryUpdate]:
self.is_stopped = False
async for batch in batches:
for point in batch:
self.joint_positions = JointPositions(values=point.positions)
yield Arm.TrajectoryUpdate()
self.is_stopped = True

async def get_3d_models(self, extra: Optional[Dict[str, Any]] = None, **kwargs) -> Mapping[str, Mesh]:
return {}

Expand Down
208 changes: 184 additions & 24 deletions src/viam/components/arm/arm.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,22 @@
import abc
from typing import Any, Dict, Final, List, Mapping, Optional, TypeAlias
from collections.abc import AsyncIterator, Mapping
from dataclasses import dataclass
from datetime import timedelta
from typing import Any, Final, TypeAlias

from google.protobuf.duration_pb2 import Duration

from viam.components import KinematicsReturn
from viam.components.component_base import ComponentBase
from viam.proto.component.arm import GetPropertiesResponse
from viam.proto.component.arm import (
GetPropertiesResponse,
JointAccelerations,
JointVelocities,
MoveThroughJointPositionsStreamedResponse,
)
from viam.proto.component.arm import (
TrajectoryPoint as TrajectoryPointPb,
)
from viam.resource.types import API, RESOURCE_NAMESPACE_RDK, RESOURCE_TYPE_COMPONENT

from . import JointPositions, Mesh, MoveOptions, Pose
Expand Down Expand Up @@ -36,12 +49,105 @@ class Arm(ComponentBase):

API: Final = API(RESOURCE_NAMESPACE_RDK, RESOURCE_TYPE_COMPONENT, "arm") # pyright: ignore [reportIncompatibleVariableOverride]

@dataclass
class KinematicConstraints:
"""
Optional per-waypoint kinematic constraints attached to a ``TrajectoryPoint``.

Velocities are required whenever constraints are present; accelerations are optional and
may only be given alongside velocities. Each list runs from the base joint out to the end
effector and must match the arm's degrees of freedom.
"""

velocities: list[float]
"""Target joint velocities at this waypoint. Rotational values in degrees per second,
translational values in mm per second."""

accelerations: list[float] | None = None
"""Optional target joint accelerations at this waypoint. Rotational values in
degrees per second squared, translational values in mm per second squared."""

@dataclass
class TrajectoryPoint:
"""
A single waypoint of a kinematized trajectory, as consumed by
``move_through_joint_positions_streamed``.

Point times must strictly increase across a stream, and the first point must have a
``time`` of zero.
"""

time: timedelta
"""Time at which this waypoint should be reached, measured from the start of the motion."""

positions: list[float]
"""Joint positions at this waypoint. Rotational values in degrees, translational values in mm."""

constraints: "Arm.KinematicConstraints | None" = None
"""Optional kinematic constraints at this waypoint."""

def to_proto(self) -> TrajectoryPointPb:
duration = Duration()
duration.FromTimedelta(self.time)
constraints_pb = None
if self.constraints is not None:
accelerations_pb = None
if self.constraints.accelerations is not None:
accelerations_pb = JointAccelerations(values=self.constraints.accelerations)
constraints_pb = TrajectoryPointPb.KinematicConstraints(
velocities=JointVelocities(values=self.constraints.velocities),
accelerations=accelerations_pb,
)
return TrajectoryPointPb(
time=duration,
positions=JointPositions(values=self.positions),
constraints=constraints_pb,
)

@classmethod
def from_proto(cls, proto: TrajectoryPointPb) -> "Arm.TrajectoryPoint":
constraints = None
if proto.HasField("constraints"):
accelerations = None
if proto.constraints.HasField("accelerations"):
accelerations = list(proto.constraints.accelerations.values)
constraints = Arm.KinematicConstraints(
velocities=list(proto.constraints.velocities.values),
accelerations=accelerations,
)
return cls(
time=proto.time.ToTimedelta(),
positions=list(proto.positions.values),
constraints=constraints,
)

@dataclass
class TrajectoryUpdate:
"""
An update reported by the arm as it executes a ``move_through_joint_positions_streamed`` trajectory.

The type is intentionally empty. The response is a ``oneof`` whose only branch today is an empty

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Per the comment, note that this type is slightly a lie: the real API Response type for streaming has a oneof message as the only field. But that oneof currently contains only one possible result: a BatchAck, which is itself empty. All of this is in furtherance of not breaking compatibility for old clients if and when additional response data types are added, which is anticipated. For now, it seems reasonable to keep the type silent, since the existence of a reply on the stream is implicitly a batch ack (it can't be anything else on the wire).

``BatchAck``, so receiving a response is itself the acknowledgment. The ``oneof`` exists so the arm's
replies can grow new branches without breaking existing clients on the wire; when a branch carries
data worth surfacing (``BatchAck``'s ``extra``, or a new branch entirely), this type grows to match.
"""

def to_proto(self) -> MoveThroughJointPositionsStreamedResponse:
# A received response is itself the acknowledgment, so send the default message and leave the
# oneof unset: the only branch is empty, nothing reads it today, and RDK and the C++ SDK send
# it unset as well. If a future branch carries data, set it here.
return MoveThroughJointPositionsStreamedResponse()

@classmethod
def from_proto(cls, proto: MoveThroughJointPositionsStreamedResponse) -> "Arm.TrajectoryUpdate":
return cls()

@abc.abstractmethod
async def get_end_position(
self,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
) -> Pose:
"""
Expand Down Expand Up @@ -69,8 +175,8 @@ async def move_to_position(
self,
pose: Pose,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
):
"""
Expand Down Expand Up @@ -101,8 +207,8 @@ async def move_to_joint_positions(
self,
positions: JointPositions,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
):
"""
Expand Down Expand Up @@ -132,11 +238,11 @@ async def move_to_joint_positions(
@abc.abstractmethod
async def move_through_joint_positions(
self,
positions: List[JointPositions],
options: Optional[MoveOptions] = None,
positions: list[JointPositions],
options: MoveOptions | None = None,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
):
"""
Expand Down Expand Up @@ -186,12 +292,66 @@ async def move_through_joint_positions(
"""
...

@abc.abstractmethod
async def move_through_joint_positions_streamed(
self,
batches: AsyncIterator[list["Arm.TrajectoryPoint"]],
*,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
) -> AsyncIterator["Arm.TrajectoryUpdate"]:
"""
Move the arm through a time-parameterized stream of joint waypoints.

The caller supplies an asynchronous iterator of batches, each batch a ``list`` of
``TrajectoryPoint``. Each list the caller yields is sent as one wire ``TrajectoryBatch``,
so the caller sets the wire cadence by choosing how many points go in each list; a caller
that wants to send one point at a time yields a single-element list. The arm's updates are
yielded back as they arrive, so iterating the return value observes execution in real time.
If the arm faults mid-trajectory, that fault arrives as a gRPC error on the iteration, so the
``async for`` raises instead of ending normally. Delivering faults mid-execution, not only at
the end, is the point of streaming this call.

The first point of the stream must have time zero, and if it carries velocity constraints
those velocities must all be zero, since the trajectory starts from rest. Point times must
strictly increase across the whole stream, not merely within a batch. A ``timeout``, if
given, bounds the entire stream, not a single message, so an open-ended trajectory should
normally leave it unset.

An implementation must yield at least one ``TrajectoryUpdate`` before returning. Besides
reporting progress, this is what makes the implementation an asynchronous generator; a
coroutine that never yields cannot be iterated as a stream and fails at runtime.

::

my_arm = Arm.from_robot(robot=machine, name="my_arm")

async def batches():
yield [
Arm.TrajectoryPoint(time=timedelta(seconds=0.0), positions=[0.0, 0.0, 0.0, 0.0, 0.0]),
Arm.TrajectoryPoint(time=timedelta(seconds=1.0), positions=[10.0, 0.0, 0.0, 0.0, 0.0]),
]

async for update in my_arm.move_through_joint_positions_streamed(batches()):
# Observe the arm's updates; a fault raises out of this iteration.
pass

Args:
batches: an asynchronous iterator of lists of ``TrajectoryPoint``. Each list becomes
one wire ``TrajectoryBatch``.

Returns:
AsyncIterator[Arm.TrajectoryUpdate]: the arm's updates, yielded as they arrive.
"""
...

@abc.abstractmethod
async def get_joint_positions(
self,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
) -> JointPositions:
"""
Expand All @@ -217,8 +377,8 @@ async def get_joint_positions(
async def stop(
self,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
):
"""
Expand Down Expand Up @@ -259,7 +419,7 @@ async def is_moving(self) -> bool:

@abc.abstractmethod
async def get_kinematics(
self, *, extra: Optional[Dict[str, Any]] = None, timeout: Optional[float] = None, **kwargs
self, *, extra: dict[str, Any] | None = None, timeout: float | None = None, **kwargs
) -> KinematicsReturn:
"""
Get the kinematics information associated with the arm.
Expand Down Expand Up @@ -291,7 +451,7 @@ async def get_kinematics(

@abc.abstractmethod
async def get_3d_models(
self, *, extra: Optional[Dict[str, Any]] = None, timeout: Optional[float] = None, **kwargs
self, *, extra: dict[str, Any] | None = None, timeout: float | None = None, **kwargs
) -> Mapping[str, Mesh]:
"""
Get the 3D models associated with the arm, keyed by name.
Expand Down Expand Up @@ -325,8 +485,8 @@ async def set_manual_mode(
manual_mode: bool,
enabled_for: int = 0,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
):
"""
Expand Down Expand Up @@ -354,8 +514,8 @@ async def set_manual_mode(
async def get_manual_mode(
self,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
) -> bool:
"""
Expand All @@ -379,8 +539,8 @@ async def get_manual_mode(
async def get_properties(
self,
*,
extra: Optional[Dict[str, Any]] = None,
timeout: Optional[float] = None,
extra: dict[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
) -> Properties:
"""
Expand Down
Loading