diff --git a/docs/examples/my_cool_arm.py b/docs/examples/my_cool_arm.py index 0fd940c86..2d40baf0e 100644 --- a/docs/examples/my_cool_arm.py +++ b/docs/examples/my_cool_arm.py @@ -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 @@ -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 {} diff --git a/examples/complex_module/src/arm/my_arm.py b/examples/complex_module/src/arm/my_arm.py index cae773c06..50362d5ff 100644 --- a/examples/complex_module/src/arm/my_arm.py +++ b/examples/complex_module/src/arm/my_arm.py @@ -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 @@ -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 {} diff --git a/examples/server/v1/components.py b/examples/server/v1/components.py index cf3b8be68..03bb7d989 100644 --- a/examples/server/v1/components.py +++ b/examples/server/v1/components.py @@ -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 {} diff --git a/src/viam/components/arm/arm.py b/src/viam/components/arm/arm.py index a7cd61d45..d506eba26 100644 --- a/src/viam/components/arm/arm.py +++ b/src/viam/components/arm/arm.py @@ -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 @@ -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 + ``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: """ @@ -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, ): """ @@ -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, ): """ @@ -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, ): """ @@ -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: """ @@ -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, ): """ @@ -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. @@ -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. @@ -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, ): """ @@ -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: """ @@ -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: """ diff --git a/src/viam/components/arm/client.py b/src/viam/components/arm/client.py index 343ab8885..174e49bfa 100644 --- a/src/viam/components/arm/client.py +++ b/src/viam/components/arm/client.py @@ -1,4 +1,5 @@ -from typing import Any, Dict, List, Mapping, Optional +import asyncio +from typing import Any, AsyncIterator, Dict, List, Mapping, Optional from grpclib.client import Channel @@ -30,6 +31,7 @@ MoveOptions, MoveThroughJointPositionsRequest, MoveToJointPositionsRequest, + MoveThroughJointPositionsStreamedRequest, MoveToPositionRequest, SetManualModeRequest, StopRequest, @@ -114,6 +116,118 @@ async def move_through_joint_positions( request = MoveThroughJointPositionsRequest(name=self.name, positions=positions, options=options, extra=dict_to_struct(extra)) await self.client.MoveThroughJointPositions(request, timeout=timeout, metadata=md) + 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]: + md = kwargs.get("metadata", self.Metadata()).proto + # A timeout, if the caller supplies one, bounds the whole stream rather than a single + # message, so it defaults to none; binding a deadline here would cancel a long but + # healthy trajectory partway through. + async with self.client.MoveThroughJointPositionsStreamed.open(timeout=timeout, metadata=md) as stream: + await stream.send_message( + MoveThroughJointPositionsStreamedRequest( + name=self.name, + init=MoveThroughJointPositionsStreamedRequest.Init( + extra=dict_to_struct(extra), + ), + ) + ) + + # Sending and receiving run concurrently as tasks: the arm can report an update or a + # fault at any point, including while the caller is still producing batches. An async + # generator cannot yield a value produced inside a task, so the receive task feeds a + # queue that this generator drains and yields from; a sentinel marks the point past + # which no more updates will arrive. Each list the caller yields becomes one wire + # TrajectoryBatch. + # + # A failure of the caller's own batch iterator is recorded separately. It is the + # caller's bug and the fault they need to see, so it wins over whatever the receive side + # reports while the stream is torn down. + producer_exception: Optional[BaseException] = None + updates: asyncio.Queue = asyncio.Queue() + end_of_updates = object() + + async def send_batches() -> None: + nonlocal producer_exception + try: + async for batch in batches: + await stream.send_message( + MoveThroughJointPositionsStreamedRequest( + batch=MoveThroughJointPositionsStreamedRequest.TrajectoryBatch( + points=[point.to_proto() for point in batch], + ) + ) + ) + # Batches exhausted cleanly; half-close so the arm knows the trajectory completed. + await stream.end() + except asyncio.CancelledError: + # Our own teardown cancelling this task, not the caller's failure. + raise + except BaseException as exc: + producer_exception = exc + raise + + async def receive_updates() -> None: + try: + while True: + update = await stream.recv_message() + if update is None: + break + updates.put_nowait(update) + finally: + updates.put_nowait(end_of_updates) + + send_task = asyncio.create_task(send_batches()) + receive_task = asyncio.create_task(receive_updates()) + + # If the producer fails, stop receiving so the queue terminates and the fault can be + # surfaced. A clean producer finish leaves the receive alone: the arm still has updates + # to send until it closes the response stream itself. + def stop_receiving_if_producer_failed(task: asyncio.Task) -> None: + if not task.cancelled() and task.exception() is not None: + receive_task.cancel() + + send_task.add_done_callback(stop_receiving_if_producer_failed) + + try: + while True: + update = await updates.get() + if update is end_of_updates: + break + yield Arm.TrajectoryUpdate.from_proto(update) + finally: + # Before the `async with` resets the stream, make sure both tasks have finished and + # their outcomes have been retrieved, so neither is parked in a read or write during + # the reset. A parked read is exactly what deadlocks a direct stream.cancel(); the + # reset that aborts the arm comes from leaving the `async with` instead. A finished + # task's result is retrieved with `.exception()` rather than by awaiting it, which + # keeps a recorded producer failure's traceback pointed at the caller's code. + for task in (send_task, receive_task): + if task.done(): + if not task.cancelled(): + task.exception() + else: + task.cancel() + try: + await task + except BaseException: + pass + + # Surface the terminal cause: the caller's producer failure first, then a fault from the + # receive side, otherwise the stream completed cleanly. Raising leaves the `async with`, + # which resets the stream so the arm sees an abort rather than a clean end. + if producer_exception is not None: + raise producer_exception + if not receive_task.cancelled(): + receive_error = receive_task.exception() + if receive_error is not None: + raise receive_error + async def stop( self, *, diff --git a/src/viam/components/arm/service.py b/src/viam/components/arm/service.py index 741711a5b..10d18a707 100644 --- a/src/viam/components/arm/service.py +++ b/src/viam/components/arm/service.py @@ -1,3 +1,7 @@ +from datetime import timedelta +from typing import AsyncIterator, List + +from grpclib import GRPCError, Status from grpclib.server import Stream from viam.proto.common import ( @@ -25,6 +29,8 @@ IsMovingResponse, MoveThroughJointPositionsRequest, MoveThroughJointPositionsResponse, + MoveThroughJointPositionsStreamedRequest, + MoveThroughJointPositionsStreamedResponse, MoveToJointPositionsRequest, MoveToJointPositionsResponse, MoveToPositionRequest, @@ -41,6 +47,51 @@ from .arm import Arm +class _TrajectoryStreamValidator: + """ + Enforces the trajectory contract across a single streamed request, one point at a time. + + The checks match what the other SDKs apply on the server side (see the C++ SDK's + TrajectoryStreamValidator): the trajectory begins at time zero and from rest, point times + strictly increase, and any constraints are dimensionally consistent with the positions. + State carries across batches, so a single validator must see every point of the stream in + order. + """ + + def __init__(self) -> None: + self._seen_first = False + self._last_time = timedelta() + + def check(self, point: Arm.TrajectoryPoint) -> None: + if not self._seen_first: + if point.time != timedelta(): + raise GRPCError(Status.INVALID_ARGUMENT, "first trajectory point must have time zero") + elif point.time <= self._last_time: + raise GRPCError(Status.INVALID_ARGUMENT, "trajectory point times must strictly increase") + + if not point.positions: + raise GRPCError(Status.INVALID_ARGUMENT, "trajectory point must carry at least one position") + + constraints = point.constraints + if constraints is not None: + if len(constraints.velocities) != len(point.positions): + raise GRPCError( + Status.INVALID_ARGUMENT, + "trajectory point must carry one velocity per position when constraints are present", + ) + # The arm has to start from a standstill, so the first point may not ask for motion. + if not self._seen_first and any(velocity != 0.0 for velocity in constraints.velocities): + raise GRPCError(Status.INVALID_ARGUMENT, "first trajectory point must start from rest (all velocities zero)") + if constraints.accelerations is not None and len(constraints.accelerations) != len(constraints.velocities): + raise GRPCError( + Status.INVALID_ARGUMENT, + "trajectory point must carry one acceleration per velocity when accelerations are present", + ) + + self._seen_first = True + self._last_time = point.time + + class ArmRPCService(UnimplementedArmServiceBase, ResourceRPCServiceBase[Arm]): """ gRPC Service for an Arm @@ -107,6 +158,55 @@ async def MoveThroughJointPositions(self, stream: Stream[MoveThroughJointPositio response = MoveThroughJointPositionsResponse() await stream.send_message(response) + async def MoveThroughJointPositionsStreamed( + self, + stream: Stream[MoveThroughJointPositionsStreamedRequest, MoveThroughJointPositionsStreamedResponse], + ) -> None: + # The stream opens with exactly one Init, which names the arm and carries the sticky extra + # arguments. The name sits at the top level of the request rather than inside Init because + # that is where the RDK server reads it. + first_request = await stream.recv_message() + if first_request is None: + raise GRPCError(Status.INVALID_ARGUMENT, "stream closed before init message") + if not first_request.HasField("init"): + raise GRPCError(Status.INVALID_ARGUMENT, "first message must be init") + + name = first_request.name + arm = self.get_resource(name) + extra = struct_to_dict(first_request.init.extra) + timeout = stream.deadline.time_remaining() if stream.deadline else None + + # Turn the rest of the request stream into the async iterator of point-lists the driver + # consumes, validating each point as it arrives. Validation lives here on the server so + # that every arm implementation gets the same contract enforcement without having to + # repeat it. An empty batch is a wire no-op and is skipped; a second Init, or any message + # that is not a batch, is a protocol violation that ends the stream with an error. + validator = _TrajectoryStreamValidator() + + async def batches() -> AsyncIterator[List[Arm.TrajectoryPoint]]: + while True: + request = await stream.recv_message() + if request is None: + return + message = request.WhichOneof("message") + if message == "init": + raise GRPCError(Status.INVALID_ARGUMENT, "init may only appear as the first message") + if message != "batch": + raise GRPCError(Status.INVALID_ARGUMENT, "expected a trajectory batch") + points = [Arm.TrajectoryPoint.from_proto(point_proto) for point_proto in request.batch.points] + for point in points: + validator.check(point) + if points: + yield points + + async for update in arm.move_through_joint_positions_streamed( # pyright: ignore [reportGeneralTypeIssues] + batches(), + extra=extra, + timeout=timeout, + metadata=stream.metadata, + ): + await stream.send_message(update.to_proto()) + async def Stop(self, stream: Stream[StopRequest, StopResponse]) -> None: request = await stream.recv_message() assert request is not None diff --git a/tests/mocks/components.py b/tests/mocks/components.py index 5c37746a1..8338d288a 100644 --- a/tests/mocks/components.py +++ b/tests/mocks/components.py @@ -75,6 +75,7 @@ def __init__(self, name: str): self.timeout: Optional[float] = None self.waypoints: List[JointPositions] = [] self.move_options: Optional[MoveOptions] = None + self.streamed_points: List[Arm.TrajectoryPoint] = [] self.models_3d = MODELS_3D super().__init__(name) @@ -128,6 +129,26 @@ async def move_through_joint_positions( self.extra = extra self.timeout = timeout + 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 + self.extra = extra + self.timeout = timeout + self.streamed_points = [] + async for batch in batches: + for point in batch: + self.streamed_points.append(point) + self.joint_positions = JointPositions(values=point.positions) + # Acknowledge each batch with one update, as a real driver would. + yield Arm.TrajectoryUpdate() + self.is_stopped = True + async def get_3d_models( self, *, extra: Optional[Dict[str, Any]] = None, timeout: Optional[float] = None, **kwargs ) -> Mapping[str, Mesh]: diff --git a/tests/test_arm.py b/tests/test_arm.py index 633f69850..09b3c9028 100644 --- a/tests/test_arm.py +++ b/tests/test_arm.py @@ -1,8 +1,11 @@ +from datetime import timedelta +from typing import AsyncIterator, List + import pytest -from grpclib import GRPCError +from grpclib import GRPCError, Status from grpclib.testing import ChannelFor -from viam.components.arm import ArmClient, KinematicsFileFormat +from viam.components.arm import Arm, ArmClient, KinematicsFileFormat from viam.components.arm.service import ArmRPCService from viam.proto.common import ( DoCommandRequest, @@ -429,3 +432,112 @@ async def test_extra(self): client = ArmClient(self.name, channel) await client.get_end_position(extra={"foo": "bar"}) assert self.arm.extra == {"foo": "bar"} + + +async def _batches(batch_lists: List[List[Arm.TrajectoryPoint]]) -> AsyncIterator[List[Arm.TrajectoryPoint]]: + for batch in batch_lists: + yield batch + + +class TestArmStreamed: + @classmethod + def setup_class(cls): + cls.name = "arm" + cls.arm = MockArm(name=cls.name) + cls.manager = ResourceManager([cls.arm]) + cls.service = ArmRPCService(cls.manager) + + async def test_streamed_round_trip(self): + async with ChannelFor([self.service]) as channel: + client = ArmClient(self.name, channel) + first = [ + Arm.TrajectoryPoint(time=timedelta(0), positions=[0.0, 0.0, 0.0]), + Arm.TrajectoryPoint(time=timedelta(seconds=1), positions=[1.0, 2.0, 3.0]), + ] + second = [Arm.TrajectoryPoint(time=timedelta(seconds=2), positions=[4.0, 5.0, 6.0])] + updates = [update async for update in client.move_through_joint_positions_streamed(_batches([first, second]))] + assert len(updates) == 2 + assert len(self.arm.streamed_points) == 3 + assert self.arm.streamed_points[-1].positions == [4.0, 5.0, 6.0] + + async def test_streamed_rejects_nonzero_first_time(self): + async with ChannelFor([self.service]) as channel: + client = ArmClient(self.name, channel) + bad = [Arm.TrajectoryPoint(time=timedelta(seconds=1), positions=[0.0])] + with pytest.raises(GRPCError) as excinfo: + async for _ in client.move_through_joint_positions_streamed(_batches([bad])): + pass + assert excinfo.value.status == Status.INVALID_ARGUMENT + + async def test_streamed_rejects_nonmonotonic_time(self): + async with ChannelFor([self.service]) as channel: + client = ArmClient(self.name, channel) + bad = [ + Arm.TrajectoryPoint(time=timedelta(0), positions=[0.0]), + Arm.TrajectoryPoint(time=timedelta(0), positions=[0.0]), + ] + with pytest.raises(GRPCError) as excinfo: + async for _ in client.move_through_joint_positions_streamed(_batches([bad])): + pass + assert excinfo.value.status == Status.INVALID_ARGUMENT + + async def test_streamed_rejects_first_point_in_motion(self): + async with ChannelFor([self.service]) as channel: + client = ArmClient(self.name, channel) + bad = [ + Arm.TrajectoryPoint( + time=timedelta(0), positions=[0.0], constraints=Arm.KinematicConstraints(velocities=[1.0]) + ) + ] + with pytest.raises(GRPCError) as excinfo: + async for _ in client.move_through_joint_positions_streamed(_batches([bad])): + pass + assert excinfo.value.status == Status.INVALID_ARGUMENT + + async def test_streamed_producer_error_surfaces(self): + async with ChannelFor([self.service]) as channel: + client = ArmClient(self.name, channel) + + class ProducerError(Exception): + pass + + async def failing_batches() -> AsyncIterator[List[Arm.TrajectoryPoint]]: + yield [Arm.TrajectoryPoint(time=timedelta(0), positions=[0.0])] + raise ProducerError() + + with pytest.raises(ProducerError): + async for _ in client.move_through_joint_positions_streamed(failing_batches()): + pass + + +class TestTrajectoryConversions: + def test_trajectory_point_round_trip_with_accelerations(self): + point = Arm.TrajectoryPoint( + time=timedelta(seconds=1, milliseconds=500), + positions=[1.0, 2.0, 3.0], + constraints=Arm.KinematicConstraints(velocities=[0.5, 0.5, 0.5], accelerations=[0.25, 0.25, 0.25]), + ) + restored = Arm.TrajectoryPoint.from_proto(point.to_proto()) + assert restored.time == point.time + assert restored.positions == point.positions + assert restored.constraints is not None + assert restored.constraints.velocities == [0.5, 0.5, 0.5] + assert restored.constraints.accelerations == [0.25, 0.25, 0.25] + + def test_trajectory_point_round_trip_without_constraints(self): + point = Arm.TrajectoryPoint(time=timedelta(0), positions=[0.0, 0.0]) + restored = Arm.TrajectoryPoint.from_proto(point.to_proto()) + assert restored.time == timedelta(0) + assert restored.positions == [0.0, 0.0] + assert restored.constraints is None + + def test_kinematic_constraints_without_accelerations(self): + point = Arm.TrajectoryPoint(time=timedelta(seconds=2), positions=[1.0], constraints=Arm.KinematicConstraints(velocities=[0.0])) + restored = Arm.TrajectoryPoint.from_proto(point.to_proto()) + assert restored.constraints is not None + assert restored.constraints.velocities == [0.0] + assert restored.constraints.accelerations is None + + def test_trajectory_update_round_trip(self): + restored = Arm.TrajectoryUpdate.from_proto(Arm.TrajectoryUpdate().to_proto()) + assert isinstance(restored, Arm.TrajectoryUpdate) diff --git a/tests/test_robot.py b/tests/test_robot.py index 207432920..1e5e2d150 100644 --- a/tests/test_robot.py +++ b/tests/test_robot.py @@ -1,5 +1,5 @@ import asyncio -from typing import Any, Dict, List, Mapping, Optional, Tuple +from typing import Any, AsyncIterator, Dict, List, Mapping, Optional, Tuple from unittest import mock import pytest @@ -507,6 +507,16 @@ async def move_through_joint_positions( ): return await self.actual_client.move_through_joint_positions(positions, options=options, extra=extra, timeout=timeout) + async def move_through_joint_positions_streamed( + self, + batches: AsyncIterator[List[Arm.TrajectoryPoint]], + *, + extra: Optional[Dict[str, Any]] = None, + timeout: Optional[float] = None, + ) -> AsyncIterator[Arm.TrajectoryUpdate]: + async for update in self.actual_client.move_through_joint_positions_streamed(batches, extra=extra, timeout=timeout): + yield update + async def get_3d_models( self, *,