From c6d9e1bed8b5656dd13ac7ba792e6cecfbd7279d Mon Sep 17 00:00:00 2001 From: andig Date: Fri, 1 Nov 2024 18:20:14 +0100 Subject: [PATCH] chore: deduplicate getter implementations (#17031) --- provider/getter.go | 68 ++++++++++++++++++++++++++++++++++++++++ provider/http.go | 62 ++++++------------------------------ provider/http_test.go | 6 ++-- provider/mqtt.go | 31 +++--------------- provider/mqtt_handler.go | 34 +------------------- provider/script.go | 55 ++------------------------------ provider/socket.go | 49 ++--------------------------- 7 files changed, 90 insertions(+), 215 deletions(-) create mode 100644 provider/getter.go diff --git a/provider/getter.go b/provider/getter.go new file mode 100644 index 000000000..5ec6f38ee --- /dev/null +++ b/provider/getter.go @@ -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 +} diff --git a/provider/http.go b/provider/http.go index 9b3ddde81..3f601386f 100644 --- a/provider/http.go +++ b/provider/http.go @@ -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 { diff --git a/provider/http_test.go b/provider/http_test.go index 7316c6d7a..417fe053a 100644 --- a/provider/http_test.go +++ b/provider/http_test.go @@ -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) diff --git a/provider/mqtt.go b/provider/mqtt.go index 6f253da36..0f3ac5556 100644 --- a/provider/mqtt.go +++ b/provider/mqtt.go @@ -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) diff --git a/provider/mqtt_handler.go b/provider/mqtt_handler.go index acc11d3e5..ebeea2440 100644 --- a/provider/mqtt_handler.go +++ b/provider/mqtt_handler.go @@ -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 -} diff --git a/provider/script.go b/provider/script.go index 85987afaa..5b3a79174 100644 --- a/provider/script.go +++ b/provider/script.go @@ -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 diff --git a/provider/socket.go b/provider/socket.go index d9ccb1902..67bbaacf9 100644 --- a/provider/socket.go +++ b/provider/socket.go @@ -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 -}