Cardata: fix mqtt channel stealing (#32244)
This commit is contained in:
parent
106d44652c
commit
614c91080b
5 changed files with 166 additions and 20 deletions
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
86
vehicle/bmw/cardata/mqtt_test.go
Normal file
86
vehicle/bmw/cardata/mqtt_test.go
Normal 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)
|
||||
}
|
||||
|
|
@ -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() {
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue