fix(stdiscosrv): must not modify database entries in-place

Signed-off-by: Jakob Borg <jakob@kastelo.net>
This commit is contained in:
Jakob Borg
2026-02-04 10:24:17 +01:00
parent b40f2acdad
commit f731cfa746
2 changed files with 41 additions and 27 deletions
+40 -26
View File
@@ -113,7 +113,7 @@ func (s *inMemoryStore) merge(key *protocol.DeviceID, addrs []*discosrv.Database
} }
if oldRec, ok := s.m.Load(*key); ok { if oldRec, ok := s.m.Load(*key); ok {
newRec = merge(oldRec, newRec) newRec = merge(newRec, oldRec)
} }
s.m.Store(*key, newRec) s.m.Store(*key, newRec)
@@ -135,7 +135,13 @@ func (s *inMemoryStore) get(key *protocol.DeviceID) (*discosrv.DatabaseRecord, e
return &discosrv.DatabaseRecord{}, nil return &discosrv.DatabaseRecord{}, nil
} }
rec.Addresses = expire(rec.Addresses, s.clock.Now()) naddresses, changed := expire(rec.Addresses, s.clock.Now())
if changed {
rec = &discosrv.DatabaseRecord{
Addresses: naddresses,
Seen: rec.Seen,
}
}
databaseOperations.WithLabelValues(dbOpGet, dbResSuccess).Inc() databaseOperations.WithLabelValues(dbOpGet, dbResSuccess).Inc()
return rec, nil return rec, nil
} }
@@ -184,12 +190,12 @@ func (s *inMemoryStore) expireAndCalculateStatistics() {
} }
n++ n++
addresses := expire(rec.Addresses, now) addresses, changed := expire(rec.Addresses, now)
if len(addresses) == 0 { if changed {
rec.Addresses = nil rec = &discosrv.DatabaseRecord{
s.m.Store(key, rec) Addresses: addresses,
} else if len(addresses) != len(rec.Addresses) { Seen: rec.Seen,
rec.Addresses = addresses }
s.m.Store(key, rec) s.m.Store(key, rec)
} }
@@ -371,9 +377,9 @@ func (s *inMemoryStore) read() (int, error) {
} }
slices.SortFunc(rec.Addresses, Cmp) slices.SortFunc(rec.Addresses, Cmp)
rec.Addresses = slices.CompactFunc(rec.Addresses, Equal) rec.Addresses, _ = expire(slices.CompactFunc(rec.Addresses, Equal), s.clock.Now())
s.m.Store(key, &discosrv.DatabaseRecord{ s.m.Store(key, &discosrv.DatabaseRecord{
Addresses: expire(rec.Addresses, s.clock.Now()), Addresses: rec.Addresses,
Seen: rec.Seen, Seen: rec.Seen,
}) })
nr++ nr++
@@ -384,7 +390,7 @@ func (s *inMemoryStore) read() (int, error) {
// merge returns the merged result of the two database records a and b. The // merge returns the merged result of the two database records a and b. The
// result is the union of the two address sets, with the newer expiry time // result is the union of the two address sets, with the newer expiry time
// chosen for any duplicates. The address list in a is overwritten and // chosen for any duplicates. The address list in a is overwritten and
// reused for the result. // reused for the result; b is not modified.
func merge(a, b *discosrv.DatabaseRecord) *discosrv.DatabaseRecord { func merge(a, b *discosrv.DatabaseRecord) *discosrv.DatabaseRecord {
// Both lists must be sorted for this to work. // Both lists must be sorted for this to work.
@@ -415,25 +421,33 @@ func merge(a, b *discosrv.DatabaseRecord) *discosrv.DatabaseRecord {
return a return a
} }
// expire returns the list of addresses after removing expired entries. // expire returns the list of addresses after removing expired entries. A
// Expiration happen in place, so the slice given as the parameter is // new slice is allocated if any changes are required, and the changed
// destroyed. Internal order is preserved. // boolean indicates whether that happened or not.
func expire(addrs []*discosrv.DatabaseAddress, now time.Time) []*discosrv.DatabaseAddress { func expire(addrs []*discosrv.DatabaseAddress, now time.Time) (result []*discosrv.DatabaseAddress, changed bool) {
cutoff := now.UnixNano() cutoff := now.UnixNano()
naddrs := addrs[:0] remains := 0
for i := range addrs { for _, a := range addrs {
if i > 0 && addrs[i].Address == addrs[i-1].Address { if a.Expires < cutoff {
// Skip duplicates changed = true
continue } else {
} remains++
if addrs[i].Expires >= cutoff {
naddrs = append(naddrs, addrs[i])
} }
} }
if len(naddrs) == 0 { if !changed {
return nil return addrs, false
} }
return naddrs if remains == 0 {
return nil, true
}
naddrs := make([]*discosrv.DatabaseAddress, 0, remains)
for _, a := range addrs {
if a.Expires >= cutoff {
naddrs = append(naddrs, a)
}
}
return naddrs, true
} }
func Cmp(d, other *discosrv.DatabaseAddress) (n int) { func Cmp(d, other *discosrv.DatabaseAddress) (n int) {
+1 -1
View File
@@ -161,7 +161,7 @@ func TestFilter(t *testing.T) {
} }
for _, tc := range cases { for _, tc := range cases {
res := expire(tc.a, time.Unix(0, 10)) res, _ := expire(tc.a, time.Unix(0, 10))
if fmt.Sprint(res) != fmt.Sprint(tc.b) { if fmt.Sprint(res) != fmt.Sprint(tc.b) {
t.Errorf("Incorrect result %v, expected %v", res, tc.b) t.Errorf("Incorrect result %v, expected %v", res, tc.b)
} }