Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 1 addition & 3 deletions cmd/server/channel_proposals_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -384,9 +384,7 @@ func TestChannelsListDoesNotMutateCaches(t *testing.T) {
// capacity, so pad the decrypted channels until it does.
var cached []map[string]interface{}
for i := 0; ; i++ {
db.channelsCacheMu.Lock()
db.channelsCacheRes = nil
db.channelsCacheMu.Unlock()
db.channelsCache.reset()
var err error
if cached, err = db.GetChannels(""); err != nil {
t.Fatal(err)
Expand Down
14 changes: 6 additions & 8 deletions cmd/server/channels_cache_append_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -91,11 +91,8 @@ func insertChannelTx(t *testing.T, db *DB, hash, channelHash, decodedJSON string
}

func resetDBChannelsCache(db *DB) {
db.channelsCacheMu.Lock()
db.channelsCacheRes = nil
db.channelsCacheKey = ""
db.channelsCacheExp = time.Time{}
db.channelsCacheMu.Unlock()
db.channelsCache.reset()
db.encChannelsCache.reset()
}

// TestChannelsCacheAppend_DBPath drives GET /api/channels through the real
Expand Down Expand Up @@ -153,9 +150,10 @@ func TestChannelsCacheAppend_DBPath(t *testing.T) {
}
}

db.channelsCacheMu.Lock()
final := db.channelsCacheRes
db.channelsCacheMu.Unlock()
var final []map[string]interface{}
if e, ok := db.channelsCache.get("", time.Now()); ok {
final = e.channels
}
if cap(final) != cap(cached) || len(final) != len(cached) {
t.Fatalf("cache slice header changed unexpectedly: before cap=%d len=%d, after cap=%d len=%d (cache may have expired/refreshed — this test assumes it didn't)", cap(cached), len(cached), cap(final), len(final))
}
Expand Down
94 changes: 94 additions & 0 deletions cmd/server/channels_key_bound_109_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
package main

import (
"strconv"
"strings"
"sync"
"testing"
)

// Every code in the region parameter becomes part of the key, so a long
// parameter must not be kept as a map key (and flight key) byte for byte.
// Long keys are replaced by a fixed-size digest; they are still cached,
// coalesced and kept apart per region, and the entry bound still holds.
func TestChannelListCacheKeyLengthIsBounded_109(t *testing.T) {
db := singleflightDB(t)
q := &queryCounter{want: 1}
db.channelsQueryHook = q.hook
many := make([]string, 300)
for i := range many {
many[i] = "R" + strconv.Itoa(i)
}
regions := []string{
strings.Join(many, ","), // many short codes
strings.Repeat("Y", 1<<20), // one 1 MiB code
strings.Repeat("Y", 1<<20) + "Z", // differs only in the last byte
"AAR",
}
for _, r := range regions {
for i := 0; i < 2; i++ { // the second call is served from the cache
if _, err := db.GetChannels(r); err != nil {
t.Fatal(err)
}
if _, err := db.GetEncryptedChannels(r); err != nil {
t.Fatal(err)
}
}
}
for _, kind := range []string{"channels", "encrypted"} {
if got := q.count(kind); got != len(regions) {
t.Errorf("%s: %d queries for %d regions called twice each, want one per region", kind, got, len(regions))
}
}
for name, c := range map[string]*channelListCache{"channels": &db.channelsCache, "encrypted": &db.encChannelsCache} {
if len(c.entries) != len(regions) {
t.Errorf("%s: %d entries for %d regions", name, len(c.entries), len(regions))
}
for k := range c.entries {
if len(k) > channelListMaxKeyBytes {
t.Errorf("%s: cache key of %d bytes kept, limit %d", name, len(k), channelListMaxKeyBytes)
}
}
if _, ok := c.entries["AAR"]; !ok {
t.Errorf("%s: a short key is not stored as itself", name)
}
}
}

// Only the cache and the flight use the bounded key: the query itself must
// still get the full normalized region, or a long region parameter would be
// looked up as the region "sha256:…" and cache an empty list.
func TestChannelListLongKeyQueriesFullRegion_109(t *testing.T) {
db := singleflightDB(t)
var mu sync.Mutex
got := map[string][]string{}
db.channelsQueryHook = func(kind, region string) error {
mu.Lock()
got[kind] = append(got[kind], region)
mu.Unlock()
return nil
}
many := make([]string, 300)
for i := range many {
many[i] = "r" + strconv.Itoa(i)
}
param := strings.Join(many, ",")
want := channelsRegionKey(param)
if len(want) <= channelListMaxKeyBytes {
t.Fatalf("test region key is %d bytes, want more than %d", len(want), channelListMaxKeyBytes)
}
if _, err := db.GetChannels(param); err != nil {
t.Fatal(err)
}
if _, err := db.GetEncryptedChannels(param); err != nil {
t.Fatal(err)
}
for _, kind := range []string{"channels", "encrypted"} {
if len(got[kind]) != 1 {
t.Fatalf("%s: %d queries, want 1", kind, len(got[kind]))
}
if got[kind][0] != want {
t.Errorf("%s: query got region %.40q…, want the full normalized key %.40q…", kind, got[kind][0], want)
}
}
}
140 changes: 140 additions & 0 deletions cmd/server/channels_list_cache.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
package main

import (
"crypto/sha256"
"encoding/hex"
"slices"
"sort"
"strings"
"sync"
"time"

"golang.org/x/sync/singleflight"
)

// channelListCache backs DB.GetChannels and DB.GetEncryptedChannels (#109).
//
// - Keyed by the normalized region (channelsRegionKey), so " aar " and
// "AAR,aar" share an entry and a flight; different regions never do.
// - Concurrent misses for one key run the query once (singleflight). The
// flight checks the cache again first: a caller that missed before
// another flight stored a result must not query again.
// - Only successful results are stored; an error is returned to the
// callers of that flight and the next call queries again.
// - Bounded: at most channelListMaxKeys entries. The region comes from a
// query parameter, so the key space is caller-controlled; when full,
// expired entries are dropped first, then the one closest to expiry.
// A key longer than channelListMaxKeyBytes is stored (and used as the
// flight key) as its SHA-256 digest, so each stored key is small too.
//
// A stored entry is never modified. Its channels slice is shared by every
// caller, so callers must not write into it (handleChannels appends with a
// full slice expression, #98).
type channelListCache struct {
mu sync.Mutex
entries map[string]*channelListEntry
flights singleflight.Group
}

const (
channelListTTL = 60 * time.Second
channelListMaxKeys = 64
// channelListMaxKeyBytes caps a stored key. A longer one becomes
// "sha256:" + 64 hex digits (71 bytes). Normalized region codes are
// upper-case (normalizeRegionCodes), so no short key starts with the
// lower-case "sha256:" and a digest never equals a short key.
channelListMaxKeyBytes = 256
)

// channelListStoreKey is the key the cache and the flight use for key.
func channelListStoreKey(key string) string {
if len(key) <= channelListMaxKeyBytes {
return key
}
sum := sha256.Sum256([]byte(key))
return "sha256:" + hex.EncodeToString(sum[:])
}

// channelsRegionKey is the cache and flight key for a region parameter: the
// normalized codes (normalizeRegionCodes), sorted and de-duplicated. The
// query runs with normalizeRegionCodes(key), the same set of codes.
func channelsRegionKey(region string) string {
codes := normalizeRegionCodes(region)
if len(codes) == 0 {
return ""
}
sort.Strings(codes)
return strings.Join(slices.Compact(codes), ",")
}

func (c *channelListCache) get(key string, now time.Time) (*channelListEntry, bool) {
c.mu.Lock()
defer c.mu.Unlock()
e := c.entries[key]
if e == nil || !now.Before(e.expires) {
return nil, false
}
return e, true
}

func (c *channelListCache) put(key string, e *channelListEntry, now time.Time) {
c.mu.Lock()
defer c.mu.Unlock()
if c.entries == nil {
c.entries = make(map[string]*channelListEntry)
}
if _, ok := c.entries[key]; !ok && len(c.entries) >= channelListMaxKeys {
for k, old := range c.entries {
if !now.Before(old.expires) {
delete(c.entries, k)
}
}
if len(c.entries) >= channelListMaxKeys {
victim, first := "", true
var soonest time.Time
for k, old := range c.entries {
if first || old.expires.Before(soonest) {
victim, soonest, first = k, old.expires, false
}
}
delete(c.entries, victim)
}
}
c.entries[key] = e
}

// reset drops every entry (tests).
func (c *channelListCache) reset() {
c.mu.Lock()
c.entries = nil
c.mu.Unlock()
}

// load returns the cached entry for key or, on a miss, the result of one
// query shared by every concurrent caller of key. missed runs after the
// first cache check fails (a test seam; nil in production).
func (c *channelListCache) load(key string, missed func(), query func() (*channelListEntry, error)) (*channelListEntry, error) {
key = channelListStoreKey(key)
if e, ok := c.get(key, time.Now()); ok {
return e, nil
}
if missed != nil {
missed()
}
v, err, _ := c.flights.Do(key, func() (interface{}, error) {
if e, ok := c.get(key, time.Now()); ok {
return e, nil
}
e, err := query()
if err != nil {
return nil, err
}
e.expires = time.Now().Add(channelListTTL)
c.put(key, e, time.Now())
return e, nil
})
if err != nil {
return nil, err
}
return v.(*channelListEntry), nil
}
Loading
Loading