chore: drop gorilla/websocket (#7768)

This commit is contained in:
andig 2023-05-01 19:55:30 +02:00 • committed by GitHub
parent b1ee684db1
commit e9ac325592
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 163 additions and 133 deletions

2
go.mod
View file

@ -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

View file

@ -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
}

View file

@ -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)
}
}

View file

@ -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

63
server/socket_helper.go Normal file
View file

@ -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()
}