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

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