MQTT: make Monitor first-read cancellable via context (#31116)

This commit is contained in:
andig 2026-06-22 11:32:56 +02:00 • committed by GitHub
parent aa7b5103dd
commit 2294066b27
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 44 additions and 3 deletions

View file

@ -12,6 +12,7 @@ import (
// Mqtt provider // Mqtt provider
type Mqtt struct { type Mqtt struct {
*getter *getter
ctx context.Context
log *util.Logger log *util.Logger
client *mqtt.Client client *mqtt.Client
topic string topic string
@ -49,7 +50,7 @@ func NewMqttPluginFromConfig(ctx context.Context, other map[string]any) (Plugin,
return nil, err return nil, err
} }
m := NewMqtt(log, client, cc.Topic, cc.Timeout).WithScale(cc.Scale).WithPayload(cc.Payload) m := NewMqtt(log, client, cc.Topic, cc.Timeout).WithContext(ctx).WithScale(cc.Scale).WithPayload(cc.Payload)
if cc.Retained { if cc.Retained {
m = m.WithRetained() m = m.WithRetained()
} }
@ -65,6 +66,7 @@ func NewMqttPluginFromConfig(ctx context.Context, other map[string]any) (Plugin,
// NewMqtt creates mqtt provider for given topic // NewMqtt creates mqtt provider for given topic
func NewMqtt(log *util.Logger, client *mqtt.Client, topic string, timeout time.Duration) *Mqtt { func NewMqtt(log *util.Logger, client *mqtt.Client, topic string, timeout time.Duration) *Mqtt {
m := &Mqtt{ m := &Mqtt{
ctx: context.Background(),
log: log, log: log,
client: client, client: client,
topic: topic, topic: topic,
@ -76,6 +78,12 @@ func NewMqtt(log *util.Logger, client *mqtt.Client, topic string, timeout time.D
return m return m
} }
// WithContext sets a context that can cancel the blocking first-read wait
func (m *Mqtt) WithContext(ctx context.Context) *Mqtt {
m.ctx = ctx
return m
}
// WithPayload adds payload for setters // WithPayload adds payload for setters
func (m *Mqtt) WithPayload(payload string) *Mqtt { func (m *Mqtt) WithPayload(payload string) *Mqtt {
m.payload = payload m.payload = payload
@ -103,6 +111,7 @@ func (p *Mqtt) WithPipeline(pipeline *pipeline.Pipeline) *Mqtt {
// newReceiver creates a msgHandler and subscribes it to the topic. // newReceiver creates a msgHandler and subscribes it to the topic.
func (m *Mqtt) newReceiver() (*msgHandler, error) { func (m *Mqtt) newReceiver() (*msgHandler, error) {
h := &msgHandler{ h := &msgHandler{
ctx: m.ctx,
topic: m.topic, topic: m.topic,
pipeline: m.pipeline, pipeline: m.pipeline,
val: util.NewMonitor[string](m.timeout), val: util.NewMonitor[string](m.timeout),

View file

@ -1,11 +1,14 @@
package plugin package plugin
import ( import (
"context"
"github.com/evcc-io/evcc/plugin/pipeline" "github.com/evcc-io/evcc/plugin/pipeline"
"github.com/evcc-io/evcc/util" "github.com/evcc-io/evcc/util"
) )
type msgHandler struct { type msgHandler struct {
ctx context.Context
topic string topic string
pipeline *pipeline.Pipeline pipeline *pipeline.Pipeline
val *util.Monitor[string] val *util.Monitor[string]
@ -17,7 +20,7 @@ func (h *msgHandler) receive(payload string) {
// hasValue returned the received and processed payload as string // hasValue returned the received and processed payload as string
func (h *msgHandler) hasValue() (string, error) { func (h *msgHandler) hasValue() (string, error) {
payload, err := h.val.Get() payload, err := h.val.GetContext(h.ctx)
if err != nil { if err != nil {
return "", err return "", err
} }

View file

@ -1,6 +1,7 @@
package util package util
import ( import (
"context"
"sync" "sync"
"time" "time"
@ -54,8 +55,14 @@ func (m *Monitor[T]) SetFunc(set func(T) T) {
// Get returns the current value or ErrOutdated if timeout exceeded // Get returns the current value or ErrOutdated if timeout exceeded
func (m *Monitor[T]) Get() (T, error) { func (m *Monitor[T]) Get() (T, error) {
return m.GetContext(context.Background())
}
// GetContext returns the current value or ErrOutdated if timeout exceeded.
// The context can cancel the blocking first-call wait.
func (m *Monitor[T]) GetContext(ctx context.Context) (T, error) {
var res T var res T
err := m.GetFunc(func(v T) { err := m.GetFuncContext(ctx, func(v T) {
res = v res = v
}) })
return res, err return res, err
@ -63,6 +70,12 @@ func (m *Monitor[T]) Get() (T, error) {
// GetFunc returns the current value or ErrOutdated if timeout exceeded while holding the lock // GetFunc returns the current value or ErrOutdated if timeout exceeded while holding the lock
func (m *Monitor[T]) GetFunc(get func(T)) error { func (m *Monitor[T]) GetFunc(get func(T)) error {
return m.GetFuncContext(context.Background(), get)
}
// GetFuncContext returns the current value or ErrOutdated if timeout exceeded while holding the lock.
// The context can cancel the blocking first-call wait.
func (m *Monitor[T]) GetFuncContext(ctx context.Context, get func(T)) error {
m.mu.RLock() m.mu.RLock()
defer m.mu.RUnlock() defer m.mu.RUnlock()
@ -94,6 +107,8 @@ func (m *Monitor[T]) GetFunc(get func(T)) error {
case <-m.done: case <-m.done:
// got value and updated timestamp // got value and updated timestamp
err = nil err = nil
case <-ctx.Done():
err = ctx.Err()
case <-m.clock.After(m.timeout): case <-m.clock.After(m.timeout):
} }

View file

@ -1,6 +1,7 @@
package util package util
import ( import (
"context"
"math/rand" "math/rand"
"testing" "testing"
"time" "time"
@ -69,3 +70,16 @@ func TestMonitorWithoutTimeout(t *testing.T) {
_, err = m.Get() _, err = m.Get()
assert.NoError(t, err) assert.NoError(t, err)
} }
func TestMonitorGetContext(t *testing.T) {
// a long timeout would block the first-read wait; a cancelled context cuts it short
m := NewMonitor[int](time.Minute)
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
defer cancel()
start := time.Now()
_, err := m.GetContext(ctx)
assert.ErrorIs(t, err, context.DeadlineExceeded)
assert.Less(t, time.Since(start), time.Second, "context must cancel the first-read wait")
}