diff --git a/.github/workflows/main.yml b/.github/workflows/main.yml index 9457511c..81603c14 100644 --- a/.github/workflows/main.yml +++ b/.github/workflows/main.yml @@ -8,6 +8,24 @@ on: branches: [master] jobs: + test: + name: Test (${{ matrix.os }}) + runs-on: ${{ matrix.os }} + strategy: + matrix: + os: [windows-latest, ubuntu-latest] + steps: + - name: Checkout + uses: actions/checkout@v4 + + - name: Setup Go + uses: actions/setup-go@v5 + with: + go-version: "1.25.7" + + - name: Run tests + run: go test ./... + build-windows: name: Build (windows) runs-on: windows-latest diff --git a/pkg/deej/obs.go b/pkg/deej/obs.go index 31fe947c..a148ae82 100644 --- a/pkg/deej/obs.go +++ b/pkg/deej/obs.go @@ -3,44 +3,52 @@ package deej import ( "errors" "fmt" - "sync" "time" "github.com/andreykaipov/goobs" "github.com/andreykaipov/goobs/api/requests/inputs" "go.uber.org/zap" + + "github.com/nik9play/deej/pkg/reconnect" ) -type OBSClient struct { - deej *Deej - logger *zap.SugaredLogger +const obsRetryDelay = 5 * time.Second +// obsConn is one OBS connection generation, carrying the config snapshot it +// was dialed with so config reloads can detect parameter changes +type obsConn struct { client *goobs.Client - lock sync.Mutex + cfg OBSConfig +} - stopChannel chan struct{} - errChannel chan error - wg sync.WaitGroup +type OBSClient struct { + deej *Deej + logger *zap.SugaredLogger - // config values at time of connection - hostConfig string - portConfig int - passwordConfig string + reconnector *reconnect.Reconnector[obsConn] } -const ( - obsRetryDelay = 5 * time.Second -) - func NewOBSClient(deej *Deej, logger *zap.SugaredLogger) *OBSClient { logger = logger.Named("obs") o := &OBSClient{ - deej: deej, - logger: logger, - errChannel: make(chan error, 1), + deej: deej, + logger: logger, } + o.reconnector = reconnect.New(reconnect.Options[obsConn]{ + Logger: logger, + Enabled: func() bool { return o.deej.config.Values().OBSConfig.Enabled }, + Dial: o.dial, + Watch: o.watch, + Close: o.close, + OnUp: o.onUp, + OnDown: func(err error) { + o.logger.Warnw("OBS connection error, reconnecting...", "error", err) + }, + Backoff: func(int) time.Duration { return obsRetryDelay }, + }) + logger.Debug("Created OBS client instance") o.setupOnConfigReload() @@ -49,40 +57,34 @@ func NewOBSClient(deej *Deej, logger *zap.SugaredLogger) *OBSClient { } func (o *OBSClient) Start() { - o.stopChannel = make(chan struct{}) o.logger.Info("OBS client starting") - - go o.managerLoop() + o.reconnector.Start() } func (o *OBSClient) Stop() { - if o.stopChannel == nil { - return - } - - close(o.stopChannel) - o.wg.Wait() - + o.reconnector.Stop() o.logger.Info("OBS client stopped") } func (o *OBSClient) IsConnected() bool { - o.lock.Lock() - defer o.lock.Unlock() - - return o.client != nil + return o.reconnector.Connected() } -func (o *OBSClient) SetInputVolume(inputName string, volume float32) error { - o.lock.Lock() - defer o.lock.Unlock() +// obsVolumeMulFromPercent converts a 0.0-1.0 slider position to an OBS volume +// multiplier using the cubic curve OBS's own mixer faders use +func obsVolumeMulFromPercent(percent float32) float64 { + p := float64(percent) + return p * p * p +} - if o.client == nil { +func (o *OBSClient) SetInputVolume(inputName string, percent float32) error { + conn, ok := o.reconnector.Current() + if !ok { return fmt.Errorf("not connected to OBS") } - vol := float64(volume) - _, err := o.client.Inputs.SetInputVolume(&inputs.SetInputVolumeParams{ + vol := obsVolumeMulFromPercent(percent) + _, err := conn.client.Inputs.SetInputVolume(&inputs.SetInputVolumeParams{ InputName: &inputName, InputVolumeMul: &vol, }) @@ -91,46 +93,25 @@ func (o *OBSClient) SetInputVolume(inputName string, volume float32) error { return err } - o.logger.Debugw("Set OBS input volume", "input", inputName, "volume", volume) + o.logger.Debugw("Set OBS input volume", "input", inputName, "volume", percent) return nil } -func (o *OBSClient) GetInputVolume(inputName string) (float32, error) { - o.lock.Lock() - defer o.lock.Unlock() - - if o.client == nil { - return 0, fmt.Errorf("not connected to OBS") - } - - resp, err := o.client.Inputs.GetInputVolume(&inputs.GetInputVolumeParams{ - InputName: &inputName, - }) - - if err != nil { - return 0, err +func (o *OBSClient) onUp(conn obsConn) { + // re-check the snapshot now that the connection is adopted + if o.deej.config.Values().OBSConfig != conn.cfg { + o.logger.Debug("OBS config changed while connecting, triggering reconnect") + o.reconnector.Reconnect(errors.New("config changed during dial")) + return } - return float32(resp.InputVolumeMul), nil -} - -func (o *OBSClient) signalError(err error) { - select { - case o.errChannel <- err: - default: - // channel full, error already pending - } + o.logger.Info("Connected to OBS") } -func (o *OBSClient) connect() error { - o.lock.Lock() - defer o.lock.Unlock() - - if o.client != nil { - return fmt.Errorf("already connected") - } - +// dial connects to OBS using the current config without touching any client +// state - the reconnector decides whether to adopt the returned connection +func (o *OBSClient) dial() (obsConn, error) { cfg := o.deej.config.Values().OBSConfig address := fmt.Sprintf("%s:%d", cfg.Host, cfg.Port) @@ -144,142 +125,31 @@ func (o *OBSClient) connect() error { client, err := goobs.New(address, opts...) if err != nil { o.logger.Debugw("Failed to connect to OBS", "error", err) - return fmt.Errorf("connect to OBS: %w", err) + return obsConn{}, fmt.Errorf("connect to OBS: %w", err) } - o.client = client - o.hostConfig = cfg.Host - o.portConfig = cfg.Port - o.passwordConfig = cfg.Password - - o.logger.Info("Connected to OBS") - - return nil + return obsConn{client: client, cfg: cfg}, nil } -func (o *OBSClient) disconnect() { - o.lock.Lock() - defer o.lock.Unlock() - - if o.client == nil { - return +// watch drains a single connection's incoming events to detect +// disconnection; Disconnect closes the events channel, which ends the loop +func (o *OBSClient) watch(conn obsConn, errChannel chan<- error) { + for range conn.client.IncomingEvents { //nolint:revive // draining only } - _ = o.client.Disconnect() - o.client = nil - - o.logger.Info("Disconnected from OBS") -} - -func (o *OBSClient) managerLoop() { - o.wg.Add(1) - defer o.wg.Done() - - obsConfig := o.deej.config.Values().OBSConfig - o.logger.Infow("Trying OBS connection", - "host", obsConfig.Host, - "port", obsConfig.Port, - ) - - for { - // check if OBS is enabled - if !o.deej.config.Values().OBSConfig.Enabled { - select { - case <-o.stopChannel: - o.logger.Debug("managerLoop: stop signal") - return - case <-time.After(obsRetryDelay): - continue - } - } - - // attempt connection in goroutine so we can respond to stop signal - connectResult := make(chan error, 1) - go func() { - connectResult <- o.connect() - }() - - // wait for connection result or stop signal - select { - case <-o.stopChannel: - o.logger.Debug("managerLoop: stop signal during connect") - // wait for connect to finish, then disconnect if it succeeded - if err := <-connectResult; err == nil { - o.disconnect() - } - return - - case err := <-connectResult: - if err != nil { - o.logger.Debugw("OBS connection error, retrying...", "error", err) - - select { - case <-o.stopChannel: - o.logger.Debug("managerLoop: stop signal") - return - case <-time.After(obsRetryDelay): - continue - } - } - } - - // re-check if OBS was disabled while connecting - if !o.deej.config.Values().OBSConfig.Enabled { - o.logger.Debug("OBS disabled while connecting, disconnecting") - o.disconnect() - continue - } - - // drain any stale errors from previous connection - select { - case <-o.errChannel: - default: - } - - // start event listener to detect disconnection - go o.eventLoop() - - select { - case <-o.stopChannel: - o.logger.Debug("managerLoop: stop signal") - o.disconnect() - return - - case err := <-o.errChannel: - o.logger.Warnw("OBS connection error, reconnecting...", "error", err) - o.disconnect() - time.Sleep(obsRetryDelay) - continue - } + select { + case errChannel <- errors.New("OBS connection closed"): + default: } } -func (o *OBSClient) eventLoop() { - o.wg.Add(1) - defer o.wg.Done() - - o.lock.Lock() - client := o.client - o.lock.Unlock() - - if client == nil { - return - } - - for { - select { - case <-o.stopChannel: - return - case _, ok := <-client.IncomingEvents: - if !ok { - // channel closed = disconnected - o.signalError(errors.New("OBS connection closed")) - return - } - } - } +func (o *OBSClient) close(conn obsConn) { + _ = conn.client.Disconnect() + o.logger.Info("Disconnected from OBS") } +// setupOnConfigReload triggers a reconnect when +// the OBS connection parameters change func (o *OBSClient) setupOnConfigReload() { configReloadedChannel := o.deej.config.SubscribeToChanges() @@ -287,25 +157,16 @@ func (o *OBSClient) setupOnConfigReload() { for { <-configReloadedChannel - // only trigger reconnect if currently connected - if !o.IsConnected() { + conn, ok := o.reconnector.Current() + if !ok { continue } cfg := o.deej.config.Values().OBSConfig - // the connection-time params are written by connect under the lock - o.lock.Lock() - host, port, password := o.hostConfig, o.portConfig, o.passwordConfig - o.lock.Unlock() - - if cfg.Host != host || - cfg.Port != port || - cfg.Password != password || - !cfg.Enabled { - + if cfg != conn.cfg { o.logger.Debug("OBS config changed, triggering reconnect") - o.signalError(errors.New("config changed")) + o.reconnector.Reconnect(errors.New("config changed")) } } }() diff --git a/pkg/deej/serial.go b/pkg/deej/serial.go index 85c4594a..0871109f 100644 --- a/pkg/deej/serial.go +++ b/pkg/deej/serial.go @@ -9,8 +9,6 @@ import ( "strconv" "strings" "sync" - "sync/atomic" - "time" "go.bug.st/serial" "go.bug.st/serial/enumerator" @@ -19,24 +17,27 @@ import ( "go.uber.org/zap" "github.com/nik9play/deej/pkg/deej/util" + "github.com/nik9play/deej/pkg/reconnect" ) +// serialConn is one serial connection generation, carrying the port name it +// resolved to and the connection parameters it was opened with so config +// reloads can detect changes +type serialConn struct { + port serial.Port + comPort string + connInfo ConnectionInfo +} + // SerialIO provides a deej-aware abstraction layer to managing serial I/O type SerialIO struct { - comPortToUse string - - lastConnectionInfo atomic.Pointer[ConnectionInfo] - deej *Deej logger *zap.SugaredLogger - stopChannel chan struct{} - wg sync.WaitGroup - port serial.Port - mode serial.Mode - connected atomic.Bool + reconnector *reconnect.Reconnector[serialConn] stateLock sync.Mutex + comPortToUse string lastKnownNumSliders int currentSliderValues []int @@ -65,11 +66,19 @@ func NewSerialIO(deej *Deej, logger *zap.SugaredLogger) (*SerialIO, error) { sio := &SerialIO{ deej: deej, logger: logger, - port: nil, sliderMoveConsumers: []chan SliderMoveEvent{}, stateChangeConsumers: []chan bool{}, } + sio.reconnector = reconnect.New(reconnect.Options[serialConn]{ + Logger: logger, + Dial: sio.dial, + Watch: sio.watch, + Close: sio.close, + OnUp: sio.onUp, + OnDown: sio.onDown, + }) + logger.Debug("Created serial i/o instance") // respond to config changes @@ -78,32 +87,26 @@ func NewSerialIO(deej *Deej, logger *zap.SugaredLogger) (*SerialIO, error) { return sio, nil } -func (sio *SerialIO) connect() error { - // don't allow multiple concurrent connections - if sio.port != nil { - sio.logger.Warn("Already connected, can't start another without closing first") - return errors.New("serial: connection already active") - } - +// dial resolves the com port (autodetecting by VID/PID if configured) and +// opens it +func (sio *SerialIO) dial() (serialConn, error) { config := sio.deej.config.Values() - sio.lastConnectionInfo.Store(&config.ConnectionInfo) - - sio.comPortToUse = config.ConnectionInfo.COMPort + comPort := config.ConnectionInfo.COMPort allowedVIDPID := config.AutoSearchVIDPID - if config.ConnectionInfo.COMPort == "auto" { + if comPort == "auto" { sio.logger.Debugw("Trying to autodetect serial port") ports, err := enumerator.GetDetailedPortsList() if err != nil { sio.logger.Errorw("Failed to enumarate serial ports, retrying", "err", err) - return ErrNoSerialPorts + return serialConn{}, ErrNoSerialPorts } if len(ports) == 0 { sio.logger.Debug("No serial ports found, retrying") - return ErrNoSerialPorts + return serialConn{}, ErrNoSerialPorts } for _, port := range ports { sio.logger.Debugf("Found port: %s", port.Name) @@ -116,35 +119,35 @@ func (sio *SerialIO) connect() error { if vid == allowedVIDPID.VID && pid == allowedVIDPID.PID { sio.logger.Debugw("Found COM port", "com", port.Name, "vid", port.VID, "pid", port.PID) - sio.comPortToUse = port.Name + comPort = port.Name break } } } - if sio.comPortToUse == "auto" { + if comPort == "auto" { sio.logger.Debug("COM port not found, retrying") - return ErrAutoPortNotFound + return serialConn{}, ErrAutoPortNotFound } } - sio.mode = serial.Mode{ + mode := serial.Mode{ BaudRate: config.ConnectionInfo.BaudRate, DataBits: 8, StopBits: serial.OneStopBit, } sio.logger.Debugw("Attempting serial connection", - "comPort", sio.comPortToUse, - "baudRate", sio.mode.BaudRate) + "comPort", comPort, + "baudRate", mode.BaudRate) - port, err := serial.Open(sio.comPortToUse, &sio.mode) + port, err := serial.Open(comPort, &mode) if err != nil { // might need a user notification here, TBD sio.logger.Debugw("Failed to open serial connection", "error", err) - return fmt.Errorf("open serial connection: %w", err) + return serialConn{}, fmt.Errorf("open serial connection: %w", err) } // actually, this sets timeout to 0x7FFFFFFE instead of 0xFFFFFFFE @@ -160,17 +163,90 @@ func (sio *SerialIO) connect() error { sio.logger.Warnw("Failed to close serial connection", "error", closeErr) } - return fmt.Errorf("set read timeout: %w", err) + return serialConn{}, fmt.Errorf("set read timeout: %w", err) } - sio.port = port - sio.connected.Store(true) + return serialConn{ + port: port, + comPort: comPort, + connInfo: config.ConnectionInfo, + }, nil +} + +func (sio *SerialIO) onUp(conn serialConn) { + sio.stateLock.Lock() + sio.comPortToUse = conn.comPort + sio.stateLock.Unlock() + + if sio.deej.config.Values().ConnectionInfo != conn.connInfo { + sio.logger.Info("Connection parameters changed while connecting, renewing connection") + sio.reconnector.Reconnect(errors.New("connection parameters changed during dial")) + return + } + + sio.sendStateChangeEvent(true) + + sio.logger.Named(strings.ToLower(conn.comPort)).Infow("Connected") + + connectedTitle := sio.deej.localizer.MustLocalize(&i18n.LocalizeConfig{ + DefaultMessage: &i18n.Message{ + ID: "ComPortConnectedNotificationTitle", + Other: "Connected to {{.ComPort}}.", + }, + TemplateData: map[string]string{ + "ComPort": conn.comPort, + }, + }) + connectedDescription := sio.deej.localizer.MustLocalize(&i18n.LocalizeConfig{ + DefaultMessage: &i18n.Message{ + ID: "ComPortConnectedNotificationDescription", + Other: "Succesfully connected to deej.", + }, + }) + sio.deej.notifier.Notify(connectedTitle, connectedDescription) +} + +func (sio *SerialIO) onDown(err error) { + sio.logger.Warnw("Read line error", "err", err) + sio.logger.Warn("Closing serial port") + + disconnectedTitle := sio.deej.localizer.MustLocalize(&i18n.LocalizeConfig{ + DefaultMessage: &i18n.Message{ + ID: "ComPortDisconnectedNotificationTitle", + Other: "Disconnected from {{.ComPort}} due to an error.", + }, + TemplateData: map[string]string{ + "ComPort": sio.CurrentComPort(), + }, + }) + disconnectedDescription := sio.deej.localizer.MustLocalize(&i18n.LocalizeConfig{ + DefaultMessage: &i18n.Message{ + ID: "ComPortDisconnectedNotificationDescription", + Other: "Trying to reconnect.", + }, + }) + sio.deej.notifier.Notify(disconnectedTitle, disconnectedDescription) +} + +func (sio *SerialIO) close(conn serialConn) { + if err := conn.port.Close(); err != nil { + sio.logger.Warnw("Failed to close serial connection", "error", err) + } else { + sio.logger.Info("Serial connection closed") + } - return nil + sio.sendStateChangeEvent(false) } func (sio *SerialIO) GetState() bool { - return sio.connected.Load() + return sio.reconnector.Connected() +} + +func (sio *SerialIO) CurrentComPort() string { + sio.stateLock.Lock() + defer sio.stateLock.Unlock() + + return sio.comPortToUse } func (sio *SerialIO) CurrentSliderValues() []int { @@ -182,22 +258,19 @@ func (sio *SerialIO) CurrentSliderValues() []int { // Start attempts to connect to our arduino chip func (sio *SerialIO) Start() { - sio.stopChannel = make(chan struct{}) - sio.logger.Info("Serial starting") + config := sio.deej.config.Values() + sio.logger.Infow("Trying serial connection", + "port", config.ConnectionInfo.COMPort, + "vid", fmt.Sprintf("%X", config.AutoSearchVIDPID.VID), + "pid", fmt.Sprintf("%X", config.AutoSearchVIDPID.PID), + ) - // wg.Add must happen before the goroutine starts, not inside it, - // so that a quick Stop can't call wg.Wait before the count is registered - sio.wg.Add(1) - go sio.managerLoop() + sio.reconnector.Start() } // Stop signals us to shut down our serial connection, if one is active func (sio *SerialIO) Stop() { - close(sio.stopChannel) - - // Wait for all goroutines to finish - sio.wg.Wait() - + sio.reconnector.Stop() sio.logger.Info("Serial stopped") } @@ -235,117 +308,39 @@ func (sio *SerialIO) setupOnConfigReload() { sio.lastKnownNumSliders = 0 sio.stateLock.Unlock() - config := sio.deej.config.Values() - lastConnectionInfo := sio.lastConnectionInfo.Load() + // if connection params have changed, ask the reconnector to + // renew the connection. when disconnected there's nothing to do: + // every dial attempt reads the current config anyway + conn, ok := sio.reconnector.Current() + if !ok { + continue + } - // if connection params have changed, attempt to stop and start the connection - if lastConnectionInfo == nil || config.ConnectionInfo != *lastConnectionInfo { + config := sio.deej.config.Values() + if config.ConnectionInfo != conn.connInfo { sio.logger.Info("Detected change in connection parameters, attempting to renew connection") - sio.Stop() - - // let the connection close - time.Sleep(2 * time.Second) - - sio.Start() + sio.reconnector.Reconnect(errors.New("connection parameters changed")) } } }() } -// manages serial connection and retries -func (sio *SerialIO) managerLoop() { - defer sio.wg.Done() - - config := sio.deej.config.Values() - sio.logger.Infow("Trying serial connection", - "port", config.ConnectionInfo.COMPort, - "vid", fmt.Sprintf("%X", config.AutoSearchVIDPID.VID), - "pid", fmt.Sprintf("%X", config.AutoSearchVIDPID.PID), - ) +// watch reads lines off a single connection until it fails; closing the port +// unblocks the pending read +func (sio *SerialIO) watch(conn serialConn, errChannel chan<- error) { + logger := sio.logger.Named(strings.ToLower(conn.comPort)) + reader := bufio.NewReader(conn.port) for { - err := sio.connect() + line, err := reader.ReadString('\n') if err != nil { - sio.logger.Debugw("Serial connection error. Trying again...", "err", err) - + // non-blocking send: Reconnect may have already filled this + // generation's buffer, and nothing drains it after teardown select { - case <-sio.stopChannel: - sio.logger.Debug("managerLoop: stop signal") - return - case <-time.After(2 * time.Second): - continue + case errChannel <- fmt.Errorf("read error: %w", err): + default: } - } - - sio.sendStateChangeEvent(true) - - namedLogger := sio.logger.Named(strings.ToLower(sio.comPortToUse)) - namedLogger.Infow("Connected") - - connectedTitle := sio.deej.localizer.MustLocalize(&i18n.LocalizeConfig{ - DefaultMessage: &i18n.Message{ - ID: "ComPortConnectedNotificationTitle", - Other: "Connected to {{.ComPort}}.", - }, - TemplateData: map[string]string{ - "ComPort": sio.comPortToUse, - }, - }) - connectedDescription := sio.deej.localizer.MustLocalize(&i18n.LocalizeConfig{ - DefaultMessage: &i18n.Message{ - ID: "ComPortConnectedNotificationDescription", - Other: "Succesfully connected to deej.", - }, - }) - sio.deej.notifier.Notify(connectedTitle, connectedDescription) - - errChannel := make(chan error, 1) - sio.wg.Add(1) - go sio.readLoop(namedLogger, errChannel) - - select { - case err := <-errChannel: - sio.logger.Warnw("Read line error", "err", err) - sio.logger.Warn("Closing serial port") - - disconnectedTitle := sio.deej.localizer.MustLocalize(&i18n.LocalizeConfig{ - DefaultMessage: &i18n.Message{ - ID: "ComPortDisconnectedNotificationTitle", - Other: "Disconnected from {{.ComPort}} due to an error.", - }, - TemplateData: map[string]string{ - "ComPort": sio.comPortToUse, - }, - }) - disconnectedDescription := sio.deej.localizer.MustLocalize(&i18n.LocalizeConfig{ - DefaultMessage: &i18n.Message{ - ID: "ComPortDisconnectedNotificationDescription", - Other: "Trying to reconnect.", - }, - }) - sio.deej.notifier.Notify(disconnectedTitle, disconnectedDescription) - - _ = sio.closePort() - time.Sleep(2 * time.Second) - continue - - case <-sio.stopChannel: - sio.logger.Debug("managerLoop: stop signal") - _ = sio.closePort() - return - } - } -} - -func (sio *SerialIO) readLoop(logger *zap.SugaredLogger, errChannel chan<- error) { - defer sio.wg.Done() - - reader := bufio.NewReader(sio.port) - for { - line, err := reader.ReadString('\n') - if err != nil { - errChannel <- fmt.Errorf("read error: %w", err) return } @@ -357,23 +352,6 @@ func (sio *SerialIO) readLoop(logger *zap.SugaredLogger, errChannel chan<- error } } -func (sio *SerialIO) closePort() error { - if sio.port == nil { - return fmt.Errorf("port is already closed") - } - - if err := sio.port.Close(); err != nil { - sio.logger.Warnw("Failed to close serial connection", "error", err) - return fmt.Errorf("close serial connection: %w", err) - } - - sio.logger.Info("Serial connection closed") - sio.port = nil - sio.connected.Store(false) - sio.sendStateChangeEvent(false) - return nil -} - func (sio *SerialIO) handleLine(logger *zap.SugaredLogger, line string) { // this function receives an unsanitized line which is guaranteed to end with LF, // but most lines will end with CRLF. it may also have garbage instead of diff --git a/pkg/deej/serial_test.go b/pkg/deej/serial_test.go new file mode 100644 index 00000000..61a6603f --- /dev/null +++ b/pkg/deej/serial_test.go @@ -0,0 +1,169 @@ +package deej + +import ( + "testing" + + "go.uber.org/zap" +) + +// newTestSerialIO builds a SerialIO with just enough wiring for handleLine, +// plus a buffered channel that captures emitted slider move events. +func newTestSerialIO(invertSliders bool, noiseReductionLevel string) (*SerialIO, chan SliderMoveEvent) { + config := &CanonicalConfig{} + config.current.Store(&ConfigValues{ + InvertSliders: invertSliders, + NoiseReductionLevel: noiseReductionLevel, + }) + + sio := &SerialIO{ + deej: &Deej{config: config}, + logger: zap.NewNop().Sugar(), + } + + events := make(chan SliderMoveEvent, 64) + sio.sliderMoveConsumers = append(sio.sliderMoveConsumers, events) + + return sio, events +} + +func drainEvents(ch chan SliderMoveEvent) []SliderMoveEvent { + var events []SliderMoveEvent + for { + select { + case e := <-ch: + events = append(events, e) + default: + return events + } + } +} + +func TestHandleLineParsesValidLine(t *testing.T) { + sio, ch := newTestSerialIO(false, "") + + sio.handleLine(sio.logger, "0|512|1023\r\n") + + events := drainEvents(ch) + if len(events) != 3 { + t.Fatalf("got %d events, expected 3", len(events)) + } + + expected := []SliderMoveEvent{ + {SliderID: 0, PercentValue: 0.0}, + {SliderID: 1, PercentValue: 0.5}, + {SliderID: 2, PercentValue: 1.0}, + } + for i, e := range events { + if e != expected[i] { + t.Errorf("event %d = %+v, expected %+v", i, e, expected[i]) + } + } +} + +func TestHandleLineIgnoresMalformedLines(t *testing.T) { + tests := []struct { + name string + line string + }{ + {"empty", ""}, + {"garbage", "hello world\r\n"}, + {"missing CR", "512|512\n"}, + {"missing line ending", "512|512"}, + {"non-numeric value", "512|abc\r\n"}, + {"negative value", "-1|512\r\n"}, + {"five digit value", "10000|512\r\n"}, + {"trailing pipe", "512|512|\r\n"}, + {"dirty first value", "4558|925|41\r\n"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + sio, ch := newTestSerialIO(false, "") + + sio.handleLine(sio.logger, tt.line) + + if events := drainEvents(ch); len(events) != 0 { + t.Errorf("line %q produced %d events, expected none", tt.line, len(events)) + } + }) + } +} + +func TestHandleLineNoiseReduction(t *testing.T) { + sio, ch := newTestSerialIO(false, "") + + // first line always emits (values initialized to an impossible -1023) + sio.handleLine(sio.logger, "512|512\r\n") + if events := drainEvents(ch); len(events) != 2 { + t.Fatalf("initial line produced %d events, expected 2", len(events)) + } + + // jitter below the default threshold of 10 must be filtered out + sio.handleLine(sio.logger, "515|509\r\n") + if events := drainEvents(ch); len(events) != 0 { + t.Errorf("noisy line produced %d events, expected none", len(events)) + } + + // a real move at/above the threshold must go through + sio.handleLine(sio.logger, "525|512\r\n") + events := drainEvents(ch) + if len(events) != 1 { + t.Fatalf("got %d events, expected 1", len(events)) + } + if events[0].SliderID != 0 { + t.Errorf("event came from slider %d, expected 0", events[0].SliderID) + } +} + +func TestHandleLineNoiseReductionNone(t *testing.T) { + sio, ch := newTestSerialIO(false, "none") + + sio.handleLine(sio.logger, "512\r\n") + drainEvents(ch) + + // with "none", even a difference of 1 is significant + sio.handleLine(sio.logger, "513\r\n") + if events := drainEvents(ch); len(events) != 1 { + t.Errorf("got %d events, expected 1", len(events)) + } +} + +func TestHandleLineEdgeSnapping(t *testing.T) { + sio, ch := newTestSerialIO(false, "") + + sio.handleLine(sio.logger, "1013\r\n") + drainEvents(ch) + + // near the top edge the threshold drops to 5, so a move of 5 emits + // even though it is below the default threshold of 10 + sio.handleLine(sio.logger, "1018\r\n") + events := drainEvents(ch) + if len(events) != 1 { + t.Fatalf("got %d events, expected 1", len(events)) + } + if events[0].PercentValue != 1.0 { + t.Errorf("edge value = %v, expected 1.0", events[0].PercentValue) + } +} + +func TestHandleLineSliderCountChange(t *testing.T) { + sio, ch := newTestSerialIO(false, "") + + sio.handleLine(sio.logger, "512\r\n") + if events := drainEvents(ch); len(events) != 1 { + t.Fatalf("got %d events, expected 1", len(events)) + } + if sio.lastKnownNumSliders != 1 { + t.Errorf("lastKnownNumSliders = %d, expected 1", sio.lastKnownNumSliders) + } + + // a different slider count resets state and re-emits everything, + // including sliders whose values did not change + sio.handleLine(sio.logger, "512|512\r\n") + if events := drainEvents(ch); len(events) != 2 { + t.Errorf("got %d events after count change, expected 2", len(events)) + } + if sio.lastKnownNumSliders != 2 { + t.Errorf("lastKnownNumSliders = %d, expected 2", sio.lastKnownNumSliders) + } +} diff --git a/pkg/deej/session_finder_linux.go b/pkg/deej/session_finder_linux.go index 57b98bdd..c756fd18 100644 --- a/pkg/deej/session_finder_linux.go +++ b/pkg/deej/session_finder_linux.go @@ -9,22 +9,55 @@ import ( "github.com/jfreymuth/pulse/proto" "go.uber.org/zap" + + "github.com/nik9play/deej/pkg/reconnect" ) const ( - reconnectDelay = 2 * time.Second - // buffer size for the event work channel paWorkChanSize = 512 + + // how often watch pings the server to catch silently dead connections + paKeepaliveInterval = 30 * time.Second ) +// paConn is one PulseAudio connection generation. +type paConn struct { + client *proto.Client + conn net.Conn + closedCh chan struct{} + closeOnce *sync.Once + logger *zap.SugaredLogger +} + +func (c paConn) forceClose() { + c.closeOnce.Do(func() { close(c.closedCh) }) +} + +// request performs one round-trip on this connection +func (c paConn) request(req proto.RequestArgs, rpl proto.Reply) error { + err := c.client.Request(req, rpl) + if err == nil { + return nil + } + + var protoErr proto.Error + if !errors.As(err, &protoErr) { + c.logger.Debugw("Transport-level request failure, marking connection closed", "error", err) + c.forceClose() + } + + return err +} + type paSessionFinder struct { logger *zap.SugaredLogger sessionLogger *zap.SugaredLogger - mu sync.RWMutex - client *proto.Client - conn net.Conn + reconnector *reconnect.Reconnector[paConn] + + // mu guards the session maps + mu sync.Mutex masterSink *masterSession masterSource *masterSession sinkInputs map[uint32]*paSession @@ -33,8 +66,7 @@ type paSessionFinder struct { namedSinks map[uint32]*masterSession namedSources map[uint32]*masterSession - // receives session events synchronously; set once by Start before any - // connection is made + // receives session events synchronously handler SessionEventHandler started bool @@ -42,8 +74,7 @@ type paSessionFinder struct { // preserving the order in which PulseAudio delivered the events workChan chan func() - reconnectCh chan struct{} - stopCh chan struct{} + stopCh chan struct{} } func newSessionFinder(logger *zap.SugaredLogger) (SessionFinder, error) { @@ -54,15 +85,23 @@ func newSessionFinder(logger *zap.SugaredLogger) (SessionFinder, error) { namedSinks: make(map[uint32]*masterSession), namedSources: make(map[uint32]*masterSession), workChan: make(chan func(), paWorkChanSize), - reconnectCh: make(chan struct{}, 1), stopCh: make(chan struct{}), } + sf.reconnector = reconnect.New(reconnect.Options[paConn]{ + Logger: sf.logger, + Dial: sf.dial, + Watch: sf.watch, + Close: sf.close, + OnUp: sf.onUp, + OnDown: sf.onDown, + }) + sf.logger.Debug("Created event-driven PA session finder") return sf, nil } -// Start begins session discovery, delivering events synchronously to handler +// Begins session discovery, delivering events synchronously to handler. func (sf *paSessionFinder) Start(handler SessionEventHandler) error { if sf.started { return errors.New("session finder already started") @@ -70,12 +109,10 @@ func (sf *paSessionFinder) Start(handler SessionEventHandler) error { sf.started = true sf.handler = handler - if err := sf.connect(); err != nil { - return err - } - + // the worker must be running before the first connection is made, since + // onUp queues the initial enumeration on it go sf.eventWorker() - go sf.connectionManager() + sf.reconnector.Start() return nil } @@ -104,41 +141,122 @@ func (sf *paSessionFinder) dispatchWork(fn func()) { } } -func (sf *paSessionFinder) connectionManager() { - for { - select { - case <-sf.stopCh: - return - case <-sf.reconnectCh: - sf.handleReconnect() +// dial connects to PulseAudio and completes the handshake (client name, +// event subscription) without touching any shared state - the reconnector +// decides whether to adopt the returned connection +func (sf *paSessionFinder) dial() (paConn, error) { + client, conn, err := proto.Connect("") + if err != nil { + return paConn{}, fmt.Errorf("connect to PulseAudio: %w", err) + } + + newConn := paConn{ + client: client, + conn: conn, + closedCh: make(chan struct{}), + closeOnce: &sync.Once{}, + logger: sf.logger, + } + + // the callback belongs to this connection only: subscription events are + // dispatched to the shared work queue, but a ConnectionClosed can only + // signal its own generation's closed channel, so a stale callback can't + // tear down a newer connection + client.Callback = func(msg any) { + switch v := msg.(type) { + case *proto.SubscribeEvent: + sf.handleSubscribeEvent(v) + case *proto.ConnectionClosed: + newConn.forceClose() } } + + if err := client.Request(&proto.SetClientName{ + Props: proto.PropList{"application.name": proto.PropListString("deej")}, + }, &proto.SetClientNameReply{}); err != nil { + _ = conn.Close() + return paConn{}, fmt.Errorf("set client name: %w", err) + } + + if err := client.Request(&proto.Subscribe{ + Mask: proto.SubscriptionMaskSinkInput | proto.SubscriptionMaskServer | proto.SubscriptionMaskSink | proto.SubscriptionMaskSource, + }, nil); err != nil { + _ = conn.Close() + return paConn{}, fmt.Errorf("subscribe to events: %w", err) + } + + return newConn, nil } -func (sf *paSessionFinder) handleReconnect() { - sf.clearSessions() +// watch blocks until this connection dies. detection has two prongs: the +// connection's callback observing ConnectionClosed (which the pulse client +// only fires on clean EOF), and transport-level request failures - including +// the periodic keepalive here - marking the connection closed via forceClose +func (sf *paSessionFinder) watch(conn paConn, errChannel chan<- error) { + keepalive := time.NewTicker(paKeepaliveInterval) + defer keepalive.Stop() for { select { - case <-sf.stopCh: + case <-conn.closedCh: + select { + case errChannel <- errors.New("PulseAudio connection closed"): + default: + } return - default: - } - sf.logger.Debug("Attempting to reconnect to PulseAudio") - if err := sf.connect(); err != nil { - sf.logger.Debugw("Reconnect failed, retrying", "error", err) - time.Sleep(reconnectDelay) - continue + case <-keepalive.C: + // a failed keepalive marks the connection closed inside + // conn.request; the next loop iteration picks that up + _ = conn.request(&proto.GetServerInfo{}, &proto.GetServerInfoReply{}) } - sf.logger.Info("Reconnected to PulseAudio") - return + } +} + +func (sf *paSessionFinder) close(conn paConn) { + _ = conn.conn.Close() + + conn.forceClose() + + sf.logger.Debug("Closed PulseAudio connection") +} + +func (sf *paSessionFinder) onUp(_ paConn) { + sf.logger.Info("Connected to PulseAudio") + + // queue the initial state sync + sf.dispatchWork(func() { + sf.refreshMaster() + sf.enumerateExistingSessions() + sf.enumerateExistingDevices() + }) +} + +func (sf *paSessionFinder) onDown(err error) { + sf.logger.Warnw("PulseAudio connection lost, clearing sessions and reconnecting", "error", err) + sf.clearSessions() +} + +// handleSubscribeEvent routes one subscription event from a connection's +// callback onto the serial work queue +func (sf *paSessionFinder) handleSubscribeEvent(v *proto.SubscribeEvent) { + eventType := v.Event.GetType() + index := v.Index + + switch v.Event.GetFacility() { + case proto.EventSinkSinkInput: + sf.dispatchWork(func() { sf.handleSinkInputEvent(eventType, index) }) + case proto.EventServer: + sf.dispatchWork(sf.refreshMaster) + case proto.EventSink: + sf.dispatchWork(func() { sf.handleSinkEvent(eventType, index) }) + case proto.EventSource: + sf.dispatchWork(func() { sf.handleSourceEvent(eventType, index) }) } } func (sf *paSessionFinder) clearSessions() { - // collect all sessions under the lock, then notify outside of it: the - // handler runs synchronously and must not be called while holding sf.mu + // collect all sessions under the lock, then notify outside of it sf.mu.Lock() removed := make([]Session, 0, len(sf.sinkInputs)+len(sf.namedSinks)+len(sf.namedSources)+2) @@ -167,12 +285,6 @@ func (sf *paSessionFinder) clearSessions() { sf.masterSource = nil } - if sf.conn != nil { - sf.conn.Close() - sf.conn = nil - } - sf.client = nil - sf.mu.Unlock() for _, s := range removed { @@ -181,76 +293,6 @@ func (sf *paSessionFinder) clearSessions() { } } -func (sf *paSessionFinder) connect() error { - client, conn, err := proto.Connect("") - if err != nil { - return fmt.Errorf("connect to PulseAudio: %w", err) - } - - client.Callback = sf.onPulseEvent - - if err := client.Request(&proto.SetClientName{ - Props: proto.PropList{"application.name": proto.PropListString("deej")}, - }, &proto.SetClientNameReply{}); err != nil { - conn.Close() - return fmt.Errorf("set client name: %w", err) - } - - sf.mu.Lock() - sf.client = client - sf.conn = conn - sf.mu.Unlock() - - // queue the initial enumeration before subscribing: subscription events go - // through the same queue, so anything that changes during or after the - // enumeration is processed strictly after it. subscribing after - // enumerating directly (the old order) would lose sessions created in - // between, with no periodic refresh to ever pick them up - sf.dispatchWork(func() { - sf.refreshMaster() - sf.enumerateExistingSessions() - sf.enumerateExistingDevices() - }) - - if err := client.Request(&proto.Subscribe{ - Mask: proto.SubscriptionMaskSinkInput | proto.SubscriptionMaskServer | proto.SubscriptionMaskSink | proto.SubscriptionMaskSource, - }, nil); err != nil { - conn.Close() - return fmt.Errorf("subscribe to events: %w", err) - } - - return nil -} - -func (sf *paSessionFinder) requestReconnect() { - select { - case sf.reconnectCh <- struct{}{}: - default: - } -} - -func (sf *paSessionFinder) onPulseEvent(msg any) { - switch v := msg.(type) { - case *proto.SubscribeEvent: - eventType := v.Event.GetType() - index := v.Index - - switch v.Event.GetFacility() { - case proto.EventSinkSinkInput: - sf.dispatchWork(func() { sf.handleSinkInputEvent(eventType, index) }) - case proto.EventServer: - sf.dispatchWork(sf.refreshMaster) - case proto.EventSink: - sf.dispatchWork(func() { sf.handleSinkEvent(eventType, index) }) - case proto.EventSource: - sf.dispatchWork(func() { sf.handleSourceEvent(eventType, index) }) - } - case *proto.ConnectionClosed: - sf.logger.Warn("PulseAudio connection closed, trying to reconnect") - sf.requestReconnect() - } -} - func (sf *paSessionFinder) handleSinkInputEvent(eventType proto.SubscriptionEventType, index uint32) { switch eventType { case proto.EventNew: @@ -266,22 +308,20 @@ func (sf *paSessionFinder) refreshMaster() { } func (sf *paSessionFinder) refreshMasterSink() { - sf.mu.RLock() - client := sf.client - sf.mu.RUnlock() - if client == nil { + conn, ok := sf.reconnector.Current() + if !ok { return } reply := proto.GetSinkInfoReply{} - if err := client.Request(&proto.GetSinkInfo{SinkIndex: proto.Undefined}, &reply); err != nil { + if err := conn.request(&proto.GetSinkInfo{SinkIndex: proto.Undefined}, &reply); err != nil { sf.logger.Debugw("Failed to get master sink info", "error", err) return } sf.mu.Lock() old := sf.masterSink - newMaster := newMasterSession(sf.sessionLogger, sf.client, reply.SinkIndex, reply.Channels, true) + newMaster := newMasterSession(sf.sessionLogger, conn, reply.SinkIndex, reply.Channels, true) sf.masterSink = newMaster sf.mu.Unlock() @@ -293,22 +333,20 @@ func (sf *paSessionFinder) refreshMasterSink() { } func (sf *paSessionFinder) refreshMasterSource() { - sf.mu.RLock() - client := sf.client - sf.mu.RUnlock() - if client == nil { + conn, ok := sf.reconnector.Current() + if !ok { return } reply := proto.GetSourceInfoReply{} - if err := client.Request(&proto.GetSourceInfo{SourceIndex: proto.Undefined}, &reply); err != nil { + if err := conn.request(&proto.GetSourceInfo{SourceIndex: proto.Undefined}, &reply); err != nil { sf.logger.Debugw("Failed to get master source info", "error", err) return } sf.mu.Lock() old := sf.masterSource - newMaster := newMasterSession(sf.sessionLogger, sf.client, reply.SourceIndex, reply.Channels, false) + newMaster := newMasterSession(sf.sessionLogger, conn, reply.SourceIndex, reply.Channels, false) sf.masterSource = newMaster sf.mu.Unlock() @@ -320,15 +358,13 @@ func (sf *paSessionFinder) refreshMasterSource() { } func (sf *paSessionFinder) enumerateExistingSessions() { - sf.mu.RLock() - client := sf.client - sf.mu.RUnlock() - if client == nil { + conn, ok := sf.reconnector.Current() + if !ok { return } reply := proto.GetSinkInputInfoListReply{} - if err := client.Request(&proto.GetSinkInputInfoList{}, &reply); err != nil { + if err := conn.request(&proto.GetSinkInputInfoList{}, &reply); err != nil { sf.logger.Errorw("Failed to enumerate sessions", "error", err) return } @@ -340,15 +376,13 @@ func (sf *paSessionFinder) enumerateExistingSessions() { } func (sf *paSessionFinder) addSinkInput(index uint32) { - sf.mu.RLock() - client := sf.client - sf.mu.RUnlock() - if client == nil { + conn, ok := sf.reconnector.Current() + if !ok { return } reply := proto.GetSinkInputInfoReply{} - if err := client.Request(&proto.GetSinkInputInfo{SinkInputIndex: index}, &reply); err != nil { + if err := conn.request(&proto.GetSinkInputInfo{SinkInputIndex: index}, &reply); err != nil { sf.logger.Debugw("Failed to get sink input info", "index", index, "error", err) return } @@ -356,6 +390,11 @@ func (sf *paSessionFinder) addSinkInput(index uint32) { } func (sf *paSessionFinder) addSinkInputFromInfo(info *proto.GetSinkInputInfoReply) { + conn, ok := sf.reconnector.Current() + if !ok { + return + } + // Try application.process.binary, then application.id, then application.name name, ok := info.Properties["application.process.binary"] if !ok { @@ -373,7 +412,7 @@ func (sf *paSessionFinder) addSinkInputFromInfo(info *proto.GetSinkInputInfoRepl sf.mu.Unlock() return } - session := newPASession(sf.sessionLogger, sf.client, info.SinkInputIndex, info.Channels, name.String()) + session := newPASession(sf.sessionLogger, conn, info.SinkInputIndex, info.Channels, name.String()) sf.sinkInputs[info.SinkInputIndex] = session sf.mu.Unlock() @@ -402,15 +441,13 @@ func (sf *paSessionFinder) enumerateExistingDevices() { } func (sf *paSessionFinder) enumerateExistingSinks() { - sf.mu.RLock() - client := sf.client - sf.mu.RUnlock() - if client == nil { + conn, ok := sf.reconnector.Current() + if !ok { return } reply := proto.GetSinkInfoListReply{} - if err := client.Request(&proto.GetSinkInfoList{}, &reply); err != nil { + if err := conn.request(&proto.GetSinkInfoList{}, &reply); err != nil { sf.logger.Errorw("Failed to enumerate sinks", "error", err) return } @@ -422,15 +459,13 @@ func (sf *paSessionFinder) enumerateExistingSinks() { } func (sf *paSessionFinder) enumerateExistingSources() { - sf.mu.RLock() - client := sf.client - sf.mu.RUnlock() - if client == nil { + conn, ok := sf.reconnector.Current() + if !ok { return } reply := proto.GetSourceInfoListReply{} - if err := client.Request(&proto.GetSourceInfoList{}, &reply); err != nil { + if err := conn.request(&proto.GetSourceInfoList{}, &reply); err != nil { sf.logger.Errorw("Failed to enumerate sources", "error", err) return } @@ -460,15 +495,13 @@ func (sf *paSessionFinder) handleSourceEvent(eventType proto.SubscriptionEventTy } func (sf *paSessionFinder) addSink(index uint32) { - sf.mu.RLock() - client := sf.client - sf.mu.RUnlock() - if client == nil { + conn, ok := sf.reconnector.Current() + if !ok { return } reply := proto.GetSinkInfoReply{} - if err := client.Request(&proto.GetSinkInfo{SinkIndex: index}, &reply); err != nil { + if err := conn.request(&proto.GetSinkInfo{SinkIndex: index}, &reply); err != nil { sf.logger.Debugw("Failed to get sink info", "index", index, "error", err) return } @@ -476,6 +509,11 @@ func (sf *paSessionFinder) addSink(index uint32) { } func (sf *paSessionFinder) addSinkFromInfo(info *proto.GetSinkInfoReply) { + conn, ok := sf.reconnector.Current() + if !ok { + return + } + // Get description from properties, fallback to sink name description := info.Device if description == "" { @@ -489,7 +527,7 @@ func (sf *paSessionFinder) addSinkFromInfo(info *proto.GetSinkInfoReply) { sf.mu.Unlock() return } - session := newNamedMasterSession(sf.sessionLogger, sf.client, info.SinkIndex, info.Channels, true, description) + session := newNamedMasterSession(sf.sessionLogger, conn, info.SinkIndex, info.Channels, true, description) session.device = true sf.namedSinks[info.SinkIndex] = session sf.mu.Unlock() @@ -499,15 +537,13 @@ func (sf *paSessionFinder) addSinkFromInfo(info *proto.GetSinkInfoReply) { } func (sf *paSessionFinder) addSource(index uint32) { - sf.mu.RLock() - client := sf.client - sf.mu.RUnlock() - if client == nil { + conn, ok := sf.reconnector.Current() + if !ok { return } reply := proto.GetSourceInfoReply{} - if err := client.Request(&proto.GetSourceInfo{SourceIndex: index}, &reply); err != nil { + if err := conn.request(&proto.GetSourceInfo{SourceIndex: index}, &reply); err != nil { sf.logger.Debugw("Failed to get source info", "index", index, "error", err) return } @@ -515,6 +551,11 @@ func (sf *paSessionFinder) addSource(index uint32) { } func (sf *paSessionFinder) addSourceFromInfo(info *proto.GetSourceInfoReply) { + conn, ok := sf.reconnector.Current() + if !ok { + return + } + // Skip monitor sources (they mirror sink outputs) if info.MonitorSourceName != "" { return @@ -533,7 +574,7 @@ func (sf *paSessionFinder) addSourceFromInfo(info *proto.GetSourceInfoReply) { sf.mu.Unlock() return } - session := newNamedMasterSession(sf.sessionLogger, sf.client, info.SourceIndex, info.Channels, false, description) + session := newNamedMasterSession(sf.sessionLogger, conn, info.SourceIndex, info.Channels, false, description) session.device = true sf.namedSources[info.SourceIndex] = session sf.mu.Unlock() @@ -580,15 +621,9 @@ func (sf *paSessionFinder) notify(event SessionEvent) { } func (sf *paSessionFinder) Release() error { + sf.reconnector.Stop() close(sf.stopCh) - sf.mu.Lock() - conn := sf.conn - sf.mu.Unlock() - - if conn != nil { - conn.Close() - } sf.logger.Debug("Released PA session finder") return nil } diff --git a/pkg/deej/session_linux.go b/pkg/deej/session_linux.go index b41457a2..07a6f0ac 100644 --- a/pkg/deej/session_linux.go +++ b/pkg/deej/session_linux.go @@ -16,7 +16,7 @@ type paSession struct { processName string - client *proto.Client + conn paConn sinkInputIndex uint32 sinkInputChannels byte @@ -25,7 +25,7 @@ type paSession struct { type masterSession struct { baseSession - client *proto.Client + conn paConn streamIndex uint32 streamChannels byte @@ -34,14 +34,14 @@ type masterSession struct { func newPASession( logger *zap.SugaredLogger, - client *proto.Client, + conn paConn, sinkInputIndex uint32, sinkInputChannels byte, processName string, ) *paSession { s := &paSession{ - client: client, + conn: conn, sinkInputIndex: sinkInputIndex, sinkInputChannels: sinkInputChannels, } @@ -59,7 +59,7 @@ func newPASession( func newMasterSession( logger *zap.SugaredLogger, - client *proto.Client, + conn paConn, streamIndex uint32, streamChannels byte, isOutput bool, @@ -71,19 +71,19 @@ func newMasterSession( key = inputSessionName } - return newNamedMasterSession(logger, client, streamIndex, streamChannels, isOutput, key) + return newNamedMasterSession(logger, conn, streamIndex, streamChannels, isOutput, key) } func newNamedMasterSession( logger *zap.SugaredLogger, - client *proto.Client, + conn paConn, streamIndex uint32, streamChannels byte, isOutput bool, name string, ) *masterSession { s := &masterSession{ - client: client, + conn: conn, streamIndex: streamIndex, streamChannels: streamChannels, isOutput: isOutput, @@ -105,13 +105,12 @@ func (s *paSession) GetVolume() float32 { } reply := proto.GetSinkInputInfoReply{} - if err := s.client.Request(&request, &reply); err != nil { + if err := s.conn.request(&request, &reply); err != nil { s.logger.Warnw("Failed to get session volume", "error", err) + return 0 } - level := parseChannelVolumes(reply.ChannelVolumes) - - return level + return parseChannelVolumes(reply.ChannelVolumes) } func (s *paSession) SetVolume(v float32) error { @@ -121,7 +120,7 @@ func (s *paSession) SetVolume(v float32) error { ChannelVolumes: volumes, } - if err := s.client.Request(&request, nil); err != nil { + if err := s.conn.request(&request, nil); err != nil { s.logger.Warnw("Failed to set session volume", "error", err) return fmt.Errorf("adjust session volume: %w", err) } @@ -148,7 +147,7 @@ func (s *masterSession) GetVolume() float32 { } reply := proto.GetSinkInfoReply{} - if err := s.client.Request(&request, &reply); err != nil { + if err := s.conn.request(&request, &reply); err != nil { s.logger.Warnw("Failed to get session volume", "error", err) return 0 } @@ -160,7 +159,7 @@ func (s *masterSession) GetVolume() float32 { } reply := proto.GetSourceInfoReply{} - if err := s.client.Request(&request, &reply); err != nil { + if err := s.conn.request(&request, &reply); err != nil { s.logger.Warnw("Failed to get session volume", "error", err) return 0 } @@ -188,7 +187,7 @@ func (s *masterSession) SetVolume(v float32) error { } } - if err := s.client.Request(request, nil); err != nil { + if err := s.conn.request(request, nil); err != nil { s.logger.Warnw("Failed to set session volume", "error", err, "volume", v) diff --git a/pkg/deej/tray.go b/pkg/deej/tray.go index 569826af..a94ade8d 100644 --- a/pkg/deej/tray.go +++ b/pkg/deej/tray.go @@ -89,7 +89,7 @@ func getStatusItemTitle(d *Deej) string { Other: "Connected to {{.ComPort}}", }, TemplateData: map[string]string{ - "ComPort": d.serial.comPortToUse, + "ComPort": d.serial.CurrentComPort(), }, }) } else { diff --git a/pkg/reconnect/reconnect.go b/pkg/reconnect/reconnect.go new file mode 100644 index 00000000..5cab785e --- /dev/null +++ b/pkg/reconnect/reconnect.go @@ -0,0 +1,317 @@ +// This package manages the lifecycle of a long-lived connection that +// must be re-established whenever it drops. +// +// Contracts between the reconnector and its callbacks: +// +// - Close must cause a blocked Watch to return. +// Watch does not need to observe any stop signal of its own. +// +// - Watch reports a failure with a single non-blocking send to errChannel +// and then returns. The channel belongs to a single connection +// generation, so a slow Watch can never tear down a newer connection. +// The send must be non-blocking because Reconnect +// shares the channel and may have already filled its buffer. +// +// - If Dial needs to remember the configuration it dialed with (for change +// detection on config reload), it must embed that snapshot in C so it +// travels with the connection it produced, rather than the client +// re-reading config later and comparing against a moving target. +package reconnect + +import ( + "sync" + "time" + + "go.uber.org/zap" +) + +// DefaultRetryDelay is the fixed delay between attempts +const DefaultRetryDelay = 2 * time.Second + +// Options configures a Reconnector. +type Options[C any] struct { + Logger *zap.SugaredLogger + + // Enabled gates connection attempts. + // When nil, the reconnector is always on. + Enabled func() bool + + // Dial establishes a new connection. It may block and is not cancellable; + // the reconnector calls it from its own goroutine so that Stop never + // waits for a dial in flight. It must be free of side effects on shared + // state - the reconnector decides whether the result is adopted. + Dial func() (C, error) + + // Watch runs for the lifetime of a single connection and returns when + // that connection is finished. When the connection fails on its own, + // Watch sends one error to errChannel before returning; when it returns + // because Close unblocked it, no send is needed. + Watch func(conn C, errChannel chan<- error) + + // Close releases a connection. It must unblock a Watch that is blocked + // on the connection. It is called at most once per connection. + Close func(conn C) + + // OnUp is called after a connection is adopted, before Watch starts. + OnUp func(conn C) + + // OnDown is called when an established connection is lost (or torn down + // via Reconnect), before Close. It is not called for failed dials. + OnDown func(err error) + + // Backoff returns the delay before the next dial, given the number of + // consecutive failures so far (0 for the first retry after a failure, + // increasing while failures continue, reset when a connection is + // adopted). When nil, a fixed DefaultRetryDelay is used. + Backoff func(attempt int) time.Duration +} + +// Reconnector maintains at most one live connection of type C, redialing +// with backoff whenever the connection drops. +// +// Start, Stop and the Reconnector's construction must happen on one +// goroutine (or be otherwise serialized); all other methods are safe to call +// from any goroutine. +type Reconnector[C any] struct { + opts Options[C] + logger *zap.SugaredLogger + + // current, hasConn and errChannel describe the current connection + // generation and are replaced together under lock on every (re)connect + lock sync.Mutex + current C + hasConn bool + errChannel chan error + + stopChannel chan struct{} + wg sync.WaitGroup +} + +// New creates a Reconnector from opts. +// Dial, Watch or Close are required +func New[C any](opts Options[C]) *Reconnector[C] { + if opts.Dial == nil || opts.Watch == nil || opts.Close == nil { + panic("reconnect: Options.Dial, Options.Watch and Options.Close are required") + } + + logger := opts.Logger + if logger == nil { + logger = zap.NewNop().Sugar() + } + + return &Reconnector[C]{ + opts: opts, + logger: logger, + } +} + +// Start launches the manager loop. +func (r *Reconnector[C]) Start() { + stopChannel := make(chan struct{}) + r.stopChannel = stopChannel + + r.wg.Go(func() { + r.managerLoop(stopChannel) + }) +} + +// Stop tears down the current connection (if any) and waits for the manager +// loop and its Watch to exit. It never waits for a dial in flight. +func (r *Reconnector[C]) Stop() { + if r.stopChannel == nil { + return + } + + close(r.stopChannel) + r.wg.Wait() + r.stopChannel = nil +} + +// Reconnect asks the manager loop to tear down the current connection and +// dial again, reporting reason to OnDown. +func (r *Reconnector[C]) Reconnect(reason error) { + r.lock.Lock() + errChannel := r.errChannel + r.lock.Unlock() + + if errChannel == nil { + return + } + + select { + case errChannel <- reason: + default: + // channel full, teardown already pending + } +} + +// Current returns the live connection, or false when disconnected +func (r *Reconnector[C]) Current() (C, bool) { + r.lock.Lock() + defer r.lock.Unlock() + + return r.current, r.hasConn +} + +// Connected reports whether a connection is currently established. +func (r *Reconnector[C]) Connected() bool { + r.lock.Lock() + defer r.lock.Unlock() + + return r.hasConn +} + +func (r *Reconnector[C]) backoffDelay(attempt int) time.Duration { + if r.opts.Backoff == nil { + return DefaultRetryDelay + } + + return r.opts.Backoff(attempt) +} + +// sleeps for delay, returning false early +// if stopChannel closes first. +func (r *Reconnector[C]) waitOrStop(stopChannel <-chan struct{}, delay time.Duration) bool { + select { + case <-stopChannel: + return false + case <-time.After(delay): + return true + } +} + +// drop clears the current connection generation and closes its connection. +func (r *Reconnector[C]) drop() { + r.lock.Lock() + + if !r.hasConn { + r.lock.Unlock() + return + } + + conn := r.current + + var zero C + r.current = zero + r.hasConn = false + r.errChannel = nil + + r.lock.Unlock() + + // close outside the lock: Close may block until Watch unblocks, and + // request paths calling Current must not stall behind it + r.opts.Close(conn) +} + +type dialResult[C any] struct { + conn C + err error +} + +func (r *Reconnector[C]) managerLoop(stopChannel <-chan struct{}) { + // consecutive failures since the last adopted connection + attempt := 0 + + for { + if r.opts.Enabled != nil && !r.opts.Enabled() { + if !r.waitOrStop(stopChannel, r.backoffDelay(0)) { + r.logger.Debug("managerLoop: stop signal") + return + } + + continue + } + + // dial in a goroutine so we can respond to the stop signal + dialDone := make(chan dialResult[C], 1) + go func() { + conn, err := r.opts.Dial() + dialDone <- dialResult[C]{conn: conn, err: err} + }() + + var result dialResult[C] + + select { + case <-stopChannel: + r.logger.Debug("managerLoop: stop signal during dial") + + // let it resolve in the background and close the connection if + // it ends up succeeding. + go func() { + if late := <-dialDone; late.err == nil { + r.opts.Close(late.conn) + } + }() + return + + case result = <-dialDone: + } + + if result.err != nil { + r.logger.Debugw("Dial failed, retrying...", "error", result.err, "attempt", attempt) + + if !r.waitOrStop(stopChannel, r.backoffDelay(attempt)) { + r.logger.Debug("managerLoop: stop signal") + return + } + + attempt++ + continue + } + + // if the connector was turned off while dialing, + // drop the connection + if r.opts.Enabled != nil && !r.opts.Enabled() { + r.logger.Debug("Disabled while dialing, dropping connection") + r.opts.Close(result.conn) + continue + } + + // publish connection together with a fresh error + // channel for this connection generation + errChannel := make(chan error, 1) + + r.lock.Lock() + r.current = result.conn + r.hasConn = true + r.errChannel = errChannel + r.lock.Unlock() + + attempt = 0 + + r.logger.Debug("Connection established") + + if r.opts.OnUp != nil { + r.opts.OnUp(result.conn) + } + + // watch this connection to detect disconnection. it receives this + // generation's connection and error channel + r.wg.Go(func() { + r.opts.Watch(result.conn, errChannel) + }) + + select { + case <-stopChannel: + r.logger.Debug("managerLoop: stop signal") + r.drop() + return + + case err := <-errChannel: + r.logger.Debugw("Connection lost, reconnecting...", "error", err) + + if r.opts.OnDown != nil { + r.opts.OnDown(err) + } + + r.drop() + + if !r.waitOrStop(stopChannel, r.backoffDelay(attempt)) { + r.logger.Debug("managerLoop: stop signal") + return + } + + attempt++ + } + } +} diff --git a/pkg/reconnect/reconnect_test.go b/pkg/reconnect/reconnect_test.go new file mode 100644 index 00000000..ac93e1ac --- /dev/null +++ b/pkg/reconnect/reconnect_test.go @@ -0,0 +1,392 @@ +package reconnect + +import ( + "errors" + "testing" + "time" +) + +const testTimeout = 5 * time.Second + +// testConn is the fake connection type driven by the fakeConnector. +// done is closed by Close, which unblocks the fakeConnector's Watch +type testConn struct { + id int + done chan struct{} +} + +type watchHandle struct { + conn testConn + errChannel chan<- error +} + +type connectResult struct { + err error +} + +// fakeConnector scripts a Reconnector +type fakeConnector struct { + r *Reconnector[testConn] + + fakeConnRes chan connectResult + dialStarted chan struct{} + watches chan watchHandle + ups chan int + downs chan error + closes chan int + backoffs chan int + + enabled chan bool // when non-nil, holds the gate's current value +} + +func newFakeConnector(backoff func(int) time.Duration, gated bool) *fakeConnector { + f := &fakeConnector{ + fakeConnRes: make(chan connectResult, 16), + dialStarted: make(chan struct{}, 16), + watches: make(chan watchHandle, 16), + ups: make(chan int, 16), + downs: make(chan error, 16), + closes: make(chan int, 16), + backoffs: make(chan int, 64), + } + + opts := Options[testConn]{ + Dial: f.dial, + Watch: f.watch, + Close: f.close, + OnUp: func(c testConn) { f.ups <- c.id }, + OnDown: func(err error) { f.downs <- err }, + Backoff: nil, + } + + if backoff != nil { + opts.Backoff = func(attempt int) time.Duration { + select { + case f.backoffs <- attempt: + default: + } + return backoff(attempt) + } + } + + if gated { + f.enabled = make(chan bool, 1) + f.enabled <- false + opts.Enabled = func() bool { + v := <-f.enabled + f.enabled <- v + return v + } + } + + f.r = New(opts) + return f +} + +func (f *fakeConnector) setEnabled(v bool) { + <-f.enabled + f.enabled <- v +} + +var nextConnID = make(chan int) + +func init() { + go func() { + for id := 1; ; id++ { + nextConnID <- id + } + }() +} + +func (f *fakeConnector) dial() (testConn, error) { + f.dialStarted <- struct{}{} + + outcome := <-f.fakeConnRes + if outcome.err != nil { + return testConn{}, outcome.err + } + + return testConn{id: <-nextConnID, done: make(chan struct{})}, nil +} + +func (f *fakeConnector) watch(conn testConn, errChannel chan<- error) { + f.watches <- watchHandle{conn: conn, errChannel: errChannel} + <-conn.done +} + +func (f *fakeConnector) close(conn testConn) { + close(conn.done) + f.closes <- conn.id +} + +func expectRecv[T any](t *testing.T, ch <-chan T, what string) T { + t.Helper() + + select { + case v := <-ch: + return v + case <-time.After(testTimeout): + t.Fatalf("timed out waiting for %s", what) + panic("unreachable") + } +} + +func expectNone[T any](t *testing.T, ch <-chan T, what string) { + t.Helper() + + select { + case v := <-ch: + t.Fatalf("unexpected %s: %v", what, v) + case <-time.After(150 * time.Millisecond): + } +} + +// stopWithDeadline fails the test if Stop blocks. +func stopWithDeadline(t *testing.T, r *Reconnector[testConn]) { + t.Helper() + + done := make(chan struct{}) + go func() { + r.Stop() + close(done) + }() + + select { + case <-done: + case <-time.After(testTimeout): + t.Fatal("Stop did not return") + } +} + +func fastBackoff(int) time.Duration { return time.Millisecond } + +func TestStopDuringDial(t *testing.T) { + f := newFakeConnector(fastBackoff, false) + + // no scripted outcome yet, so the dial stays in flight + f.r.Start() + expectRecv(t, f.dialStarted, "dial start") + + // Stop must not wait for the un-cancellable dial + stopWithDeadline(t, f.r) + + // when the dial eventually succeeds, the never-adopted connection must + // be closed by the detached cleanup + f.fakeConnRes <- connectResult{} + expectRecv(t, f.closes, "close of never-adopted connection") + expectNone(t, f.ups, "OnUp") +} + +func TestStopDuringRetryWait(t *testing.T) { + f := newFakeConnector(func(int) time.Duration { return time.Hour }, false) + + f.fakeConnRes <- connectResult{err: errors.New("dial failed")} + f.r.Start() + + // wait until the loop is inside its hour-long retry wait, then make + // sure Stop interrupts it + expectRecv(t, f.backoffs, "backoff call") + time.Sleep(50 * time.Millisecond) + + stopWithDeadline(t, f.r) +} + +func TestDialFailRetrySuccess(t *testing.T) { + f := newFakeConnector(fastBackoff, false) + + f.fakeConnRes <- connectResult{err: errors.New("fail 1")} + f.fakeConnRes <- connectResult{err: errors.New("fail 2")} + f.fakeConnRes <- connectResult{} + + f.r.Start() + defer stopWithDeadline(t, f.r) + + upID := expectRecv(t, f.ups, "OnUp") + + if !f.r.Connected() { + t.Error("Connected() = false after OnUp") + } + if conn, ok := f.r.Current(); !ok || conn.id != upID { + t.Errorf("Current() = (%v, %v), expected id %d", conn.id, ok, upID) + } + + // consecutive dial failures must see increasing attempt counts + if a := expectRecv(t, f.backoffs, "first backoff"); a != 0 { + t.Errorf("first backoff attempt = %d, expected 0", a) + } + if a := expectRecv(t, f.backoffs, "second backoff"); a != 1 { + t.Errorf("second backoff attempt = %d, expected 1", a) + } +} + +func TestWatchErrorTriggersReconnect(t *testing.T) { + f := newFakeConnector(fastBackoff, false) + + f.fakeConnRes <- connectResult{} + f.fakeConnRes <- connectResult{} + + f.r.Start() + defer stopWithDeadline(t, f.r) + + firstID := expectRecv(t, f.ups, "first OnUp") + handle := expectRecv(t, f.watches, "first watch") + + watchErr := errors.New("connection lost") + handle.errChannel <- watchErr + + if got := expectRecv(t, f.downs, "OnDown"); !errors.Is(got, watchErr) { + t.Errorf("OnDown error = %v, expected %v", got, watchErr) + } + if closedID := expectRecv(t, f.closes, "close"); closedID != firstID { + t.Errorf("closed connection %d, expected %d", closedID, firstID) + } + + secondID := expectRecv(t, f.ups, "second OnUp") + if secondID == firstID { + t.Error("reconnect did not produce a new connection") + } +} + +func TestReconnectTriggersTeardown(t *testing.T) { + f := newFakeConnector(fastBackoff, false) + + f.fakeConnRes <- connectResult{} + f.fakeConnRes <- connectResult{} + + f.r.Start() + defer stopWithDeadline(t, f.r) + + firstID := expectRecv(t, f.ups, "first OnUp") + + reason := errors.New("config changed") + f.r.Reconnect(reason) + + if got := expectRecv(t, f.downs, "OnDown"); !errors.Is(got, reason) { + t.Errorf("OnDown error = %v, expected %v", got, reason) + } + if closedID := expectRecv(t, f.closes, "close"); closedID != firstID { + t.Errorf("closed connection %d, expected %d", closedID, firstID) + } + + expectRecv(t, f.ups, "second OnUp") +} + +func TestStaleWatchErrorIgnored(t *testing.T) { + f := newFakeConnector(fastBackoff, false) + + f.fakeConnRes <- connectResult{} + f.fakeConnRes <- connectResult{} + + f.r.Start() + defer stopWithDeadline(t, f.r) + + expectRecv(t, f.ups, "first OnUp") + staleHandle := expectRecv(t, f.watches, "first watch") + + staleHandle.errChannel <- errors.New("gen 1 lost") + _ = expectRecv(t, f.downs, "OnDown") + expectRecv(t, f.closes, "close of gen 1") + + secondID := expectRecv(t, f.ups, "second OnUp") + expectRecv(t, f.watches, "second watch") + + // a late error from the first generation's watch must not tear down + // the second generation's connection + staleHandle.errChannel <- errors.New("gen 1 late error") + + expectNone(t, f.downs, "OnDown from stale generation") + + if conn, ok := f.r.Current(); !ok || conn.id != secondID { + t.Errorf("Current() = (%v, %v), expected id %d", conn.id, ok, secondID) + } +} + +func TestEnabledGate(t *testing.T) { + f := newFakeConnector(fastBackoff, true) + + f.r.Start() + defer stopWithDeadline(t, f.r) + + // while disabled, the reconnector must not dial at all + expectNone(t, f.dialStarted, "dial while disabled") + + f.setEnabled(true) + expectRecv(t, f.dialStarted, "dial after enabling") + + f.fakeConnRes <- connectResult{} + expectRecv(t, f.ups, "OnUp") +} + +func TestDisabledWhileDialing(t *testing.T) { + f := newFakeConnector(fastBackoff, true) + + f.setEnabled(true) + f.r.Start() + defer stopWithDeadline(t, f.r) + + expectRecv(t, f.dialStarted, "dial start") + + f.setEnabled(false) + f.fakeConnRes <- connectResult{} + + expectRecv(t, f.closes, "close of non-adopted connection") + expectNone(t, f.ups, "OnUp while disabled") + + if f.r.Connected() { + t.Error("Connected() = true for a dropped connection") + } +} + +func TestBackoffResetAfterConnection(t *testing.T) { + f := newFakeConnector(fastBackoff, false) + + f.fakeConnRes <- connectResult{err: errors.New("fail 1")} + f.fakeConnRes <- connectResult{err: errors.New("fail 2")} + f.fakeConnRes <- connectResult{} + + f.r.Start() + defer stopWithDeadline(t, f.r) + + expectRecv(t, f.ups, "first OnUp") + handle := expectRecv(t, f.watches, "first watch") + + if a := expectRecv(t, f.backoffs, "backoff 1"); a != 0 { + t.Errorf("backoff attempt = %d, expected 0", a) + } + if a := expectRecv(t, f.backoffs, "backoff 2"); a != 1 { + t.Errorf("backoff attempt = %d, expected 1", a) + } + + // losing an established connection must restart the attempt count + f.fakeConnRes <- connectResult{} + handle.errChannel <- errors.New("connection lost") + + _ = expectRecv(t, f.downs, "OnDown") + if a := expectRecv(t, f.backoffs, "backoff after loss"); a != 0 { + t.Errorf("backoff attempt after loss = %d, expected 0", a) + } + + expectRecv(t, f.ups, "second OnUp") +} + +func TestRestartAfterStop(t *testing.T) { + f := newFakeConnector(fastBackoff, false) + + f.fakeConnRes <- connectResult{} + f.r.Start() + firstID := expectRecv(t, f.ups, "first OnUp") + + stopWithDeadline(t, f.r) + if closedID := expectRecv(t, f.closes, "close on stop"); closedID != firstID { + t.Errorf("closed connection %d, expected %d", closedID, firstID) + } + if f.r.Connected() { + t.Error("Connected() = true after Stop") + } + + f.fakeConnRes <- connectResult{} + f.r.Start() + expectRecv(t, f.ups, "OnUp after restart") + + stopWithDeadline(t, f.r) +}