diff --git a/server/socket.go b/server/socket.go index f096f9d38..40d042051 100644 --- a/server/socket.go +++ b/server/socket.go @@ -145,13 +145,8 @@ func (h *SocketHub) broadcast(p util.Param) { // Sharder splits data into chunks if sp, ok := (p.Val).(util.Sharder); ok { - shards := sp.Shards() - if len(shards) == 0 { - return // nothing changed, skip broadcast - } - - for _, shard := range shards { - msg[k+"."+shard.Key] = json.RawMessage(socketEncode(shard.Value)) + for key, val := range sp.ModifiedShards() { + msg[k+"."+key] = json.RawMessage(socketEncode(val)) } } else { msg[k] = json.RawMessage(socketEncode(p.Val)) diff --git a/util/param_shard.go b/util/param_shard.go index 41411682b..4ce0fe9d7 100644 --- a/util/param_shard.go +++ b/util/param_shard.go @@ -4,6 +4,7 @@ import ( "crypto/sha256" "encoding/json" "fmt" + "iter" "strings" "sync" @@ -13,7 +14,8 @@ import ( // Sharder splits data into chunks, omitting unmodified chunks type Sharder interface { - Shards() []Shard + ModifiedShards() iter.Seq2[string, any] + AllShards() iter.Seq2[string, any] } type Shard struct { @@ -36,41 +38,60 @@ func (s *sharderImpl) MarshalJSON() ([]byte, error) { return json.Marshal(s.struc) } -func (s *sharderImpl) Shards() []Shard { - ff := structs.Fields(s.struc) - res := make([]Shard, 0, len(ff)) +func (s *sharderImpl) AllShards() iter.Seq2[string, any] { + return s.shards(false) +} - shardMu.Lock() - defer shardMu.Unlock() +func (s *sharderImpl) ModifiedShards() iter.Seq2[string, any] { + return s.shards(true) +} - for _, f := range ff { - key := f.Name() - if t := f.Tag("json"); t != "" { - if n := strings.Split(t, ",")[0]; n != "" { - key = n - } - } - - // Use JSON for stable hashing (fmt.Append includes pointer addresses) - b, err := json.Marshal(f.Value()) - if err != nil { - // Fallback to fmt.Append if JSON fails - b = fmt.Append(nil, f.Value()) - } - - hash := sha256.Sum256(b) - if cached, ok := shardCache[s.prefix+key]; ok && hash == cached { - continue - } - shardCache[s.prefix+key] = hash - - res = append(res, Shard{ - Key: key, - Value: f.Value(), - }) +func (s *sharderImpl) shards(useCache bool) iter.Seq2[string, any] { + if useCache { + shardMu.Lock() + defer shardMu.Unlock() } - return res + return func(yield func(string, any) bool) { + for _, f := range structs.Fields(s.struc) { + key := jsonKey(f) + if useCache && s.skipCachedShard(key, f.Value()) { + continue + } + if !yield(key, f.Value()) { + break + } + } + } +} + +func jsonKey(f *structs.Field) string { + key := f.Name() + if t := f.Tag("json"); t != "" { + if n := strings.Split(t, ",")[0]; n != "" { + key = n + } + } + return key +} + +func (s *sharderImpl) skipCachedShard(key string, value any) bool { + // Use JSON for stable hashing (fmt.Append includes pointer addresses) + b, err := json.Marshal(value) + if err != nil { + // Fallback to fmt.Append if JSON fails + b = fmt.Append(nil, value) + } + + hash := sha256.Sum256(b) + cacheKey := s.prefix + key + + if cached, ok := shardCache[cacheKey]; ok && hash == cached { + return true + } + + shardCache[cacheKey] = hash + return false } var _ api.StructMarshaler = (*sharderImpl)(nil) diff --git a/util/param_shard_test.go b/util/param_shard_test.go new file mode 100644 index 000000000..de305306a --- /dev/null +++ b/util/param_shard_test.go @@ -0,0 +1,34 @@ +package util + +import ( + "maps" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestSharder(t *testing.T) { + type S struct { + A, B string + } + + s := NewSharder("foo", S{"a", "b"}) + + assert.Equal(t, map[string]any{ + "A": "a", + "B": "b", + }, maps.Collect(s.AllShards()), "non-cached") + + assert.Equal(t, map[string]any{ + "A": "a", + "B": "b", + }, maps.Collect(s.ModifiedShards()), "cache cold") + + assert.Equal(t, map[string]any{}, maps.Collect(s.ModifiedShards()), "cache warm") + + s = NewSharder("foo", S{"a", "c"}) + + assert.Equal(t, map[string]any{ + "B": "c", + }, maps.Collect(s.ModifiedShards()), "cache modied") +}