package main import ( "bufio" "bytes" "encoding/json" "errors" "fmt" "io" "io/fs" "log" "math" "net" "net/http" "net/http/httptest" "net/url" "os" "path/filepath" "strings" "sync" "sync/atomic" "syscall" "testing" "time" "github.com/gorilla/websocket" ) func TestGeneratedRelayProtocolVersionsMatchSpec(t *testing.T) { data, err := os.ReadFile(filepath.Join("..", "relay_protocol.json")) if err != nil { t.Fatalf("read relay protocol spec: %v", err) } var spec struct { ProtocolVersion int `json:"protocolVersion"` LegacyProtocolVersion int `json:"legacyProtocolVersion"` } if err := json.Unmarshal(data, &spec); err != nil { t.Fatalf("decode relay protocol spec: %v", err) } if relayProtocolVersion != spec.ProtocolVersion || legacyRelayProtocolVersion != spec.LegacyProtocolVersion { t.Fatalf( "generated versions=(%d,%d), spec=(%d,%d)", relayProtocolVersion, legacyRelayProtocolVersion, spec.ProtocolVersion, spec.LegacyProtocolVersion, ) } } // 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()), posters: newPosterStore(t.TempDir(), maxPosterStoreSize, posterMaxAge), posterUploads: newPosterUploadLimiter(posterPerIPRateBurst, posterPerIPRateSustained, posterGlobalRateBurst, posterGlobalRateSustained, maxConcurrentPosterUploads, time.Now()), conns: newConnTracker(), clientIPs: newClientIPResolver(nil), } s.snap = newSnapshotter(stateFile, s.buildSnapshot) return s } func mustReconnectToken(t *testing.T) (string, reconnectVerifier) { t.Helper() token, verifier, err := mintReconnectToken() if err != nil { t.Fatalf("mint reconnect token: %v", err) } return token, verifier } func makeRoomSnapshots(count int, maximumLengthIDs bool, now time.Time) []roomSnapshot { rooms := make([]roomSnapshot, 0, count) for i := range count { suffix := fmt.Sprintf("%04d", i) sessionID := "S" + suffix hostPeerID := "H" if maximumLengthIDs { sessionID = suffix + strings.Repeat("S", maxSessionIDLength-len(suffix)) hostPeerID = strings.Repeat("H", maxPeerIDLength) } rooms = append(rooms, roomSnapshot{ SessionID: sessionID, HostPeerID: hostPeerID, HostReconnectVerifier: encodeReconnectVerifier(reconnectVerifier{1}), CreatedAt: now.Add(-time.Minute), LastActivityAt: now, }) } return rooms } 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) _, verifier1 := mustReconnectToken(t) _, verifier2 := mustReconnectToken(t) s.rooms["ABC12"] = &Room{ SessionID: "ABC12", HostPeerID: "host-1", hostVerifier: verifier1, peerVerifiers: make(map[string]reconnectVerifier), Peers: map[string]*Client{}, CreatedAt: now.Add(-time.Minute), LastActivityAt: now, quotaOwnerKey: "203.0.113.44", } s.rooms["XYZ99"] = &Room{ SessionID: "XYZ99", HostPeerID: "host-2", hostVerifier: verifier2, peerVerifiers: make(map[string]reconnectVerifier), 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) } data, err := os.ReadFile(path) if err != nil { t.Fatalf("read snapshot: %v", err) } if bytes.Contains(data, []byte("203.0.113.44")) || bytes.Contains(data, []byte("quotaOwner")) { t.Fatalf("snapshot persisted process-local client identity: %s", data) } // 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 !reconnectVerifierMatches(r.hostVerifier, orig.hostVerifier) { t.Errorf("%s: host reconnect verifier did not round-trip", id) } 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.quotaOwnerKey != "" { t.Errorf("%s: quotaOwnerKey=%q after reload, want empty", id, r.quotaOwnerKey) } 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)) } } s2.conns.mu.Lock() defer s2.conns.mu.Unlock() if len(s2.conns.roomsPerIP) != 0 { t.Fatalf("reload restored process-local room quota: %v", s2.conns.roomsPerIP) } } func TestLoadSkipsExpired(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "rooms.json") now := time.Now() _, verifier := mustReconnectToken(t) encodedVerifier := encodeReconnectVerifier(verifier) snap := stateSnapshot{ Version: snapshotFormatVersion, SavedAt: now, Rooms: []roomSnapshot{ {SessionID: "FRESH", HostPeerID: "h", HostReconnectVerifier: encodedVerifier, CreatedAt: now.Add(-time.Minute), LastActivityAt: now.Add(-30 * time.Second)}, {SessionID: "OLD24", HostPeerID: "h", HostReconnectVerifier: encodedVerifier, CreatedAt: now.Add(-25 * time.Hour), LastActivityAt: now.Add(-time.Second)}, {SessionID: "IDLE6", HostPeerID: "h", HostReconnectVerifier: encodedVerifier, CreatedAt: now.Add(-2 * time.Hour), LastActivityAt: now.Add(-6 * time.Minute)}, {SessionID: "", HostPeerID: "h", HostReconnectVerifier: encodedVerifier, CreatedAt: now, LastActivityAt: now}, {SessionID: "NOHOS", HostPeerID: "", HostReconnectVerifier: encodedVerifier, 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 TestLoadRejectsSnapshotsWithoutHostAuthority(t *testing.T) { for _, version := range []int{1, 99} { t.Run(fmt.Sprintf("version_%d", version), func(t *testing.T) { path := filepath.Join(t.TempDir(), "rooms.json") data := []byte(fmt.Sprintf(`{"version":%d,"rooms":[{"sessionId":"X","hostPeerId":"H"}]}`, version)) 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) != 0 { t.Fatalf("expected empty rooms for version %d, got %d", version, 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 TestRemoveRoomLockedReleasesOwnedQuotaExactlyOnce(t *testing.T) { s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) ownerKey := "203.0.113.8" room := &Room{ SessionID: "OWNED", HostPeerID: "H", Peers: map[string]*Client{}, quotaOwnerKey: ownerKey, } s.rooms[room.SessionID] = room if !s.conns.tryCreateRoom(ownerKey) { t.Fatal("reserve room quota") } s.mu.Lock() firstRemoval := s.removeRoomLocked(room.SessionID, room) secondRemoval := s.removeRoomLocked(room.SessionID, room) s.mu.Unlock() if !firstRemoval || secondRemoval { t.Fatalf("first removal=%v second removal=%v, want true then false", firstRemoval, secondRemoval) } s.conns.mu.Lock() remaining := s.conns.roomsPerIP[ownerKey] s.conns.mu.Unlock() if remaining != 0 { t.Fatalf("quota after repeated removal=%d, want 0", remaining) } replacement := &Room{SessionID: room.SessionID, Peers: map[string]*Client{}} s.mu.Lock() s.rooms[room.SessionID] = replacement staleRemoval := s.removeRoomLocked(room.SessionID, room) authoritative := s.rooms[room.SessionID] s.mu.Unlock() if staleRemoval || authoritative != replacement { t.Fatal("stale pointer removed the authoritative replacement") } } 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 TestSnapshotWriteRejectsOversizeAndPreservesLastValidFile(t *testing.T) { path := filepath.Join(t.TempDir(), "rooms.json") s := newTestServer(t, path) now := time.Now() s.rooms["ORIG"] = &Room{ SessionID: "ORIG", HostPeerID: "H", Peers: map[string]*Client{}, CreatedAt: now, LastActivityAt: now, } if err := s.snap.write(); err != nil { t.Fatalf("write valid snapshot: %v", err) } original, err := os.ReadFile(path) if err != nil { t.Fatalf("read valid snapshot: %v", err) } s.snap.build = func() stateSnapshot { return stateSnapshot{ Version: snapshotFormatVersion, SavedAt: now, Rooms: []roomSnapshot{{ SessionID: strings.Repeat("S", snapshotMaxFileSize), HostPeerID: "H", CreatedAt: now, LastActivityAt: now, }}, } } if err := s.snap.write(); err == nil { t.Fatal("oversized snapshot write succeeded") } after, err := os.ReadFile(path) if err != nil { t.Fatalf("read preserved snapshot: %v", err) } if !bytes.Equal(after, original) { t.Fatal("oversized write replaced the last valid snapshot") } reloaded := newTestServer(t, path) if err := reloaded.loadSnapshot(path); err != nil { t.Fatalf("load preserved snapshot: %v", err) } if _, ok := reloaded.rooms["ORIG"]; !ok { t.Fatal("last valid snapshot did not survive oversized write") } } func TestSnapshotAtRetainedRoomCapFitsAndReloads(t *testing.T) { path := filepath.Join(t.TempDir(), "rooms.json") s := newTestServer(t, path) now := time.Now().UTC() for _, room := range makeRoomSnapshots(maxRetainedRooms, true, now) { s.rooms[room.SessionID] = &Room{ SessionID: room.SessionID, HostPeerID: room.HostPeerID, Peers: map[string]*Client{}, CreatedAt: room.CreatedAt, LastActivityAt: room.LastActivityAt, } } snapshot := s.buildSnapshot() data, err := json.Marshal(snapshot) if err != nil { t.Fatalf("marshal maximum snapshot: %v", err) } if len(data) > snapshotMaxFileSize { t.Fatalf("maximum admitted snapshot is %d bytes, exceeds %d", len(data), snapshotMaxFileSize) } if err := s.snap.write(); err != nil { t.Fatalf("write maximum snapshot: %v", err) } reloaded := newTestServer(t, path) if err := reloaded.loadSnapshot(path); err != nil { t.Fatalf("load maximum snapshot: %v", err) } if got := len(reloaded.rooms); got != maxRetainedRooms { t.Fatalf("reloaded rooms=%d, want %d", got, maxRetainedRooms) } for _, expected := range snapshot.Rooms { room := reloaded.rooms[expected.SessionID] if room == nil || room.Peers == nil { t.Fatalf("room %q is not available for joins after reload", expected.SessionID) } } } func TestLoadRejectsSnapshotOverRetainedRoomCapWithoutPartialState(t *testing.T) { path := filepath.Join(t.TempDir(), "rooms.json") now := time.Now().UTC() snapshot := stateSnapshot{ Version: snapshotFormatVersion, SavedAt: now, Rooms: makeRoomSnapshots(maxRetainedRooms+1, false, now), } data, err := json.Marshal(snapshot) if err != nil { t.Fatalf("marshal over-count snapshot: %v", err) } if len(data) > snapshotMaxFileSize { t.Fatalf("over-count fixture is %d bytes, must exercise count limit below %d", len(data), snapshotMaxFileSize) } if err := os.WriteFile(path, data, 0644); err != nil { t.Fatalf("write over-count snapshot: %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("over-count snapshot partially loaded %d rooms", len(s.rooms)) } s.conns.mu.Lock() defer s.conns.mu.Unlock() if len(s.conns.roomsPerIP) != 0 { t.Fatalf("over-count snapshot changed process quota state: %v", s.conns.roomsPerIP) } if _, err := os.Stat(path); err != nil { t.Fatalf("rejected snapshot was not preserved: %v", err) } } func TestSnapshotDebounceCoalesces(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "rooms.json") var ( buildCount int countMu sync.Mutex ) built := make(chan struct{}, 1) sn := newSnapshotter(path, func() stateSnapshot { countMu.Lock() buildCount++ countMu.Unlock() select { case built <- struct{}{}: default: } 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() } select { case <-built: case <-time.After(2 * time.Second): t.Fatal("timed out waiting for debounced snapshot build") } quiet := time.NewTimer(5 * snapshotDebounce) defer quiet.Stop() select { case <-built: t.Fatal("debounced burst produced an unexpected trailing snapshot build") case <-quiet.C: } countMu.Lock() got := buildCount countMu.Unlock() if got != 1 { t.Fatalf("expected 1 build from burst, got %d", got) } } // ====================================================================== // Integration harness — boots a real Server behind httptest with the full // HTTP mux. Each dial sets X-Forwarded-For so tests control the perceived // client IP independently of the rate limiters. // ====================================================================== type relayHarness struct { srv *Server httpSrv *httptest.Server wsURL string baseURL string } func mustClientIPResolver(t *testing.T, cidrs string) clientIPResolver { t.Helper() prefixes, err := parseTrustedProxyCIDRs(cidrs) if err != nil { t.Fatalf("parse trusted proxies: %v", err) } return newClientIPResolver(prefixes) } func newRelayHarness(t *testing.T) *relayHarness { t.Helper() tmpDir := t.TempDir() return newRelayHarnessAt(t, tmpDir, filepath.Join(tmpDir, "rooms.json")) } func newRelayHarnessNoTrust(t *testing.T) *relayHarness { t.Helper() tmpDir := t.TempDir() return newRelayHarnessAtWithResolver( t, tmpDir, filepath.Join(tmpDir, "rooms.json"), newClientIPResolver(nil), ) } // newRelayHarnessAt lets a test control the stateFile path so two harnesses // can share a snapshot across a simulated restart. func newRelayHarnessAt(t *testing.T, logDir, stateFile string) *relayHarness { t.Helper() return newRelayHarnessAtWithResolver(t, logDir, stateFile, mustClientIPResolver(t, "127.0.0.0/8")) } func newRelayHarnessAtWithResolver( t *testing.T, logDir, stateFile string, clientIPs clientIPResolver, ) *relayHarness { t.Helper() srv := newServer(logDir, stateFile, filepath.Join(t.TempDir(), "posters"), clientIPs) return newRelayHarnessWithServer(t, srv, true) } func newStorageHarness(t *testing.T, logs *logStore, posters *posterStore) *relayHarness { t.Helper() srv := &Server{ rooms: make(map[string]*Room), logs: logs, posters: posters, posterUploads: newPosterUploadLimiter(posterPerIPRateBurst, posterPerIPRateSustained, posterGlobalRateBurst, posterGlobalRateSustained, maxConcurrentPosterUploads, time.Now()), logLookups: make(chan struct{}, maxConcurrentLogLookups), conns: newConnTracker(), clientIPs: mustClientIPResolver(t, "127.0.0.0/8"), } return newRelayHarnessWithServer(t, srv, false) } func newRelayHarnessWithServer(t *testing.T, srv *Server, stopSnapshot bool) *relayHarness { t.Helper() mux := http.NewServeMux() mux.HandleFunc("/relay", srv.handleWS) mux.HandleFunc("/logs", srv.handlePostLogs) mux.HandleFunc("/logs/", srv.handleGetLogs) mux.HandleFunc("/posters", srv.handlePostPosters) mux.HandleFunc("/posters/", srv.handleGetPosters) httpSrv := httptest.NewServer(mux) t.Cleanup(func() { httpSrv.Close() if stopSnapshot { _ = srv.snap.flushAndStop(time.Second) } }) u, _ := url.Parse(httpSrv.URL) wsURL := "ws://" + u.Host + "/relay" return &relayHarness{srv: srv, httpSrv: httpSrv, wsURL: wsURL, baseURL: httpSrv.URL} } func createModernRoomWithGuest( t *testing.T, h *relayHarness, sessionID, hostIP, guestIP string, ) (host, guest *testConn, hostToken, guestToken string) { t.Helper() hostToken, _ = mustReconnectToken(t) host = h.dial(t, hostIP) host.send(clientMsg{ Type: relayTypeCreate, SessionID: sessionID, PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") guestToken, _ = mustReconnectToken(t) guest = h.dial(t, guestIP) guest.send(clientMsg{ Type: relayTypeJoin, SessionID: sessionID, PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) return host, guest, hostToken, guestToken } func injectSnapshotPersistenceFailure(t *testing.T, sn *snapshotter, injected error) { t.Helper() if err := sn.write(); err != nil { t.Fatalf("persist pre-failure baseline: %v", err) } sn.writeMu.Lock() original := sn.persist sn.persist = func([]byte) error { return injected } sn.writeMu.Unlock() t.Cleanup(func() { sn.writeMu.Lock() sn.persist = original sn.writeMu.Unlock() }) } func copySnapshotForRestart(t *testing.T, source string) string { t.Helper() data, err := os.ReadFile(source) if err != nil { t.Fatalf("read crash-window snapshot: %v", err) } path := filepath.Join(t.TempDir(), "rooms.json") if err := os.WriteFile(path, data, 0644); err != nil { t.Fatalf("copy crash-window snapshot: %v", err) } return path } type deterministicRemover struct { mu sync.Mutex failures map[string]error calls map[string]int } func newDeterministicRemover() *deterministicRemover { return &deterministicRemover{ failures: make(map[string]error), calls: make(map[string]int), } } func (r *deterministicRemover) remove(path string) error { r.mu.Lock() r.calls[path]++ err := r.failures[path] r.mu.Unlock() if err != nil { return err } return os.Remove(path) } func (r *deterministicRemover) fail(path string, err error) { r.mu.Lock() defer r.mu.Unlock() r.failures[path] = err } func (r *deterministicRemover) recover(path string) { r.mu.Lock() defer r.mu.Unlock() delete(r.failures, path) } func (r *deterministicRemover) callCount(path string) int { r.mu.Lock() defer r.mu.Unlock() return r.calls[path] } func (h *relayHarness) dial(t *testing.T, ip string) *testConn { t.Helper() headers := http.Header{} if ip != "" { headers.Set("X-Forwarded-For", ip) } conn, _, err := websocket.DefaultDialer.Dial(h.wsURL, headers) if err != nil { t.Fatalf("dial (ip=%s): %v", ip, err) } tc := &testConn{t: t, conn: conn} t.Cleanup(func() { conn.Close() }) return tc } func (h *relayHarness) dialRaw(ip string) (*websocket.Conn, error) { headers := http.Header{} if ip != "" { headers.Set("X-Forwarded-For", ip) } conn, _, err := websocket.DefaultDialer.Dial(h.wsURL, headers) return conn, err } func newWebSocketPair(t *testing.T) (*websocket.Conn, *websocket.Conn) { t.Helper() serverConnCh := make(chan *websocket.Conn, 1) upgradeErrCh := make(chan error, 1) httpServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := upgrader.Upgrade(w, r, nil) if err != nil { upgradeErrCh <- err return } serverConnCh <- conn })) u, err := url.Parse(httpServer.URL) if err != nil { httpServer.Close() t.Fatalf("parse websocket pair URL: %v", err) } peerConn, _, err := websocket.DefaultDialer.Dial("ws://"+u.Host, nil) if err != nil { httpServer.Close() t.Fatalf("dial websocket pair: %v", err) } var serverConn *websocket.Conn select { case serverConn = <-serverConnCh: case err := <-upgradeErrCh: peerConn.Close() httpServer.Close() t.Fatalf("upgrade websocket pair: %v", err) case <-time.After(2 * time.Second): peerConn.Close() httpServer.Close() t.Fatal("timed out waiting for websocket pair upgrade") } t.Cleanup(func() { serverConn.Close() peerConn.Close() httpServer.Close() }) return serverConn, peerConn } func (h *relayHarness) dialWithHeaders(headers http.Header) (*websocket.Conn, *http.Response, error) { return websocket.DefaultDialer.Dial(h.wsURL, headers) } func (h *relayHarness) waitRoomPeers(t *testing.T, sessionID string, want int) { t.Helper() deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { h.srv.mu.RLock() room := h.srv.rooms[sessionID] h.srv.mu.RUnlock() if room != nil { room.mu.RLock() got := len(room.Peers) room.mu.RUnlock() if got == want { return } } time.Sleep(20 * time.Millisecond) } t.Fatalf("room %s never reached %d peers within 2s", sessionID, want) } func (h *relayHarness) waitIPConnections(t *testing.T, ip string, want int) { t.Helper() deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { h.srv.conns.mu.Lock() got := h.srv.conns.perIP[ip] h.srv.conns.mu.Unlock() if got == want { return } time.Sleep(20 * time.Millisecond) } t.Fatalf("IP %s never reached %d connections within 2s", ip, want) } type testConn struct { t *testing.T conn *websocket.Conn } func (c *testConn) send(msg clientMsg) { c.t.Helper() data, err := json.Marshal(msg) if err != nil { c.t.Fatalf("marshal: %v", err) } if err := c.conn.WriteMessage(websocket.TextMessage, data); err != nil { c.t.Fatalf("write: %v", err) } } func (c *testConn) sendRaw(data []byte) { c.t.Helper() if err := c.conn.WriteMessage(websocket.TextMessage, data); err != nil { c.t.Fatalf("write raw: %v", err) } } func (c *testConn) recv() serverMsg { c.t.Helper() c.conn.SetReadDeadline(time.Now().Add(2 * time.Second)) _, data, err := c.conn.ReadMessage() if err != nil { c.t.Fatalf("read: %v", err) } var m serverMsg if err := json.Unmarshal(data, &m); err != nil { c.t.Fatalf("unmarshal %q: %v", data, err) } return m } func (c *testConn) expect(typ string) serverMsg { c.t.Helper() m := c.recv() if m.Type != typ { c.t.Fatalf("expected type=%s, got type=%s code=%s message=%s", typ, m.Type, m.Code, m.Message) } return m } func (c *testConn) expectError(code string) serverMsg { c.t.Helper() m := c.expect("error") if m.Code != code { c.t.Fatalf("expected code=%s, got code=%s message=%s", code, m.Code, m.Message) } return m } func (c *testConn) expectAuthority(typ, hostPeerID string) serverMsg { c.t.Helper() message := c.expect(typ) if message.HostPeerID != hostPeerID { c.t.Fatalf("%s hostPeerId=%q, want %q", typ, message.HostPeerID, hostPeerID) } if _, ok := reconnectVerifierFromToken(message.ReconnectToken); !ok { c.t.Fatalf("%s reconnectToken has invalid shape", typ) } return message } // recvNothing asserts no message arrives within the given window. Used to // verify silent paths (sender not receiving own broadcast, stale-peer skip). func (c *testConn) recvNothing(within time.Duration) { c.t.Helper() c.conn.SetReadDeadline(time.Now().Add(within)) _, data, err := c.conn.ReadMessage() if err == nil { c.t.Fatalf("expected no message within %v, got %s", within, data) } if ne, ok := err.(net.Error); !ok || !ne.Timeout() { c.t.Fatalf("expected read timeout, got %v", err) } } // recvUntilClosed consumes any frames already queued on the wire and requires // a permanent terminal read error before the absolute deadline. func (c *testConn) recvUntilClosed(within time.Duration) ([]serverMsg, error) { c.t.Helper() if err := c.conn.SetReadDeadline(time.Now().Add(within)); err != nil { return nil, fmt.Errorf("set close-read deadline: %w", err) } var messages []serverMsg for { messageType, data, err := c.conn.ReadMessage() if err != nil { if ne, ok := err.(net.Error); ok && ne.Timeout() { return messages, fmt.Errorf("terminal closure not observed within %v: %w", within, err) } return messages, nil } if messageType != websocket.TextMessage { return messages, fmt.Errorf("unexpected websocket message type %d before closure", messageType) } var message serverMsg if err := json.Unmarshal(data, &message); err != nil { return messages, fmt.Errorf("decode frame before closure %q: %w", data, err) } messages = append(messages, message) } } func requireClientClosed(t *testing.T, client *Client) { t.Helper() select { case <-client.done: case <-time.After(2 * time.Second): t.Fatal("client did not close") } } func requirePeerClosed(t *testing.T, peer *websocket.Conn) { t.Helper() testPeer := &testConn{t: t, conn: peer} if messages, err := testPeer.recvUntilClosed(2 * time.Second); err != nil { t.Fatalf("peer remained open after client failure (messages=%v): %v", messages, err) } } func TestClientQueueOverflowClosesConnection(t *testing.T) { serverConn, peerConn := newWebSocketPair(t) client := &Client{ conn: serverConn, send: make(chan outboundFrame, 1), done: make(chan struct{}), } if !client.enqueue([]byte(`{"sequence":1}`)) { t.Fatal("first frame was not accepted") } if client.enqueue([]byte(`{"sequence":2}`)) { t.Fatal("overflowing frame was accepted") } requireClientClosed(t, client) requirePeerClosed(t, peerConn) if client.enqueue([]byte(`{"sequence":3}`)) { t.Fatal("frame was accepted after terminal close") } client.close() } func TestClientQueueOverflowBroadcastKeepsHealthyRecipient(t *testing.T) { slowServerConn, slowPeerConn := newWebSocketPair(t) slow := &Client{ conn: slowServerConn, send: make(chan outboundFrame, 1), done: make(chan struct{}), } if !slow.enqueue([]byte(`{"sequence":1}`)) { t.Fatal("failed to prime slow client queue") } healthyServerConn, healthyPeerConn := newWebSocketPair(t) healthy := newClient(healthyServerConn) t.Cleanup(healthy.close) room := &Room{ Peers: map[string]*Client{ "slow": slow, "healthy": healthy, }, } payload := json.RawMessage(`{"sequence":2}`) room.broadcastExcept("sender", serverMsg{ Type: relayTypeMessage, From: "sender", Payload: payload, }) requireClientClosed(t, slow) requirePeerClosed(t, slowPeerConn) received := (&testConn{t: t, conn: healthyPeerConn}).expect(relayTypeMessage) if received.From != "sender" { t.Fatalf("healthy recipient sender=%q, want sender", received.From) } if string(received.Payload) != string(payload) { t.Fatalf("healthy recipient payload=%s, want %s", received.Payload, payload) } } func TestClientQueueOverflowDirectedTargetStillExists(t *testing.T) { serverConn, peerConn := newWebSocketPair(t) target := &Client{ conn: serverConn, send: make(chan outboundFrame, 1), done: make(chan struct{}), } if !target.enqueue([]byte(`{"sequence":1}`)) { t.Fatal("failed to prime directed target queue") } sender := &Client{} room := &Room{Peers: map[string]*Client{"sender": sender, "target": target}} if result := room.sendFrom("sender", sender, "target", serverMsg{ Type: relayTypeMessage, From: "sender", Payload: json.RawMessage(`{"sequence":2}`), }); result != directedTargetFound { t.Fatalf("full existing target result=%v, want directedTargetFound", result) } requireClientClosed(t, target) requirePeerClosed(t, peerConn) if result := room.sendFrom("sender", sender, "missing", serverMsg{Type: relayTypeMessage}); result != directedTargetMissing { t.Fatalf("missing target result=%v, want directedTargetMissing", result) } } func TestClientWriteFailureClosesConnection(t *testing.T) { serverConn, _ := newWebSocketPair(t) client := &Client{ conn: serverConn, send: make(chan outboundFrame, 1), done: make(chan struct{}), } if err := serverConn.Close(); err != nil { t.Fatalf("close writer connection: %v", err) } client.send <- outboundFrame{data: []byte(`{"type":"queued"}`)} exited := make(chan struct{}) go func() { client.writePump() close(exited) }() requireClientClosed(t, client) select { case <-exited: case <-time.After(2 * time.Second): t.Fatal("write pump did not exit after write failure") } client.close() } // ====================================================================== // Unit tests — pure logic // ====================================================================== func TestRateLimiterBurstExhausts(t *testing.T) { rl := newRateLimiter(5, 10) for i := 0; i < 5; i++ { if !rl.allow() { t.Fatalf("allow %d: expected true", i) } } if rl.allow() { t.Fatal("allow 6: expected false (burst exhausted)") } } func TestRateLimiterRefillsOverTime(t *testing.T) { rl := newRateLimiter(5, 10) // 10 tokens/sec for i := 0; i < 5; i++ { rl.allow() } if rl.allow() { t.Fatal("burst should be exhausted before sleep") } time.Sleep(1200 * time.Millisecond) count := 0 for rl.allow() { count++ } if count < 1 { t.Fatalf("expected at least 1 token after 1.2s refill, got %d", count) } if count > 5 { t.Fatalf("expected at most burst=5 after refill, got %d", count) } } func TestRateLimiterAllowRace(t *testing.T) { rl := newRateLimiter(100, 1000) var wg sync.WaitGroup var successes atomic.Int64 for i := 0; i < 10; i++ { wg.Add(1) go func() { defer wg.Done() for j := 0; j < 50; j++ { if rl.allow() { successes.Add(1) } } }() } wg.Wait() // Real assertion is that -race finds no data race. Spot-check the // result is within plausible bounds. if got := successes.Load(); got <= 0 || got > 500 { t.Fatalf("unexpected successes count %d (want 1..500)", got) } } func TestRateLimiterReclaimableOnlyAfterFullRefill(t *testing.T) { now := time.Now() limiter := &rateLimiter{ tokens: 0, maxTokens: 5, refillRate: 1, lastTime: now, } if limiter.reclaimable(now.Add(4 * time.Second)) { t.Fatal("partially refilled limiter must retain its effective state") } if !limiter.reclaimable(now.Add(5 * time.Second)) { t.Fatal("fully refilled limiter should be reclaimable") } } func TestCleanupRateWindowsUsesWindowBoundary(t *testing.T) { now := time.Now() windows := map[string]time.Time{ "active": now.Add(-logRateInterval + time.Nanosecond), "expired": now.Add(-logRateInterval), } cleanupRateWindows(windows, now, logRateInterval) if _, ok := windows["active"]; !ok { t.Fatal("active fixed-window limiter was removed early") } if _, ok := windows["expired"]; ok { t.Fatal("expired fixed-window limiter was retained") } } // ====================================================================== // connTracker unit tests // ====================================================================== func TestConnTrackerPerIPLimit(t *testing.T) { ct := newConnTracker() for i := 0; i < maxConnsPerIP; i++ { if !ct.tryConnect("10.0.0.1") { t.Fatalf("tryConnect %d: expected true", i) } } if ct.tryConnect("10.0.0.1") { t.Fatalf("tryConnect %d from same IP: expected false", maxConnsPerIP+1) } } func TestConnTrackerGlobalLimit(t *testing.T) { ct := newConnTracker() for i := 0; i < maxGlobalConns; i++ { ip := fmt.Sprintf("10.0.%d.%d", i/256, i%256) if !ct.tryConnect(ip) { t.Fatalf("tryConnect %d (ip=%s): expected true", i, ip) } } if ct.tryConnect("10.99.99.99") { t.Fatal("tryConnect should fail once globalCount hits max") } } func TestConnTrackerDisconnectFrees(t *testing.T) { ct := newConnTracker() ip := "10.0.0.2" for i := 0; i < 5; i++ { ct.tryConnect(ip) } for i := 0; i < 5; i++ { ct.disconnect(ip) } ct.mu.Lock() if _, ok := ct.perIP[ip]; ok { t.Error("perIP entry should be deleted when count reaches 0") } if ct.globalCount != 0 { t.Errorf("globalCount=%d, want 0", ct.globalCount) } ct.mu.Unlock() // Extra disconnect is a no-op (doesn't panic). ct.disconnect(ip) } func TestConnTrackerRoomQuota(t *testing.T) { ct := newConnTracker() ip := "10.0.0.3" for i := 0; i < maxRoomsPerIP; i++ { if !ct.tryCreateRoom(ip) { t.Fatalf("tryCreateRoom %d: expected true", i) } } if ct.tryCreateRoom(ip) { t.Fatalf("tryCreateRoom %d: expected false (quota)", maxRoomsPerIP+1) } ct.releaseRoom(ip) if !ct.tryCreateRoom(ip) { t.Fatal("tryCreateRoom after release: expected true") } } func TestConnTrackerCleanupPreservesEffectiveRateLimits(t *testing.T) { ct := newConnTracker() ip := "10.0.1.1" for i := range connRateBurst { if !ct.tryConnect(ip) { t.Fatalf("tryConnect %d: expected true", i) } } for range connRateBurst { ct.disconnect(ip) } ct.cleanup(time.Now()) if ct.tryConnect(ip) { t.Fatal("cleanup reset a connection rate limit that was still effective") } ct.cleanup(time.Now().Add(10 * time.Second)) if !ct.tryConnect(ip) { t.Fatal("fully refilled limiter should be reclaimable") } ct.cleanup(time.Now().Add(10 * time.Second)) ct.mu.Lock() _, retainedWhileConnected := ct.ipRate[ip] ct.mu.Unlock() if !retainedWhileConnected { t.Fatal("cleanup removed a limiter with an active connection") } } func TestConnTrackerConnectRateLimit(t *testing.T) { ct := newConnTracker() ip := "10.0.0.4" for i := range connRateBurst { if !ct.tryConnect(ip) { t.Fatalf("warmup tryConnect %d: expected true", i) } } // Free one slot so the perIP check won't be what rejects us. ct.disconnect(ip) // Rate-limit bucket is empty now; this should be the denial path. if ct.tryConnect(ip) { t.Fatal("expected false from rate-limit bucket, not per-IP cap") } } func TestPosterUploadLimiterAdmissionPolicy(t *testing.T) { now := time.Unix(1_700_000_000, 0) t.Run("per IP burst and independent clients", func(t *testing.T) { limiter := newPosterUploadLimiter(2, 1, 10, 1, 10, now) for range 2 { if !limiter.tryStart("203.0.113.1", now) { t.Fatal("per-IP burst rejected early") } limiter.finish() } if limiter.tryStart("203.0.113.1", now) { t.Fatal("request beyond per-IP burst succeeded") } if !limiter.tryStart("203.0.113.2", now) { t.Fatal("independent IP was denied") } limiter.finish() }) t.Run("global burst spans distinct clients", func(t *testing.T) { limiter := newPosterUploadLimiter(10, 1, 2, 1, 10, now) for _, ip := range []string{"203.0.113.1", "203.0.113.2"} { if !limiter.tryStart(ip, now) { t.Fatalf("%s rejected before global burst exhausted", ip) } limiter.finish() } if limiter.tryStart("203.0.113.3", now) { t.Fatal("request beyond global burst succeeded") } if len(limiter.perIP) != 2 { t.Fatalf("globally denied request allocated per-IP state: %d buckets", len(limiter.perIP)) } }) t.Run("concurrency denial consumes no tokens", func(t *testing.T) { limiter := newPosterUploadLimiter(1, 0, 2, 0, 1, now) if !limiter.tryStart("203.0.113.1", now) { t.Fatal("first upload denied") } if limiter.tryStart("203.0.113.2", now) { t.Fatal("upload above concurrency limit succeeded") } limiter.finish() if !limiter.tryStart("203.0.113.2", now) { t.Fatal("concurrency denial consumed admission tokens") } limiter.finish() }) t.Run("per IP denial refunds global token", func(t *testing.T) { limiter := newPosterUploadLimiter(1, 0, 2, 0, 2, now) if !limiter.tryStart("203.0.113.1", now) { t.Fatal("first upload denied") } limiter.finish() if limiter.tryStart("203.0.113.1", now) { t.Fatal("exhausted IP unexpectedly admitted") } if !limiter.tryStart("203.0.113.2", now) { t.Fatal("refunded global token was unavailable to another IP") } limiter.finish() }) t.Run("finish restores only concurrency and time restores rate", func(t *testing.T) { limiter := newPosterUploadLimiter(1, 1, 1, 1, 1, now) if !limiter.tryStart("203.0.113.1", now) { t.Fatal("first upload denied") } limiter.finish() if limiter.active != 0 { t.Fatalf("active=%d, want 0", limiter.active) } if limiter.tryStart("203.0.113.1", now) { t.Fatal("finish incorrectly refunded rate tokens") } if !limiter.tryStart("203.0.113.1", now.Add(time.Second)) { t.Fatal("sustained refill did not restore capacity") } limiter.finish() }) t.Run("cleanup retains effective buckets then reclaims full ones", func(t *testing.T) { limiter := newPosterUploadLimiter(2, 1, 10, 1, 2, now) if !limiter.tryStart("203.0.113.1", now) { t.Fatal("first upload denied") } limiter.finish() limiter.cleanup(now) if _, ok := limiter.perIP["203.0.113.1"]; !ok { t.Fatal("cleanup removed effective per-IP limiter") } limiter.cleanup(now.Add(time.Second)) if _, ok := limiter.perIP["203.0.113.1"]; ok { t.Fatal("cleanup retained fully refilled per-IP limiter") } }) } // ====================================================================== // clientIPResolver unit tests // ====================================================================== func TestClientIPResolverTrustChains(t *testing.T) { tests := []struct { name string trusted string remote string headers []string want string wantErr bool }{ {name: "absent forwarding header", remote: "127.0.0.1:12345", want: "127.0.0.1"}, {name: "untrusted peer ignores spoof", remote: "198.51.100.10:12345", headers: []string{"203.0.113.5"}, want: "198.51.100.10"}, {name: "one trusted proxy", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"203.0.113.5"}, want: "203.0.113.5"}, {name: "append chain ignores forged leftmost", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"198.51.100.99, 203.0.113.5"}, want: "203.0.113.5"}, {name: "two trusted proxies", trusted: "10.0.0.0/8, 192.0.2.0/24", remote: "10.0.0.2:8080", headers: []string{"203.0.113.5, 192.0.2.10"}, want: "203.0.113.5"}, {name: "untrusted intermediate is client boundary", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"203.0.113.5, 198.51.100.7"}, want: "198.51.100.7"}, {name: "repeated header lines preserve chain", trusted: "10.0.0.0/8, 192.0.2.0/24", remote: "10.0.0.2:8080", headers: []string{"203.0.113.5", "192.0.2.10"}, want: "203.0.113.5"}, {name: "trusted proxy without forwarding header", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", want: "10.0.0.2"}, {name: "IPv4 mapped peer is unmapped", remote: "[::ffff:192.0.2.4]:8080", want: "192.0.2.4"}, {name: "IPv4 mapped forwarded address is unmapped", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"::ffff:203.0.113.5"}, want: "203.0.113.5"}, {name: "native IPv6 client is grouped to 64", trusted: "10.0.0.0/8", remote: "10.0.0.2:8080", headers: []string{"2001:db8:85a3:12::abcd"}, want: "2001:db8:85a3:12::"}, {name: "untrusted malformed header is ignored", remote: "198.51.100.10:12345", headers: []string{"bad,,host:123"}, want: "198.51.100.10"}, {name: "malformed immediate peer", remote: "not-an-address", wantErr: true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { resolver := mustClientIPResolver(t, tt.trusted) req := &http.Request{RemoteAddr: tt.remote, Header: make(http.Header)} for _, value := range tt.headers { req.Header.Add("X-Forwarded-For", value) } got, err := resolver.resolve(req) if tt.wantErr { if err == nil { t.Fatalf("resolve()=%q, want error", got) } return } if err != nil { t.Fatalf("resolve(): %v", err) } if got != tt.want { t.Fatalf("resolve()=%q, want %q", got, tt.want) } }) } } func TestClientIPResolverRejectsMalformedTrustedChains(t *testing.T) { resolver := mustClientIPResolver(t, "10.0.0.0/8") for _, value := range []string{ "", "203.0.113.5,", "203.0.113.5,,192.0.2.1", "not-an-ip", "203.0.113.5:1234", "fe80::1%eth0", } { t.Run(fmt.Sprintf("%q", value), func(t *testing.T) { req := &http.Request{ RemoteAddr: "10.0.0.2:8080", Header: http.Header{"X-Forwarded-For": []string{value}}, } if got, err := resolver.resolve(req); err == nil { t.Fatalf("resolve()=%q, want error", got) } }) } } func TestParseTrustedProxyCIDRs(t *testing.T) { prefixes, err := parseTrustedProxyCIDRs(" ") if err != nil || len(prefixes) != 0 { t.Fatalf("empty config = %v, %v; want no prefixes", prefixes, err) } prefixes, err = parseTrustedProxyCIDRs(" 10.1.2.3/8, ::ffff:192.0.2.12/120, 2001:db8::1/32 ") if err != nil { t.Fatalf("valid config: %v", err) } got := make([]string, len(prefixes)) for i, prefix := range prefixes { got[i] = prefix.String() } want := []string{"10.0.0.0/8", "192.0.2.0/24", "2001:db8::/32"} if fmt.Sprint(got) != fmt.Sprint(want) { t.Fatalf("prefixes=%v, want %v", got, want) } for _, value := range []string{"10.0.0.0/8,", "10.0.0.0/8,garbage", "::ffff:192.0.2.1/64"} { if prefixes, err := parseTrustedProxyCIDRs(value); err == nil || prefixes != nil { t.Fatalf("parseTrustedProxyCIDRs(%q)=(%v, %v), want nil error result", value, prefixes, err) } } } // ====================================================================== // generateLogID // ====================================================================== func TestGenerateLogIDShape(t *testing.T) { if entropyBits := float64(logIDLength) * math.Log2(float64(len(idChars))); entropyBits < 128 { t.Fatalf("log capability entropy=%f bits, want at least 128", entropyBits) } seen := map[string]struct{}{} for i := 0; i < 200; i++ { id := generateLogID() if len(id) != logIDLength { t.Fatalf("len=%d want %d (id=%q)", len(id), logIDLength, id) } for _, c := range id { if !strings.ContainsRune(idChars, c) { t.Fatalf("id %q has unexpected char %q", id, c) } } if _, dup := seen[id]; dup { t.Fatalf("duplicate id %q after %d calls", id, i) } seen[id] = struct{}{} } } // ====================================================================== // handleWS — create case // ====================================================================== func TestCreateSucceeds(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "1.1.1.1") c.send(clientMsg{Type: relayTypeCreate, SessionID: "ROOM1", PeerID: "host-a"}) message := c.expectAuthority(relayTypeCreated, "host-a") if message.SessionID != "ROOM1" { t.Errorf("SessionID=%q want ROOM1", message.SessionID) } h.waitRoomPeers(t, "ROOM1", 1) } func TestCreateMissingSessionIDRejected(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "1.1.1.2") c.send(clientMsg{Type: "create", PeerID: "host-a"}) c.expectError("invalid_message") } func TestCreateMissingPeerIDRejected(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "1.1.1.3") c.send(clientMsg{Type: "create", SessionID: "ROOM1"}) c.expectError("invalid_message") } func TestCreateDuplicateReturnsRoomExists(t *testing.T) { h := newRelayHarness(t) c1 := h.dial(t, "1.1.1.4") c1.send(clientMsg{Type: "create", SessionID: "SAME", PeerID: "host-1"}) c1.expect("created") // Different IP to avoid the per-IP rooms quota interfering. c2 := h.dial(t, "1.1.1.5") c2.send(clientMsg{Type: "create", SessionID: "SAME", PeerID: "host-2"}) c2.expectError("room_exists") } func TestCreateNegotiatesModernProtocolWithClientKnownToken(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) host := h.dial(t, "1.1.1.40") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "MODERN_CREATE", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion + 1, }) mismatch := host.expectError(relayErrorProtocolMismatch) if mismatch.ProtocolVersion != relayProtocolVersion { t.Fatalf("protocol mismatch advertised version=%d, want %d", mismatch.ProtocolVersion, relayProtocolVersion) } host.send(clientMsg{ Type: relayTypeCreate, SessionID: "MODERN_CREATE", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: mismatch.ProtocolVersion, }) created := host.expectAuthority(relayTypeCreated, "H") if created.ReconnectToken != hostToken { t.Fatal("modern create rotated the client-known reconnect token") } if created.ProtocolVersion != relayProtocolVersion { t.Fatalf("created protocolVersion=%d, want %d", created.ProtocolVersion, relayProtocolVersion) } } func TestModernCreateRetryAfterLostSetupResponseIsIdempotent(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) create := clientMsg{ Type: relayTypeCreate, SessionID: "CREATE_RETRY", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, } first := h.dial(t, "1.1.1.41") first.send(create) if err := first.conn.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil { t.Fatalf("set discarded response deadline: %v", err) } messageType, _, err := first.conn.ReadMessage() if err != nil { t.Fatalf("read discarded setup response: %v", err) } if messageType != websocket.TextMessage { t.Fatalf("discarded setup response type=%d, want text", messageType) } retry := h.dial(t, "1.1.1.42") retry.send(create) created := retry.expectAuthority(relayTypeCreated, "H") if created.ReconnectToken != hostToken || created.ProtocolVersion != relayProtocolVersion { t.Fatalf("retry authority changed: tokenMatch=%v protocol=%d", created.ReconnectToken == hostToken, created.ProtocolVersion) } if len(created.Peers) != 0 { t.Fatalf("retry reported unexpected peers: %v", created.Peers) } if messages, err := first.recvUntilClosed(2 * time.Second); err != nil { t.Fatalf("superseded create connection remained open: %v (frames=%v)", err, messages) } h.srv.mu.RLock() roomCount := len(h.srv.rooms) room := h.srv.rooms["CREATE_RETRY"] h.srv.mu.RUnlock() if roomCount != 1 || room == nil { t.Fatalf("idempotent retry retained rooms=%d targetPresent=%v", roomCount, room != nil) } } func TestIdempotentCreateReannouncesPreviouslyAbsentHost(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) create := clientMsg{ Type: relayTypeCreate, SessionID: "CREATE_REANNOUNCE", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, } host := h.dial(t, "1.1.1.43") host.send(create) host.expectAuthority(relayTypeCreated, "H") guestToken, _ := mustReconnectToken(t) guest := h.dial(t, "1.1.1.44") guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "CREATE_REANNOUNCE", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) if err := host.conn.Close(); err != nil { t.Fatalf("close original host: %v", err) } left := guest.expect(relayTypePeerLeft) if left.PeerID != "H" { t.Fatalf("disconnected host event peerId=%q, want H", left.PeerID) } returning := h.dial(t, "1.1.1.45") returning.send(create) returning.expectAuthority(relayTypeCreated, "H") reannounced := guest.expect(relayTypePeerJoined) if reannounced.PeerID != "H" { t.Fatalf("returning host event peerId=%q, want H", reannounced.PeerID) } } func TestCreateCannotReclaimReservedEmptyRoom(t *testing.T) { h := newRelayHarness(t) hostToken, hostVerifier := mustReconnectToken(t) original := &Room{ SessionID: "STALE", HostPeerID: "old-host", hostVerifier: hostVerifier, peerVerifiers: make(map[string]reconnectVerifier), Peers: map[string]*Client{}, CreatedAt: time.Now().Add(-time.Minute), LastActivityAt: time.Now(), } h.srv.mu.Lock() h.srv.rooms["STALE"] = original h.srv.mu.Unlock() creator := h.dial(t, "1.1.1.6") creator.send(clientMsg{Type: relayTypeCreate, SessionID: "STALE", PeerID: "new-host"}) creator.expectError(relayErrorRoomExists) unproved := h.dial(t, "1.1.1.60") unproved.send(clientMsg{Type: relayTypeJoin, SessionID: "STALE", PeerID: "old-host"}) unproved.expectError(relayErrorPeerIdUnavailable) reconnected := h.dial(t, "1.1.1.61") reconnected.send(clientMsg{ Type: relayTypeJoin, SessionID: "STALE", PeerID: "old-host", ReconnectToken: hostToken, }) reconnected.expectAuthority(relayTypeJoined, "old-host") h.srv.mu.RLock() current := h.srv.rooms["STALE"] h.srv.mu.RUnlock() if current != original { t.Fatal("reserved room identity was replaced") } } func TestCreateReclaimsOwnedEmptyRoomWithoutDoubleCharging(t *testing.T) { h := newRelayHarness(t) ownerKey := "1.1.1.61" now := time.Now() replacementToken, replacementVerifier := mustReconnectToken(t) h.srv.mu.Lock() for i := range maxRoomsPerIP { sessionID := fmt.Sprintf("OWNED%d", i) h.srv.rooms[sessionID] = &Room{ SessionID: sessionID, HostPeerID: "old-host", hostVerifier: replacementVerifier, peerVerifiers: make(map[string]reconnectVerifier), Peers: map[string]*Client{}, quotaOwnerKey: ownerKey, CreatedAt: now.Add(-time.Hour), LastActivityAt: now, } if !h.srv.conns.tryCreateRoom(ownerKey) { h.srv.mu.Unlock() t.Fatal("failed to seed retained quota") } } h.srv.mu.Unlock() c := h.dial(t, ownerKey) c.send(clientMsg{ Type: relayTypeCreate, SessionID: "OWNED0", PeerID: "old-host", ReconnectToken: replacementToken, }) c.expect(relayTypeCreated) h.srv.conns.mu.Lock() quota := h.srv.conns.roomsPerIP[ownerKey] h.srv.conns.mu.Unlock() if quota != maxRoomsPerIP { t.Fatalf("replacement quota=%d, want %d", quota, maxRoomsPerIP) } h.srv.mu.RLock() roomCount := len(h.srv.rooms) h.srv.mu.RUnlock() if roomCount != maxRoomsPerIP { t.Fatalf("replacement room count=%d, want %d", roomCount, maxRoomsPerIP) } } func TestCreateHitsRoomsPerIPLimit(t *testing.T) { h := newRelayHarness(t) ip := "1.1.1.7" for i := 0; i < maxRoomsPerIP; i++ { c := h.dial(t, ip) c.send(clientMsg{Type: "create", SessionID: fmt.Sprintf("R%d", i), PeerID: "host"}) c.expect("created") } // 4th create from same IP exceeds the quota. c := h.dial(t, ip) c.send(clientMsg{Type: "create", SessionID: "ROVERFLOW", PeerID: "host"}) c.expectError("rate_limited") } func TestRetainedRoomQuotaSurvivesDisconnectAndReturnsOnRemoval(t *testing.T) { h := newRelayHarness(t) ownerKey := "1.1.1.70" sessionIDs := make([]string, 0, maxRoomsPerIP) for i := range maxRoomsPerIP { sessionID := fmt.Sprintf("RETAIN%d", i) host := h.dial(t, ownerKey) host.send(clientMsg{Type: relayTypeCreate, SessionID: sessionID, PeerID: "H"}) host.expect(relayTypeCreated) host.conn.Close() h.waitRoomPeers(t, sessionID, 0) sessionIDs = append(sessionIDs, sessionID) } for i, sessionID := range sessionIDs { guest := h.dial(t, fmt.Sprintf("1.1.2.%d", i+1)) guest.send(clientMsg{Type: relayTypeJoin, SessionID: sessionID, PeerID: "G"}) guest.expect(relayTypeJoined) guest.conn.Close() h.waitRoomPeers(t, sessionID, 0) } blocked := h.dial(t, ownerKey) blocked.send(clientMsg{Type: relayTypeCreate, SessionID: "RETAINX", PeerID: "H"}) blocked.expectError(relayErrorRateLimited) h.srv.conns.mu.Lock() retainedQuota := h.srv.conns.roomsPerIP[ownerKey] h.srv.conns.mu.Unlock() if retainedQuota != maxRoomsPerIP { t.Fatalf("retained quota=%d, want %d after creator disconnects", retainedQuota, maxRoomsPerIP) } now := time.Now() h.srv.mu.RLock() idleRoom := h.srv.rooms[sessionIDs[0]] h.srv.mu.RUnlock() idleRoom.mu.Lock() idleRoom.LastActivityAt = now.Add(-emptyRoomMaxAge - time.Second) idleRoom.mu.Unlock() h.srv.runCleanupStep(now) h.srv.conns.mu.Lock() afterIdleRemoval := h.srv.conns.roomsPerIP[ownerKey] h.srv.conns.mu.Unlock() if afterIdleRemoval != maxRoomsPerIP-1 { t.Fatalf("quota after idle removal=%d, want %d", afterIdleRemoval, maxRoomsPerIP-1) } blocked.send(clientMsg{Type: relayTypeCreate, SessionID: "RETAIN3", PeerID: "H"}) blocked.expect(relayTypeCreated) occupied := h.dial(t, "1.1.2.99") occupied.send(clientMsg{Type: relayTypeJoin, SessionID: sessionIDs[1], PeerID: "G"}) occupied.expect(relayTypeJoined) h.srv.mu.RLock() expiringRoom := h.srv.rooms[sessionIDs[1]] h.srv.mu.RUnlock() expiringRoom.mu.Lock() expiringRoom.CreatedAt = now.Add(-roomMaxAge - time.Second) expiringRoom.mu.Unlock() h.srv.runCleanupStep(now) if _, err := occupied.recvUntilClosed(2 * time.Second); err != nil { t.Fatalf("occupied expired room client remained connected: %v", err) } h.srv.conns.mu.Lock() afterHardExpiry := h.srv.conns.roomsPerIP[ownerKey] h.srv.conns.mu.Unlock() if afterHardExpiry != maxRoomsPerIP-1 { t.Fatalf("quota after hard expiry=%d, want %d", afterHardExpiry, maxRoomsPerIP-1) } recovered := h.dial(t, ownerKey) recovered.send(clientMsg{Type: relayTypeCreate, SessionID: "RETAIN4", PeerID: "H"}) recovered.expect(relayTypeCreated) h.srv.conns.mu.Lock() finalQuota := h.srv.conns.roomsPerIP[ownerKey] h.srv.conns.mu.Unlock() if finalQuota != maxRoomsPerIP { t.Fatalf("final retained quota=%d, want %d", finalQuota, maxRoomsPerIP) } } func TestGlobalRetainedRoomCapBlocksCreateButPreservesJoin(t *testing.T) { h := newRelayHarness(t) now := time.Now() h.srv.mu.Lock() for _, persisted := range makeRoomSnapshots(maxRetainedRooms, false, now) { h.srv.rooms[persisted.SessionID] = &Room{ SessionID: persisted.SessionID, HostPeerID: persisted.HostPeerID, Peers: map[string]*Client{}, CreatedAt: persisted.CreatedAt, LastActivityAt: persisted.LastActivityAt, } } h.srv.mu.Unlock() ownerKey := "1.1.3.1" client := h.dial(t, ownerKey) client.send(clientMsg{Type: relayTypeCreate, SessionID: "OVERGLOBAL", PeerID: "H"}) client.expectError(relayErrorRateLimited) h.srv.mu.RLock() roomCount := len(h.srv.rooms) h.srv.mu.RUnlock() snapshotRoomCount := len(h.srv.buildSnapshot().Rooms) if roomCount != maxRetainedRooms || snapshotRoomCount != maxRetainedRooms { t.Fatalf("rejected create mutated retained state: rooms=%d snapshot=%d", roomCount, snapshotRoomCount) } h.srv.conns.mu.Lock() reservation := h.srv.conns.roomsPerIP[ownerKey] h.srv.conns.mu.Unlock() if reservation != 0 { t.Fatalf("global rejection reserved per-source quota: %d", reservation) } client.send(clientMsg{Type: relayTypeJoin, SessionID: "S0000", PeerID: "G"}) client.expect(relayTypeJoined) client.conn.Close() h.waitRoomPeers(t, "S0000", 0) h.srv.mu.RLock() expired := h.srv.rooms["S0001"] h.srv.mu.RUnlock() expired.mu.Lock() expired.LastActivityAt = now.Add(-emptyRoomMaxAge - time.Second) expired.mu.Unlock() h.srv.runCleanupStep(now) creator := h.dial(t, ownerKey) creator.send(clientMsg{Type: relayTypeCreate, SessionID: "AFTERGLOBAL", PeerID: "H"}) creator.expect(relayTypeCreated) h.srv.mu.RLock() roomCount = len(h.srv.rooms) h.srv.mu.RUnlock() if roomCount != maxRetainedRooms { t.Fatalf("room count after cleanup and create=%d, want %d", roomCount, maxRetainedRooms) } } func TestConcurrentCreatesCannotExceedGlobalRetainedRoomCap(t *testing.T) { for iteration := range 8 { t.Run(fmt.Sprintf("iteration_%d", iteration), func(t *testing.T) { h := newRelayHarness(t) now := time.Now() h.srv.mu.Lock() for _, persisted := range makeRoomSnapshots(maxRetainedRooms-1, false, now) { h.srv.rooms[persisted.SessionID] = &Room{ SessionID: persisted.SessionID, HostPeerID: persisted.HostPeerID, Peers: map[string]*Client{}, CreatedAt: persisted.CreatedAt, LastActivityAt: persisted.LastActivityAt, } } h.srv.mu.Unlock() type createResult struct { ownerKey string message serverMsg err error } connections := make([]*websocket.Conn, 2) for i := range connections { conn, err := h.dialRaw(fmt.Sprintf("1.1.4.%d", i+1)) if err != nil { t.Fatalf("dial create contender %d: %v", i, err) } connections[i] = conn t.Cleanup(func() { conn.Close() }) } start := make(chan struct{}) results := make(chan createResult, len(connections)) for i, conn := range connections { ownerKey := fmt.Sprintf("1.1.4.%d", i+1) go func(conn *websocket.Conn, ownerKey string, index int) { <-start err := conn.WriteJSON(clientMsg{ Type: relayTypeCreate, SessionID: fmt.Sprintf("RACE%d", index), PeerID: "H", }) var message serverMsg if err == nil { err = conn.ReadJSON(&message) } results <- createResult{ownerKey: ownerKey, message: message, err: err} }(conn, ownerKey, i) } close(start) created, rejected := 0, 0 acceptedOwner := "" for range connections { result := <-results if result.err != nil { t.Fatalf("concurrent create failed: %v", result.err) } switch { case result.message.Type == relayTypeCreated: created++ acceptedOwner = result.ownerKey case result.message.Type == relayTypeError && result.message.Code == relayErrorRateLimited: rejected++ default: t.Fatalf("unexpected concurrent result: %+v", result.message) } } if created != 1 || rejected != 1 { t.Fatalf("created=%d rejected=%d, want one each", created, rejected) } h.srv.mu.RLock() roomCount := len(h.srv.rooms) h.srv.mu.RUnlock() if roomCount != maxRetainedRooms { t.Fatalf("room count=%d, want %d", roomCount, maxRetainedRooms) } h.srv.conns.mu.Lock() acceptedQuota := h.srv.conns.roomsPerIP[acceptedOwner] totalQuota := 0 for _, count := range h.srv.conns.roomsPerIP { totalQuota += count } h.srv.conns.mu.Unlock() if acceptedQuota != 1 || totalQuota != 1 { t.Fatalf("accepted quota=%d total quota=%d, want 1 and 1", acceptedQuota, totalQuota) } }) } } func TestRelayUntrustedXFFCannotRotateConnectionOrRoomIdentity(t *testing.T) { t.Run("connections", func(t *testing.T) { h := newRelayHarnessNoTrust(t) var conns []*websocket.Conn for i := 0; i < maxConnsPerIP; i++ { conn, err := h.dialRaw(fmt.Sprintf("203.0.113.%d", i+1)) if err != nil { t.Fatalf("dial %d: %v", i, err) } conns = append(conns, conn) } t.Cleanup(func() { for _, conn := range conns { conn.Close() } }) headers := http.Header{"X-Forwarded-For": []string{"198.51.100.200"}} conn, resp, err := h.dialWithHeaders(headers) if conn != nil { conn.Close() t.Fatal("connection above direct peer limit unexpectedly succeeded") } if err == nil || resp == nil { t.Fatalf("dial error=%v response=%v, want HTTP 429", err, resp) } defer resp.Body.Close() if resp.StatusCode != http.StatusTooManyRequests { t.Fatalf("status=%d, want 429", resp.StatusCode) } }) t.Run("rooms", func(t *testing.T) { h := newRelayHarnessNoTrust(t) for i := 0; i <= maxRoomsPerIP; i++ { client := h.dial(t, fmt.Sprintf("203.0.113.%d", i+1)) client.send(clientMsg{Type: "create", SessionID: fmt.Sprintf("SPOOF%d", i), PeerID: "host"}) if i < maxRoomsPerIP { client.expect("created") } else { client.expectError("rate_limited") } } }) } func TestRelayTrustedClientsHaveIndependentConnectionBuckets(t *testing.T) { h := newRelayHarness(t) var conns []*websocket.Conn for i := 0; i < maxConnsPerIP; i++ { conn, err := h.dialRaw("203.0.113.10") if err != nil { t.Fatalf("client A dial %d: %v", i, err) } conns = append(conns, conn) } conn, err := h.dialRaw("203.0.113.11") if err != nil { t.Fatalf("client B should have an independent bucket: %v", err) } conns = append(conns, conn) t.Cleanup(func() { for _, conn := range conns { conn.Close() } }) } func TestRelayMalformedTrustedChainDoesNotMutateAdmission(t *testing.T) { h := newRelayHarness(t) headers := http.Header{"X-Forwarded-For": []string{"203.0.113.5,"}} conn, resp, err := h.dialWithHeaders(headers) if conn != nil { conn.Close() t.Fatal("malformed trusted chain unexpectedly upgraded") } if err == nil || resp == nil { t.Fatalf("dial error=%v response=%v, want HTTP 400", err, resp) } defer resp.Body.Close() if resp.StatusCode != http.StatusBadRequest { t.Fatalf("status=%d, want 400", resp.StatusCode) } h.srv.conns.mu.Lock() defer h.srv.conns.mu.Unlock() if h.srv.conns.globalCount != 0 || len(h.srv.conns.perIP) != 0 || len(h.srv.conns.ipRate) != 0 { t.Fatalf("malformed chain mutated connection admission: %+v", h.srv.conns) } } func TestConnectionCannotRetainMultipleRoomMemberships(t *testing.T) { h := newRelayHarness(t) ip := "1.1.1.8" client := h.dial(t, ip) client.send(clientMsg{Type: "create", SessionID: "PRIMARY", PeerID: "host"}) client.expect("created") otherHost := h.dial(t, "1.1.1.9") otherHost.send(clientMsg{Type: "create", SessionID: "OTHER", PeerID: "other-host"}) otherHost.expect("created") client.send(clientMsg{Type: "join", SessionID: "OTHER", PeerID: "ghost"}) client.expectError("already_in_room") client.send(clientMsg{Type: "create", SessionID: "EXTRA", PeerID: "extra-host"}) client.expectError("already_in_room") h.waitRoomPeers(t, "PRIMARY", 1) h.waitRoomPeers(t, "OTHER", 1) h.srv.mu.RLock() _, extraExists := h.srv.rooms["EXTRA"] h.srv.mu.RUnlock() if extraExists { t.Fatal("rejected create retained an extra room") } h.srv.conns.mu.Lock() roomsForIP := h.srv.conns.roomsPerIP[ip] h.srv.conns.mu.Unlock() if roomsForIP != 1 { t.Fatalf("roomsPerIP[%q]=%d, want 1", ip, roomsForIP) } } // ====================================================================== // handleWS — join case // ====================================================================== func TestJoinSucceedsAndBroadcastsPeerJoined(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "2.0.0.1") host.send(clientMsg{Type: relayTypeCreate, SessionID: "J1", PeerID: "H"}) host.expectAuthority(relayTypeCreated, "H") guest := h.dial(t, "2.0.0.2") guest.send(clientMsg{Type: relayTypeJoin, SessionID: "J1", PeerID: "G"}) joined := guest.expectAuthority(relayTypeJoined, "H") if joined.SessionID != "J1" { t.Errorf("SessionID=%q want J1", joined.SessionID) } if len(joined.Peers) != 1 || joined.Peers[0] != "H" { t.Errorf("Peers=%v, want [H]", joined.Peers) } peerJoined := host.expect(relayTypePeerJoined) if peerJoined.PeerID != "G" { t.Errorf("peerJoined.PeerID=%q want G", peerJoined.PeerID) } } func TestHostIdentityClaimsRequireReconnectCapability(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "2.0.1.1") host.send(clientMsg{Type: relayTypeCreate, SessionID: "AUTH", PeerID: "HOST"}) created := host.expectAuthority(relayTypeCreated, "HOST") guest := h.dial(t, "2.0.1.2") guest.send(clientMsg{Type: relayTypeJoin, SessionID: "AUTH", PeerID: "GUEST"}) guest.expectAuthority(relayTypeJoined, "HOST") host.expect(relayTypePeerJoined) attacker := h.dial(t, "2.0.1.3") attacker.send(clientMsg{Type: relayTypeJoin, SessionID: "AUTH", PeerID: "HOST"}) attacker.expectError(relayErrorPeerIdUnavailable) wrongToken, _ := mustReconnectToken(t) attacker.send(clientMsg{ Type: relayTypeJoin, SessionID: "AUTH", PeerID: "HOST", ReconnectToken: wrongToken, }) attacker.expectError(relayErrorPeerIdUnavailable) attacker.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"forged":true}`)}) attacker.expectError(relayErrorNotInRoom) host.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"real":true}`)}) message := guest.expect(relayTypeMessage) if message.From != "HOST" { t.Fatalf("message sender=%q, want HOST", message.From) } hostVerifier, ok := reconnectVerifierFromToken(created.ReconnectToken) if !ok { t.Fatal("created reconnect token became invalid") } h.srv.mu.RLock() room := h.srv.rooms["AUTH"] h.srv.mu.RUnlock() room.mu.RLock() currentVerifier := room.hostVerifier currentHost := room.Peers["HOST"] room.mu.RUnlock() if !reconnectVerifierMatches(hostVerifier, currentVerifier) || currentHost == nil { t.Fatal("failed claims mutated host authority") } } func TestJoinMissingFieldsRejected(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "2.0.0.3") c.send(clientMsg{Type: "join"}) c.expectError("invalid_message") } func TestRelayIdentifiersRejectUnsafeOrOversizedValues(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "2.0.0.30") invalid := []string{"has space", "has/slash", strings.Repeat("x", maxSessionIDLength+1)} for _, sessionID := range invalid { c.send(clientMsg{Type: "create", SessionID: sessionID, PeerID: "H"}) c.expectError("invalid_message") } c.send(clientMsg{Type: "create", SessionID: "SAFE_ID-1", PeerID: "H"}) c.expect("created") c.send(clientMsg{Type: "sendTo", To: "bad target", Payload: json.RawMessage(`{}`)}) c.expectError("invalid_message") } func TestJoinUnknownRoomFails(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "2.0.0.4") c.send(clientMsg{Type: "join", SessionID: "NOPE", PeerID: "G"}) c.expectError("room_not_found") } func TestReleasedGuestProbesDoNotConsumeDurableCapacity(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) host := h.dial(t, "2.0.0.40") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "PROBE_CAPACITY", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") for i := range maxRoomSize * 2 { peerID := fmt.Sprintf("P%d", i) probeToken, _ := mustReconnectToken(t) probe := h.dial(t, fmt.Sprintf("2.0.1.%d", i+1)) probe.send(clientMsg{ Type: relayTypeJoin, SessionID: "PROBE_CAPACITY", PeerID: peerID, ReconnectToken: probeToken, ProtocolVersion: relayProtocolVersion, }) probe.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) probe.send(clientMsg{ Type: relayTypeLeave, ReconnectToken: probeToken, ProtocolVersion: relayProtocolVersion, }) probe.expect(relayTypeLeft) left := host.expect(relayTypePeerLeft) if left.PeerID != peerID { t.Fatalf("released probe event peerId=%q, want %q", left.PeerID, peerID) } } h.srv.mu.RLock() room := h.srv.rooms["PROBE_CAPACITY"] h.srv.mu.RUnlock() room.mu.RLock() reservations := len(room.peerVerifiers) room.mu.RUnlock() if reservations != 0 { t.Fatalf("released probes retained %d durable reservations", reservations) } for i := range maxRoomSize - 1 { peerID := fmt.Sprintf("G%d", i) token, _ := mustReconnectToken(t) guest := h.dial(t, fmt.Sprintf("2.0.2.%d", i+1)) guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "PROBE_CAPACITY", PeerID: peerID, ReconnectToken: token, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) } } func TestFullRoomAllowsOnlyAuthenticatedLiveReplacements(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) host := h.dial(t, "2.1.0.1") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "FULL", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) created := host.expectAuthority(relayTypeCreated, "H") guests := make(map[string]*testConn) guestTokens := make(map[string]string) for i := 1; i < maxRoomSize; i++ { peerID := fmt.Sprintf("G%d", i) guestToken, _ := mustReconnectToken(t) guest := h.dial(t, fmt.Sprintf("2.1.0.%d", 100+i)) guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "FULL", PeerID: peerID, ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) joined := guest.expectAuthority(relayTypeJoined, "H") if joined.ReconnectToken != guestToken { t.Fatalf("%s join rotated its client-known token", peerID) } guests[peerID] = guest guestTokens[peerID] = guestToken } for range maxRoomSize - 1 { host.expect(relayTypePeerJoined) } for range maxRoomSize - 2 { guests["G1"].expect(relayTypePeerJoined) } overflowToken, _ := mustReconnectToken(t) overflow := h.dial(t, "2.1.0.250") overflow.send(clientMsg{ Type: relayTypeJoin, SessionID: "FULL", PeerID: "LATE", ReconnectToken: overflowToken, ProtocolVersion: relayProtocolVersion, }) overflow.expectError(relayErrorRoomFull) overflow.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{}`)}) overflow.expectError(relayErrorNotInRoom) unprovedHost := h.dial(t, "2.1.0.251") unprovedHost.send(clientMsg{ Type: relayTypeJoin, SessionID: "FULL", PeerID: "H", ProtocolVersion: relayProtocolVersion, }) unprovedHost.expectError(relayErrorPeerIdUnavailable) unprovedGuest := h.dial(t, "2.1.0.252") unprovedGuest.send(clientMsg{ Type: relayTypeJoin, SessionID: "FULL", PeerID: "G1", ProtocolVersion: relayProtocolVersion, }) unprovedGuest.expectError(relayErrorPeerIdUnavailable) wrongGuestToken, _ := mustReconnectToken(t) unprovedGuest.send(clientMsg{ Type: relayTypeJoin, SessionID: "FULL", PeerID: "G1", ReconnectToken: wrongGuestToken, ProtocolVersion: relayProtocolVersion, }) unprovedGuest.expectError(relayErrorPeerIdUnavailable) newHost := h.dial(t, "2.1.0.253") newHost.send(clientMsg{ Type: relayTypeJoin, SessionID: "FULL", PeerID: "H", ReconnectToken: created.ReconnectToken, ProtocolVersion: relayProtocolVersion, }) hostJoined := newHost.expectAuthority(relayTypeJoined, "H") if len(hostJoined.Peers) != maxRoomSize-1 { t.Fatalf("replacement host peers=%v, want %d peers", hostJoined.Peers, maxRoomSize-1) } if messages, err := host.recvUntilClosed(2 * time.Second); err != nil { t.Fatalf("displaced host did not close: %v (frames=%v)", err, messages) } hostReturn := guests["G1"].expect(relayTypePeerJoined) if hostReturn.PeerID != "H" { t.Fatalf("host replacement event peerId=%q", hostReturn.PeerID) } newGuest := h.dial(t, "2.1.0.254") newGuest.send(clientMsg{ Type: relayTypeJoin, SessionID: "FULL", PeerID: "G1", ReconnectToken: guestTokens["G1"], ProtocolVersion: relayProtocolVersion, }) newGuest.expectAuthority(relayTypeJoined, "H") if messages, err := guests["G1"].recvUntilClosed(2 * time.Second); err != nil { t.Fatalf("displaced guest did not close: %v (frames=%v)", err, messages) } guestReturn := newHost.expect(relayTypePeerJoined) if guestReturn.PeerID != "G1" { t.Fatalf("guest replacement event peerId=%q", guestReturn.PeerID) } newHost.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"state":"current"}`)}) message := newGuest.expect(relayTypeMessage) if message.From != "H" { t.Fatalf("replacement sender=%q, want H", message.From) } } func TestDisconnectedHostKeepsAReservedRoomSlot(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "2.1.1.1") host.send(clientMsg{Type: relayTypeCreate, SessionID: "HOST_SLOT", PeerID: "H"}) created := host.expectAuthority(relayTypeCreated, "H") for i := 1; i < maxRoomSize; i++ { guest := h.dial(t, fmt.Sprintf("2.1.1.%d", 100+i)) guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "HOST_SLOT", PeerID: fmt.Sprintf("G%d", i), }) guest.expectAuthority(relayTypeJoined, "H") } h.waitRoomPeers(t, "HOST_SLOT", maxRoomSize) if err := host.conn.Close(); err != nil { t.Fatalf("close host: %v", err) } h.waitRoomPeers(t, "HOST_SLOT", maxRoomSize-1) lateGuest := h.dial(t, "2.1.1.250") lateGuest.send(clientMsg{Type: relayTypeJoin, SessionID: "HOST_SLOT", PeerID: "LATE"}) lateGuest.expectError(relayErrorRoomFull) returningHost := h.dial(t, "2.1.1.251") returningHost.send(clientMsg{ Type: relayTypeJoin, SessionID: "HOST_SLOT", PeerID: "H", ReconnectToken: created.ReconnectToken, }) joined := returningHost.expectAuthority(relayTypeJoined, "H") if len(joined.Peers) != maxRoomSize-1 { t.Fatalf("returning host peers=%v, want %d peers", joined.Peers, maxRoomSize-1) } h.waitRoomPeers(t, "HOST_SLOT", maxRoomSize) } func TestLegacySameSourceHostReconnectAndModernTokenEnforcement(t *testing.T) { t.Run("legacy same-source tokenless reconnect", func(t *testing.T) { h := newRelayHarness(t) source := "2.1.2.1" host := h.dial(t, source) host.send(clientMsg{Type: relayTypeCreate, SessionID: "LEGACY_RECONNECT", PeerID: "H"}) host.expectAuthority(relayTypeCreated, "H") guest := h.dial(t, "2.1.2.2") guest.send(clientMsg{Type: relayTypeJoin, SessionID: "LEGACY_RECONNECT", PeerID: "G"}) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) if err := host.conn.Close(); err != nil { t.Fatalf("close legacy host: %v", err) } left := guest.expect(relayTypePeerLeft) if left.PeerID != "H" { t.Fatalf("legacy disconnect peerId=%q, want H", left.PeerID) } returning := h.dial(t, source) returning.send(clientMsg{Type: relayTypeJoin, SessionID: "LEGACY_RECONNECT", PeerID: "H"}) joined := returning.expect(relayTypeJoined) if joined.HostPeerID != "H" || joined.ReconnectToken != "" || joined.ProtocolVersion != legacyRelayProtocolVersion { t.Fatalf("legacy reconnect authority=%+v", joined) } rejoined := guest.expect(relayTypePeerJoined) if rejoined.PeerID != "H" { t.Fatalf("legacy reconnect event peerId=%q, want H", rejoined.PeerID) } }) t.Run("modern same-source reconnect requires token", func(t *testing.T) { h := newRelayHarness(t) source := "2.1.3.1" hostToken, _ := mustReconnectToken(t) host := h.dial(t, source) host.send(clientMsg{ Type: relayTypeCreate, SessionID: "MODERN_RECONNECT", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") guestToken, _ := mustReconnectToken(t) guest := h.dial(t, "2.1.3.2") guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "MODERN_RECONNECT", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) if err := host.conn.Close(); err != nil { t.Fatalf("close modern host: %v", err) } left := guest.expect(relayTypePeerLeft) if left.PeerID != "H" { t.Fatalf("modern disconnect peerId=%q, want H", left.PeerID) } unproved := h.dial(t, source) unproved.send(clientMsg{ Type: relayTypeJoin, SessionID: "MODERN_RECONNECT", PeerID: "H", ProtocolVersion: relayProtocolVersion, }) unproved.expectError(relayErrorPeerIdUnavailable) returning := h.dial(t, source) returning.send(clientMsg{ Type: relayTypeJoin, SessionID: "MODERN_RECONNECT", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) joined := returning.expectAuthority(relayTypeJoined, "H") if joined.ReconnectToken != hostToken || joined.ProtocolVersion != relayProtocolVersion { t.Fatalf("modern reconnect authority changed: %+v", joined) } }) } func TestJoinAdmissionIsAtomicWithEmptyRoomCleanup(t *testing.T) { h := newRelayHarness(t) _, hostVerifier := mustReconnectToken(t) now := time.Now() room := &Room{ SessionID: "ATOMIC_CLEANUP", HostPeerID: "H", hostVerifier: hostVerifier, peerVerifiers: make(map[string]reconnectVerifier), Peers: make(map[string]*Client), CreatedAt: now.Add(-time.Hour), LastActivityAt: now.Add(-emptyRoomMaxAge - time.Second), } h.srv.mu.Lock() h.srv.rooms[room.SessionID] = room h.srv.mu.Unlock() reached := make(chan struct{}) release := make(chan struct{}) var once sync.Once h.srv.beforeJoinRoomLock = func() { once.Do(func() { close(reached) <-release }) } joiner := h.dial(t, "2.2.0.1") joiner.send(clientMsg{Type: relayTypeJoin, SessionID: room.SessionID, PeerID: "G1"}) <-reached cleanupDone := make(chan struct{}) go func() { h.srv.runCleanupStep(now) close(cleanupDone) }() select { case <-cleanupDone: t.Fatal("cleanup passed a join that still owns the server read lock") case <-time.After(100 * time.Millisecond): } close(release) joiner.expectAuthority(relayTypeJoined, "H") <-cleanupDone h.srv.mu.RLock() current := h.srv.rooms[room.SessionID] h.srv.mu.RUnlock() if current != room { t.Fatal("successful join committed to a detached room") } room.mu.RLock() _, present := room.Peers["G1"] room.mu.RUnlock() if !present { t.Fatal("joining peer missing from authoritative room") } second := h.dial(t, "2.2.0.2") second.send(clientMsg{Type: relayTypeJoin, SessionID: room.SessionID, PeerID: "G2"}) second.expectAuthority(relayTypeJoined, "H") joiner.expect(relayTypePeerJoined) second.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"atomic":true}`)}) if message := joiner.expect(relayTypeMessage); message.From != "G2" { t.Fatalf("message sender=%q, want G2", message.From) } } func TestJoinAdmissionIsAtomicWithReservedRoomCreate(t *testing.T) { h := newRelayHarness(t) _, hostVerifier := mustReconnectToken(t) now := time.Now() room := &Room{ SessionID: "ATOMIC_CREATE", HostPeerID: "H", hostVerifier: hostVerifier, peerVerifiers: make(map[string]reconnectVerifier), Peers: make(map[string]*Client), CreatedAt: now, LastActivityAt: now, } h.srv.mu.Lock() h.srv.rooms[room.SessionID] = room h.srv.mu.Unlock() reached := make(chan struct{}) release := make(chan struct{}) var once sync.Once h.srv.beforeJoinRoomLock = func() { once.Do(func() { close(reached) <-release }) } joiner := h.dial(t, "2.3.0.1") joiner.send(clientMsg{Type: relayTypeJoin, SessionID: room.SessionID, PeerID: "G"}) <-reached creator := h.dial(t, "2.3.0.2") creator.send(clientMsg{Type: relayTypeCreate, SessionID: room.SessionID, PeerID: "OTHER"}) type readResult struct { message serverMsg err error } createResult := make(chan readResult, 1) go func() { creator.conn.SetReadDeadline(time.Now().Add(2 * time.Second)) _, data, err := creator.conn.ReadMessage() if err != nil { createResult <- readResult{err: err} return } var message serverMsg err = json.Unmarshal(data, &message) createResult <- readResult{message: message, err: err} }() select { case result := <-createResult: t.Fatalf("create completed before join admission committed: message=%+v err=%v", result.message, result.err) case <-time.After(100 * time.Millisecond): } close(release) joiner.expectAuthority(relayTypeJoined, "H") result := <-createResult if result.err != nil { t.Fatalf("read create result: %v", result.err) } if result.message.Type != relayTypeError || result.message.Code != relayErrorRoomExists { t.Fatalf("create result=%+v, want room_exists", result.message) } h.srv.mu.RLock() current := h.srv.rooms[room.SessionID] h.srv.mu.RUnlock() if current != room { t.Fatal("reserved room was replaced during admission") } } // ====================================================================== // handleWS — broadcast / sendTo // ====================================================================== func TestBroadcastDeliversToOthersNotSender(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "3.0.0.1") host.send(clientMsg{Type: "create", SessionID: "B1", PeerID: "H"}) host.expect("created") g1 := h.dial(t, "3.0.0.2") g1.send(clientMsg{Type: "join", SessionID: "B1", PeerID: "G1"}) g1.expect("joined") host.expect("peerJoined") g2 := h.dial(t, "3.0.0.3") g2.send(clientMsg{Type: "join", SessionID: "B1", PeerID: "G2"}) g2.expect("joined") host.expect("peerJoined") g1.expect("peerJoined") payload := json.RawMessage(`{"hello":"world"}`) g1.send(clientMsg{Type: "broadcast", Payload: payload}) hostMsg := host.expect("message") if hostMsg.From != "G1" { t.Errorf("host From=%q want G1", hostMsg.From) } if string(hostMsg.Payload) != string(payload) { t.Errorf("host payload=%s want %s", hostMsg.Payload, payload) } g2Msg := g2.expect("message") if g2Msg.From != "G1" { t.Errorf("g2 From=%q want G1", g2Msg.From) } // Sender should not receive its own broadcast. g1.recvNothing(200 * time.Millisecond) } func TestBroadcastNotInRoomRejected(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "3.0.0.4") c.send(clientMsg{Type: "broadcast", Payload: json.RawMessage(`{}`)}) c.expectError("not_in_room") } func TestSendToDeliversToTargetOnly(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "4.0.0.1") host.send(clientMsg{Type: "create", SessionID: "S1", PeerID: "H"}) host.expect("created") g1 := h.dial(t, "4.0.0.2") g1.send(clientMsg{Type: "join", SessionID: "S1", PeerID: "G1"}) g1.expect("joined") host.expect("peerJoined") g2 := h.dial(t, "4.0.0.3") g2.send(clientMsg{Type: "join", SessionID: "S1", PeerID: "G2"}) g2.expect("joined") host.expect("peerJoined") g1.expect("peerJoined") payload := json.RawMessage(`{"direct":true}`) host.send(clientMsg{Type: "sendTo", To: "G1", Payload: payload}) m := g1.expect("message") if m.From != "H" { t.Errorf("From=%q want H", m.From) } if string(m.Payload) != string(payload) { t.Errorf("payload mismatch: %s", m.Payload) } g2.recvNothing(200 * time.Millisecond) } func TestRelayMessagesRefreshRoomActivity(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "4.0.0.7") host.send(clientMsg{Type: "create", SessionID: "ACTIVE", PeerID: "H"}) host.expect("created") guest := h.dial(t, "4.0.0.8") guest.send(clientMsg{Type: "join", SessionID: "ACTIVE", PeerID: "G"}) guest.expect("joined") host.expect("peerJoined") h.srv.mu.RLock() room := h.srv.rooms["ACTIVE"] h.srv.mu.RUnlock() old := time.Now().Add(-time.Hour) room.mu.Lock() room.LastActivityAt = old room.mu.Unlock() host.send(clientMsg{Type: "broadcast", Payload: json.RawMessage(`{"broadcast":true}`)}) guest.expect("message") room.mu.RLock() broadcastActivity := room.LastActivityAt room.mu.RUnlock() if !broadcastActivity.After(old) { t.Fatalf("broadcast activity=%v, want after %v", broadcastActivity, old) } room.mu.Lock() room.LastActivityAt = old room.mu.Unlock() host.send(clientMsg{Type: "sendTo", To: "G", Payload: json.RawMessage(`{"direct":true}`)}) guest.expect("message") room.mu.RLock() directActivity := room.LastActivityAt room.mu.RUnlock() if !directActivity.After(old) { t.Fatalf("sendTo activity=%v, want after %v", directActivity, old) } } func TestSendToUnknownTargetRejected(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "4.0.0.4") host.send(clientMsg{Type: "create", SessionID: "S2", PeerID: "H"}) host.expect("created") host.send(clientMsg{Type: "sendTo", To: "ghost", Payload: json.RawMessage(`{}`)}) host.expectError("not_in_room") } func TestSendToMissingToRejected(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "4.0.0.5") host.send(clientMsg{Type: "create", SessionID: "S3", PeerID: "H"}) host.expect("created") host.send(clientMsg{Type: "sendTo", Payload: json.RawMessage(`{}`)}) host.expectError("invalid_message") } func TestSendToNotInRoomRejected(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "4.0.0.6") c.send(clientMsg{Type: "sendTo", To: "anyone", Payload: json.RawMessage(`{}`)}) c.expectError("not_in_room") } // ====================================================================== // handleWS — ping / misc / rate limits // ====================================================================== func TestPingReturnsPong(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "5.0.0.1") c.send(clientMsg{Type: "ping"}) c.expect("pong") } func TestUnknownTypeRejected(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "5.0.0.2") c.send(clientMsg{Type: "nope"}) c.expectError("invalid_message") } func TestInvalidJSONRejected(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "5.0.0.3") c.sendRaw([]byte("not json {{{")) c.expectError("invalid_message") } func TestPerConnectionMessageRateLimit(t *testing.T) { h := newRelayHarness(t) c := h.dial(t, "5.0.0.4") c.send(clientMsg{Type: "create", SessionID: "RL", PeerID: "H"}) c.expect("created") // The per-connection bucket is rateBurst=30. After ~30 pings we start seeing rate_limited. sawRateLimit := false for i := 0; i < rateBurst+10; i++ { c.send(clientMsg{Type: "ping"}) } for i := 0; i < rateBurst+10; i++ { m := c.recv() if m.Code == "rate_limited" { sawRateLimit = true break } } if !sawRateLimit { t.Fatal("expected to hit rate_limited within burst+10 messages") } } // ====================================================================== // handleWS — disconnect lifecycle // ====================================================================== func TestDisconnectBroadcastsPeerLeft(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "6.0.0.1") host.send(clientMsg{Type: "create", SessionID: "D1", PeerID: "H"}) host.expect("created") guest := h.dial(t, "6.0.0.2") guest.send(clientMsg{Type: "join", SessionID: "D1", PeerID: "G"}) guest.expect("joined") host.expect("peerJoined") guest.conn.Close() left := host.expect("peerLeft") if left.PeerID != "G" { t.Errorf("PeerID=%q, want G", left.PeerID) } } func TestStalePeerSkipsCleanupBroadcast(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) host := h.dial(t, "6.1.0.1") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "D2", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") guestToken, _ := mustReconnectToken(t) g1 := h.dial(t, "6.1.0.2") g1.send(clientMsg{ Type: relayTypeJoin, SessionID: "D2", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) g1.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) g2 := h.dial(t, "6.1.0.3") g2.send(clientMsg{ Type: relayTypeJoin, SessionID: "D2", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) g2.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) if messages, err := g1.recvUntilClosed(2 * time.Second); err != nil { t.Fatalf("displaced guest did not close: %v (frames=%v)", err, messages) } host.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"after":"replacement"}`)}) message := g2.expect(relayTypeMessage) if message.From != "H" { t.Fatalf("post-replacement sender=%q, want H", message.From) } } func TestDisconnectedModernGuestIdentityRejectsTheftAndAcceptsRightfulReconnect(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) host := h.dial(t, "6.1.0.10") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "GUEST_RECONNECT", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") guestToken, _ := mustReconnectToken(t) guest := h.dial(t, "6.1.0.11") guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "GUEST_RECONNECT", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) if err := guest.conn.Close(); err != nil { t.Fatalf("close guest: %v", err) } left := host.expect(relayTypePeerLeft) if left.PeerID != "G" { t.Fatalf("disconnected peerId=%q, want G", left.PeerID) } thiefToken, _ := mustReconnectToken(t) thief := h.dial(t, "6.1.0.12") thief.send(clientMsg{ Type: relayTypeJoin, SessionID: "GUEST_RECONNECT", PeerID: "G", ReconnectToken: thiefToken, ProtocolVersion: relayProtocolVersion, }) thief.expectError(relayErrorPeerIdUnavailable) thief.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"forged":true}`)}) thief.expectError(relayErrorNotInRoom) rightful := h.dial(t, "6.1.0.13") rightful.send(clientMsg{ Type: relayTypeJoin, SessionID: "GUEST_RECONNECT", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) joined := rightful.expectAuthority(relayTypeJoined, "H") if joined.ReconnectToken != guestToken { t.Fatal("rightful reconnect rotated the retained guest token") } rejoined := host.expect(relayTypePeerJoined) if rejoined.PeerID != "G" { t.Fatalf("rightful reconnect event peerId=%q, want G", rejoined.PeerID) } } func TestTerminalSuccessFramesFollowCrashDurableSnapshot(t *testing.T) { t.Run("guest leave releases persisted reservation", func(t *testing.T) { root := t.TempDir() statePath := filepath.Join(root, "rooms.json") h := newRelayHarnessAt(t, filepath.Join(root, "logs"), statePath) terminalReady := make(chan struct{}) releaseTerminal := make(chan struct{}) var releaseOnce sync.Once h.srv.beforeTerminalDelivery = func() { close(terminalReady) <-releaseTerminal } t.Cleanup(func() { releaseOnce.Do(func() { close(releaseTerminal) }) }) host, guest, _, guestToken := createModernRoomWithGuest( t, h, "DURABLE_LEAVE", "6.1.0.20", "6.1.0.21", ) guest.send(clientMsg{ Type: relayTypeLeave, ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) select { case <-terminalReady: case <-time.After(2 * time.Second): t.Fatal("leave did not reach the post-persistence delivery barrier") } restartPath := copySnapshotForRestart(t, statePath) restarted := newRelayHarnessAt(t, t.TempDir(), restartPath) replacementToken, _ := mustReconnectToken(t) replacement := restarted.dial(t, "6.1.0.22") replacement.send(clientMsg{ Type: relayTypeJoin, SessionID: "DURABLE_LEAVE", PeerID: "G", ReconnectToken: replacementToken, ProtocolVersion: relayProtocolVersion, }) replacement.expectAuthority(relayTypeJoined, "H") releaseOnce.Do(func() { close(releaseTerminal) }) guest.expect(relayTypeLeft) left := host.expect(relayTypePeerLeft) if left.PeerID != "G" { t.Fatalf("released peer event peerId=%q, want G", left.PeerID) } }) t.Run("host end removes persisted room", func(t *testing.T) { root := t.TempDir() statePath := filepath.Join(root, "rooms.json") h := newRelayHarnessAt(t, filepath.Join(root, "logs"), statePath) terminalReady := make(chan struct{}) releaseTerminal := make(chan struct{}) var releaseOnce sync.Once h.srv.beforeTerminalDelivery = func() { close(terminalReady) <-releaseTerminal } t.Cleanup(func() { releaseOnce.Do(func() { close(releaseTerminal) }) }) host, guest, hostToken, _ := createModernRoomWithGuest( t, h, "DURABLE_END", "6.1.0.23", "6.1.0.24", ) host.send(clientMsg{ Type: relayTypeEndSession, ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) select { case <-terminalReady: case <-time.After(2 * time.Second): t.Fatal("end did not reach the post-persistence delivery barrier") } restartPath := copySnapshotForRestart(t, statePath) restarted := newRelayHarnessAt(t, t.TempDir(), restartPath) probe := restarted.dial(t, "6.1.0.25") probe.send(clientMsg{ Type: relayTypeJoin, SessionID: "DURABLE_END", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) probe.expectError(relayErrorRoomNotFound) releaseOnce.Do(func() { close(releaseTerminal) }) host.expect(relayTypeEnded) messages, err := guest.recvUntilClosed(2 * time.Second) if err != nil { t.Fatalf("guest remained connected after durable end: %v (frames=%v)", err, messages) } if len(messages) != 1 || messages[0].Type != relayTypeEnded { t.Fatalf("guest terminal frames=%+v, want one ended notification", messages) } }) } func TestTerminalPersistenceFailureSuppressesSuccess(t *testing.T) { injectedErr := errors.New("injected snapshot persistence failure") t.Run("guest leave", func(t *testing.T) { root := t.TempDir() statePath := filepath.Join(root, "rooms.json") h := newRelayHarnessAt(t, filepath.Join(root, "logs"), statePath) _, guest, _, guestToken := createModernRoomWithGuest( t, h, "FAILED_LEAVE", "6.1.0.26", "6.1.0.27", ) injectSnapshotPersistenceFailure(t, h.srv.snap, injectedErr) guest.send(clientMsg{ Type: relayTypeLeave, ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) failure := guest.expectError(relayErrorInvalidMessage) if !strings.Contains(failure.Message, "persist") { t.Fatalf("leave persistence error message=%q", failure.Message) } restartPath := copySnapshotForRestart(t, statePath) restarted := newRelayHarnessAt(t, t.TempDir(), restartPath) replacementToken, _ := mustReconnectToken(t) replacement := restarted.dial(t, "6.1.0.28") replacement.send(clientMsg{ Type: relayTypeJoin, SessionID: "FAILED_LEAVE", PeerID: "G", ReconnectToken: replacementToken, ProtocolVersion: relayProtocolVersion, }) replacement.expectError(relayErrorPeerIdUnavailable) }) t.Run("host end", func(t *testing.T) { root := t.TempDir() statePath := filepath.Join(root, "rooms.json") h := newRelayHarnessAt(t, filepath.Join(root, "logs"), statePath) host, guest, hostToken, _ := createModernRoomWithGuest( t, h, "FAILED_END", "6.1.0.29", "6.1.0.30", ) injectSnapshotPersistenceFailure(t, h.srv.snap, injectedErr) host.send(clientMsg{ Type: relayTypeEndSession, ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) failure := host.expectError(relayErrorInvalidMessage) if !strings.Contains(failure.Message, "persist") { t.Fatalf("end persistence error message=%q", failure.Message) } messages, err := guest.recvUntilClosed(2 * time.Second) if err != nil { t.Fatalf("guest remained connected after failed end persistence: %v (frames=%v)", err, messages) } if len(messages) != 0 { t.Fatalf("guest received success frames after persistence failure: %+v", messages) } restartPath := copySnapshotForRestart(t, statePath) restarted := newRelayHarnessAt(t, t.TempDir(), restartPath) reconnected := restarted.dial(t, "6.1.0.31") reconnected.send(clientMsg{ Type: relayTypeJoin, SessionID: "FAILED_END", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) reconnected.expectAuthority(relayTypeJoined, "H") }) } func TestAuthenticatedModernGuestLeaveReleasesIdentity(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) host := h.dial(t, "6.1.0.20") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "GUEST_LEAVE", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") guestToken, _ := mustReconnectToken(t) guest := h.dial(t, "6.1.0.21") guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "GUEST_LEAVE", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) wrongToken, _ := mustReconnectToken(t) guest.send(clientMsg{ Type: relayTypeLeave, ReconnectToken: wrongToken, ProtocolVersion: relayProtocolVersion, }) guest.expectError(relayErrorPeerIdUnavailable) guest.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"still":"joined"}`)}) stillJoined := host.expect(relayTypeMessage) if stillJoined.From != "G" { t.Fatalf("message after rejected leave came from %q, want G", stillJoined.From) } guest.send(clientMsg{ Type: relayTypeLeave, ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) leftAck := guest.expect(relayTypeLeft) if leftAck.SessionID != "GUEST_LEAVE" || leftAck.PeerID != "G" || leftAck.ProtocolVersion != relayProtocolVersion { t.Fatalf("left acknowledgement=%+v", leftAck) } leftEvent := host.expect(relayTypePeerLeft) if leftEvent.PeerID != "G" { t.Fatalf("released peer event peerId=%q, want G", leftEvent.PeerID) } replacementToken, _ := mustReconnectToken(t) replacement := h.dial(t, "6.1.0.22") replacement.send(clientMsg{ Type: relayTypeJoin, SessionID: "GUEST_LEAVE", PeerID: "G", ReconnectToken: replacementToken, ProtocolVersion: relayProtocolVersion, }) joined := replacement.expectAuthority(relayTypeJoined, "H") if joined.ReconnectToken != replacementToken { t.Fatal("released guest identity retained its old verifier") } rejoined := host.expect(relayTypePeerJoined) if rejoined.PeerID != "G" { t.Fatalf("replacement event peerId=%q, want G", rejoined.PeerID) } } func TestAuthenticatedModernHostEndDeletesRoomAndIsRetrySafe(t *testing.T) { h := newRelayHarness(t) hostToken, _ := mustReconnectToken(t) host := h.dial(t, "6.1.0.30") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "HOST_END", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") guestToken, _ := mustReconnectToken(t) guest := h.dial(t, "6.1.0.31") guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "HOST_END", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) wrongToken, _ := mustReconnectToken(t) host.send(clientMsg{ Type: relayTypeEndSession, ReconnectToken: wrongToken, ProtocolVersion: relayProtocolVersion, }) host.expectError(relayErrorPeerIdUnavailable) host.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"room":"live"}`)}) stillLive := guest.expect(relayTypeMessage) if stillLive.From != "H" { t.Fatalf("message after rejected end came from %q, want H", stillLive.From) } host.send(clientMsg{ Type: relayTypeEndSession, ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) ended := host.expect(relayTypeEnded) if ended.SessionID != "HOST_END" || ended.ProtocolVersion != relayProtocolVersion { t.Fatalf("ended acknowledgement=%+v", ended) } messages, err := guest.recvUntilClosed(2 * time.Second) if err != nil { t.Fatalf("guest remained connected after room end: %v (frames=%v)", err, messages) } if len(messages) != 1 || messages[0].Type != relayTypeEnded || messages[0].SessionID != "HOST_END" || messages[0].ProtocolVersion != relayProtocolVersion { t.Fatalf("guest terminal frames=%+v, want one ended notification", messages) } host.send(clientMsg{ Type: relayTypeEndSession, ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectError(relayErrorNotInRoom) retry := h.dial(t, "6.1.0.32") retry.send(clientMsg{ Type: relayTypeJoin, SessionID: "HOST_END", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) retry.expectError(relayErrorRoomNotFound) } func TestHostEndDeliversEndedAfterConcurrentGuestTraffic(t *testing.T) { h := newRelayHarness(t) endDeliveryReady := make(chan struct{}) releaseEndDelivery := make(chan struct{}) var releaseOnce sync.Once h.srv.beforeTerminalDelivery = func() { close(endDeliveryReady) <-releaseEndDelivery } t.Cleanup(func() { releaseOnce.Do(func() { close(releaseEndDelivery) }) }) hostToken, _ := mustReconnectToken(t) host := h.dial(t, "6.1.0.33") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "END_RACE", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") guestToken, _ := mustReconnectToken(t) guest := h.dial(t, "6.1.0.34") guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "END_RACE", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) host.send(clientMsg{ Type: relayTypeEndSession, ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) select { case <-endDeliveryReady: case <-time.After(2 * time.Second): t.Fatal("host end did not reach the terminal-delivery barrier") } h.srv.mu.RLock() _, discoverable := h.srv.rooms["END_RACE"] h.srv.mu.RUnlock() if discoverable { t.Fatal("ending room remained discoverable before terminal delivery") } // WebSocket frames are processed in order. Receiving pong proves the // preceding membership-sensitive traffic was handled while ended delivery // was blocked, without closing the guest as a stale client. guest.send(clientMsg{ Type: relayTypeBroadcast, Payload: json.RawMessage(`{"during":"end"}`), }) guest.send(clientMsg{ Type: relayTypeSendTo, To: "H", Payload: json.RawMessage(`{"also":"during-end"}`), }) guest.send(clientMsg{Type: relayTypePing}) guest.expect(relayTypePong) releaseOnce.Do(func() { close(releaseEndDelivery) }) endedAck := host.expect(relayTypeEnded) if endedAck.SessionID != "END_RACE" || endedAck.ProtocolVersion != relayProtocolVersion { t.Fatalf("host ended acknowledgement=%+v", endedAck) } messages, err := guest.recvUntilClosed(2 * time.Second) if err != nil { t.Fatalf("guest remained connected after terminal delivery: %v (frames=%v)", err, messages) } if len(messages) != 1 || messages[0].Type != relayTypeEnded || messages[0].SessionID != "END_RACE" || messages[0].ProtocolVersion != relayProtocolVersion { t.Fatalf("guest terminal frames=%+v, want one ended notification", messages) } } func TestHostReconnectReplacesStaleConnectionWithoutLeaving(t *testing.T) { h := newRelayHarness(t) oldHostIP := "6.1.1.1" oldHost := h.dial(t, oldHostIP) oldHost.send(clientMsg{Type: relayTypeCreate, SessionID: "REJOIN", PeerID: "H"}) created := oldHost.expectAuthority(relayTypeCreated, "H") guest := h.dial(t, "6.1.1.2") guest.send(clientMsg{Type: relayTypeJoin, SessionID: "REJOIN", PeerID: "G"}) guest.expectAuthority(relayTypeJoined, "H") oldHost.expect(relayTypePeerJoined) newHost := h.dial(t, "6.1.1.3") newHost.send(clientMsg{ Type: relayTypeJoin, SessionID: "REJOIN", PeerID: "H", ReconnectToken: created.ReconnectToken, }) joined := newHost.expectAuthority(relayTypeJoined, "H") if len(joined.Peers) != 1 || joined.Peers[0] != "G" { t.Fatalf("reconnected host peers=%v, want [G]", joined.Peers) } guest.expect(relayTypePeerJoined) if messages, err := oldHost.recvUntilClosed(2 * time.Second); err != nil { t.Fatalf("displaced host did not close: %v (frames=%v)", err, messages) } h.waitIPConnections(t, oldHostIP, 0) newHost.send(clientMsg{Type: relayTypeBroadcast, Payload: json.RawMessage(`{"state":"ready"}`)}) message := guest.expect(relayTypeMessage) if message.From != "H" { t.Fatalf("message sender=%q, want H", message.From) } } func TestEmptyRoomRequiresHostProofUntilExpiryThenSupportsFallbackCreate(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "6.1.2.1") host.send(clientMsg{Type: relayTypeCreate, SessionID: "EMPTY", PeerID: "H"}) created := host.expectAuthority(relayTypeCreated, "H") host.conn.Close() h.waitRoomPeers(t, "EMPTY", 0) unproved := h.dial(t, "6.1.2.20") unproved.send(clientMsg{Type: relayTypeJoin, SessionID: "EMPTY", PeerID: "H"}) unproved.expectError(relayErrorPeerIdUnavailable) reconnected := h.dial(t, "6.1.2.2") reconnected.send(clientMsg{ Type: relayTypeJoin, SessionID: "EMPTY", PeerID: "H", ReconnectToken: created.ReconnectToken, }) joined := reconnected.expectAuthority(relayTypeJoined, "H") if joined.ReconnectToken != created.ReconnectToken { t.Fatal("host reconnect rotated its capability") } if len(joined.Peers) != 0 { t.Fatalf("empty-room reconnect peers=%v, want none", joined.Peers) } reconnected.conn.Close() h.waitRoomPeers(t, "EMPTY", 0) now := time.Now() h.srv.mu.RLock() room := h.srv.rooms["EMPTY"] h.srv.mu.RUnlock() room.mu.Lock() room.LastActivityAt = now.Add(-emptyRoomMaxAge - time.Second) room.mu.Unlock() h.srv.runCleanupStep(now) fallback := h.dial(t, "6.1.2.3") fallback.send(clientMsg{ Type: relayTypeJoin, SessionID: "EMPTY", PeerID: "H", ReconnectToken: created.ReconnectToken, }) fallback.expectError(relayErrorRoomNotFound) fallback.send(clientMsg{ Type: relayTypeCreate, SessionID: "EMPTY", PeerID: "H", ReconnectToken: created.ReconnectToken, }) recreated := fallback.expectAuthority(relayTypeCreated, "H") if recreated.ReconnectToken != created.ReconnectToken { t.Fatal("fallback create rotated its retained capability") } } func TestCleanupDisconnectsPeersBeforeRemovingExpiredOccupiedRoom(t *testing.T) { h := newRelayHarness(t) host := h.dial(t, "6.2.0.1") host.send(clientMsg{Type: "create", SessionID: "EXPIRED", PeerID: "H"}) host.expect("created") guest := h.dial(t, "6.2.0.2") guest.send(clientMsg{Type: "join", SessionID: "EXPIRED", PeerID: "G"}) guest.expect("joined") host.expect("peerJoined") now := time.Now() h.srv.mu.RLock() room := h.srv.rooms["EXPIRED"] h.srv.mu.RUnlock() room.mu.Lock() room.CreatedAt = now.Add(-roomMaxAge - time.Second) room.mu.Unlock() h.srv.runCleanupStep(now) h.srv.mu.RLock() _, exists := h.srv.rooms["EXPIRED"] h.srv.mu.RUnlock() if exists { t.Fatal("expired room still exists after cleanup") } room.mu.RLock() closing := room.closing remainingPeers := len(room.Peers) room.mu.RUnlock() if !closing { t.Error("expired room was not marked closing") } if remainingPeers != 0 { t.Errorf("expired room retained %d peers after cleanup", remainingPeers) } for name, connection := range map[string]*testConn{"host": host, "guest": guest} { messages, err := connection.recvUntilClosed(2 * time.Second) if err != nil { t.Errorf("%s did not reach terminal closure: %v", name, err) } for _, message := range messages { if message.Type == relayTypePeerLeft || message.Type == relayTypePeerJoined { t.Errorf("%s received %s during expired-room teardown", name, message.Type) } } } } // ====================================================================== // Logs endpoints // ====================================================================== func postLog(t *testing.T, baseURL, ip string, body []byte) *http.Response { t.Helper() req, err := http.NewRequest(http.MethodPost, baseURL+"/logs", bytes.NewReader(body)) if err != nil { t.Fatalf("new request: %v", err) } if ip != "" { req.Header.Set("X-Forwarded-For", ip) } resp, err := http.DefaultClient.Do(req) if err != nil { t.Fatalf("post: %v", err) } return resp } func getLog(t *testing.T, baseURL, ip, id string) *http.Response { t.Helper() req, err := http.NewRequest(http.MethodGet, baseURL+"/logs/"+id, nil) if err != nil { t.Fatalf("new request: %v", err) } if ip != "" { req.Header.Set("X-Forwarded-For", ip) } resp, err := http.DefaultClient.Do(req) if err != nil { t.Fatalf("get: %v", err) } return resp } // postLogAndGetID uploads a log and returns the generated id, asserting the // POST succeeded. func postLogAndGetID(t *testing.T, baseURL, ip string, body []byte) string { t.Helper() resp := postLog(t, baseURL, ip, body) defer resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("post status=%d", resp.StatusCode) } var out struct { ID string `json:"id"` } if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { t.Fatalf("decode: %v", err) } if len(out.ID) != logIDLength { t.Fatalf("id=%q len=%d want %d", out.ID, len(out.ID), logIDLength) } return out.ID } type posterUploadResponse struct { ID string `json:"id"` URL string `json:"url"` ExpiresIn int `json:"expiresIn"` } func postPoster(t *testing.T, baseURL, ip string, body []byte) *http.Response { t.Helper() req, err := http.NewRequest(http.MethodPost, baseURL+"/posters", bytes.NewReader(body)) if err != nil { t.Fatalf("new request: %v", err) } if ip != "" { req.Header.Set("X-Forwarded-For", ip) } resp, err := http.DefaultClient.Do(req) if err != nil { t.Fatalf("post: %v", err) } return resp } func postPosterAndDecode(t *testing.T, baseURL, ip string, body []byte) posterUploadResponse { t.Helper() resp := postPoster(t, baseURL, ip, body) defer resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("post status=%d", resp.StatusCode) } var out posterUploadResponse if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { t.Fatalf("decode: %v", err) } if len(out.ID) != posterIDLength { t.Fatalf("id=%q len=%d want %d", out.ID, len(out.ID), posterIDLength) } if out.URL == "" { t.Fatal("empty poster url") } return out } type blockingLogResponseWriter struct { header http.Header writeStarted chan struct{} releaseWrite chan struct{} writeErr error startOnce sync.Once mu sync.Mutex deadlines []time.Time } func newBlockingLogResponseWriter(writeErr error) *blockingLogResponseWriter { return &blockingLogResponseWriter{ header: make(http.Header), writeStarted: make(chan struct{}), releaseWrite: make(chan struct{}), writeErr: writeErr, } } func (w *blockingLogResponseWriter) Header() http.Header { return w.header } func (w *blockingLogResponseWriter) WriteHeader(int) {} func (w *blockingLogResponseWriter) Write(p []byte) (int, error) { w.startOnce.Do(func() { close(w.writeStarted) }) <-w.releaseWrite if w.writeErr != nil { return 0, w.writeErr } return len(p), nil } func (w *blockingLogResponseWriter) SetWriteDeadline(deadline time.Time) error { w.mu.Lock() w.deadlines = append(w.deadlines, deadline) w.mu.Unlock() return nil } func TestLogResponseTransmissionReleasesLookupSlotAndUsesDeadline(t *testing.T) { logs := newLogStore(t.TempDir()) id, _, err := logs.store([]byte("diagnostic"), time.Now()) if err != nil { t.Fatalf("store log: %v", err) } srv := &Server{ logs: logs, logLookups: make(chan struct{}, 1), clientIPs: newClientIPResolver(nil), } writer := newBlockingLogResponseWriter(errors.New("synthetic write failure: " + id)) request := httptest.NewRequest(http.MethodGet, "/logs/"+id, nil) var output bytes.Buffer previousOutput := log.Writer() previousFlags := log.Flags() previousPrefix := log.Prefix() log.SetOutput(&output) log.SetFlags(0) log.SetPrefix("") t.Cleanup(func() { log.SetOutput(previousOutput) log.SetFlags(previousFlags) log.SetPrefix(previousPrefix) }) done := make(chan struct{}) go func() { srv.handleGetLogs(writer, request) close(done) }() select { case <-writer.writeStarted: case <-time.After(time.Second): t.Fatal("log response did not reach blocked write") } if occupied := len(srv.logLookups); occupied != 0 { t.Fatalf("blocked response retained %d lookup slots", occupied) } second := httptest.NewRecorder() srv.handleGetLogs(second, httptest.NewRequest(http.MethodGet, "/logs/"+id, nil)) if second.Code != http.StatusOK { t.Fatalf("lookup while first response blocked status=%d, want 200", second.Code) } writer.mu.Lock() deadlines := append([]time.Time(nil), writer.deadlines...) writer.mu.Unlock() if len(deadlines) == 0 || deadlines[0].IsZero() { t.Fatalf("response write deadline not applied: %v", deadlines) } remaining := time.Until(deadlines[0]) if remaining <= 0 || remaining > httpResponseWriteTimeout { t.Fatalf("response write deadline remaining=%v, want within (0, %v]", remaining, httpResponseWriteTimeout) } close(writer.releaseWrite) select { case <-done: case <-time.After(time.Second): t.Fatal("log handler did not finish after writer release") } if !strings.Contains(output.String(), "logs: response write failed") { t.Fatalf("write failure was not observed: %q", output.String()) } if strings.Contains(output.String(), id) { t.Fatalf("write failure leaked log capability %q", id) } } func TestHTTPServerWriteTimeoutCoversOAuthResultLongPoll(t *testing.T) { server := newHTTPServer("127.0.0.1:0", http.NewServeMux()) if server.WriteTimeout != httpResponseWriteTimeout || server.WriteTimeout <= 0 { t.Fatalf("WriteTimeout=%v, want bounded timeout %v", server.WriteTimeout, httpResponseWriteTimeout) } if server.WriteTimeout <= oauthResultWait { t.Fatalf("WriteTimeout=%v must exceed oauthResultWait=%v", server.WriteTimeout, oauthResultWait) } if margin := server.WriteTimeout - oauthResultWait; margin < httpResponseWriteMargin { t.Fatalf("WriteTimeout margin=%v, want at least %v", margin, httpResponseWriteMargin) } } func TestLogsRoundTrip(t *testing.T) { h := newRelayHarness(t) payload := []byte("hello log world") id := postLogAndGetID(t, h.baseURL, "7.0.0.1", payload) getResp, err := http.Get(h.baseURL + "/logs/" + id) if err != nil { t.Fatalf("get: %v", err) } defer getResp.Body.Close() if getResp.StatusCode != http.StatusOK { t.Fatalf("get status=%d", getResp.StatusCode) } got, _ := io.ReadAll(getResp.Body) if !bytes.Equal(got, payload) { t.Fatalf("round-tripped bytes mismatch: got %q want %q", got, payload) } if ct := getResp.Header.Get("Content-Type"); !strings.HasPrefix(ct, "text/plain") { t.Errorf("Content-Type=%q", ct) } } func TestLogsUploadDoesNotWriteCapabilityToOperationalLog(t *testing.T) { h := newRelayHarness(t) var output bytes.Buffer previousOutput := log.Writer() previousFlags := log.Flags() previousPrefix := log.Prefix() log.SetOutput(&output) log.SetFlags(0) log.SetPrefix("") t.Cleanup(func() { log.SetOutput(previousOutput) log.SetFlags(previousFlags) log.SetPrefix(previousPrefix) }) id := postLogAndGetID(t, h.baseURL, "203.0.113.40", []byte("safe diagnostic")) if strings.Contains(output.String(), id) { t.Fatalf("operational log retained bearer capability %q", id) } if !strings.Contains(output.String(), "logs: stored 15 bytes from 203.0.113.40") { t.Fatalf("successful upload was not observable: %q", output.String()) } } func TestLogStoreRetiresLegacyCapabilitiesOnStartup(t *testing.T) { dir := t.TempDir() now := time.Now().Add(-time.Minute) legacyID := "abcde" currentID := strings.Repeat("a", logIDLength) legacyPath := filepath.Join(dir, legacyID+".log") currentPath := filepath.Join(dir, currentID+".log") for path, body := range map[string]string{legacyPath: "legacy", currentPath: "current"} { if err := os.WriteFile(path, []byte(body), 0o644); err != nil { t.Fatalf("seed %s: %v", path, err) } if err := os.Chtimes(path, now, now); err != nil { t.Fatalf("chtimes %s: %v", path, err) } } store := newLogStore(dir) if _, err := os.Stat(legacyPath); !errors.Is(err, os.ErrNotExist) { t.Fatalf("legacy capability file still exists: %v", err) } if _, ok, err := store.lookup(legacyID, time.Now()); err != nil || ok { t.Fatalf("legacy capability lookup=(ok=%v, err=%v), want absent", ok, err) } if _, ok, err := store.lookup(currentID, time.Now()); err != nil || !ok { t.Fatalf("current capability lookup=(ok=%v, err=%v), want indexed", ok, err) } restarted := newLogStore(dir) if _, ok, err := restarted.lookup(currentID, time.Now()); err != nil || !ok { t.Fatalf("restarted capability lookup=(ok=%v, err=%v), want indexed", ok, err) } } func TestLogsFailedLookupsAreBoundedButValidCapabilitiesRemainAvailable(t *testing.T) { h := newRelayHarness(t) payload := []byte("retrievable") validID := postLogAndGetID(t, h.baseURL, "203.0.113.1", payload) source := "203.0.113.50" for i := range logLookupRateBurst { unknownID := strings.Repeat("z", logIDLength-2) + fmt.Sprintf("%02d", i) resp := getLog(t, h.baseURL, source, unknownID) resp.Body.Close() if resp.StatusCode != http.StatusNotFound { t.Fatalf("failed lookup %d status=%d, want 404", i, resp.StatusCode) } if got := resp.Header.Get("Cache-Control"); got != "private, no-store" { t.Fatalf("failed lookup Cache-Control=%q", got) } } throttled := getLog(t, h.baseURL, source, strings.Repeat("y", logIDLength)) throttled.Body.Close() if throttled.StatusCode != http.StatusTooManyRequests { t.Fatalf("exhausted lookup status=%d, want 429", throttled.StatusCode) } if got := throttled.Header.Get("Cache-Control"); got != "private, no-store" { t.Fatalf("throttled Cache-Control=%q", got) } success := getLog(t, h.baseURL, source, validID) defer success.Body.Close() if success.StatusCode != http.StatusOK { t.Fatalf("valid capability after exhausted failures status=%d", success.StatusCode) } got, err := io.ReadAll(success.Body) if err != nil || !bytes.Equal(got, payload) { t.Fatalf("valid body=%q err=%v, want %q", got, err, payload) } if cache := success.Header.Get("Cache-Control"); cache != "private, no-store" { t.Fatalf("success Cache-Control=%q", cache) } independent := getLog(t, h.baseURL, "203.0.113.51", strings.Repeat("x", logIDLength)) independent.Body.Close() if independent.StatusCode != http.StatusNotFound { t.Fatalf("independent source status=%d, want 404", independent.StatusCode) } } func TestLogFailedLookupCleanupIsDeterministic(t *testing.T) { store := newLogStore(t.TempDir()) now := time.Unix(1_700_000_000, 0) id, _, err := store.store([]byte("keep"), now) if err != nil { t.Fatalf("store: %v", err) } for range logLookupRateBurst { if !store.allowFailedLookup("203.0.113.1", now) { t.Fatal("burst rejected early") } } if store.allowFailedLookup("203.0.113.1", now) { t.Fatal("lookup beyond burst unexpectedly allowed") } store.cleanup(now) if _, ok := store.failedLookupRate["203.0.113.1"]; !ok { t.Fatal("cleanup removed an effective limiter") } store.cleanup(now.Add(time.Duration(logLookupRateBurst) * time.Second)) if _, ok := store.failedLookupRate["203.0.113.1"]; ok { t.Fatal("cleanup retained a fully refilled limiter") } if _, ok := store.entries[id]; !ok { t.Fatal("limiter cleanup removed stored log") } } func TestLogStorePersistsAcrossRestartAndAvoidsIDCollisions(t *testing.T) { dir := t.TempDir() now := time.Now().Add(-time.Second) first := newLogStore(dir) first.generateID = func() string { return strings.Repeat("a", logIDLength) } firstID, _, err := first.store([]byte("original"), now) if err != nil { t.Fatalf("store original: %v", err) } restarted := newLogStore(dir) if _, ok, err := restarted.lookup(firstID, time.Now()); err != nil || !ok { t.Fatal("stored log was not restored after restart") } secondWant := strings.Repeat("b", logIDLength) ids := []string{firstID, secondWant} restarted.generateID = func() string { id := ids[0] ids = ids[1:] return id } secondID, _, err := restarted.store([]byte("second"), time.Now()) if err != nil { t.Fatalf("store after restart: %v", err) } if secondID != secondWant { t.Fatalf("collision generated id %q, want %q", secondID, secondWant) } original, err := os.ReadFile(restarted.filePath(firstID)) if err != nil { t.Fatalf("read original: %v", err) } if string(original) != "original" { t.Fatalf("colliding store overwrote original: %q", original) } } func TestLogsUploadRateLimitedPerIP(t *testing.T) { h := newRelayHarness(t) r1 := postLog(t, h.baseURL, "7.0.0.2", []byte("first")) r1.Body.Close() if r1.StatusCode != http.StatusOK { t.Fatalf("first post status=%d", r1.StatusCode) } r2 := postLog(t, h.baseURL, "7.0.0.2", []byte("second")) r2.Body.Close() if r2.StatusCode != http.StatusTooManyRequests { t.Fatalf("second post status=%d want 429", r2.StatusCode) } } func TestLogsUseTrustedCanonicalClientIdentity(t *testing.T) { t.Run("untrusted spoofing shares direct peer bucket", func(t *testing.T) { h := newRelayHarnessNoTrust(t) first := postLog(t, h.baseURL, "203.0.113.1", []byte("first")) first.Body.Close() if first.StatusCode != http.StatusOK { t.Fatalf("first status=%d", first.StatusCode) } second := postLog(t, h.baseURL, "203.0.113.2", []byte("second")) second.Body.Close() if second.StatusCode != http.StatusTooManyRequests { t.Fatalf("rotated spoof status=%d, want 429", second.StatusCode) } }) t.Run("trusted clients have independent buckets", func(t *testing.T) { h := newRelayHarness(t) for _, ip := range []string{"203.0.113.1", "203.0.113.2"} { resp := postLog(t, h.baseURL, ip, []byte(ip)) resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("client %s status=%d, want 200", ip, resp.StatusCode) } } }) t.Run("malformed trusted chain mutates no log state", func(t *testing.T) { h := newRelayHarness(t) post := postLog(t, h.baseURL, "203.0.113.1,", []byte("body")) post.Body.Close() if post.StatusCode != http.StatusBadRequest { t.Fatalf("post status=%d, want 400", post.StatusCode) } get := getLog(t, h.baseURL, "203.0.113.1,", strings.Repeat("a", logIDLength)) get.Body.Close() if get.StatusCode != http.StatusBadRequest { t.Fatalf("get status=%d, want 400", get.StatusCode) } h.srv.logs.mu.RLock() defer h.srv.logs.mu.RUnlock() if len(h.srv.logs.entries) != 0 || len(h.srv.logs.rateLimit) != 0 || len(h.srv.logs.failedLookupRate) != 0 { t.Fatalf("malformed chain mutated log state: entries=%d uploads=%d failures=%d", len(h.srv.logs.entries), len(h.srv.logs.rateLimit), len(h.srv.logs.failedLookupRate)) } }) t.Run("untrusted malformed header is ignored", func(t *testing.T) { h := newRelayHarnessNoTrust(t) resp := postLog(t, h.baseURL, "bad,", []byte("body")) resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("status=%d, want 200", resp.StatusCode) } }) } func TestLogsUploadTooLargeRejected(t *testing.T) { h := newRelayHarness(t) body := make([]byte, maxLogSize+1) resp := postLog(t, h.baseURL, "7.0.0.3", body) resp.Body.Close() if resp.StatusCode != http.StatusRequestEntityTooLarge { t.Fatalf("status=%d want 413", resp.StatusCode) } } func TestLogsUploadEmptyRejected(t *testing.T) { h := newRelayHarness(t) resp := postLog(t, h.baseURL, "7.0.0.4", nil) resp.Body.Close() if resp.StatusCode != http.StatusBadRequest { t.Fatalf("status=%d want 400", resp.StatusCode) } } func TestLogsUploadStoreFull(t *testing.T) { h := newRelayHarness(t) // Saturate the store with distinct IPs so per-IP rate limit doesn't bite. for i := 0; i < maxLogEntries; i++ { ip := fmt.Sprintf("7.1.%d.%d", i/256, i%256) resp := postLog(t, h.baseURL, ip, []byte("x")) resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("warmup %d (ip=%s) status=%d", i, ip, resp.StatusCode) } } resp := postLog(t, h.baseURL, "7.2.0.1", []byte("overflow")) resp.Body.Close() if resp.StatusCode != http.StatusServiceUnavailable { t.Fatalf("overflow status=%d want 503", resp.StatusCode) } } func TestLogsGetUnknownIDIs404(t *testing.T) { h := newRelayHarness(t) resp, err := http.Get(h.baseURL + "/logs/" + strings.Repeat("c", logIDLength)) if err != nil { t.Fatalf("get: %v", err) } resp.Body.Close() if resp.StatusCode != http.StatusNotFound { t.Fatalf("status=%d want 404", resp.StatusCode) } } func TestLogsGetMalformedIDIs404(t *testing.T) { h := newRelayHarness(t) for _, id := range []string{"", "abc", "abcde", strings.Repeat("a", logIDLength+1), strings.Repeat("!", logIDLength)} { resp, err := http.Get(h.baseURL + "/logs/" + id) if err != nil { t.Fatalf("get %q: %v", id, err) } resp.Body.Close() if resp.StatusCode != http.StatusNotFound { t.Errorf("id=%q status=%d want 404", id, resp.StatusCode) } if got := resp.Header.Get("Cache-Control"); got != "private, no-store" { t.Errorf("id=%q Cache-Control=%q", id, got) } } } func TestLogsGetExpiredIs404(t *testing.T) { h := newRelayHarness(t) id := postLogAndGetID(t, h.baseURL, "7.3.0.1", []byte("temp")) // Poison the entry's ExpiresAt into the past. h.srv.logs.mu.Lock() entry := h.srv.logs.entries[id] entry.ExpiresAt = time.Now().Add(-time.Minute) h.srv.logs.entries[id] = entry h.srv.logs.mu.Unlock() get, err := http.Get(h.baseURL + "/logs/" + id) if err != nil { t.Fatalf("get: %v", err) } get.Body.Close() if get.StatusCode != http.StatusNotFound { t.Fatalf("status=%d want 404", get.StatusCode) } } func TestLogsMethodNotAllowed(t *testing.T) { h := newRelayHarness(t) resp, err := http.Get(h.baseURL + "/logs") if err != nil { t.Fatalf("get: %v", err) } resp.Body.Close() if resp.StatusCode != http.StatusMethodNotAllowed { t.Errorf("GET /logs status=%d want 405", resp.StatusCode) } } // ====================================================================== // Poster endpoints // ====================================================================== var minimalPNG = []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a, 0x01, 0x02, 0x03} type countingReadCloser struct { reader *bytes.Reader reads atomic.Int32 } func newCountingReadCloser(data []byte) *countingReadCloser { return &countingReadCloser{reader: bytes.NewReader(data)} } func (r *countingReadCloser) Read(p []byte) (int, error) { r.reads.Add(1) return r.reader.Read(p) } func (r *countingReadCloser) Close() error { return nil } type blockingReadCloser struct { data []byte offset int started chan struct{} release <-chan struct{} once sync.Once } func (r *blockingReadCloser) Read(p []byte) (int, error) { r.once.Do(func() { close(r.started) }) <-r.release if r.offset == len(r.data) { return 0, io.EOF } n := copy(p, r.data[r.offset:]) r.offset += n return n, nil } func (r *blockingReadCloser) Close() error { return nil } type deadlineBlockingReadCloser struct { started chan struct{} closed chan struct{} startOnce sync.Once closeOnce sync.Once } func newDeadlineBlockingReadCloser() *deadlineBlockingReadCloser { return &deadlineBlockingReadCloser{ started: make(chan struct{}), closed: make(chan struct{}), } } func (r *deadlineBlockingReadCloser) Read([]byte) (int, error) { r.startOnce.Do(func() { close(r.started) }) <-r.closed return 0, errors.New("body closed") } func (r *deadlineBlockingReadCloser) Close() error { r.closeOnce.Do(func() { close(r.closed) }) return nil } type failingReadCloser struct{} func (failingReadCloser) Read([]byte) (int, error) { return 0, errors.New("read failed") } func (failingReadCloser) Close() error { return nil } func servePosterUpload(s *Server, body io.ReadCloser, xff string) *httptest.ResponseRecorder { req := httptest.NewRequest(http.MethodPost, "/posters", body) req.RemoteAddr = "198.51.100.10:1234" if xff != "" { req.Header.Set("X-Forwarded-For", xff) } recorder := httptest.NewRecorder() s.handlePostPosters(recorder, req) return recorder } func snapshotPosterStore(t *testing.T, store *posterStore) (int, int64, []string) { t.Helper() store.mu.RLock() entryCount := len(store.entries) totalBytes := store.totalBytes store.mu.RUnlock() files, err := os.ReadDir(store.dir) if err != nil { t.Fatalf("read poster dir: %v", err) } names := make([]string, len(files)) for i, file := range files { names[i] = file.Name() } return entryCount, totalBytes, names } func TestPosterHandlerRejectsRateLimitedRequestBeforeReadingOrStoring(t *testing.T) { s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) now := time.Now() s.posterUploads = newPosterUploadLimiter(1, 0, 10, 0, 2, now) first := servePosterUpload(s, io.NopCloser(bytes.NewReader(minimalPNG)), "203.0.113.1") if first.Code != http.StatusOK { t.Fatalf("first status=%d", first.Code) } beforeEntries, beforeBytes, beforeFiles := snapshotPosterStore(t, s.posters) rejectedBody := newCountingReadCloser(minimalPNG) rejected := servePosterUpload(s, rejectedBody, "203.0.113.2") if rejected.Code != http.StatusTooManyRequests { t.Fatalf("rejected status=%d, want 429", rejected.Code) } if rejectedBody.reads.Load() != 0 { t.Fatalf("rate-limited body read %d times", rejectedBody.reads.Load()) } afterEntries, afterBytes, afterFiles := snapshotPosterStore(t, s.posters) if beforeEntries != afterEntries || beforeBytes != afterBytes || fmt.Sprint(beforeFiles) != fmt.Sprint(afterFiles) { t.Fatalf("denial mutated poster store: before=(%d,%d,%v) after=(%d,%d,%v)", beforeEntries, beforeBytes, beforeFiles, afterEntries, afterBytes, afterFiles) } } func TestPosterHandlerConcurrencyRejectsBeforeReadAndRecovers(t *testing.T) { s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) s.posterUploads = newPosterUploadLimiter(20, 0, 20, 0, 2, time.Now()) release := make(chan struct{}) recorders := make(chan *httptest.ResponseRecorder, 2) for range 2 { body := &blockingReadCloser{ data: minimalPNG, started: make(chan struct{}), release: release, } go func() { recorders <- servePosterUpload(s, body, "") }() select { case <-body.started: case <-time.After(time.Second): t.Fatal("admitted body was not read") } } extraBody := newCountingReadCloser(minimalPNG) extra := servePosterUpload(s, extraBody, "") if extra.Code != http.StatusTooManyRequests { t.Fatalf("extra status=%d, want 429", extra.Code) } if extraBody.reads.Load() != 0 { t.Fatalf("concurrency-rejected body read %d times", extraBody.reads.Load()) } close(release) for range 2 { select { case recorder := <-recorders: if recorder.Code != http.StatusOK { t.Fatalf("admitted status=%d", recorder.Code) } case <-time.After(time.Second): t.Fatal("admitted upload did not finish") } } recovered := servePosterUpload(s, io.NopCloser(bytes.NewReader(minimalPNG)), "") if recovered.Code != http.StatusOK { t.Fatalf("post-completion status=%d, want 200", recovered.Code) } s.posterUploads.mu.Lock() active := s.posterUploads.active s.posterUploads.mu.Unlock() if active != 0 { t.Fatalf("active=%d after completion, want 0", active) } } func TestPosterHandlerDeadlineReleasesStalledUploadSlot(t *testing.T) { s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) s.posterUploads = newPosterUploadLimiter(10, 0, 10, 0, 1, time.Now()) s.posterBodyReadTimeout = 20 * time.Millisecond stalled := newDeadlineBlockingReadCloser() result := make(chan *httptest.ResponseRecorder, 1) go func() { result <- servePosterUpload(s, stalled, "") }() select { case <-stalled.started: case <-time.After(time.Second): t.Fatal("stalled body was not read") } var timedOut *httptest.ResponseRecorder select { case timedOut = <-result: case <-time.After(time.Second): t.Fatal("stalled upload did not honor body deadline") } if timedOut.Code != http.StatusRequestTimeout { t.Fatalf("stalled status=%d, want 408", timedOut.Code) } s.posterUploads.mu.Lock() active := s.posterUploads.active s.posterUploads.mu.Unlock() if active != 0 { t.Fatalf("active=%d after body timeout, want 0", active) } recovered := servePosterUpload(s, io.NopCloser(bytes.NewReader(minimalPNG)), "") if recovered.Code != http.StatusOK { t.Fatalf("post-timeout status=%d, want 200", recovered.Code) } } func TestPosterHandlerSlowChunkedBodyDeadline(t *testing.T) { s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) s.posterUploads = newPosterUploadLimiter(10, 0, 10, 0, 1, time.Now()) s.posterBodyReadTimeout = 30 * time.Millisecond httpServer := httptest.NewServer(http.HandlerFunc(s.handlePostPosters)) t.Cleanup(httpServer.Close) address := strings.TrimPrefix(httpServer.URL, "http://") conn, err := net.DialTimeout("tcp", address, time.Second) if err != nil { t.Fatalf("dial: %v", err) } if _, err := fmt.Fprintf( conn, "POST /posters HTTP/1.1\r\nHost: %s\r\nTransfer-Encoding: chunked\r\n\r\n1\r\nx\r\n", address, ); err != nil { conn.Close() t.Fatalf("write partial chunked request: %v", err) } if err := conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil { conn.Close() t.Fatalf("set response deadline: %v", err) } response, err := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodPost}) if err != nil { conn.Close() t.Fatalf("read timeout response: %v", err) } response.Body.Close() conn.Close() if response.StatusCode != http.StatusRequestTimeout { t.Fatalf("slow chunked status=%d, want 408", response.StatusCode) } recovered, err := http.Post(httpServer.URL, "image/png", bytes.NewReader(minimalPNG)) if err != nil { t.Fatalf("post after timeout: %v", err) } recovered.Body.Close() if recovered.StatusCode != http.StatusOK { t.Fatalf("post-timeout status=%d, want 200", recovered.StatusCode) } } func TestPosterHandlerReleasesConcurrencyOnEveryExit(t *testing.T) { tests := []struct { name string body func() io.ReadCloser wantStatus int storeFailure bool }{ {name: "read error", body: func() io.ReadCloser { return failingReadCloser{} }, wantStatus: http.StatusBadRequest}, {name: "oversized", body: func() io.ReadCloser { return io.NopCloser(bytes.NewReader(make([]byte, maxPosterSize+1))) }, wantStatus: http.StatusRequestEntityTooLarge}, {name: "empty", body: func() io.ReadCloser { return io.NopCloser(bytes.NewReader(nil)) }, wantStatus: http.StatusBadRequest}, {name: "unsupported", body: func() io.ReadCloser { return io.NopCloser(strings.NewReader("not an image")) }, wantStatus: http.StatusUnsupportedMediaType}, {name: "store failure", body: func() io.ReadCloser { return io.NopCloser(bytes.NewReader(minimalPNG)) }, wantStatus: http.StatusInternalServerError, storeFailure: true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) s.posterUploads = newPosterUploadLimiter(10, 0, 10, 0, 1, time.Now()) originalDir := s.posters.dir if tt.storeFailure { s.posters.dir = filepath.Join(t.TempDir(), "missing", "posters") } failed := servePosterUpload(s, tt.body(), "") if failed.Code != tt.wantStatus { t.Fatalf("status=%d, want %d", failed.Code, tt.wantStatus) } s.posters.dir = originalDir recovery := servePosterUpload(s, io.NopCloser(bytes.NewReader(minimalPNG)), "") if recovery.Code != http.StatusOK { t.Fatalf("recovery status=%d, want 200", recovery.Code) } s.posterUploads.mu.Lock() active := s.posterUploads.active s.posterUploads.mu.Unlock() if active != 0 { t.Fatalf("active=%d, want 0", active) } }) } } func TestPosterHandlerUsesTrustedCanonicalIdentityAndGlobalBudget(t *testing.T) { t.Run("untrusted XFF rotation cannot bypass per-IP limit", func(t *testing.T) { h := newRelayHarnessNoTrust(t) h.srv.posterUploads = newPosterUploadLimiter(3, 0, 20, 0, 4, time.Now()) for i := range 3 { resp := postPoster(t, h.baseURL, fmt.Sprintf("203.0.113.%d", i+1), minimalPNG) resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("upload %d status=%d", i, resp.StatusCode) } } beforeEntries, beforeBytes, beforeFiles := snapshotPosterStore(t, h.srv.posters) denied := postPoster(t, h.baseURL, "203.0.113.99", minimalPNG) denied.Body.Close() if denied.StatusCode != http.StatusTooManyRequests { t.Fatalf("rotated spoof status=%d, want 429", denied.StatusCode) } afterEntries, afterBytes, afterFiles := snapshotPosterStore(t, h.srv.posters) if beforeEntries != afterEntries || beforeBytes != afterBytes || fmt.Sprint(beforeFiles) != fmt.Sprint(afterFiles) { t.Fatal("per-IP denial mutated poster store") } }) t.Run("trusted clients are independent but share global budget", func(t *testing.T) { h := newRelayHarness(t) h.srv.posterUploads = newPosterUploadLimiter(3, 0, 8, 0, 4, time.Now()) for range 3 { resp := postPoster(t, h.baseURL, "203.0.113.1", minimalPNG) resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("client A status=%d", resp.StatusCode) } } perIPDenied := postPoster(t, h.baseURL, "203.0.113.1", minimalPNG) perIPDenied.Body.Close() if perIPDenied.StatusCode != http.StatusTooManyRequests { t.Fatalf("client A overflow status=%d, want 429", perIPDenied.StatusCode) } for _, ip := range []string{"203.0.113.2", "203.0.113.2", "203.0.113.2", "203.0.113.3", "203.0.113.3"} { resp := postPoster(t, h.baseURL, ip, minimalPNG) resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("client %s status=%d before global exhaustion", ip, resp.StatusCode) } } beforeEntries, beforeBytes, beforeFiles := snapshotPosterStore(t, h.srv.posters) globalDenied := postPoster(t, h.baseURL, "203.0.113.4", minimalPNG) globalDenied.Body.Close() if globalDenied.StatusCode != http.StatusTooManyRequests { t.Fatalf("global overflow status=%d, want 429", globalDenied.StatusCode) } afterEntries, afterBytes, afterFiles := snapshotPosterStore(t, h.srv.posters) if beforeEntries != afterEntries || beforeBytes != afterBytes || fmt.Sprint(beforeFiles) != fmt.Sprint(afterFiles) { t.Fatal("global denial mutated poster store") } }) } func TestPosterHandlerMalformedTrustedChainMutatesNothing(t *testing.T) { s := newTestServer(t, filepath.Join(t.TempDir(), "rooms.json")) s.clientIPs = mustClientIPResolver(t, "10.0.0.0/8") body := newCountingReadCloser(minimalPNG) req := httptest.NewRequest(http.MethodPost, "/posters", body) req.RemoteAddr = "10.0.0.2:1234" req.Header.Set("X-Forwarded-For", "203.0.113.1,") recorder := httptest.NewRecorder() s.handlePostPosters(recorder, req) if recorder.Code != http.StatusBadRequest { t.Fatalf("status=%d, want 400", recorder.Code) } if body.reads.Load() != 0 { t.Fatalf("malformed-chain body read %d times", body.reads.Load()) } s.posterUploads.mu.Lock() active := s.posterUploads.active perIP := len(s.posterUploads.perIP) s.posterUploads.global.mu.Lock() globalTokens := s.posterUploads.global.tokens s.posterUploads.global.mu.Unlock() s.posterUploads.mu.Unlock() if active != 0 || perIP != 0 || globalTokens != posterGlobalRateBurst { t.Fatalf("malformed chain mutated admission: active=%d perIP=%d global=%v", active, perIP, globalTokens) } entries, total, files := snapshotPosterStore(t, s.posters) if entries != 0 || total != 0 || len(files) != 0 { t.Fatalf("malformed chain mutated store: entries=%d total=%d files=%v", entries, total, files) } } func TestPosterHandlerIgnoresMalformedHeaderFromUntrustedPeer(t *testing.T) { h := newRelayHarnessNoTrust(t) resp := postPoster(t, h.baseURL, "bad,", minimalPNG) resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Fatalf("status=%d, want 200", resp.StatusCode) } } func TestPostersRoundTrip(t *testing.T) { h := newRelayHarness(t) payload := []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a, 0x01, 0x02, 0x03} out := postPosterAndDecode(t, h.baseURL, "9.0.0.1", payload) if out.ExpiresIn != int(posterMaxAge.Seconds()) { t.Fatalf("expiresIn=%d want %d", out.ExpiresIn, int(posterMaxAge.Seconds())) } if !strings.HasPrefix(out.URL, "/posters/") || !strings.HasSuffix(out.URL, ".png") { t.Fatalf("url=%q should be a relative png poster path", out.URL) } if strings.Contains(out.URL, "://") { t.Fatalf("url=%q should be relative", out.URL) } getResp, err := http.Get(h.baseURL + out.URL) if err != nil { t.Fatalf("get: %v", err) } defer getResp.Body.Close() if getResp.StatusCode != http.StatusOK { t.Fatalf("get status=%d", getResp.StatusCode) } got, _ := io.ReadAll(getResp.Body) if !bytes.Equal(got, payload) { t.Fatalf("round-tripped bytes mismatch: got %v want %v", got, payload) } if ct := getResp.Header.Get("Content-Type"); !strings.HasPrefix(ct, "image/png") { t.Errorf("Content-Type=%q", ct) } } func TestPostersRejectInvalidAndOversizedUploads(t *testing.T) { h := newRelayHarness(t) invalid := postPoster(t, h.baseURL, "9.0.0.2", []byte("not an image")) invalid.Body.Close() if invalid.StatusCode != http.StatusUnsupportedMediaType { t.Fatalf("invalid status=%d want 415", invalid.StatusCode) } oversized := postPoster(t, h.baseURL, "9.0.0.3", make([]byte, maxPosterSize+1)) oversized.Body.Close() if oversized.StatusCode != http.StatusRequestEntityTooLarge { t.Fatalf("oversized status=%d want 413", oversized.StatusCode) } } func TestPosterStoreEvictsOldestOverQuota(t *testing.T) { ps := newPosterStore(t.TempDir(), 12, time.Hour) now := time.Now() payload := []byte{1, 2, 3, 4, 5, 6, 7} id1, entry1, err := ps.store(payload, "image/png", now) if err != nil { t.Fatalf("store first: %v", err) } id2, entry2, err := ps.store(payload, "image/png", now.Add(time.Minute)) if err != nil { t.Fatalf("store second: %v", err) } ps.mu.RLock() _, hasFirst := ps.entries[id1] _, hasSecond := ps.entries[id2] total := ps.totalBytes ps.mu.RUnlock() if hasFirst { t.Fatal("oldest poster should have been evicted") } if !hasSecond { t.Fatal("newest poster should remain") } if total != int64(len(payload)) { t.Fatalf("totalBytes=%d want %d", total, len(payload)) } if _, err := os.Stat(ps.filePath(entry1.Filename)); !os.IsNotExist(err) { t.Fatalf("oldest file still exists or stat failed unexpectedly: %v", err) } if _, err := os.Stat(ps.filePath(entry2.Filename)); err != nil { t.Fatalf("newest file missing: %v", err) } } func TestPosterStoreCleanupExpiresOldPosters(t *testing.T) { ps := newPosterStore(t.TempDir(), 1024, time.Hour) now := time.Now() id, entry, err := ps.store([]byte{1, 2, 3}, "image/png", now.Add(-2*time.Hour)) if err != nil { t.Fatalf("store: %v", err) } if err := ps.cleanup(now); err != nil { t.Fatalf("cleanup: %v", err) } ps.mu.RLock() _, exists := ps.entries[id] total := ps.totalBytes ps.mu.RUnlock() if exists { t.Fatal("expired poster should have been removed") } if total != 0 { t.Fatalf("totalBytes=%d want 0", total) } if _, err := os.Stat(ps.filePath(entry.Filename)); !os.IsNotExist(err) { t.Fatalf("expired file still exists or stat failed unexpectedly: %v", err) } } func regularFileBytes(t *testing.T, dir string) int64 { t.Helper() files, err := os.ReadDir(dir) if err != nil { t.Fatalf("read directory: %v", err) } var total int64 for _, file := range files { info, err := file.Info() if err != nil { t.Fatalf("stat %s: %v", file.Name(), err) } if info.Mode().IsRegular() { total += info.Size() } } return total } func TestLogStoreRemovalFailureRetainsEntryUntilRetry(t *testing.T) { dir := t.TempDir() remover := newDeterministicRemover() ls := newLogStoreWithRemover(dir, remover.remove) ls.generateID = func() string { return strings.Repeat("a", logIDLength) } now := time.Now() id, _, err := ls.store([]byte("retained"), now) if err != nil { t.Fatalf("store: %v", err) } path := ls.filePath(id) ls.mu.Lock() entry := ls.entries[id] entry.ExpiresAt = now.Add(-time.Minute) ls.entries[id] = entry ls.mu.Unlock() remover.fail(path, fs.ErrPermission) if _, ok, err := ls.lookup(id, now); !errors.Is(err, fs.ErrPermission) || ok { t.Fatalf("lookup=(ok=%v, err=%v), want unavailable permission error", ok, err) } ls.mu.RLock() _, indexed := ls.entries[id] artifacts := ls.artifactCountLocked() ls.mu.RUnlock() if !indexed || artifacts != 1 { t.Fatalf("failed removal changed metadata: indexed=%v artifacts=%d", indexed, artifacts) } if _, err := os.Stat(path); err != nil { t.Fatalf("failed removal lost file: %v", err) } remover.recover(path) if err := ls.cleanup(now); err != nil { t.Fatalf("retry cleanup: %v", err) } if remover.callCount(path) != 2 { t.Fatalf("remove calls=%d want 2", remover.callCount(path)) } if _, err := os.Stat(path); !errors.Is(err, fs.ErrNotExist) { t.Fatalf("file remains after retry: %v", err) } if err := ls.cleanup(now); err != nil { t.Fatalf("idempotent cleanup: %v", err) } if remover.callCount(path) != 2 { t.Fatalf("already committed entry removed again: calls=%d", remover.callCount(path)) } } func TestRemovalFailureDoesNotBlockUploadsWhileCapacityRemains(t *testing.T) { t.Run("logs", func(t *testing.T) { dir := t.TempDir() remover := newDeterministicRemover() store := newLogStoreWithRemover(dir, remover.remove) ids := []string{strings.Repeat("a", logIDLength), strings.Repeat("b", logIDLength)} nextID := 0 store.generateID = func() string { id := ids[nextID] nextID++ return id } now := time.Now() firstID, _, err := store.store([]byte("expired"), now) if err != nil { t.Fatalf("store expired log: %v", err) } store.mu.Lock() entry := store.entries[firstID] entry.ExpiresAt = now.Add(-time.Minute) store.entries[firstID] = entry store.mu.Unlock() remover.fail(store.filePath(firstID), fs.ErrPermission) secondID, _, err := store.store([]byte("new"), now) if err != nil { t.Fatalf("unrelated removal failure blocked log upload: %v", err) } store.mu.RLock() _, firstRetained := store.entries[firstID] _, secondStored := store.entries[secondID] store.mu.RUnlock() if !firstRetained || !secondStored { t.Fatalf("log accounting lost entries: first=%v second=%v", firstRetained, secondStored) } }) t.Run("posters", func(t *testing.T) { dir := t.TempDir() remover := newDeterministicRemover() store := newPosterStoreWithRemover(dir, 1024, time.Hour, remover.remove) now := time.Now() firstID, first, err := store.store([]byte{1, 2, 3}, "image/png", now.Add(-2*time.Hour)) if err != nil { t.Fatalf("store expired poster: %v", err) } remover.fail(store.filePath(first.Filename), fs.ErrPermission) secondID, _, err := store.store([]byte{4, 5, 6}, "image/png", now) if err != nil { t.Fatalf("unrelated removal failure blocked poster upload: %v", err) } store.mu.RLock() _, firstRetained := store.entries[firstID] _, secondStored := store.entries[secondID] store.mu.RUnlock() if !firstRetained || !secondStored { t.Fatalf("poster accounting lost entries: first=%v second=%v", firstRetained, secondStored) } }) } func TestLogStoreErrNotExistCommitsDeletionOnce(t *testing.T) { remover := newDeterministicRemover() ls := newLogStoreWithRemover(t.TempDir(), remover.remove) ls.generateID = func() string { return strings.Repeat("a", logIDLength) } now := time.Now() id, _, err := ls.store([]byte("gone"), now) if err != nil { t.Fatalf("store: %v", err) } path := ls.filePath(id) if err := os.Remove(path); err != nil { t.Fatalf("external remove: %v", err) } ls.mu.Lock() entry := ls.entries[id] entry.ExpiresAt = now.Add(-time.Minute) ls.entries[id] = entry ls.mu.Unlock() if _, ok, err := ls.lookup(id, now); err != nil || ok { t.Fatalf("lookup=(ok=%v, err=%v), want clean miss", ok, err) } if err := ls.cleanup(now); err != nil { t.Fatalf("repeat cleanup: %v", err) } if remover.callCount(path) != 1 { t.Fatalf("remove calls=%d want 1", remover.callCount(path)) } } func TestLogStoreTracksFailedTempCleanup(t *testing.T) { dir := t.TempDir() remover := newDeterministicRemover() ls := newLogStoreWithRemover(dir, remover.remove) logID := strings.Repeat("a", logIDLength) ls.generateID = func() string { return logID } tmpPath := ls.filePath(logID) + ".tmp" if err := os.Mkdir(tmpPath, 0755); err != nil { t.Fatalf("seed temp directory: %v", err) } cleanupErr := errors.New("synthetic temp removal failure") remover.fail(tmpPath, cleanupErr) if _, _, err := ls.store([]byte("payload"), time.Now()); err == nil { t.Fatal("store succeeded despite temp write failure") } ls.mu.RLock() _, pending := ls.pendingRemovals[filepath.Base(tmpPath)] artifacts := ls.artifactCountLocked() ls.mu.RUnlock() if !pending || artifacts != 1 { t.Fatalf("temp cleanup not tracked: pending=%v artifacts=%d", pending, artifacts) } remover.recover(tmpPath) if err := ls.cleanup(time.Now()); err != nil { t.Fatalf("retry temp cleanup: %v", err) } if _, err := os.Stat(tmpPath); !errors.Is(err, fs.ErrNotExist) { t.Fatalf("temp artifact remains: %v", err) } } func TestLogStoreStartupReconcilesLiveAndPendingRemovals(t *testing.T) { dir := t.TempDir() now := time.Now() expiredID := strings.Repeat("a", logIDLength) expiredPath := filepath.Join(dir, expiredID+".log") tempPath := filepath.Join(dir, "upload.log.tmp") malformedPath := filepath.Join(dir, "malformed") for path, data := range map[string][]byte{ expiredPath: []byte("expired"), tempPath: []byte("partial"), malformedPath: []byte("invalid"), } { if err := os.WriteFile(path, data, 0644); err != nil { t.Fatalf("seed %s: %v", filepath.Base(path), err) } } old := now.Add(-logMaxAge - time.Hour) if err := os.Chtimes(expiredPath, old, old); err != nil { t.Fatalf("age expired log: %v", err) } remover := newDeterministicRemover() for _, path := range []string{expiredPath, tempPath, malformedPath} { remover.fail(path, fs.ErrPermission) } ls := newLogStoreWithRemover(dir, remover.remove) if ls.startupErr == nil { t.Fatal("startup removal failures were not reported") } ls.mu.RLock() _, live := ls.entries[expiredID] pending := len(ls.pendingRemovals) artifacts := ls.artifactCountLocked() ls.mu.RUnlock() if !live || pending != 2 || artifacts != 3 { t.Fatalf("startup accounting: live=%v pending=%d artifacts=%d", live, pending, artifacts) } newID, _, err := ls.store([]byte("new"), now) if err != nil { t.Fatalf("startup cleanup failure blocked new log: %v", err) } for _, path := range []string{expiredPath, tempPath, malformedPath} { remover.recover(path) } if err := ls.cleanup(now); err != nil { t.Fatalf("startup retry cleanup: %v", err) } restarted := newLogStore(dir) restarted.mu.RLock() restartedArtifacts := restarted.artifactCountLocked() _, newLogRestored := restarted.entries[newID] restarted.mu.RUnlock() if restartedArtifacts != 1 || !newLogRestored { t.Fatalf("restart reconstructed %d artifacts, new log restored=%v", restartedArtifacts, newLogRestored) } } func TestStoresReconcileConfinedNonEmptyStaleDirectories(t *testing.T) { t.Run("logs", func(t *testing.T) { dir := t.TempDir() staleDir := filepath.Join(dir, "abandoned.log.tmp") if err := os.MkdirAll(filepath.Join(staleDir, "nested"), 0755); err != nil { t.Fatalf("seed stale log directory: %v", err) } if err := os.WriteFile(filepath.Join(staleDir, "nested", "partial"), []byte("stale"), 0644); err != nil { t.Fatalf("seed stale log payload: %v", err) } store := newLogStore(dir) if store.startupErr != nil { t.Fatalf("startup reconciliation: %v", store.startupErr) } if _, err := os.Stat(staleDir); !errors.Is(err, fs.ErrNotExist) { t.Fatalf("stale log directory remains: %v", err) } store.generateID = func() string { return strings.Repeat("a", logIDLength) } if _, _, err := store.store([]byte("new log"), time.Now()); err != nil { t.Fatalf("store after reconciliation: %v", err) } }) t.Run("posters", func(t *testing.T) { dir := t.TempDir() staleDir := filepath.Join(dir, "abandoned.tmp") if err := os.MkdirAll(filepath.Join(staleDir, "nested"), 0755); err != nil { t.Fatalf("seed stale poster directory: %v", err) } if err := os.WriteFile(filepath.Join(staleDir, "nested", "partial"), []byte("stale"), 0644); err != nil { t.Fatalf("seed stale poster payload: %v", err) } store := newPosterStore(dir, 1024, time.Hour) if store.startupErr != nil { t.Fatalf("startup reconciliation: %v", store.startupErr) } if _, err := os.Stat(staleDir); !errors.Is(err, fs.ErrNotExist) { t.Fatalf("stale poster directory remains: %v", err) } if _, _, err := store.store([]byte{1, 2, 3}, "image/png", time.Now()); err != nil { t.Fatalf("store after reconciliation: %v", err) } }) } func TestRecursiveArtifactRemovalRejectsOutsideStore(t *testing.T) { root := t.TempDir() outside := t.TempDir() nested := filepath.Join(outside, "nested") if err := os.Mkdir(nested, 0755); err != nil { t.Fatalf("seed outside directory: %v", err) } if err := os.WriteFile(filepath.Join(nested, "keep"), []byte("keep"), 0644); err != nil { t.Fatalf("seed outside payload: %v", err) } err := removeArtifact(os.Remove, root, nested) if !errors.Is(err, errArtifactOutsideStore) { t.Fatalf("outside removal error=%v, want confinement error", err) } if _, err := os.Stat(filepath.Join(nested, "keep")); err != nil { t.Fatalf("outside artifact was removed: %v", err) } } func TestPosterQuotaRemovalFailureDoesNotReclaimAccounting(t *testing.T) { dir := t.TempDir() remover := newDeterministicRemover() ps := newPosterStoreWithRemover(dir, 12, time.Hour, remover.remove) now := time.Now() payload := []byte{1, 2, 3, 4, 5, 6, 7} oldID, oldEntry, err := ps.store(payload, "image/png", now) if err != nil { t.Fatalf("store oldest: %v", err) } oldPath := ps.filePath(oldEntry.Filename) remover.fail(oldPath, fs.ErrPermission) newID, newEntry, err := ps.store(payload, "image/png", now.Add(time.Minute)) if !errors.Is(err, fs.ErrPermission) { t.Fatalf("quota store error=%v want permission error", err) } if newID != "" || newEntry != (posterEntry{}) { t.Fatalf("failed store returned success values: id=%q entry=%+v", newID, newEntry) } ps.mu.RLock() _, retained := ps.entries[oldID] total := ps.totalBytes pending := ps.pendingBytes accounted := ps.accountedBytesLocked() ps.mu.RUnlock() if !retained || total != int64(len(payload)) || pending != 0 { t.Fatalf("failed eviction accounting: retained=%v total=%d pending=%d", retained, total, pending) } if physical := regularFileBytes(t, dir); physical != accounted { t.Fatalf("accounted bytes=%d physical bytes=%d", accounted, physical) } remover.recover(oldPath) retryID, retryEntry, err := ps.store(payload, "image/png", now.Add(time.Minute)) if err != nil { t.Fatalf("retry store: %v", err) } if retryID == "" || retryEntry.Size != int64(len(payload)) { t.Fatalf("retry result: id=%q entry=%+v", retryID, retryEntry) } if remover.callCount(oldPath) != 2 { t.Fatalf("old poster remove calls=%d want 2", remover.callCount(oldPath)) } ps.mu.RLock() accounted = ps.accountedBytesLocked() total = ps.totalBytes ps.mu.RUnlock() if total != int64(len(payload)) || regularFileBytes(t, dir) != accounted { t.Fatalf("retry accounting: total=%d accounted=%d physical=%d", total, accounted, regularFileBytes(t, dir)) } } func TestPosterExpiredRemovalFailureAndErrNotExistAreExactOnce(t *testing.T) { t.Run("failure retains accounting for retry", func(t *testing.T) { remover := newDeterministicRemover() ps := newPosterStoreWithRemover(t.TempDir(), 1024, time.Hour, remover.remove) now := time.Now() id, entry, err := ps.store([]byte{1, 2, 3}, "image/png", now) if err != nil { t.Fatalf("store: %v", err) } path := ps.filePath(entry.Filename) ps.mu.Lock() expired := ps.entries[id] expired.ExpiresAt = now.Add(-time.Minute) ps.entries[id] = expired ps.mu.Unlock() remover.fail(path, fs.ErrPermission) if _, ok, err := ps.lookup(entry.Filename, now); !errors.Is(err, fs.ErrPermission) || ok { t.Fatalf("lookup=(ok=%v, err=%v), want unavailable permission error", ok, err) } ps.mu.RLock() _, retained := ps.entries[id] total := ps.totalBytes ps.mu.RUnlock() if !retained || total != entry.Size { t.Fatalf("failed expiry accounting: retained=%v total=%d", retained, total) } remover.recover(path) if err := ps.cleanup(now); err != nil { t.Fatalf("retry cleanup: %v", err) } if err := ps.cleanup(now); err != nil { t.Fatalf("repeat cleanup: %v", err) } if remover.callCount(path) != 2 { t.Fatalf("remove calls=%d want 2", remover.callCount(path)) } ps.mu.RLock() total = ps.totalBytes ps.mu.RUnlock() if total != 0 { t.Fatalf("totalBytes=%d want 0", total) } }) t.Run("not exist commits once", func(t *testing.T) { remover := newDeterministicRemover() ps := newPosterStoreWithRemover(t.TempDir(), 1024, time.Hour, remover.remove) now := time.Now() id, entry, err := ps.store([]byte{1, 2, 3}, "image/png", now) if err != nil { t.Fatalf("store: %v", err) } path := ps.filePath(entry.Filename) if err := os.Remove(path); err != nil { t.Fatalf("external remove: %v", err) } ps.mu.Lock() expired := ps.entries[id] expired.ExpiresAt = now.Add(-time.Minute) ps.entries[id] = expired ps.mu.Unlock() if _, ok, err := ps.lookup(entry.Filename, now); err != nil || ok { t.Fatalf("lookup=(ok=%v, err=%v), want clean miss", ok, err) } if err := ps.cleanup(now); err != nil { t.Fatalf("repeat cleanup: %v", err) } if remover.callCount(path) != 1 { t.Fatalf("remove calls=%d want 1", remover.callCount(path)) } ps.mu.RLock() total := ps.totalBytes ps.mu.RUnlock() if total != 0 { t.Fatalf("totalBytes=%d want 0", total) } }) } func TestPosterStoreKnownCleanupDebtConsumesCapacityAndRetries(t *testing.T) { dir := t.TempDir() stalePath := filepath.Join(dir, "poster.tmp") if err := os.WriteFile(stalePath, []byte("1234"), 0644); err != nil { t.Fatalf("seed stale poster: %v", err) } remover := newDeterministicRemover() remover.fail(stalePath, fs.ErrPermission) ps := newPosterStoreWithRemover(dir, 5, time.Hour, remover.remove) if ps.startupErr == nil { t.Fatal("startup removal failure was not reported") } if _, _, err := ps.store([]byte{1, 2}, "image/png", time.Now()); err == nil { t.Fatal("upload exceeded capacity after known stale bytes were accounted") } ps.mu.RLock() pendingBytes := ps.pendingBytes accountedBytes := ps.accountedBytesLocked() ps.mu.RUnlock() if pendingBytes != 4 || accountedBytes != 4 { t.Fatalf("known debt accounting: pending=%d accounted=%d, want 4", pendingBytes, accountedBytes) } if calls := remover.callCount(stalePath); calls != 2 { t.Fatalf("known debt remove calls=%d, want startup plus upload retry", calls) } remover.recover(stalePath) if _, entry, err := ps.store([]byte{1, 2}, "image/png", time.Now()); err != nil { t.Fatalf("store after known debt recovery: %v", err) } else if entry.Size != 2 { t.Fatalf("stored entry size=%d, want 2", entry.Size) } ps.mu.RLock() pendingBytes = ps.pendingBytes ps.mu.RUnlock() if pendingBytes != 0 { t.Fatalf("known debt remained after successful retry: %d bytes", pendingBytes) } } func TestPosterStoreUnknownCleanupDebtDoesNotBlockUploadAndRecovers(t *testing.T) { dir := t.TempDir() unknownPath := filepath.Join(dir, "unknown-dir") if err := os.Mkdir(unknownPath, 0755); err != nil { t.Fatalf("seed unknown artifact: %v", err) } remover := newDeterministicRemover() remover.fail(unknownPath, fs.ErrPermission) ps := newPosterStoreWithRemover(dir, 5, time.Hour, remover.remove) if ps.startupErr == nil { t.Fatal("startup removal failure was not reported") } if calls := remover.callCount(unknownPath); calls != 1 { t.Fatalf("startup remove calls=%d, want 1", calls) } if _, entry, err := ps.store([]byte{1, 2, 3}, "image/png", time.Now()); err != nil { t.Fatalf("capacity-safe upload blocked by unknown artifact: %v", err) } else if entry.Size != 3 { t.Fatalf("stored entry size=%d, want 3", entry.Size) } if calls := remover.callCount(unknownPath); calls != 1 { t.Fatalf("upload retried permanent unknown debt: calls=%d", calls) } if err := ps.cleanup(time.Now()); !errors.Is(err, fs.ErrPermission) { t.Fatalf("failed cleanup error=%v, want permission error", err) } ps.mu.RLock() unknownPending := ps.unknownPending ps.mu.RUnlock() if unknownPending != 1 { t.Fatalf("unknown debt count=%d, want 1", unknownPending) } remover.recover(unknownPath) if err := ps.cleanup(time.Now()); err != nil { t.Fatalf("cleanup after recovery: %v", err) } ps.mu.RLock() unknownPending = ps.unknownPending pendingCount := len(ps.pendingRemovals) ps.mu.RUnlock() if unknownPending != 0 || pendingCount != 0 { t.Fatalf("recovered debt remains: unknown=%d pending=%d", unknownPending, pendingCount) } if _, err := os.Stat(unknownPath); !errors.Is(err, fs.ErrNotExist) { t.Fatalf("unknown artifact remains after recovery: %v", err) } } func TestStorageHandlersReturnGenericErrorsForRemovalFailures(t *testing.T) { remover := newDeterministicRemover() logs := newLogStoreWithRemover(t.TempDir(), remover.remove) logs.generateID = func() string { return strings.Repeat("a", logIDLength) } posters := newPosterStoreWithRemover(t.TempDir(), 1024, time.Hour, remover.remove) h := newStorageHarness(t, logs, posters) now := time.Now() logID, _, err := logs.store([]byte("expired log"), now) if err != nil { t.Fatalf("store log: %v", err) } posterID, poster, err := posters.store([]byte{1, 2, 3}, "image/png", now) if err != nil { t.Fatalf("store poster: %v", err) } logs.mu.Lock() logEntry := logs.entries[logID] logEntry.ExpiresAt = now.Add(-time.Minute) logs.entries[logID] = logEntry logs.mu.Unlock() posters.mu.Lock() posterEntry := posters.entries[posterID] posterEntry.ExpiresAt = now.Add(-time.Minute) posters.entries[posterID] = posterEntry posters.mu.Unlock() logPath := logs.filePath(logID) posterPath := posters.filePath(poster.Filename) remover.fail(logPath, fs.ErrPermission) remover.fail(posterPath, fs.ErrPermission) for name, target := range map[string]string{ "log": h.baseURL + "/logs/" + logID, "poster": h.baseURL + "/posters/" + poster.Filename, } { resp, err := http.Get(target) if err != nil { t.Fatalf("%s get: %v", name, err) } body, readErr := io.ReadAll(resp.Body) resp.Body.Close() if readErr != nil { t.Fatalf("%s response body: %v", name, readErr) } if resp.StatusCode != http.StatusInternalServerError { t.Fatalf("%s status=%d want 500", name, resp.StatusCode) } want := "Failed to retrieve " + name + "\n" if string(body) != want { t.Fatalf("%s response=%q want %q", name, body, want) } } remover.recover(logPath) remover.recover(posterPath) for name, target := range map[string]string{ "log": h.baseURL + "/logs/" + logID, "poster": h.baseURL + "/posters/" + poster.Filename, } { resp, err := http.Get(target) if err != nil { t.Fatalf("%s recovery get: %v", name, err) } resp.Body.Close() if resp.StatusCode != http.StatusNotFound { t.Fatalf("%s recovery status=%d want 404", name, resp.StatusCode) } } } func TestPosterHandlerRejectsUploadWhenQuotaRemovalFails(t *testing.T) { remover := newDeterministicRemover() logs := newLogStoreWithRemover(t.TempDir(), remover.remove) payload := []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a, 1, 2, 3} posters := newPosterStoreWithRemover(t.TempDir(), int64(len(payload)+1), time.Hour, remover.remove) now := time.Now() oldID, oldEntry, err := posters.store(payload, "image/png", now) if err != nil { t.Fatalf("store old poster: %v", err) } oldPath := posters.filePath(oldEntry.Filename) remover.fail(oldPath, fs.ErrPermission) h := newStorageHarness(t, logs, posters) resp := postPoster(t, h.baseURL, "9.9.9.9", payload) body, readErr := io.ReadAll(resp.Body) resp.Body.Close() if readErr != nil { t.Fatalf("read failed upload response: %v", readErr) } if resp.StatusCode != http.StatusInternalServerError || string(body) != "Failed to store poster\n" { t.Fatalf("failed upload status=%d body=%q", resp.StatusCode, body) } posters.mu.RLock() _, retained := posters.entries[oldID] total := posters.totalBytes posters.mu.RUnlock() if !retained || total != int64(len(payload)) { t.Fatalf("failed upload changed old poster: retained=%v total=%d", retained, total) } get, err := http.Get(h.baseURL + "/posters/" + oldEntry.Filename) if err != nil { t.Fatalf("get retained poster: %v", err) } get.Body.Close() if get.StatusCode != http.StatusOK { t.Fatalf("retained poster status=%d want 200", get.StatusCode) } } func TestCleanupStepContinuesAfterRemovalFailureAndThrottlesLogging(t *testing.T) { remover := newDeterministicRemover() logs := newLogStoreWithRemover(t.TempDir(), remover.remove) logs.generateID = func() string { return strings.Repeat("a", logIDLength) } posters := newPosterStoreWithRemover(t.TempDir(), 1024, time.Hour, remover.remove) now := time.Now() logID, _, err := logs.store([]byte("expired"), now) if err != nil { t.Fatalf("store log: %v", err) } posterID, poster, err := posters.store([]byte{1, 2, 3}, "image/png", now) if err != nil { t.Fatalf("store poster: %v", err) } logs.mu.Lock() logEntry := logs.entries[logID] logEntry.ExpiresAt = now.Add(-time.Minute) logs.entries[logID] = logEntry logs.mu.Unlock() posters.mu.Lock() posterEntry := posters.entries[posterID] posterEntry.ExpiresAt = now.Add(-time.Minute) posters.entries[posterID] = posterEntry posters.mu.Unlock() logPath := logs.filePath(logID) remover.fail(logPath, fs.ErrPermission) srv := &Server{ rooms: make(map[string]*Room), logs: logs, posters: posters, posterUploads: newPosterUploadLimiter(posterPerIPRateBurst, posterPerIPRateSustained, posterGlobalRateBurst, posterGlobalRateSustained, maxConcurrentPosterUploads, now), conns: newConnTracker(), } srv.runCleanupStep(now) logs.mu.RLock() _, logRetained := logs.entries[logID] logs.mu.RUnlock() posters.mu.RLock() _, posterRetained := posters.entries[posterID] posters.mu.RUnlock() if !logRetained || posterRetained { t.Fatalf("cleanup continuation: log retained=%v poster retained=%v", logRetained, posterRetained) } if _, err := os.Stat(posters.filePath(poster.Filename)); !errors.Is(err, fs.ErrNotExist) { t.Fatalf("poster cleanup did not continue: %v", err) } srv.removalErrors.mu.Lock() firstLog := srv.removalErrors.lastLog["logs:cleanup"] srv.removalErrors.mu.Unlock() if firstLog.IsZero() { t.Fatal("cleanup removal failure was not made operationally visible") } srv.runCleanupStep(now.Add(time.Minute)) srv.removalErrors.mu.Lock() secondLog := srv.removalErrors.lastLog["logs:cleanup"] srv.removalErrors.mu.Unlock() if !secondLog.Equal(firstLog) { t.Fatalf("persistent cleanup failure was not throttled: first=%v second=%v", firstLog, secondLog) } if remover.callCount(logPath) != 2 { t.Fatalf("cleanup retry calls=%d want 2", remover.callCount(logPath)) } } func TestRemovalFailureLogDoesNotExposeCapabilityPath(t *testing.T) { dir := t.TempDir() remover := newDeterministicRemover() store := newLogStoreWithRemover(dir, remover.remove) id := strings.Repeat("c", logIDLength) store.generateID = func() string { return id } now := time.Now() if _, _, err := store.store([]byte("sensitive"), now); err != nil { t.Fatalf("store: %v", err) } path := store.filePath(id) store.mu.Lock() entry := store.entries[id] entry.ExpiresAt = now.Add(-time.Minute) store.entries[id] = entry store.mu.Unlock() remover.fail(path, &os.PathError{Op: "remove", Path: path, Err: syscall.EACCES}) removalErr := store.cleanup(now) if removalErr == nil { t.Fatal("cleanup unexpectedly succeeded") } var output bytes.Buffer previousOutput := log.Writer() previousFlags := log.Flags() previousPrefix := log.Prefix() log.SetOutput(&output) log.SetFlags(0) log.SetPrefix("") t.Cleanup(func() { log.SetOutput(previousOutput) log.SetFlags(previousFlags) log.SetPrefix(previousPrefix) }) srv := &Server{} srv.logRemovalError("logs", "cleanup", removalErr) message := output.String() if strings.Contains(message, id) || strings.Contains(message, path) { t.Fatalf("removal log exposed capability path: %q", message) } want := fmt.Sprintf("logs: cleanup removal failed: category=permission errno=%d", syscall.EACCES) if !strings.Contains(message, want) { t.Fatalf("removal log=%q, want sanitized context %q", message, want) } } // ====================================================================== // End-to-end: rooms survive a process restart // ====================================================================== func TestSnapshotSurvivesRestartWithHostAuthority(t *testing.T) { stateFile := filepath.Join(t.TempDir(), "rooms.json") hA := newRelayHarnessAt(t, t.TempDir(), stateFile) host := hA.dial(t, "8.0.0.1") host.send(clientMsg{Type: relayTypeCreate, SessionID: "RESUM", PeerID: "H"}) created := host.expectAuthority(relayTypeCreated, "H") if err := hA.srv.snap.flushAndStop(2 * time.Second); err != nil { t.Fatalf("flushAndStop: %v", err) } snapshotBytes, err := os.ReadFile(stateFile) if err != nil { t.Fatalf("read snapshot: %v", err) } if bytes.Contains(snapshotBytes, []byte(created.ReconnectToken)) { t.Fatal("snapshot persisted the raw reconnect capability") } if !bytes.Contains(snapshotBytes, []byte(`"hostReconnectVerifier"`)) { t.Fatalf("snapshot omitted host verifier: %s", snapshotBytes) } hB := newRelayHarnessAt(t, t.TempDir(), stateFile) hB.srv.mu.RLock() _, reloaded := hB.srv.rooms["RESUM"] hB.srv.mu.RUnlock() if !reloaded { t.Fatal("room RESUM was not reloaded from snapshot") } unproved := hB.dial(t, "8.0.0.2") unproved.send(clientMsg{Type: relayTypeJoin, SessionID: "RESUM", PeerID: "H"}) unproved.expectError(relayErrorPeerIdUnavailable) reconnected := hB.dial(t, "8.0.0.3") reconnected.send(clientMsg{ Type: relayTypeJoin, SessionID: "RESUM", PeerID: "H", ReconnectToken: created.ReconnectToken, }) joined := reconnected.expectAuthority(relayTypeJoined, "H") if joined.ReconnectToken != created.ReconnectToken { t.Fatal("restored host capability changed") } duplicateCreate := hB.dial(t, "8.0.0.4") duplicateCreate.send(clientMsg{Type: relayTypeCreate, SessionID: "RESUM", PeerID: "OTHER"}) duplicateCreate.expectError(relayErrorRoomExists) } func TestSnapshotV3RetainsModernHostAndGuestVerifiersAcrossRestart(t *testing.T) { stateFile := filepath.Join(t.TempDir(), "rooms.json") hA := newRelayHarnessAt(t, t.TempDir(), stateFile) hostToken, _ := mustReconnectToken(t) host := hA.dial(t, "8.0.1.1") host.send(clientMsg{ Type: relayTypeCreate, SessionID: "V3_RESTART", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) host.expectAuthority(relayTypeCreated, "H") guestToken, _ := mustReconnectToken(t) guest := hA.dial(t, "8.0.1.2") guest.send(clientMsg{ Type: relayTypeJoin, SessionID: "V3_RESTART", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guest.expectAuthority(relayTypeJoined, "H") host.expect(relayTypePeerJoined) if err := guest.conn.Close(); err != nil { t.Fatalf("close guest before snapshot: %v", err) } left := host.expect(relayTypePeerLeft) if left.PeerID != "G" { t.Fatalf("pre-snapshot disconnect peerId=%q, want G", left.PeerID) } if err := hA.srv.snap.flushAndStop(2 * time.Second); err != nil { t.Fatalf("flush snapshot v3: %v", err) } snapshotBytes, err := os.ReadFile(stateFile) if err != nil { t.Fatalf("read snapshot v3: %v", err) } if !bytes.Contains(snapshotBytes, []byte(`"version":3`)) || !bytes.Contains(snapshotBytes, []byte(`"peerReconnectVerifiers"`)) { t.Fatalf("snapshot omitted v3 guest verifier state: %s", snapshotBytes) } if bytes.Contains(snapshotBytes, []byte(hostToken)) || bytes.Contains(snapshotBytes, []byte(guestToken)) { t.Fatal("snapshot persisted a raw reconnect capability") } hB := newRelayHarnessAt(t, t.TempDir(), stateFile) wrongHostToken, _ := mustReconnectToken(t) hostThief := hB.dial(t, "8.0.1.3") hostThief.send(clientMsg{ Type: relayTypeJoin, SessionID: "V3_RESTART", PeerID: "H", ReconnectToken: wrongHostToken, ProtocolVersion: relayProtocolVersion, }) hostThief.expectError(relayErrorPeerIdUnavailable) restartedHost := hB.dial(t, "8.0.1.4") restartedHost.send(clientMsg{ Type: relayTypeJoin, SessionID: "V3_RESTART", PeerID: "H", ReconnectToken: hostToken, ProtocolVersion: relayProtocolVersion, }) hostJoined := restartedHost.expectAuthority(relayTypeJoined, "H") if hostJoined.ReconnectToken != hostToken || hostJoined.ProtocolVersion != relayProtocolVersion { t.Fatalf("restored host authority changed: %+v", hostJoined) } wrongGuestToken, _ := mustReconnectToken(t) guestThief := hB.dial(t, "8.0.1.5") guestThief.send(clientMsg{ Type: relayTypeJoin, SessionID: "V3_RESTART", PeerID: "G", ReconnectToken: wrongGuestToken, ProtocolVersion: relayProtocolVersion, }) guestThief.expectError(relayErrorPeerIdUnavailable) restartedGuest := hB.dial(t, "8.0.1.6") restartedGuest.send(clientMsg{ Type: relayTypeJoin, SessionID: "V3_RESTART", PeerID: "G", ReconnectToken: guestToken, ProtocolVersion: relayProtocolVersion, }) guestJoined := restartedGuest.expectAuthority(relayTypeJoined, "H") if guestJoined.ReconnectToken != guestToken || guestJoined.ProtocolVersion != relayProtocolVersion { t.Fatalf("restored guest authority changed: %+v", guestJoined) } rejoined := restartedHost.expect(relayTypePeerJoined) if rejoined.PeerID != "G" { t.Fatalf("restored guest event peerId=%q, want G", rejoined.PeerID) } } func TestLoadedRoomsConsumeGlobalCapacityWithoutRestoringSourceQuota(t *testing.T) { stateFile := filepath.Join(t.TempDir(), "rooms.json") now := time.Now().UTC() snapshot := stateSnapshot{ Version: snapshotFormatVersion, SavedAt: now, Rooms: makeRoomSnapshots(maxRetainedRooms, false, now), } data, err := json.Marshal(snapshot) if err != nil { t.Fatalf("marshal full snapshot: %v", err) } if err := os.WriteFile(stateFile, data, 0644); err != nil { t.Fatalf("write full snapshot: %v", err) } h := newRelayHarnessAt(t, t.TempDir(), stateFile) h.srv.conns.mu.Lock() restoredQuotaEntries := len(h.srv.conns.roomsPerIP) h.srv.conns.mu.Unlock() if restoredQuotaEntries != 0 { t.Fatalf("restart restored %d process-local quota entries", restoredQuotaEntries) } client := h.dial(t, "8.0.0.4") client.send(clientMsg{Type: relayTypeCreate, SessionID: "RESTARTOVER", PeerID: "H"}) client.expectError(relayErrorRateLimited) client.send(clientMsg{Type: relayTypeJoin, SessionID: "S0000", PeerID: "G"}) client.expect(relayTypeJoined) }