diff --git a/provider/http.go b/provider/http.go index decba6a86..31dd43bf9 100644 --- a/provider/http.go +++ b/provider/http.go @@ -148,7 +148,7 @@ func (p *HTTP) WithAuth(typ, user, password string) (*HTTP, error) { } // request executes the configured request or returns the cached value -func (p *HTTP) request(body ...string) ([]byte, error) { +func (p *HTTP) request(url string, body ...string) ([]byte, error) { if time.Since(p.updated) >= p.cache { var b io.Reader if len(body) == 1 { @@ -156,7 +156,7 @@ func (p *HTTP) request(body ...string) ([]byte, error) { } // empty method becomes GET - req, err := request.New(strings.ToUpper(p.method), p.url, b, p.headers) + req, err := request.New(strings.ToUpper(p.method), url, b, p.headers) if err != nil { return []byte{}, err } @@ -200,7 +200,7 @@ func (p *HTTP) IntGetter() func() (int64, error) { // StringGetter sends string request func (p *HTTP) StringGetter() func() (string, error) { return func() (string, error) { - b, err := p.request(p.body) + b, err := p.request(p.url, p.body) if err == nil && p.pipeline != nil { b, err = p.pipeline.Process(b) @@ -221,12 +221,18 @@ func (p *HTTP) BoolGetter() func() (bool, error) { } func (p *HTTP) set(param string, val interface{}) error { - body, err := setFormattedValue(p.body, param, val) - - if err == nil { - _, err = p.request(body) + url, err := setFormattedValue(p.url, param, val) + if err != nil { + return err } + body, err := setFormattedValue(p.body, param, val) + if err != nil { + return err + } + + _, err = p.request(url, body) + return err } diff --git a/provider/http_test.go b/provider/http_test.go new file mode 100644 index 000000000..d8e0b6615 --- /dev/null +++ b/provider/http_test.go @@ -0,0 +1,55 @@ +package provider + +import ( + "net/http" + "net/http/httptest" + "net/url" + "testing" + + "github.com/evcc-io/evcc/util" + "github.com/samber/lo" + "github.com/stretchr/testify/assert" +) + +type httpHandler struct { + val string + req *http.Request +} + +func (h *httpHandler) ServeHTTP(w http.ResponseWriter, req *http.Request) { + h.req = req + h.val = lo.RandomString(16, lo.LettersCharset) + _, _ = w.Write([]byte(h.val)) +} + +func TestHttpGet(t *testing.T) { + h := new(httpHandler) + srv := httptest.NewServer(h) + defer srv.Close() + + uri := srv.URL + "/foo/bar" + p := NewHTTP(util.NewLogger("foo"), http.MethodGet, uri, false, 1, 0) + + uriUrl, _ := url.Parse(uri) + + res, err := p.StringGetter()() + assert.NoError(t, err) + assert.Equal(t, uriUrl.Path, h.req.URL.Path) + assert.Equal(t, h.val, res) +} + +func TestHttpSet(t *testing.T) { + h := new(httpHandler) + srv := httptest.NewServer(h) + defer srv.Close() + + uri := srv.URL + "/foo/bar?baz={{.baz}}" + p := NewHTTP(util.NewLogger("foo"), http.MethodGet, uri, false, 1, 0) + + uriUrl, _ := url.Parse(uri) + + err := p.StringSetter("baz")("4711") + assert.NoError(t, err) + assert.Equal(t, uriUrl.Path, h.req.URL.Path) + assert.Equal(t, "baz=4711", h.req.URL.RawQuery) +}