chore: deduplicate getter implementations (#17031)

This commit is contained in:
andig 2024-11-01 18:20:14 +01:00 • committed by GitHub
parent 4c70d5e2c7
commit c6d9e1bed8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 90 additions and 215 deletions

68
provider/getter.go Normal file
View file

@ -0,0 +1,68 @@
package provider
import (
"strconv"
"github.com/spf13/cast"
)
type getter struct {
sp StringProvider
scale float64
}
func defaultGetters(sp StringProvider, scale float64) *getter {
return &getter{
sp: sp,
scale: scale,
}
}
var _ FloatProvider = (*getter)(nil)
// FloatGetter parses float from exec result
func (p *getter) FloatGetter() (func() (float64, error), error) {
g, err := p.sp.StringGetter()
return func() (float64, error) {
s, err := g()
if err != nil {
return 0, err
}
f, err := strconv.ParseFloat(s, 64)
if err != nil {
return 0, err
}
return f * p.scale, nil
}, err
}
var _ IntProvider = (*getter)(nil)
// IntGetter parses int64 from exec result
func (p *getter) IntGetter() (func() (int64, error), error) {
g, err := p.FloatGetter()
return func() (int64, error) {
f, err := g()
return int64(f), err
}, err
}
var _ BoolProvider = (*getter)(nil)
// BoolGetter parses bool from exec result. "on", "true" and 1 are considered truish.
func (p *getter) BoolGetter() (func() (bool, error), error) {
g, err := p.sp.StringGetter()
return func() (bool, error) {
s, err := g()
if err != nil {
return false, err
}
return cast.ToBoolE(s)
}, err
}

View file

@ -4,9 +4,7 @@ import (
"context"
"fmt"
"io"
"math"
"net/http"
"strconv"
"strings"
"time"
@ -20,11 +18,11 @@ import (
// HTTP implements HTTP request provider
type HTTP struct {
*getter
*request.Helper
url, method string
headers map[string]string
body string
scale float64
cache time.Duration
updated time.Time
pipeline *pipeline.Pipeline
@ -65,40 +63,40 @@ func NewHTTPProviderFromConfig(ctx context.Context, other map[string]interface{}
}
log := contextLogger(ctx, util.NewLogger("http"))
http := NewHTTP(
p := NewHTTP(
log,
strings.ToUpper(cc.Method),
cc.URI,
cc.Insecure,
cc.Scale,
cc.Cache,
).
WithHeaders(cc.Headers).
WithBody(cc.Body)
http.Client.Timeout = cc.Timeout
p.Client.Timeout = cc.Timeout
p.getter = defaultGetters(p, cc.Scale)
var err error
if cc.Auth.Type != "" {
_, err = http.WithAuth(cc.Auth.Type, cc.Auth.User, cc.Auth.Password)
_, err = p.WithAuth(cc.Auth.Type, cc.Auth.User, cc.Auth.Password)
}
if err == nil {
var pipe *pipeline.Pipeline
pipe, err = pipeline.New(log, cc.Settings)
http = http.WithPipeline(pipe)
p = p.WithPipeline(pipe)
}
return http, err
return p, err
}
// NewHTTP create HTTP provider
func NewHTTP(log *util.Logger, method, uri string, insecure bool, scale float64, cache time.Duration) *HTTP {
func NewHTTP(log *util.Logger, method, uri string, insecure bool, cache time.Duration) *HTTP {
p := &HTTP{
Helper: request.NewHelper(log),
url: uri,
method: method,
scale: scale,
cache: cache,
}
@ -200,48 +198,6 @@ func (p *HTTP) StringGetter() (func() (string, error), error) {
}, nil
}
var _ FloatProvider = (*HTTP)(nil)
// FloatGetter parses float from request
func (p *HTTP) FloatGetter() (func() (float64, error), error) {
g, err := p.StringGetter()
return func() (float64, error) {
s, err := g()
if err != nil {
return 0, err
}
f, err := strconv.ParseFloat(s, 64)
return f * p.scale, err
}, err
}
var _ IntProvider = (*HTTP)(nil)
// IntGetter parses int64 from request
func (p *HTTP) IntGetter() (func() (int64, error), error) {
g, err := p.FloatGetter()
return func() (int64, error) {
f, err := g()
return int64(math.Round(f)), err
}, err
}
var _ BoolProvider = (*HTTP)(nil)
// BoolGetter parses bool from request
func (p *HTTP) BoolGetter() (func() (bool, error), error) {
g, err := p.StringGetter()
return func() (bool, error) {
s, err := g()
return util.Truish(s), err
}, err
}
func (p *HTTP) set(param string, val interface{}) error {
url, err := setFormattedValue(p.url, param, val)
if err != nil {

View file

@ -42,7 +42,7 @@ func (suite *httpTestSuite) TearDown() {
func (suite *httpTestSuite) TestGet() {
uri := suite.srv.URL + "/foo/bar{{\"/baz\"}}"
p := NewHTTP(util.NewLogger("foo"), http.MethodGet, uri, false, 1, 0)
p := NewHTTP(util.NewLogger("foo"), http.MethodGet, uri, false, 0)
g, err := p.StringGetter()
suite.Require().NoError(err)
@ -55,7 +55,7 @@ func (suite *httpTestSuite) TestGet() {
func (suite *httpTestSuite) TestSetQuery() {
uri := suite.srv.URL + "/foo/bar?baz={{.baz}}"
p := NewHTTP(util.NewLogger("foo"), http.MethodGet, uri, false, 1, 0)
p := NewHTTP(util.NewLogger("foo"), http.MethodGet, uri, false, 0)
s, err := p.StringSetter("baz")
suite.Require().NoError(err)
@ -65,7 +65,7 @@ func (suite *httpTestSuite) TestSetQuery() {
func (suite *httpTestSuite) TestSetPath() {
uri := suite.srv.URL + "/foo/bar/{{.baz}}"
p := NewHTTP(util.NewLogger("foo"), http.MethodGet, uri, false, 1, 0)
p := NewHTTP(util.NewLogger("foo"), http.MethodGet, uri, false, 0)
s, err := p.StringSetter("baz")
suite.Require().NoError(err)

View file

@ -11,12 +11,12 @@ import (
// Mqtt provider
type Mqtt struct {
*getter
log *util.Logger
client *mqtt.Client
topic string
retained bool
payload string
scale float64
timeout time.Duration
pipeline *pipeline.Pipeline
}
@ -54,6 +54,8 @@ func NewMqttFromConfig(ctx context.Context, other map[string]interface{}) (Provi
m = m.WithRetained()
}
m.getter = defaultGetters(m, cc.Scale)
pipe, err := pipeline.New(log, cc.Settings)
if err == nil {
m = m.WithPipeline(pipe)
@ -68,7 +70,6 @@ func NewMqtt(log *util.Logger, client *mqtt.Client, topic string, timeout time.D
log: log,
client: client,
topic: topic,
scale: 1,
timeout: timeout,
}
@ -112,36 +113,12 @@ func (m *Mqtt) newReceiver() (*msgHandler, error) {
return h, err
}
var _ FloatProvider = (*Mqtt)(nil)
// FloatGetter creates handler for float64 from MQTT topic that returns cached value
func (m *Mqtt) FloatGetter() (func() (float64, error), error) {
h, err := m.newReceiver()
return h.floatGetter, err
}
var _ IntProvider = (*Mqtt)(nil)
// IntGetter creates handler for int64 from MQTT topic that returns cached value
func (m *Mqtt) IntGetter() (func() (int64, error), error) {
h, err := m.newReceiver()
return h.intGetter, err
}
var _ StringProvider = (*Mqtt)(nil)
// StringGetter creates handler for string from MQTT topic that returns cached value
func (m *Mqtt) StringGetter() (func() (string, error), error) {
h, err := m.newReceiver()
return h.stringGetter, err
}
var _ BoolProvider = (*Mqtt)(nil)
// BoolGetter creates handler for string from MQTT topic that returns cached value
func (m *Mqtt) BoolGetter() (func() (bool, error), error) {
h, err := m.newReceiver()
return h.boolGetter, err
return h.value, err
}
var _ SetIntProvider = (*Mqtt)(nil)

View file

@ -1,10 +1,6 @@
package provider
import (
"fmt"
"math"
"strconv"
"github.com/evcc-io/evcc/provider/pipeline"
"github.com/evcc-io/evcc/util"
)
@ -39,26 +35,7 @@ func (h *msgHandler) hasValue() (string, error) {
return payload, nil
}
func (h *msgHandler) floatGetter() (float64, error) {
v, err := h.hasValue()
if err != nil {
return 0, err
}
f, err := strconv.ParseFloat(v, 64)
if err != nil {
return 0, fmt.Errorf("%s invalid: '%s'", h.topic, v)
}
return f * h.scale, nil
}
func (h *msgHandler) intGetter() (int64, error) {
f, err := h.floatGetter()
return int64(math.Round(f)), err
}
func (h *msgHandler) stringGetter() (string, error) {
func (h *msgHandler) value() (string, error) {
v, err := h.hasValue()
if err != nil {
return "", err
@ -66,12 +43,3 @@ func (h *msgHandler) stringGetter() (string, error) {
return v, nil
}
func (h *msgHandler) boolGetter() (bool, error) {
v, err := h.hasValue()
if err != nil {
return false, err
}
return util.Truish(v), nil
}

View file

@ -3,9 +3,7 @@ package provider
import (
"context"
"errors"
"math"
"os/exec"
"strconv"
"strings"
"time"
@ -17,6 +15,7 @@ import (
// Script implements shell script-based providers and setters
type Script struct {
*getter
log *util.Logger
script string
timeout time.Duration
@ -25,7 +24,6 @@ type Script struct {
val string
err error
pipeline *pipeline.Pipeline
scale float64
}
func init() {
@ -50,6 +48,7 @@ func NewScriptProviderFromConfig(other map[string]interface{}) (Provider, error)
}
p, err := NewScriptProvider(cc.Cmd, cc.Timeout, cc.Scale, cc.Cache)
p.getter = defaultGetters(p, cc.Scale)
if err == nil {
var pipe *pipeline.Pipeline
@ -71,7 +70,6 @@ func NewScriptProvider(script string, timeout time.Duration, scale float64, cach
log: util.NewLogger("script"),
script: script,
timeout: timeout,
scale: scale,
cache: cache,
}
@ -128,55 +126,6 @@ func (p *Script) StringGetter() (func() (string, error), error) {
}, nil
}
var _ FloatProvider = (*Script)(nil)
// FloatGetter parses float from exec result
func (p *Script) FloatGetter() (func() (float64, error), error) {
g, err := p.StringGetter()
return func() (float64, error) {
s, err := g()
if err != nil {
return 0, err
}
f, err := strconv.ParseFloat(s, 64)
if err == nil {
f *= p.scale
}
return f, err
}, err
}
var _ IntProvider = (*Script)(nil)
// IntGetter parses int64 from exec result
func (p *Script) IntGetter() (func() (int64, error), error) {
g, err := p.FloatGetter()
return func() (int64, error) {
f, err := g()
return int64(math.Round(f)), err
}, err
}
var _ BoolProvider = (*Script)(nil)
// BoolGetter parses bool from exec result. "on", "true" and 1 are considered truish.
func (p *Script) BoolGetter() (func() (bool, error), error) {
g, err := p.StringGetter()
return func() (bool, error) {
s, err := g()
if err != nil {
return false, err
}
return util.Truish(s), nil
}, err
}
var _ SetIntProvider = (*Script)(nil)
// IntSetter invokes script with parameter replaced by int value

View file

@ -2,9 +2,7 @@ package provider
import (
"context"
"math"
"net/http"
"strconv"
"sync"
"time"
@ -20,11 +18,11 @@ const retryDelay = 5 * time.Second
// Socket implements websocket request provider
type Socket struct {
*getter
*request.Helper
log *util.Logger
url string
headers map[string]string
scale float64
pipeline *pipeline.Pipeline
val *util.Monitor[[]byte]
}
@ -65,10 +63,11 @@ func NewSocketProviderFromConfig(other map[string]interface{}) (Provider, error)
Helper: request.NewHelper(log),
url: url,
headers: cc.Headers,
scale: cc.Scale,
val: util.NewMonitor[[]byte](cc.Timeout),
}
p.getter = defaultGetters(p, cc.Scale)
// handle basic auth
if cc.Auth.Type != "" {
basicAuth := transport.BasicAuthHeader(cc.Auth.User, cc.Auth.Password)
@ -163,45 +162,3 @@ func (p *Socket) StringGetter() (func() (string, error), error) {
return string(val), nil
}, nil
}
var _ FloatProvider = (*Socket)(nil)
// FloatGetter parses float from string getter
func (p *Socket) FloatGetter() (func() (float64, error), error) {
g, err := p.StringGetter()
return func() (float64, error) {
s, err := g()
if err != nil {
return 0, err
}
f, err := strconv.ParseFloat(s, 64)
return f * p.scale, err
}, err
}
var _ IntProvider = (*Socket)(nil)
// IntGetter parses int64 from float getter
func (p *Socket) IntGetter() (func() (int64, error), error) {
g, err := p.FloatGetter()
return func() (int64, error) {
f, err := g()
return int64(math.Round(f)), err
}, err
}
var _ BoolProvider = (*Socket)(nil)
// BoolGetter parses bool from string getter
func (p *Socket) BoolGetter() (func() (bool, error), error) {
g, err := p.StringGetter()
return func() (bool, error) {
s, err := g()
return util.Truish(s), err
}, err
}