From 4642a135db65c4234d05b856014b3ce9afd941f6 Mon Sep 17 00:00:00 2001 From: andig Date: Sat, 25 Feb 2023 13:09:18 +0100 Subject: [PATCH 1/2] Protect client map by mutex --- server/socket.go | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/server/socket.go b/server/socket.go index 1c83dcb94..c7d70c03e 100644 --- a/server/socket.go +++ b/server/socket.go @@ -6,6 +6,7 @@ import ( "math" "net/http" "strings" + "sync" "time" "github.com/evcc-io/evcc/util" @@ -68,6 +69,8 @@ func ServeWebsocket(hub *SocketHub, w http.ResponseWriter, r *http.Request) { // 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 @@ -133,7 +136,9 @@ func kv(p util.Param) string { } func (h *SocketHub) welcome(client *SocketClient, params []util.Param) { + h.mu.Lock() h.clients[client] = true + h.mu.Unlock() var msg strings.Builder msg.WriteString("{") @@ -153,6 +158,9 @@ func (h *SocketHub) welcome(client *SocketClient, params []util.Param) { } func (h *SocketHub) broadcast(p util.Param) { + h.mu.RLock() + defer h.mu.RUnlock() + if len(h.clients) > 0 { msg := "{" + kv(p) + "}" @@ -173,10 +181,12 @@ func (h *SocketHub) Run(in <-chan util.Param, cache *util.Cache) { 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 From db98a34e234eb82299feaa9111601b9436fc4d84 Mon Sep 17 00:00:00 2001 From: andig Date: Sat, 25 Feb 2023 13:09:37 +0100 Subject: [PATCH 2/2] Improve encoding tests --- server/socket.go | 10 ++++++---- server/socket_test.go | 22 ++++++++++------------ 2 files changed, 16 insertions(+), 16 deletions(-) diff --git a/server/socket.go b/server/socket.go index c7d70c03e..466cd2e8c 100644 --- a/server/socket.go +++ b/server/socket.go @@ -95,11 +95,11 @@ func encode(v interface{}) (string, error) { var s string switch val := v.(type) { case time.Time: - var b []byte - if !val.IsZero() { - b, _ = val.Round(time.Second).Local().MarshalText() + if val.IsZero() { + s = "null" + } else { + s = fmt.Sprintf(`"%s"`, val.Format(time.RFC3339)) } - s = fmt.Sprintf(`"%s"`, string(b)) case time.Duration: // must be before stringer to convert to seconds instead of string s = fmt.Sprintf("%d", int64(val.Seconds())) @@ -112,6 +112,8 @@ func encode(v interface{}) (string, error) { default: if b, err := json.Marshal(v); err == nil { s = string(b) + } else { + return "", err } } return s, nil diff --git a/server/socket_test.go b/server/socket_test.go index f6abef3f2..ffe7368f0 100644 --- a/server/socket_test.go +++ b/server/socket_test.go @@ -2,14 +2,18 @@ package server import ( "math" - "reflect" "testing" "time" + + "github.com/stretchr/testify/assert" ) func TestEncode(t *testing.T) { + now := time.Now() + tc := []struct { - in, out interface{} + in interface{} + out string }{ {int64(1), "1"}, {math.NaN(), "null"}, @@ -17,19 +21,13 @@ func TestEncode(t *testing.T) { {"1.2345", "\"1.2345\""}, {time.Hour, "3600"}, {"minpv", "\"minpv\""}, + {time.Time{}, "null"}, + {now, "\"" + now.Format(time.RFC3339) + "\""}, } for _, tc := range tc { - t.Logf("%+v", tc) out, err := encode(tc.in) - if err != nil { - t.Error(err) - } - - if out != tc.out { - t.Errorf("expected %v (string), got %v (%s)", - tc.out, out, reflect.TypeOf(out).Kind(), - ) - } + assert.NoError(t, err) + assert.Equal(t, tc.out, out) } }