diff --git a/go.mod b/go.mod index 92468fb46..9a95bf303 100644 --- a/go.mod +++ b/go.mod @@ -42,7 +42,6 @@ require ( github.com/google/uuid v1.3.0 github.com/gorilla/handlers v1.5.1 github.com/gorilla/mux v1.8.0 - github.com/gorilla/websocket v1.5.0 github.com/gregdel/pushover v1.1.0 github.com/grid-x/modbus v0.0.0-20230412120016-b867b6a2a633 github.com/hashicorp/go-version v1.6.0 @@ -139,6 +138,7 @@ require ( github.com/google/s2a-go v0.1.1 // indirect github.com/googleapis/enterprise-certificate-proxy v0.2.3 // indirect github.com/googleapis/gax-go/v2 v2.8.0 // indirect + github.com/gorilla/websocket v1.5.0 // indirect github.com/grid-x/serial v0.0.0-20211107191517-583c7356b3aa // indirect github.com/hashicorp/hcl v1.0.0 // indirect github.com/holoplot/go-avahi v1.0.1 // indirect diff --git a/provider/socket.go b/provider/socket.go index 6e57bcc54..ce294978d 100644 --- a/provider/socket.go +++ b/provider/socket.go @@ -1,6 +1,7 @@ package provider import ( + "context" "fmt" "math" "net/http" @@ -12,7 +13,7 @@ import ( "github.com/evcc-io/evcc/util" "github.com/evcc-io/evcc/util/request" "github.com/evcc-io/evcc/util/transport" - "github.com/gorilla/websocket" + "nhooyr.io/websocket" ) const retryDelay = 5 * time.Second @@ -99,13 +100,15 @@ func (p *Socket) listen() { headers.Set(k, v) } - dialer := &websocket.Dialer{ - Proxy: http.ProxyFromEnvironment, - HandshakeTimeout: request.Timeout, + opts := &websocket.DialOptions{ + HTTPHeader: headers, } for { - client, _, err := dialer.Dial(p.url, headers) + ctx, cancel := context.WithTimeout(context.Background(), request.Timeout) + conn, _, err := websocket.Dial(ctx, p.url, opts) + cancel() + if err != nil { p.log.ERROR.Println(err) time.Sleep(retryDelay) @@ -113,10 +116,10 @@ func (p *Socket) listen() { } for { - _, b, err := client.ReadMessage() + _, b, err := conn.Read(context.Background()) if err != nil { p.log.TRACE.Println("read:", err) - _ = client.Close() + _ = conn.Close(websocket.StatusAbnormalClosure, "done") break } diff --git a/server/socket.go b/server/socket.go index 0c9577575..4a5e48040 100644 --- a/server/socket.go +++ b/server/socket.go @@ -1,6 +1,7 @@ package server import ( + "context" "encoding/json" "fmt" "math" @@ -10,8 +11,8 @@ import ( "time" "github.com/evcc-io/evcc/util" - "github.com/gorilla/websocket" "github.com/kr/pretty" + "nhooyr.io/websocket" ) const ( @@ -19,12 +20,6 @@ const ( socketWriteTimeout = 10 * time.Second ) -var upgrader = websocket.Upgrader{ - ReadBufferSize: 1024, - WriteBufferSize: 1024, - CheckOrigin: func(r *http.Request) bool { return true }, -} - // SocketClient is a middleman between the websocket connection and the hub. type SocketClient struct { hub *SocketHub @@ -39,15 +34,16 @@ type SocketClient struct { // writePump pumps messages from the hub to the websocket connection. func (c *SocketClient) writePump() { defer func() { - c.conn.Close() + c.conn.Close(websocket.StatusNormalClosure, "done") c.hub.unregister <- c }() for msg := range c.send { - if err := c.conn.SetWriteDeadline(time.Now().Add(socketWriteTimeout)); err != nil { - return - } - if err := c.conn.WriteMessage(websocket.TextMessage, msg); err != nil { + ctx, cancel := context.WithTimeout(context.Background(), socketWriteTimeout) + err := c.conn.Write(ctx, websocket.MessageText, msg) + cancel() + + if err != nil { return } } @@ -55,7 +51,7 @@ func (c *SocketClient) writePump() { // ServeWebsocket handles websocket requests from the peer. func ServeWebsocket(hub *SocketHub, w http.ResponseWriter, r *http.Request) { - conn, err := upgrader.Upgrade(w, r, nil) + conn, err := websocket.Accept(w, r, nil) if err != nil { log.ERROR.Println(err) return