Merge pull request #6434 from evcc-io/chore/socket-encoding
Improve websocket handling
This commit is contained in:
commit
4af116a532
2 changed files with 26 additions and 16 deletions
|
|
@ -6,6 +6,7 @@ import (
|
|||
"math"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/evcc-io/evcc/util"
|
||||
|
|
@ -69,6 +70,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
|
||||
|
||||
|
|
@ -93,11 +96,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()))
|
||||
|
|
@ -110,6 +113,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
|
||||
|
|
@ -139,7 +144,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("{")
|
||||
|
|
@ -159,6 +166,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) + "}"
|
||||
|
||||
|
|
@ -179,10 +189,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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue