MQTT: make Monitor first-read cancellable via context (#31116)
This commit is contained in:
parent
aa7b5103dd
commit
2294066b27
4 changed files with 44 additions and 3 deletions
|
|
@ -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),
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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):
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue