From 2548bf552872bfa128307cd240bb06184be35940 Mon Sep 17 00:00:00 2001 From: HipsterBrown Date: Fri, 10 Jul 2026 22:13:47 -0400 Subject: [PATCH] fix(controller): propagate close-state to all shared-controller holders The ref-counted controller registry handed each consumer (arm + gripper on one SO-101 share a serial port) its own *SafeSoArmController struct copy wrapping the same bus. When one component closed first, ReleaseController closed the shared bus but the other holder's struct had no way to observe it, so its next call raced against / hit a closed bus. Fix, extracted from the lifecycle-correctness work in #11: - Add a `closed atomic.Bool` to SafeSoArmController plus an ErrControllerClosed sentinel; every public method checks it up-front and returns the sentinel instead of touching a closed bus. Close() is now idempotent (sets closed). - getExistingController / createNewController return the *cached* controller pointer instead of a fresh struct copy, so all holders share one closed flag. ReleaseController and ForceCloseController set closed=true before closing the bus, so surviving holders observe ErrControllerClosed atomically. checkClosed is best-effort (documented): callers must hold a registry refcount for the duration of a call. The runtime.Caller-based release tracking is left intact here; replacing it was a separate change in #11 and is out of scope. Tests: add a scripted mock-bus harness (no hardware) and cover the gated-method sentinel across all 9 methods, same-pointer-per-port, and release-closes-all- consumers. Full package passes under -race. Co-Authored-By: Claude Opus 4.8 (1M context) Claude-Session: https://claude.ai/code/session_01JNvy6dGUT6wBHLNBv8YPYv --- lifecycle_test.go | 132 ++++++++++++++++++++++++++++++++++++++++++ manager.go | 46 +++++++++++++++ mock_bus_test.go | 143 ++++++++++++++++++++++++++++++++++++++++++++++ registry.go | 28 ++++----- 4 files changed, 332 insertions(+), 17 deletions(-) create mode 100644 lifecycle_test.go create mode 100644 mock_bus_test.go 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