ghosteebus: propagate context to REST API calls during init (#28951)
This commit is contained in:
parent
462dbd0def
commit
71185c5a7f
3 changed files with 75 additions and 9 deletions
|
|
@ -73,7 +73,7 @@ func NewGhostEEBus(ctx context.Context, ski, ip, user, password string, hasMeter
|
|||
var phasesG func() (int, error)
|
||||
|
||||
if ip != "" && user != "" && password != "" {
|
||||
ts, err := ghostone.TokenSource(log, wb.uri, user, password)
|
||||
ts, err := ghostone.TokenSource(ctx, log, wb.uri, user, password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -85,13 +85,13 @@ func NewGhostEEBus(ctx context.Context, ski, ip, user, password string, hasMeter
|
|||
|
||||
// warn if PV optimization is active
|
||||
var pvMode ghostone.PvOptimizationMode
|
||||
if err := wb.GetJSON(wb.uri+"/charging/pvoptimization/mode", &pvMode); err == nil && pvMode.Value != ghostone.PvModeNone {
|
||||
if err := wb.getJSONCtx(ctx, wb.uri+"/charging/pvoptimization/mode", &pvMode); err == nil && pvMode.Value != ghostone.PvModeNone {
|
||||
log.WARN.Printf("wallbox PV optimization is active (%s), should be disabled when using evcc", pvMode.Value)
|
||||
}
|
||||
|
||||
// warn if phase switching is disabled
|
||||
var relaisEnabled ghostone.Enabled
|
||||
if err := wb.GetJSON(wb.uri+"/system/relais-switch/enabled", &relaisEnabled); err == nil && !relaisEnabled.Enabled {
|
||||
if err := wb.getJSONCtx(ctx, wb.uri+"/system/relais-switch/enabled", &relaisEnabled); err == nil && !relaisEnabled.Enabled {
|
||||
log.WARN.Println("phase switching is disabled, enable it in the wallbox settings to use 1p3p switching")
|
||||
}
|
||||
|
||||
|
|
@ -130,6 +130,15 @@ func (wb *GhostEEBus) Identify() (string, error) {
|
|||
return wb.EEBus.Identify()
|
||||
}
|
||||
|
||||
// getJSONCtx executes a context-aware GET request and decodes the JSON response.
|
||||
func (wb *GhostEEBus) getJSONCtx(ctx context.Context, url string, res any) error {
|
||||
req, err := request.New(http.MethodGet, url, nil, request.AcceptJSON)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return wb.DoJSON(req.WithContext(ctx), res)
|
||||
}
|
||||
|
||||
// putJSON sends a PUT request with JSON body to the REST API.
|
||||
func (wb *GhostEEBus) putJSON(url string, data any) error {
|
||||
req, err := request.New(http.MethodPut, url, request.MarshalJSON(data), request.JSONEncoding)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
package ghostone
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/url"
|
||||
|
|
@ -24,7 +25,7 @@ type tokenSource struct {
|
|||
}
|
||||
|
||||
// TokenSource creates a JWT token source for the ghost REST API
|
||||
func TokenSource(log *util.Logger, uri, user, password string) (oauth2.TokenSource, error) {
|
||||
func TokenSource(ctx context.Context, log *util.Logger, uri, user, password string) (oauth2.TokenSource, error) {
|
||||
c := &tokenSource{
|
||||
Helper: request.NewHelper(log),
|
||||
uri: uri + "/jwt/login",
|
||||
|
|
@ -34,7 +35,7 @@ func TokenSource(log *util.Logger, uri, user, password string) (oauth2.TokenSour
|
|||
|
||||
c.Client.Transport = transport.Insecure()
|
||||
|
||||
token, err := c.login()
|
||||
token, err := c.login(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -44,7 +45,7 @@ func TokenSource(log *util.Logger, uri, user, password string) (oauth2.TokenSour
|
|||
return c, nil
|
||||
}
|
||||
|
||||
func (c *tokenSource) login() (*oauth2.Token, error) {
|
||||
func (c *tokenSource) login(ctx context.Context) (*oauth2.Token, error) {
|
||||
data := url.Values{
|
||||
"user": {c.user},
|
||||
"pass": {c.password},
|
||||
|
|
@ -55,7 +56,7 @@ func (c *tokenSource) login() (*oauth2.Token, error) {
|
|||
return nil, err
|
||||
}
|
||||
|
||||
resp, err := c.Do(req)
|
||||
resp, err := c.Do(req.WithContext(ctx))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -88,6 +89,6 @@ func (c *tokenSource) login() (*oauth2.Token, error) {
|
|||
}
|
||||
|
||||
func (c *tokenSource) refresh(_ *oauth2.Token) (*oauth2.Token, error) {
|
||||
// re-login to get new token
|
||||
return c.login()
|
||||
// re-login to get new token (no context needed- refresh runs in steady-state, not during init)
|
||||
return c.login(context.Background())
|
||||
}
|
||||
|
|
|
|||
56
charger/ghostone/identity_test.go
Normal file
56
charger/ghostone/identity_test.go
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
package ghostone
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/evcc-io/evcc/util"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestTokenSource_ContextCancellation(t *testing.T) {
|
||||
// server that blocks -- simulates slow/unreachable wallbox
|
||||
unblock := make(chan struct{})
|
||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
<-unblock
|
||||
}))
|
||||
defer func() {
|
||||
close(unblock)
|
||||
srv.Close()
|
||||
}()
|
||||
|
||||
log := util.NewLogger("test")
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
start := time.Now()
|
||||
_, err := TokenSource(ctx, log, srv.URL, "user", "pass")
|
||||
elapsed := time.Since(start)
|
||||
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
assert.Less(t, elapsed, 2*time.Second, "TokenSource should return promptly when context is cancelled")
|
||||
}
|
||||
|
||||
func TestTokenSource_Success(t *testing.T) {
|
||||
// server that returns a valid JWT-style token
|
||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Authorization", fmt.Sprintf("Bearer %s",
|
||||
"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwiZXhwIjo5OTk5OTk5OTk5fQ.signature"))
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
log := util.NewLogger("test")
|
||||
|
||||
ts, err := TokenSource(context.Background(), log, srv.URL, "user", "pass")
|
||||
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, ts)
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue