diff --git a/lifecycle_test.go b/lifecycle_test.go new file mode 100644 index 0000000..da007c4 --- /dev/null +++ b/lifecycle_test.go @@ -0,0 +1,132 @@ +package so_arm + +import ( + "context" + "errors" + "testing" +) + +// TestController_PostCloseReturnsSentinel verifies that every gated controller +// method returns ErrControllerClosed when the controller's closed flag is set, +// rather than panicking or hitting the (closed) bus. Table-driven so adding a +// new method to the gated list without an entry here is a visible omission. +func TestController_PostCloseReturnsSentinel(t *testing.T) { + ctx := context.Background() + cases := []struct { + name string + call func(*SafeSoArmController) error + }{ + {"Ping", func(c *SafeSoArmController) error { return c.Ping(ctx) }}, + {"SetTorqueEnable", func(c *SafeSoArmController) error { return c.SetTorqueEnable(ctx, true) }}, + {"Stop", func(c *SafeSoArmController) error { return c.Stop(ctx) }}, + {"MoveToJointPositions", func(c *SafeSoArmController) error { + return c.MoveToJointPositions(ctx, []float64{0, 0, 0, 0, 0}, 0, 0) + }}, + {"MoveServosToPositions", func(c *SafeSoArmController) error { + return c.MoveServosToPositions(ctx, []int{1}, []float64{0}, 0, 0) + }}, + {"WriteServoRegister", func(c *SafeSoArmController) error { + return c.WriteServoRegister(ctx, 1, "goal_position", []byte{0, 0}) + }}, + {"SetCalibration", func(c *SafeSoArmController) error { + return c.SetCalibration(SO101FullCalibration{}) + }}, + {"GetJointPositions", func(c *SafeSoArmController) error { + _, err := c.GetJointPositions(ctx) + return err + }}, + {"GetJointPositionsForServos", func(c *SafeSoArmController) error { + _, err := c.GetJointPositionsForServos(ctx, []int{1}) + return err + }}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + bus, _ := newMockBus(t) + ctrl := &SafeSoArmController{ + bus: bus, + logger: newTestLogger(t), + } + ctrl.closed.Store(true) + + if err := tc.call(ctrl); !errors.Is(err, ErrControllerClosed) { + t.Errorf("%s after close: expected ErrControllerClosed, got %v", tc.name, err) + } + }) + } +} + +// TestRegistry_SamePointerForSamePort verifies that two callers acquiring a +// controller for the same port receive the *same* *SafeSoArmController, so +// that close-state propagates correctly across all consumers. +func TestRegistry_SamePointerForSamePort(t *testing.T) { + registry := NewControllerRegistry() + port := "/dev/test-port" + cfg := testConfig(port) + + // Inject a pre-built entry so we don't need a real bus. + bus, _ := newMockBus(t) + ctrl := &SafeSoArmController{ + bus: bus, + logger: cfg.Logger, + } + registry.entries[port] = &ControllerEntry{ + controller: ctrl, + config: cfg, + calibration: DefaultSO101FullCalibration, + refCount: 0, + } + + first, err := registry.GetController(port, cfg, DefaultSO101FullCalibration, false) + if err != nil { + t.Fatalf("first GetController: %v", err) + } + second, err := registry.GetController(port, cfg, DefaultSO101FullCalibration, false) + if err != nil { + t.Fatalf("second GetController: %v", err) + } + + if first != second { + t.Errorf("expected same pointer for same port; got %p and %p", first, second) + } + if first != ctrl { + t.Errorf("expected cached controller pointer to be returned") + } +} + +// TestRegistry_ReleaseClosesAllConsumers verifies that ReleaseController +// at refcount zero closes the bus and sets the closed flag on the shared +// controller, so other holders observe ErrControllerClosed on next call. +func TestRegistry_ReleaseClosesAllConsumers(t *testing.T) { + registry := NewControllerRegistry() + port := "/dev/test-port" + cfg := testConfig(port) + + bus, _ := newMockBus(t) + ctrl := &SafeSoArmController{ + bus: bus, + logger: cfg.Logger, + } + registry.entries[port] = &ControllerEntry{ + controller: ctrl, + config: cfg, + calibration: DefaultSO101FullCalibration, + refCount: 2, // simulate arm + gripper both holding + } + + // First release: refcount drops to 1, controller stays alive. + registry.ReleaseController(port) + if ctrl.closed.Load() { + t.Fatalf("controller closed prematurely at refcount > 0") + } + + // Second release: refcount drops to 0, controller closes. + registry.ReleaseController(port) + if !ctrl.closed.Load() { + t.Errorf("expected controller.closed=true after final release") + } + if err := ctrl.Ping(context.Background()); !errors.Is(err, ErrControllerClosed) { + t.Errorf("Ping after final release: expected ErrControllerClosed, got %v", err) + } +} diff --git a/manager.go b/manager.go index 298f5d2..d860f35 100644 --- a/manager.go +++ b/manager.go @@ -2,6 +2,7 @@ package so_arm import ( "context" + "errors" "fmt" "math" "sync" @@ -13,6 +14,11 @@ import ( "go.viam.com/rdk/utils" ) +// ErrControllerClosed is returned by SafeSoArmController methods after the +// underlying bus has been closed via the registry. Callers holding a stale +// reference should treat this as a permanent failure for that controller. +var ErrControllerClosed = errors.New("so101: controller is closed") + // isGripperServo checks if a servo ID is the gripper (servo 6) func isGripperServo(servoID int) bool { return servoID == 6 @@ -27,6 +33,18 @@ type SafeSoArmController struct { logger logging.Logger calibration SO101FullCalibration mu sync.RWMutex + closed atomic.Bool +} + +// checkClosed returns ErrControllerClosed if the controller has been released. +// checkClosed is best-effort: callers must hold a registry refcount for the +// duration of any controller call. A concurrent ReleaseController can race the +// unlocked Load() with an in-flight method. +func (s *SafeSoArmController) checkClosed() error { + if s.closed.Load() { + return ErrControllerClosed + } + return nil } // writePositions writes goal positions to the servo group. When speed > 0 it commands @@ -45,6 +63,9 @@ func (s *SafeSoArmController) writePositions(ctx context.Context, rawPositions f } func (s *SafeSoArmController) MoveToJointPositions(ctx context.Context, jointAngles []float64, speed, acc int) error { + if err := s.checkClosed(); err != nil { + return err + } s.mu.Lock() defer s.mu.Unlock() @@ -73,6 +94,9 @@ func (s *SafeSoArmController) MoveToJointPositions(ctx context.Context, jointAng } func (s *SafeSoArmController) MoveServosToPositions(ctx context.Context, servoIDs []int, jointAngles []float64, speed, acc int) error { + if err := s.checkClosed(); err != nil { + return err + } s.mu.Lock() defer s.mu.Unlock() @@ -106,6 +130,9 @@ func (s *SafeSoArmController) MoveServosToPositions(ctx context.Context, servoID } func (s *SafeSoArmController) GetJointPositions(ctx context.Context) ([]float64, error) { + if err := s.checkClosed(); err != nil { + return nil, err + } s.mu.RLock() defer s.mu.RUnlock() @@ -144,6 +171,9 @@ func (s *SafeSoArmController) GetJointPositions(ctx context.Context) ([]float64, } func (s *SafeSoArmController) GetJointPositionsForServos(ctx context.Context, servoIDs []int) ([]float64, error) { + if err := s.checkClosed(); err != nil { + return nil, err + } s.mu.RLock() defer s.mu.RUnlock() @@ -173,6 +203,9 @@ func (s *SafeSoArmController) GetJointPositionsForServos(ctx context.Context, se } func (s *SafeSoArmController) SetTorqueEnable(ctx context.Context, enable bool) error { + if err := s.checkClosed(); err != nil { + return err + } s.mu.Lock() defer s.mu.Unlock() @@ -189,6 +222,9 @@ func (s *SafeSoArmController) SetTorqueEnable(ctx context.Context, enable bool) } func (s *SafeSoArmController) Stop(ctx context.Context) error { + if err := s.checkClosed(); err != nil { + return err + } s.mu.Lock() defer s.mu.Unlock() @@ -260,6 +296,7 @@ func (s *SafeSoArmController) Close() error { s.mu.Lock() defer s.mu.Unlock() + s.closed.Store(true) if s.bus != nil { return s.bus.Close() } @@ -267,6 +304,9 @@ func (s *SafeSoArmController) Close() error { } func (s *SafeSoArmController) Ping(ctx context.Context) error { + if err := s.checkClosed(); err != nil { + return err + } s.mu.RLock() defer s.mu.RUnlock() @@ -280,6 +320,9 @@ func (s *SafeSoArmController) Ping(ctx context.Context) error { // WriteServoRegister writes to a specific servo register by name func (s *SafeSoArmController) WriteServoRegister(ctx context.Context, servoID int, registerName string, data []byte) error { + if err := s.checkClosed(); err != nil { + return err + } s.mu.Lock() defer s.mu.Unlock() @@ -292,6 +335,9 @@ func (s *SafeSoArmController) WriteServoRegister(ctx context.Context, servoID in } func (s *SafeSoArmController) SetCalibration(calibration SO101FullCalibration) error { + if err := s.checkClosed(); err != nil { + return err + } s.mu.Lock() defer s.mu.Unlock() diff --git a/mock_bus_test.go b/mock_bus_test.go new file mode 100644 index 0000000..d283ec2 --- /dev/null +++ b/mock_bus_test.go @@ -0,0 +1,143 @@ +package so_arm + +import ( + "io" + "sync" + "testing" + "time" + + "github.com/hipsterbrown/feetech-servo/feetech" + "go.viam.com/rdk/logging" +) + +// scriptedMockTransport is a custom Transport (not feetech.MockTransport) whose +// Read responses are queued per-request. We use a custom type rather than the +// upstream MockTransport because tests in PR2-PR5 need to script multiple +// round-trips with different responses, and MockTransport's single ReadData +// buffer doesn't support that pattern cleanly. +type scriptedMockTransport struct { + mu sync.Mutex + written []byte + responses [][]byte + closed bool +} + +func (m *scriptedMockTransport) Read(p []byte) (int, error) { + m.mu.Lock() + defer m.mu.Unlock() + if len(m.responses) == 0 { + // Match feetech.MockTransport semantics: returning io.EOF on an empty + // queue lets feetech.Bus.readRawBytesLocked fall into its 1ms-sleep + // retry path instead of busy-spinning under the bus mutex. Important + // for PR5's planned 100Hz calibration-sensor reader. + return 0, io.EOF + } + resp := m.responses[0] + n := copy(p, resp) + if n >= len(resp) { + m.responses = m.responses[1:] + } else { + m.responses[0] = resp[n:] + } + return n, nil +} + +func (m *scriptedMockTransport) Write(p []byte) (int, error) { + m.mu.Lock() + defer m.mu.Unlock() + m.written = append(m.written, p...) + return len(p), nil +} + +func (m *scriptedMockTransport) Close() error { + m.mu.Lock() + defer m.mu.Unlock() + m.closed = true + return nil +} + +func (m *scriptedMockTransport) SetReadTimeout(timeout time.Duration) error { + return nil +} + +func (m *scriptedMockTransport) Flush() error { + // No-op: tests that need to drop unconsumed responses should clear + // m.responses explicitly. Mirroring SerialTransport.Flush would risk + // silently swallowing scripted frames between operations. + return nil +} + +// queueResponse appends a raw frame to the mock's response queue. +func (m *scriptedMockTransport) queueResponse(b []byte) { + m.mu.Lock() + defer m.mu.Unlock() + m.responses = append(m.responses, b) +} + +// reset clears the written-data buffer (responses queue is preserved). +func (m *scriptedMockTransport) resetWritten() { + m.mu.Lock() + defer m.mu.Unlock() + m.written = nil +} + +// snapshotWritten returns a copy of all bytes written since the last reset. +func (m *scriptedMockTransport) snapshotWritten() []byte { + m.mu.Lock() + defer m.mu.Unlock() + out := make([]byte, len(m.written)) + copy(out, m.written) + return out +} + +// pingResponse builds the ping ACK frame for a given servo ID using the STS protocol. +// Frame: 0xFF 0xFF 0x02 0x00 +func pingResponse(servoID byte) []byte { + checksum := ^(servoID + 0x02 + 0x00) + return []byte{0xFF, 0xFF, servoID, 0x02, 0x00, checksum} +} + +// newMockBus returns a *feetech.Bus backed by a scriptedMockTransport. +// Caller is responsible for queuing responses before issuing bus operations. +func newMockBus(t *testing.T) (*feetech.Bus, *scriptedMockTransport) { + t.Helper() + mock := &scriptedMockTransport{} + bus, err := feetech.NewBus(feetech.BusConfig{ + Transport: mock, + Protocol: feetech.ProtocolSTS, + Timeout: 100 * time.Millisecond, + }) + if err != nil { + t.Fatalf("newMockBus: NewBus failed: %v", err) + } + t.Cleanup(func() { _ = bus.Close() }) + return bus, mock +} + +// newTestLogger returns a logger that discards output unless the test fails. +func newTestLogger(t *testing.T) logging.Logger { + t.Helper() + return logging.NewTestLogger(t) +} + +func TestMockBus_PingRoundTrip(t *testing.T) { + bus, mock := newMockBus(t) + // Bus.Ping does two round trips: a ping ACK, then a model-number register read. + // Queue both responses so the call completes. + mock.queueResponse(pingResponse(1)) + // Model number 777 (0x0309) read response: FF FF + mock.queueResponse([]byte{0xFF, 0xFF, 0x01, 0x04, 0x00, 0x09, 0x03, 0xEE}) + + if _, err := bus.Ping(t.Context(), 1); err != nil { + t.Fatalf("ping failed: %v", err) + } + + written := mock.snapshotWritten() + if len(written) < 6 { + t.Fatalf("expected ping packet to be written, got %d bytes: %X", len(written), written) + } + // Ping packet: 0xFF 0xFF + if written[0] != 0xFF || written[1] != 0xFF || written[2] != 0x01 { + t.Errorf("malformed ping packet: %X", written[:6]) + } +} diff --git a/registry.go b/registry.go index 2493377..d38b1c9 100644 --- a/registry.go +++ b/registry.go @@ -95,13 +95,9 @@ func (r *ControllerRegistry) getExistingController(entry *ControllerEntry, confi atomic.AddInt64(&entry.refCount, 1) r.trackCaller(entry.config.Port) - return &SafeSoArmController{ - bus: entry.controller.bus, - group: entry.controller.group, - calibratedServos: entry.controller.calibratedServos, - logger: config.Logger, - calibration: entry.calibration, - }, nil + // Return the cached pointer so that all consumers (arm + gripper sharing a + // port) observe close-state atomically via the shared closed flag. + return entry.controller, nil } func (r *ControllerRegistry) createNewController(portPath string, config *SoArm101Config, calibration SO101FullCalibration, fromFile bool) (*SafeSoArmController, error) { @@ -217,13 +213,7 @@ func (r *ControllerRegistry) createNewController(portPath string, config *SoArm1 config.Logger.Debugf("Created new feetech servo bus with %d servos for port %s", len(calibratedServos), portPath) } - return &SafeSoArmController{ - bus: bus, - group: group, - calibratedServos: calibratedServos, - logger: config.Logger, - calibration: finalCalibration, - }, nil + return entry.controller, nil } func (r *ControllerRegistry) ReleaseController(portPath string) { @@ -240,9 +230,12 @@ func (r *ControllerRegistry) ReleaseController(portPath string) { currentRefCount := atomic.AddInt64(&entry.refCount, -1) if currentRefCount <= 0 { - if entry.controller != nil && entry.controller.bus != nil { - if err := entry.controller.bus.Close(); err != nil && entry.config != nil && entry.config.Logger != nil { - entry.config.Logger.Warnf("error closing shared controller for port %s: %v", portPath, err) + if entry.controller != nil { + entry.controller.closed.Store(true) + if entry.controller.bus != nil { + if err := entry.controller.bus.Close(); err != nil && entry.config != nil && entry.config.Logger != nil { + entry.config.Logger.Warnf("error closing shared controller for port %s: %v", portPath, err) + } } } @@ -275,6 +268,7 @@ func (r *ControllerRegistry) ForceCloseController(portPath string) error { var err error if entry.controller != nil { + entry.controller.closed.Store(true) err = entry.controller.bus.Close() entry.controller = nil entry.config = nil