From e9ac325592c187395794f343ea4fa741b8dd5863 Mon Sep 17 00:00:00 2001 From: andig Date: Mon, 1 May 2023 19:55:30 +0200 Subject: [PATCH] chore: drop gorilla/websocket (#7768) --- go.mod | 2 +- provider/socket.go | 17 ++-- server/http_handler.go | 2 +- server/socket.go | 212 +++++++++++++++++----------------------- server/socket_helper.go | 63 ++++++++++++ 5 files changed, 163 insertions(+), 133 deletions(-) create mode 100644 server/socket_helper.go diff --git a/go.mod b/go.mod index b7ffc0a1e..29f60ce9c 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/http_handler.go b/server/http_handler.go index d0128d02f..f1eec3b00 100644 --- a/server/http_handler.go +++ b/server/http_handler.go @@ -402,6 +402,6 @@ func planHandler(lp loadpoint.API) http.HandlerFunc { // socketHandler attaches websocket handler to uri func socketHandler(hub *SocketHub) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - ServeWebsocket(hub, w, r) + hub.ServeWebsocket(w, r) } } diff --git a/server/socket.go b/server/socket.go index 0c9577575..7ecf3a4a2 100644 --- a/server/socket.go +++ b/server/socket.go @@ -1,17 +1,15 @@ package server import ( - "encoding/json" - "fmt" - "math" + "context" + "errors" "net/http" "strings" "sync" "time" "github.com/evcc-io/evcc/util" - "github.com/gorilla/websocket" - "github.com/kr/pretty" + "nhooyr.io/websocket" ) const ( @@ -19,135 +17,111 @@ const ( socketWriteTimeout = 10 * time.Second ) -var upgrader = websocket.Upgrader{ - ReadBufferSize: 1024, - WriteBufferSize: 1024, - CheckOrigin: func(r *http.Request) bool { return true }, +// socketSubscriber is a middleman between the websocket connection and the hub. +type socketSubscriber struct { + send chan []byte + closeSlow func() } -// SocketClient is a middleman between the websocket connection and the hub. -type SocketClient struct { - hub *SocketHub - - // The websocket connection. - conn *websocket.Conn - - // Buffered channel of outbound messages. - send chan []byte -} - -// writePump pumps messages from the hub to the websocket connection. -func (c *SocketClient) writePump() { - defer func() { - c.conn.Close() - 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 { - return - } - } -} - -// ServeWebsocket handles websocket requests from the peer. -func ServeWebsocket(hub *SocketHub, w http.ResponseWriter, r *http.Request) { - conn, err := upgrader.Upgrade(w, r, nil) - if err != nil { - log.ERROR.Println(err) - return - } - client := &SocketClient{hub: hub, conn: conn, send: make(chan []byte, 256)} - client.hub.register <- client - - // run writing to client in goroutine - go client.writePump() +func writeTimeout(ctx context.Context, timeout time.Duration, c *websocket.Conn, msg []byte) error { + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + return c.Write(ctx, websocket.MessageText, msg) } // SocketHub maintains the set of active clients and broadcasts messages to the // clients. type SocketHub struct { - mu sync.RWMutex - - // Registered clients. - clients map[*SocketClient]bool - - // Register requests from the clients. - register chan *SocketClient - - // Unregister requests from clients. - unregister chan *SocketClient + mu sync.RWMutex + register chan *socketSubscriber + subscribers map[*socketSubscriber]struct{} } // NewSocketHub creates a web socket hub that distributes meter status and // query results for the ui or other clients func NewSocketHub() *SocketHub { return &SocketHub{ - register: make(chan *SocketClient), - unregister: make(chan *SocketClient), - clients: make(map[*SocketClient]bool), + register: make(chan *socketSubscriber, 1), + subscribers: make(map[*socketSubscriber]struct{}), } } -func encode(v interface{}) (string, error) { - var s string - switch val := v.(type) { - case time.Time: - if val.IsZero() { - s = "null" - } else { - s = fmt.Sprintf(`"%s"`, val.Format(time.RFC3339)) - } - case time.Duration: - // must be before stringer to convert to seconds instead of string - s = fmt.Sprintf("%d", int64(val.Seconds())) - case float64: - if math.IsNaN(val) { - s = "null" - } else { - s = fmt.Sprintf("%.5g", val) - } - default: - if b, err := json.Marshal(v); err == nil { - s = string(b) - } else { - return "", err - } +// ServeWebsocket handles websocket requests from the peer. +func (h *SocketHub) ServeWebsocket(w http.ResponseWriter, r *http.Request) { + acceptOptions := &websocket.AcceptOptions{ + InsecureSkipVerify: true, } - return s, nil -} -func kv(p util.Param) string { - val, err := encode(p.Val) + // https://github.com/nhooyr/websocket/issues/218 + ua := strings.ToLower(r.Header.Get("User-Agent")) + if strings.Contains(ua, "safari") && !strings.Contains(ua, "chrome") && !strings.Contains(ua, "android") { + acceptOptions.CompressionMode = websocket.CompressionDisabled + } + + conn, err := websocket.Accept(w, r, acceptOptions) if err != nil { - panic(err) + log.ERROR.Println(err) + return } + defer conn.Close(websocket.StatusInternalError, "") - if p.Key == "" && val == "" { - log.ERROR.Printf("invalid key/val for %+v %# v, please report to https://github.com/evcc-io/evcc/issues/6439", p, pretty.Formatter(p.Val)) - return "\"foo\":\"bar\"" + err = h.subscribe(r.Context(), conn) + + if errors.Is(err, context.Canceled) { + return } - - var msg strings.Builder - msg.WriteString("\"") - if p.Loadpoint != nil { - msg.WriteString(fmt.Sprintf("loadpoints.%d.", *p.Loadpoint)) + if cs := websocket.CloseStatus(err); cs == websocket.StatusNormalClosure || cs == websocket.StatusGoingAway { + return + } + if err != nil { + log.ERROR.Println(err) + return } - msg.WriteString(p.Key) - msg.WriteString("\":") - msg.WriteString(val) - - return msg.String() } -func (h *SocketHub) welcome(client *SocketClient, params []util.Param) { - h.mu.Lock() - h.clients[client] = true - h.mu.Unlock() +func (h *SocketHub) subscribe(ctx context.Context, conn *websocket.Conn) error { + ctx = conn.CloseRead(ctx) + s := &socketSubscriber{ + send: make(chan []byte, 1024), + closeSlow: func() { + conn.Close(websocket.StatusPolicyViolation, "connection too slow to keep up with messages") + }, + } + + h.addSubscriber(s) + defer h.deleteSubscriber(s) + + // send welcome message + h.register <- s + + for { + select { + case msg := <-s.send: + if err := writeTimeout(ctx, socketWriteTimeout, conn, msg); err != nil { + return err + } + case <-ctx.Done(): + return ctx.Err() + } + } +} + +// addSubscriber registers a subscriber. +func (h *SocketHub) addSubscriber(s *socketSubscriber) { + h.mu.Lock() + h.subscribers[s] = struct{}{} + h.mu.Unlock() +} + +// deleteSubscriber deletes the given subscriber. +func (h *SocketHub) deleteSubscriber(s *socketSubscriber) { + h.mu.Lock() + delete(h.subscribers, s) + h.mu.Unlock() +} + +func (h *SocketHub) welcome(subscriber *socketSubscriber, params []util.Param) { var msg strings.Builder msg.WriteString("{") for _, p := range params { @@ -158,25 +132,22 @@ func (h *SocketHub) welcome(client *SocketClient, params []util.Param) { } msg.WriteString("}") - select { - case client.send <- []byte(msg.String()): - default: - close(client.send) - } + // should not block + subscriber.send <- []byte(msg.String()) } func (h *SocketHub) broadcast(p util.Param) { h.mu.RLock() defer h.mu.RUnlock() - if len(h.clients) > 0 { + if len(h.subscribers) > 0 { msg := "{" + kv(p) + "}" - for client := range h.clients { + for s := range h.subscribers { select { - case client.send <- []byte(msg): + case s.send <- []byte(msg): default: - h.unregister <- client + s.closeSlow() } } } @@ -188,13 +159,6 @@ func (h *SocketHub) Run(in <-chan util.Param, cache *util.Cache) { select { case client := <-h.register: h.welcome(client, cache.All()) - case client := <-h.unregister: - h.mu.Lock() - if _, ok := h.clients[client]; ok { - close(client.send) - delete(h.clients, client) - } - h.mu.Unlock() case msg, ok := <-in: if !ok { return // break if channel closed diff --git a/server/socket_helper.go b/server/socket_helper.go new file mode 100644 index 000000000..f1a8c2e99 --- /dev/null +++ b/server/socket_helper.go @@ -0,0 +1,63 @@ +package server + +import ( + "encoding/json" + "fmt" + "math" + "strings" + "time" + + "github.com/evcc-io/evcc/util" + "github.com/kr/pretty" +) + +func encode(v interface{}) (string, error) { + var s string + switch val := v.(type) { + case time.Time: + if val.IsZero() { + s = "null" + } else { + s = fmt.Sprintf(`"%s"`, val.Format(time.RFC3339)) + } + case time.Duration: + // must be before stringer to convert to seconds instead of string + s = fmt.Sprintf("%d", int64(val.Seconds())) + case float64: + if math.IsNaN(val) { + s = "null" + } else { + s = fmt.Sprintf("%.5g", val) + } + default: + if b, err := json.Marshal(v); err == nil { + s = string(b) + } else { + return "", err + } + } + return s, nil +} + +func kv(p util.Param) string { + val, err := encode(p.Val) + if err != nil { + panic(err) + } + + if p.Key == "" && val == "" { + log.ERROR.Printf("invalid key/val for %+v %# v, please report to https://github.com/evcc-io/evcc/issues/6439", p, pretty.Formatter(p.Val)) + return "\"foo\":\"bar\"" + } + + var msg strings.Builder + msg.WriteString("\"") + if p.Loadpoint != nil { + msg.WriteString(fmt.Sprintf("loadpoints.%d.", *p.Loadpoint)) + } + msg.WriteString(p.Key) + msg.WriteString("\":") + msg.WriteString(val) + + return msg.String() +}