chore: drop gorilla/websocket (#7768)
This commit is contained in:
parent
b1ee684db1
commit
e9ac325592
5 changed files with 163 additions and 133 deletions
2
go.mod
2
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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
212
server/socket.go
212
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
|
||||
|
|
|
|||
63
server/socket_helper.go
Normal file
63
server/socket_helper.go
Normal 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()
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue