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"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
@ -20,7 +22,48 @@ import (
|
||||||
type MqttConnector struct {
|
type MqttConnector struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
log *util.Logger
|
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 (
|
var (
|
||||||
|
|
@ -38,7 +81,7 @@ func NewMqttConnector(ctx context.Context, log *util.Logger, clientID string, ts
|
||||||
|
|
||||||
v := &MqttConnector{
|
v := &MqttConnector{
|
||||||
log: log,
|
log: log,
|
||||||
subscriptions: make(map[string]chan StreamingMessage),
|
subscriptions: make(map[string][]subscription),
|
||||||
}
|
}
|
||||||
|
|
||||||
if !testing.Testing() {
|
if !testing.Testing() {
|
||||||
|
|
@ -54,19 +97,37 @@ func (v *MqttConnector) Subscribe(vin string) <-chan StreamingMessage {
|
||||||
v.mu.Lock()
|
v.mu.Lock()
|
||||||
defer v.mu.Unlock()
|
defer v.mu.Unlock()
|
||||||
|
|
||||||
ch := make(chan StreamingMessage, 1)
|
vin = strings.ToUpper(vin)
|
||||||
v.subscriptions[vin] = ch
|
|
||||||
|
|
||||||
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()
|
v.mu.Lock()
|
||||||
defer v.mu.Unlock()
|
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)
|
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()
|
v.mu.RLock()
|
||||||
defer v.mu.RUnlock()
|
defer v.mu.RUnlock()
|
||||||
|
|
||||||
if ch, ok := v.subscriptions[res.Vin]; ok {
|
for _, sub := range v.subscriptions[strings.ToUpper(res.Vin)] {
|
||||||
ch <- res
|
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() {
|
go func() {
|
||||||
<-ctx.Done()
|
<-ctx.Done()
|
||||||
mqtt.Unsubscribe(vin)
|
mqtt.Unsubscribe(vin, recvC)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
|
|
|
||||||
|
|
@ -30,8 +30,8 @@ func TestCardataStreaming(t *testing.T) {
|
||||||
require.Equal(t, 42.0, soc)
|
require.Equal(t, 42.0, soc)
|
||||||
|
|
||||||
mqtt := mqttConnections["client"]
|
mqtt := mqttConnections["client"]
|
||||||
dataC := mqtt.subscriptions["vin"]
|
require.Len(t, mqtt.subscriptions["VIN"], 1, "streaming channel")
|
||||||
require.NotNil(t, dataC, "streaming channel")
|
dataC := mqtt.subscriptions["VIN"][0].in
|
||||||
|
|
||||||
dataC <- StreamingMessage{
|
dataC <- StreamingMessage{
|
||||||
Vin: "vin",
|
Vin: "vin",
|
||||||
|
|
@ -40,13 +40,10 @@ func TestCardataStreaming(t *testing.T) {
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
// process first message
|
require.Eventually(t, func() bool {
|
||||||
dataC <- StreamingMessage{}
|
soc, err := p.Soc()
|
||||||
dataC <- StreamingMessage{}
|
return err == nil && soc == 47.0
|
||||||
|
}, time.Second, 10*time.Millisecond, "streaming update not applied")
|
||||||
soc, err = p.Soc()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, 47.0, soc)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSocFallback(t *testing.T) {
|
func TestSocFallback(t *testing.T) {
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"slices"
|
"slices"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/evcc-io/evcc/api"
|
"github.com/evcc-io/evcc/api"
|
||||||
|
|
@ -36,6 +37,7 @@ func NewCardataFromConfig(ctx context.Context, other map[string]any) (api.Vehicl
|
||||||
if cc.VIN == "" {
|
if cc.VIN == "" {
|
||||||
return nil, errors.New("missing vin")
|
return nil, errors.New("missing vin")
|
||||||
}
|
}
|
||||||
|
cc.VIN = strings.ToUpper(cc.VIN)
|
||||||
|
|
||||||
if cc.ClientID == "" {
|
if cc.ClientID == "" {
|
||||||
return nil, api.ErrMissingCredentials
|
return nil, api.ErrMissingCredentials
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue