Cardata: fix mqtt channel stealing (#32244)

This commit is contained in:
mfuchs1984 2026-07-28 21:38:01 +02:00 • committed by GitHub
parent 106d44652c
commit 614c91080b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 166 additions and 20 deletions

View file

@ -5,6 +5,8 @@ import (
"encoding/json"
"errors"
"fmt"
"slices"
"strings"
"sync"
"testing"
"time"
@ -20,7 +22,48 @@ import (
type MqttConnector struct {
mu sync.RWMutex
log *util.Logger
subscriptions map[string]chan StreamingMessage
subscriptions map[string][]subscription
}
// subscription connects the handler to a single receiver through an unbounded queue
type subscription struct {
in chan StreamingMessage
out <-chan StreamingMessage
}
// queued forwards in to the returned channel through an unbounded fifo queue, so that
// a busy receiver cannot block the mqtt handler. Closing in drains the queue, then closes out.
func queued(in <-chan StreamingMessage) <-chan StreamingMessage {
out := make(chan StreamingMessage)
go func() {
defer close(out)
var queue []StreamingMessage
for in != nil || len(queue) > 0 {
var send chan<- StreamingMessage
var next StreamingMessage
if len(queue) > 0 {
send, next = out, queue[0]
}
select {
case msg, ok := <-in:
if !ok {
in = nil
continue
}
queue = append(queue, msg)
case send <- next:
queue = queue[1:]
}
}
}()
return out
}
var (
@ -38,7 +81,7 @@ func NewMqttConnector(ctx context.Context, log *util.Logger, clientID string, ts
v := &MqttConnector{
log: log,
subscriptions: make(map[string]chan StreamingMessage),
subscriptions: make(map[string][]subscription),
}
if !testing.Testing() {
@ -54,19 +97,37 @@ func (v *MqttConnector) Subscribe(vin string) <-chan StreamingMessage {
v.mu.Lock()
defer v.mu.Unlock()
ch := make(chan StreamingMessage, 1)
v.subscriptions[vin] = ch
vin = strings.ToUpper(vin)
return ch
sub := subscription{in: make(chan StreamingMessage)}
sub.out = queued(sub.in)
v.subscriptions[vin] = append(v.subscriptions[vin], sub)
v.log.DEBUG.Printf("mqtt subscribe: %s (%d active subscribers)", vin, len(v.subscriptions[vin]))
return sub.out
}
func (v *MqttConnector) Unsubscribe(vin string) {
func (v *MqttConnector) Unsubscribe(vin string, ch <-chan StreamingMessage) {
v.mu.Lock()
defer v.mu.Unlock()
if ch, ok := v.subscriptions[vin]; ok {
vin = strings.ToUpper(vin)
subs := v.subscriptions[vin]
i := slices.IndexFunc(subs, func(sub subscription) bool { return sub.out == ch })
if i < 0 {
return
}
// queued messages are delivered before out is closed
close(subs[i].in)
if subs = slices.Delete(subs, i, i+1); len(subs) == 0 {
delete(v.subscriptions, vin)
close(ch)
} else {
v.subscriptions[vin] = subs
}
}
@ -149,7 +210,7 @@ func (v *MqttConnector) handler(_ mqtt.Client, m mqtt.Message) {
v.mu.RLock()
defer v.mu.RUnlock()
if ch, ok := v.subscriptions[res.Vin]; ok {
ch <- res
for _, sub := range v.subscriptions[strings.ToUpper(res.Vin)] {
sub.in <- res
}
}

View file

@ -0,0 +1,86 @@
package cardata
import (
"context"
"encoding/json"
"testing"
"time"
mqtt "github.com/eclipse/paho.mqtt.golang"
"github.com/evcc-io/evcc/util"
"github.com/stretchr/testify/require"
)
type mockMessage struct {
mqtt.Message
payload []byte
}
func (m mockMessage) Topic() string { return "test" }
func (m mockMessage) Payload() []byte { return m.payload }
func message(t *testing.T, vin string) mockMessage {
t.Helper()
b, err := json.Marshal(StreamingMessage{Vin: vin})
require.NoError(t, err)
return mockMessage{payload: b}
}
func recv(t *testing.T, ch <-chan StreamingMessage) StreamingMessage {
t.Helper()
select {
case msg := <-ch:
return msg
case <-time.After(time.Second):
t.Fatal("message not delivered")
return StreamingMessage{}
}
}
func TestMqttMultiSubscribe(t *testing.T) {
conn := NewMqttConnector(context.TODO(), util.NewLogger("foo"), t.Name(), nil)
// loadpoint and config page subscribing to the same vehicle, vin entered in mixed case
ch1 := conn.Subscribe("wba12345")
ch2 := conn.Subscribe("WBA12345")
conn.handler(nil, message(t, "WBA12345"))
for _, ch := range []<-chan StreamingMessage{ch1, ch2} {
require.Equal(t, "WBA12345", recv(t, ch).Vin)
}
// unsubscribing one subscriber must not steal the other's channel
conn.Unsubscribe("WBA12345", ch1)
conn.handler(nil, message(t, "WBA12345"))
_, ok := <-ch1
require.False(t, ok, "ch1 not closed")
require.Equal(t, "WBA12345", recv(t, ch2).Vin)
conn.Unsubscribe("wba12345", ch2)
require.Empty(t, conn.subscriptions)
}
func TestMqttUnboundedQueue(t *testing.T) {
conn := NewMqttConnector(context.TODO(), util.NewLogger("foo"), t.Name(), nil)
ch := conn.Subscribe("WBA12345")
// handler must not block on a receiver that is not reading
for range 100 {
conn.handler(nil, message(t, "WBA12345"))
}
// queue is drained before the channel is closed
conn.Unsubscribe("WBA12345", ch)
var count int
for range ch {
count++
}
require.Equal(t, 100, count)
}

View file

@ -50,7 +50,7 @@ func NewProvider(ctx context.Context, log *util.Logger, api *API, ts oauth2.Toke
go func() {
<-ctx.Done()
mqtt.Unsubscribe(vin)
mqtt.Unsubscribe(vin, recvC)
}()
go func() {

View file

@ -30,8 +30,8 @@ func TestCardataStreaming(t *testing.T) {
require.Equal(t, 42.0, soc)
mqtt := mqttConnections["client"]
dataC := mqtt.subscriptions["vin"]
require.NotNil(t, dataC, "streaming channel")
require.Len(t, mqtt.subscriptions["VIN"], 1, "streaming channel")
dataC := mqtt.subscriptions["VIN"][0].in
dataC <- StreamingMessage{
Vin: "vin",
@ -40,13 +40,10 @@ func TestCardataStreaming(t *testing.T) {
},
}
// process first message
dataC <- StreamingMessage{}
dataC <- StreamingMessage{}
soc, err = p.Soc()
require.NoError(t, err)
require.Equal(t, 47.0, soc)
require.Eventually(t, func() bool {
soc, err := p.Soc()
return err == nil && soc == 47.0
}, time.Second, 10*time.Millisecond, "streaming update not applied")
}
func TestSocFallback(t *testing.T) {

View file

@ -4,6 +4,7 @@ import (
"context"
"errors"
"slices"
"strings"
"time"
"github.com/evcc-io/evcc/api"
@ -36,6 +37,7 @@ func NewCardataFromConfig(ctx context.Context, other map[string]any) (api.Vehicl
if cc.VIN == "" {
return nil, errors.New("missing vin")
}
cc.VIN = strings.ToUpper(cc.VIN)
if cc.ClientID == "" {
return nil, api.ErrMissingCredentials