From 74c6af2980291508591781eb0c4a4ea4891249ac Mon Sep 17 00:00:00 2001 From: edde746 <86283021+edde746@users.noreply.github.com> Date: Wed, 22 Apr 2026 14:57:41 +0200 Subject: [PATCH] feat(server): persist rooms across restarts --- server/Dockerfile | 2 +- server/docker-compose.yml | 3 +- server/main.go | 367 ++++++++++++++++++++++++++++++++++---- server/main_test.go | 290 ++++++++++++++++++++++++++++++ 4 files changed, 626 insertions(+), 36 deletions(-) create mode 100644 server/main_test.go diff --git a/server/Dockerfile b/server/Dockerfile index 179682b8..4602d39b 100644 --- a/server/Dockerfile +++ b/server/Dockerfile @@ -7,7 +7,7 @@ RUN CGO_ENABLED=0 go build -ldflags="-s -w" -o /relay . FROM scratch COPY --from=build /relay /relay -VOLUME /data/logs +VOLUME /data EXPOSE 8080 ENV GOMEMLIMIT=64MiB ENTRYPOINT ["/relay"] diff --git a/server/docker-compose.yml b/server/docker-compose.yml index caf64339..2e012ff3 100644 --- a/server/docker-compose.yml +++ b/server/docker-compose.yml @@ -3,8 +3,9 @@ services: build: . restart: unless-stopped mem_limit: 128m + stop_grace_period: 15s volumes: - - logs:/data/logs + - logs:/data expose: - "8080" logging: diff --git a/server/main.go b/server/main.go index de761b77..8bc94537 100644 --- a/server/main.go +++ b/server/main.go @@ -1,19 +1,24 @@ package main import ( + "context" "crypto/rand" "encoding/json" + "errors" "flag" "io" + "io/fs" "log" "math/big" "net" "net/http" "os" + "os/signal" "path/filepath" "strconv" "strings" "sync" + "syscall" "time" "github.com/gorilla/websocket" @@ -40,6 +45,11 @@ const ( maxRoomsPerIP = 3 connRateBurst = 5 connRateSustained = 1 + + snapshotFormatVersion = 1 + snapshotDebounce = 100 * time.Millisecond + snapshotFlushTimeout = 5 * time.Second + snapshotMaxFileSize = 1 * 1024 * 1024 ) var upgrader = websocket.Upgrader{ @@ -254,11 +264,27 @@ func (c *Client) close() { // --- Room --- type Room struct { - SessionID string - HostPeerID string - Peers map[string]*Client - mu sync.RWMutex - CreatedAt time.Time + SessionID string + HostPeerID string + Peers map[string]*Client `json:"-"` + mu sync.RWMutex `json:"-"` + CreatedAt time.Time + LastActivityAt time.Time +} + +// --- Snapshot types (on-disk JSON format) --- + +type roomSnapshot struct { + SessionID string `json:"sessionId"` + HostPeerID string `json:"hostPeerId"` + CreatedAt time.Time `json:"createdAt"` + LastActivityAt time.Time `json:"lastActivityAt"` +} + +type stateSnapshot struct { + Version int `json:"version"` + SavedAt time.Time `json:"savedAt"` + Rooms []roomSnapshot `json:"rooms"` } func (r *Room) peerIDs() []string { @@ -367,49 +393,284 @@ func (ls *logStore) cleanup() { } } +// --- Snapshotter (single-writer, debounced, atomic disk persistence) --- + +type snapshotter struct { + path string + dir string + trigger chan struct{} + flush chan chan error + done chan struct{} + exited chan struct{} + build func() stateSnapshot + writeMu sync.Mutex + + errMu sync.Mutex + lastErrLog time.Time +} + +func newSnapshotter(path string, build func() stateSnapshot) *snapshotter { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0755); err != nil { + log.Printf("snapshot: mkdir %s: %v", dir, err) + } + return &snapshotter{ + path: path, + dir: dir, + trigger: make(chan struct{}, 1), + flush: make(chan chan error), + done: make(chan struct{}), + exited: make(chan struct{}), + build: build, + } +} + +func (sn *snapshotter) schedule() { + select { + case sn.trigger <- struct{}{}: + default: + } +} + +func (sn *snapshotter) run() { + defer close(sn.exited) + for { + select { + case <-sn.trigger: + time.Sleep(snapshotDebounce) + // Drain triggers that piled up during the sleep window. + select { + case <-sn.trigger: + default: + } + sn.writeAndLog() + case reply := <-sn.flush: + reply <- sn.writeAndLog() + case <-sn.done: + // flushAndStop is the expected caller and has already written. + return + } + } +} + +func (sn *snapshotter) writeAndLog() error { + err := sn.write() + if err != nil { + sn.logWriteErr(err) + } + return err +} + +func (sn *snapshotter) write() error { + sn.writeMu.Lock() + defer sn.writeMu.Unlock() + + data, err := json.Marshal(sn.build()) + if err != nil { + return err + } + + tmpPath := sn.path + ".tmp" + f, err := os.OpenFile(tmpPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0644) + if err != nil { + return err + } + if _, err := f.Write(data); err != nil { + f.Close() + os.Remove(tmpPath) + return err + } + if err := f.Sync(); err != nil { + f.Close() + os.Remove(tmpPath) + return err + } + if err := f.Close(); err != nil { + os.Remove(tmpPath) + return err + } + if err := os.Rename(tmpPath, sn.path); err != nil { + os.Remove(tmpPath) + return err + } + // Best-effort dir fsync so the rename is durable after host crash. + if d, err := os.Open(sn.dir); err == nil { + d.Sync() + d.Close() + } + return nil +} + +func (sn *snapshotter) flushAndStop(timeout time.Duration) error { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + reply := make(chan error, 1) + select { + case sn.flush <- reply: + case <-ctx.Done(): + return errors.New("snapshot flush: timed out sending flush signal") + } + var flushErr error + select { + case flushErr = <-reply: + case <-ctx.Done(): + return errors.New("snapshot flush: timed out waiting for write") + } + close(sn.done) + select { + case <-sn.exited: + case <-ctx.Done(): + } + return flushErr +} + +// logWriteErr throttles snapshot-write error spam to at most once per hour. +func (sn *snapshotter) logWriteErr(err error) { + sn.errMu.Lock() + defer sn.errMu.Unlock() + if time.Since(sn.lastErrLog) < time.Hour { + return + } + sn.lastErrLog = time.Now() + log.Printf("snapshot: write failed: %v", err) +} + // --- Server --- type Server struct { rooms map[string]*Room logs *logStore conns *connTracker + snap *snapshotter mu sync.RWMutex } -func newServer(logDir string) *Server { +func newServer(logDir, stateFile string) *Server { s := &Server{rooms: make(map[string]*Room), logs: newLogStore(logDir), conns: newConnTracker()} + s.snap = newSnapshotter(stateFile, s.buildSnapshot) + if err := s.loadSnapshot(stateFile); err != nil { + log.Printf("snapshot: load error: %v", err) + } + go s.snap.run() go s.cleanupLoop() return s } +// buildSnapshot copies room identity into a serializable value with no locks held during marshal. +// Lock order: s.mu before room.mu, matching cleanupLoop. +func (s *Server) buildSnapshot() stateSnapshot { + s.mu.RLock() + defer s.mu.RUnlock() + snap := stateSnapshot{ + Version: snapshotFormatVersion, + SavedAt: time.Now(), + Rooms: make([]roomSnapshot, 0, len(s.rooms)), + } + for _, room := range s.rooms { + room.mu.RLock() + snap.Rooms = append(snap.Rooms, roomSnapshot{ + SessionID: room.SessionID, + HostPeerID: room.HostPeerID, + CreatedAt: room.CreatedAt, + LastActivityAt: room.LastActivityAt, + }) + room.mu.RUnlock() + } + return snap +} + +// loadSnapshot restores rooms from disk on startup. Missing/corrupt files log +// and return nil so the server always starts; only unexpected I/O paths bubble up. +func (s *Server) loadSnapshot(path string) error { + data, err := os.ReadFile(path) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + log.Printf("snapshot: no file at %s, starting fresh", path) + return nil + } + log.Printf("snapshot: read error, starting fresh: %v", err) + return nil + } + if len(data) > snapshotMaxFileSize { + log.Printf("snapshot: file too large (%d bytes), starting fresh", len(data)) + return nil + } + var snap stateSnapshot + if err := json.Unmarshal(data, &snap); err != nil { + log.Printf("snapshot: corrupt file at %s, starting fresh: %v", path, err) + return nil + } + if snap.Version != snapshotFormatVersion { + log.Printf("snapshot: unknown version %d, starting fresh", snap.Version) + return nil + } + now := time.Now() + loaded, skipped := 0, 0 + s.mu.Lock() + for _, r := range snap.Rooms { + if r.SessionID == "" || r.HostPeerID == "" { + skipped++ + continue + } + if now.Sub(r.CreatedAt) > roomMaxAge { + skipped++ + continue + } + if now.Sub(r.LastActivityAt) > emptyRoomMaxAge { + skipped++ + continue + } + s.rooms[r.SessionID] = &Room{ + SessionID: r.SessionID, + HostPeerID: r.HostPeerID, + Peers: make(map[string]*Client), + CreatedAt: r.CreatedAt, + LastActivityAt: r.LastActivityAt, + } + loaded++ + } + s.mu.Unlock() + log.Printf("snapshot: loaded %d rooms, skipped %d expired", loaded, skipped) + return nil +} + func (s *Server) cleanupLoop() { ticker := time.NewTicker(cleanupInterval) defer ticker.Stop() for range ticker.C { - s.mu.Lock() - now := time.Now() - for id, room := range s.rooms { - room.mu.RLock() - empty := len(room.Peers) == 0 - age := now.Sub(room.CreatedAt) - room.mu.RUnlock() - - if (empty && age > emptyRoomMaxAge) || age > roomMaxAge { - log.Printf("cleanup: removing room %s (empty=%v, age=%v)", id, empty, age) - delete(s.rooms, id) - } - } - s.mu.Unlock() - s.logs.cleanup() - s.conns.cleanup() - - s.conns.mu.Lock() - log.Printf("stats: conns=%d ips=%d rooms=%d", - s.conns.globalCount, len(s.conns.perIP), len(s.rooms)) - s.conns.mu.Unlock() + s.runCleanupStep(time.Now()) } } +func (s *Server) runCleanupStep(now time.Time) { + s.mu.Lock() + changed := false + for id, room := range s.rooms { + room.mu.RLock() + empty := len(room.Peers) == 0 + age := now.Sub(room.CreatedAt) + idle := now.Sub(room.LastActivityAt) + room.mu.RUnlock() + + if (empty && idle > emptyRoomMaxAge) || age > roomMaxAge { + log.Printf("cleanup: removing room %s (empty=%v, idle=%v, age=%v)", id, empty, idle, age) + delete(s.rooms, id) + changed = true + } + } + s.mu.Unlock() + if changed { + s.snap.schedule() + } + s.logs.cleanup() + s.conns.cleanup() + + s.conns.mu.Lock() + log.Printf("stats: conns=%d ips=%d rooms=%d", + s.conns.globalCount, len(s.conns.perIP), len(s.rooms)) + s.conns.mu.Unlock() +} + func clientIP(r *http.Request) string { var raw string if fwd := r.Header.Get("X-Forwarded-For"); fwd != "" { @@ -562,6 +823,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { stale := currentRoom.Peers[currentPeerID] != client if !stale { delete(currentRoom.Peers, currentPeerID) + currentRoom.LastActivityAt = time.Now() } currentRoom.mu.Unlock() if !stale { @@ -569,6 +831,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { Type: "peerLeft", PeerID: currentPeerID, }) + s.snap.schedule() } if isHost { s.conns.releaseRoom(ip) @@ -621,11 +884,13 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { // Empty stale room — reclaim the ID delete(s.rooms, msg.SessionID) } + now := time.Now() room := &Room{ - SessionID: msg.SessionID, - HostPeerID: msg.PeerID, - Peers: map[string]*Client{msg.PeerID: client}, - CreatedAt: time.Now(), + SessionID: msg.SessionID, + HostPeerID: msg.PeerID, + Peers: map[string]*Client{msg.PeerID: client}, + CreatedAt: now, + LastActivityAt: now, } s.rooms[msg.SessionID] = room s.mu.Unlock() @@ -634,6 +899,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { isHost = true log.Printf("room %s created by %s", msg.SessionID, msg.PeerID) client.sendJSON(serverMsg{Type: "created", SessionID: msg.SessionID}) + s.snap.schedule() case "join": if msg.SessionID == "" || msg.PeerID == "" { @@ -654,6 +920,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { continue } room.Peers[msg.PeerID] = client + room.LastActivityAt = time.Now() peers := room.peerIDs() room.mu.Unlock() currentRoom = room @@ -669,6 +936,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { } client.sendJSON(serverMsg{Type: "joined", SessionID: msg.SessionID, Peers: existingPeers}) room.broadcastExcept(msg.PeerID, serverMsg{Type: "peerJoined", PeerID: msg.PeerID}) + s.snap.schedule() case "broadcast": if currentRoom == nil { @@ -710,9 +978,10 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) { func main() { addr := flag.String("addr", ":8080", "Listen address") logDir := flag.String("log-dir", "/data/logs", "Directory for log file storage") + stateFile := flag.String("state-file", "/data/rooms.json", "Path to room snapshot file") flag.Parse() - srv := newServer(*logDir) + srv := newServer(*logDir, *stateFile) mux := http.NewServeMux() mux.HandleFunc("/relay", srv.handleWS) @@ -723,6 +992,36 @@ func main() { mux.HandleFunc("/logs", srv.handlePostLogs) mux.HandleFunc("/logs/", srv.handleGetLogs) - log.Printf("Starting relay server on %s", *addr) - log.Fatal(http.ListenAndServe(*addr, mux)) + httpSrv := &http.Server{Addr: *addr, Handler: mux} + + serveErr := make(chan error, 1) + go func() { + log.Printf("Starting relay server on %s", *addr) + if err := httpSrv.ListenAndServe(); err != nil && err != http.ErrServerClosed { + serveErr <- err + } + close(serveErr) + }() + + sig := make(chan os.Signal, 1) + signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM) + + select { + case err, ok := <-serveErr: + if ok { + log.Fatalf("listen: %v", err) + } + case s := <-sig: + log.Printf("shutdown signal received (%s), draining...", s) + } + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + if err := httpSrv.Shutdown(ctx); err != nil { + log.Printf("http shutdown: %v", err) + } + if err := srv.snap.flushAndStop(snapshotFlushTimeout); err != nil { + log.Printf("snapshot flush: %v", err) + } + log.Printf("shutdown complete") } diff --git a/server/main_test.go b/server/main_test.go new file mode 100644 index 00000000..76bb4a49 --- /dev/null +++ b/server/main_test.go @@ -0,0 +1,290 @@ +package main + +import ( + "encoding/json" + "os" + "path/filepath" + "sync" + "testing" + "time" +) + +// newTestServer builds a Server wired for tests: no goroutines, no network, +// logs scratched to a per-test temp dir. The snapshotter is constructed but +// its goroutine is NOT started — tests drive it synchronously via write() +// or call schedule() and then call write() themselves. +func newTestServer(t *testing.T, stateFile string) *Server { + t.Helper() + s := &Server{ + rooms: make(map[string]*Room), + logs: newLogStore(t.TempDir()), + conns: newConnTracker(), + } + s.snap = newSnapshotter(stateFile, s.buildSnapshot) + return s +} + +func TestSnapshotRoundTrip(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "rooms.json") + s := newTestServer(t, path) + + now := time.Now().UTC().Truncate(time.Second) + s.rooms["ABC12"] = &Room{ + SessionID: "ABC12", + HostPeerID: "host-1", + Peers: map[string]*Client{}, + CreatedAt: now.Add(-time.Minute), + LastActivityAt: now, + } + s.rooms["XYZ99"] = &Room{ + SessionID: "XYZ99", + HostPeerID: "host-2", + Peers: map[string]*Client{}, + CreatedAt: now.Add(-time.Hour), + LastActivityAt: now.Add(-time.Second), + } + + if err := s.snap.write(); err != nil { + t.Fatalf("write: %v", err) + } + + // Reconstruct into a fresh Server and verify identity. + s2 := newTestServer(t, path) + if err := s2.loadSnapshot(path); err != nil { + t.Fatalf("loadSnapshot: %v", err) + } + if got := len(s2.rooms); got != 2 { + t.Fatalf("expected 2 rooms after reload, got %d", got) + } + for _, id := range []string{"ABC12", "XYZ99"} { + r, ok := s2.rooms[id] + if !ok { + t.Fatalf("room %s missing after reload", id) + } + orig := s.rooms[id] + if r.HostPeerID != orig.HostPeerID { + t.Errorf("%s: HostPeerID=%q want %q", id, r.HostPeerID, orig.HostPeerID) + } + if !r.CreatedAt.Equal(orig.CreatedAt) { + t.Errorf("%s: CreatedAt=%v want %v", id, r.CreatedAt, orig.CreatedAt) + } + if !r.LastActivityAt.Equal(orig.LastActivityAt) { + t.Errorf("%s: LastActivityAt=%v want %v", id, r.LastActivityAt, orig.LastActivityAt) + } + if r.Peers == nil { + t.Errorf("%s: Peers map nil after reload", id) + } + if len(r.Peers) != 0 { + t.Errorf("%s: expected empty Peers, got %d", id, len(r.Peers)) + } + } +} + +func TestLoadSkipsExpired(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "rooms.json") + now := time.Now() + + snap := stateSnapshot{ + Version: snapshotFormatVersion, + SavedAt: now, + Rooms: []roomSnapshot{ + {SessionID: "FRESH", HostPeerID: "h", CreatedAt: now.Add(-time.Minute), LastActivityAt: now.Add(-30 * time.Second)}, + {SessionID: "OLD24", HostPeerID: "h", CreatedAt: now.Add(-25 * time.Hour), LastActivityAt: now.Add(-time.Second)}, + {SessionID: "IDLE6", HostPeerID: "h", CreatedAt: now.Add(-2 * time.Hour), LastActivityAt: now.Add(-6 * time.Minute)}, + {SessionID: "", HostPeerID: "h", CreatedAt: now, LastActivityAt: now}, + {SessionID: "NOHOS", HostPeerID: "", CreatedAt: now, LastActivityAt: now}, + }, + } + data, err := json.Marshal(snap) + if err != nil { + t.Fatalf("marshal: %v", err) + } + if err := os.WriteFile(path, data, 0644); err != nil { + t.Fatalf("write: %v", err) + } + + s := newTestServer(t, path) + if err := s.loadSnapshot(path); err != nil { + t.Fatalf("loadSnapshot: %v", err) + } + if len(s.rooms) != 1 { + t.Fatalf("expected 1 room after load, got %d: %v", len(s.rooms), s.rooms) + } + if _, ok := s.rooms["FRESH"]; !ok { + t.Fatalf("FRESH room should have loaded") + } +} + +func TestLoadHandlesCorrupt(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "rooms.json") + if err := os.WriteFile(path, []byte("not valid json {{{"), 0644); err != nil { + t.Fatalf("write: %v", err) + } + s := newTestServer(t, path) + if err := s.loadSnapshot(path); err != nil { + t.Fatalf("loadSnapshot returned error: %v", err) + } + if len(s.rooms) != 0 { + t.Fatalf("expected empty rooms after corrupt load, got %d", len(s.rooms)) + } + // File should be preserved for debugging. + if _, err := os.Stat(path); err != nil { + t.Fatalf("corrupt file should NOT be deleted: %v", err) + } +} + +func TestLoadHandlesMissing(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "does-not-exist.json") + s := newTestServer(t, path) + if err := s.loadSnapshot(path); err != nil { + t.Fatalf("loadSnapshot: %v", err) + } + if len(s.rooms) != 0 { + t.Fatalf("expected empty rooms, got %d", len(s.rooms)) + } +} + +func TestLoadHandlesUnknownVersion(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "rooms.json") + if err := os.WriteFile(path, []byte(`{"version":99,"rooms":[{"sessionId":"X"}]}`), 0644); err != nil { + t.Fatalf("write: %v", err) + } + s := newTestServer(t, path) + if err := s.loadSnapshot(path); err != nil { + t.Fatalf("loadSnapshot: %v", err) + } + if len(s.rooms) != 0 { + t.Fatalf("expected empty rooms for unknown version, got %d", len(s.rooms)) + } +} + +func TestCleanupUsesIdleNotAge(t *testing.T) { + s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) + now := time.Now() + + // 2h-old room that has activity 1min ago — must NOT be cleaned up. + s.rooms["KEEP"] = &Room{ + SessionID: "KEEP", + HostPeerID: "h", + Peers: map[string]*Client{}, + CreatedAt: now.Add(-2 * time.Hour), + LastActivityAt: now.Add(-1 * time.Minute), + } + // 2h-old room that emptied 10min ago — MUST be cleaned up. + s.rooms["GONE"] = &Room{ + SessionID: "GONE", + HostPeerID: "h", + Peers: map[string]*Client{}, + CreatedAt: now.Add(-2 * time.Hour), + LastActivityAt: now.Add(-10 * time.Minute), + } + // 25h-old room — absolute TTL nukes it even if recently active. + s.rooms["OLD"] = &Room{ + SessionID: "OLD", + HostPeerID: "h", + Peers: map[string]*Client{}, // empty anyway + CreatedAt: now.Add(-25 * time.Hour), + LastActivityAt: now.Add(-10 * time.Second), + } + + s.runCleanupStep(now) + + if _, ok := s.rooms["KEEP"]; !ok { + t.Errorf("KEEP should still exist (recent activity)") + } + if _, ok := s.rooms["GONE"]; ok { + t.Errorf("GONE should have been cleaned (idle>5min)") + } + if _, ok := s.rooms["OLD"]; ok { + t.Errorf("OLD should have been cleaned (age>24h)") + } +} + +func TestSnapshotAtomicWriteSurvivesRenameFailure(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "rooms.json") + s := newTestServer(t, path) + + // Seed a valid snapshot on disk. + s.rooms["ORIG"] = &Room{ + SessionID: "ORIG", + HostPeerID: "h", + Peers: map[string]*Client{}, + CreatedAt: time.Now(), + LastActivityAt: time.Now(), + } + if err := s.snap.write(); err != nil { + t.Fatalf("first write: %v", err) + } + origBytes, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read orig: %v", err) + } + + // Force a second write to fail at the tmp-file create step by making the + // snapshot directory unwritable. The rename therefore never runs, so the + // existing file must be untouched. + if err := os.Chmod(dir, 0555); err != nil { + t.Fatalf("chmod: %v", err) + } + t.Cleanup(func() { os.Chmod(dir, 0755) }) + + delete(s.rooms, "ORIG") + s.rooms["NEW"] = &Room{ + SessionID: "NEW", + HostPeerID: "h", + Peers: map[string]*Client{}, + CreatedAt: time.Now(), + LastActivityAt: time.Now(), + } + if err := s.snap.write(); err == nil { + t.Fatalf("expected write to fail with dir read-only") + } + + // Restore permissions so we can read the file back. + os.Chmod(dir, 0755) + nowBytes, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read after failed write: %v", err) + } + if string(origBytes) != string(nowBytes) { + t.Fatalf("snapshot file was corrupted after failed write:\nbefore: %s\nafter: %s", origBytes, nowBytes) + } +} + +func TestSnapshotDebounceCoalesces(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "rooms.json") + + var ( + buildCount int + countMu sync.Mutex + ) + sn := newSnapshotter(path, func() stateSnapshot { + countMu.Lock() + buildCount++ + countMu.Unlock() + return stateSnapshot{Version: snapshotFormatVersion, SavedAt: time.Now(), Rooms: nil} + }) + go sn.run() + t.Cleanup(func() { _ = sn.flushAndStop(time.Second) }) + + // Fire a burst — should collapse into one write due to debounce. + for i := 0; i < 20; i++ { + sn.schedule() + } + // Give the debounce window + a small buffer to actually run. + time.Sleep(snapshotDebounce + 50*time.Millisecond) + + countMu.Lock() + got := buildCount + countMu.Unlock() + if got != 1 { + t.Fatalf("expected 1 build from burst, got %d", got) + } +}