diff --git a/Server/db/message_queries.go b/Server/db/message_queries.go index 473c7186..3565f57f 100644 --- a/Server/db/message_queries.go +++ b/Server/db/message_queries.go @@ -20,8 +20,9 @@ func sanitizeFTSQuery(q string) string { } result := strings.TrimSpace(sb.String()) // Enforce a maximum query length to bound FTS processing. - if len(result) > 200 { - result = result[:200] + // Use rune count to avoid splitting multi-byte characters. + if runes := []rune(result); len(runes) > 200 { + result = string(runes[:200]) } return result } diff --git a/Server/db/sanitize_test.go b/Server/db/sanitize_test.go new file mode 100644 index 00000000..c0f2ca4e --- /dev/null +++ b/Server/db/sanitize_test.go @@ -0,0 +1,49 @@ +package db + +import ( + "strings" + "testing" + "unicode/utf8" +) + +func TestSanitizeFTSQuery_UTF8Truncation(t *testing.T) { + // BUG-090: sanitizeFTSQuery truncates at byte boundary, producing + // invalid UTF-8 when the input contains multi-byte runes (CJK, emoji). + + // Build a 210-rune CJK string. Each CJK rune is 3 bytes → 630 bytes. + input := strings.Repeat("漢", 210) + + got := sanitizeFTSQuery(input) + + if !utf8.ValidString(got) { + t.Fatal("sanitizeFTSQuery produced invalid UTF-8 after truncation") + } + + runeCount := utf8.RuneCountInString(got) + if runeCount > 200 { + t.Fatalf("expected at most 200 runes, got %d", runeCount) + } + if runeCount != 200 { + t.Fatalf("expected exactly 200 runes for 210-rune input, got %d", runeCount) + } +} + +func TestSanitizeFTSQuery_ASCIIUnchanged(t *testing.T) { + // ASCII-only input under 200 chars should pass through unchanged. + input := "hello world search query" + got := sanitizeFTSQuery(input) + if got != input { + t.Errorf("expected %q, got %q", input, got) + } +} + +func TestSanitizeFTSQuery_StripsOperators(t *testing.T) { + input := `hello "world" AND (test) NOT foo*` + got := sanitizeFTSQuery(input) + // Should only contain letters, digits, spaces, hyphens. + for _, r := range got { + if r == '"' || r == '(' || r == ')' || r == '*' { + t.Errorf("operator character %q not stripped", r) + } + } +} diff --git a/Server/updater/updater.go b/Server/updater/updater.go index 5ccdd18f..209f6825 100644 --- a/Server/updater/updater.go +++ b/Server/updater/updater.go @@ -407,7 +407,12 @@ func (u *Updater) downloadFile(ctx context.Context, url, destPath string) error if err != nil { return fmt.Errorf("creating destination file: %w", err) } - defer f.Close() //nolint:errcheck + closed := false + defer func() { + if !closed { + _ = f.Close() + } + }() // Cap download at 500 MiB to prevent unbounded disk usage from a // malicious or corrupted release asset. @@ -417,6 +422,7 @@ func (u *Updater) downloadFile(ctx context.Context, url, destPath string) error n, err := io.Copy(f, limitedReader) if err != nil { _ = f.Close() + closed = true _ = os.Remove(destPath) return fmt.Errorf("writing downloaded file: %w", err) } @@ -425,6 +431,7 @@ func (u *Updater) downloadFile(ctx context.Context, url, destPath string) error var probe [1]byte if extra, _ := resp.Body.Read(probe[:]); extra > 0 { _ = f.Close() + closed = true _ = os.Remove(destPath) return fmt.Errorf("downloaded file exceeds maximum size of %d bytes", maxBinarySize) } diff --git a/Server/ws/hub.go b/Server/ws/hub.go index 95a8f83a..54afbe36 100644 --- a/Server/ws/hub.go +++ b/Server/ws/hub.go @@ -24,19 +24,20 @@ type broadcastMsg struct { // Hub manages all active WebSocket clients and routes messages between them. // All exported methods are safe to call from multiple goroutines. type Hub struct { - clients map[int64]*Client - mu syncutil.RWMutex - db *db.DB - limiter *auth.RateLimiter - broadcast chan broadcastMsg - register chan *Client - unregister chan *Client - stop chan struct{} - stopOnce sync.Once - livekit *LiveKitClient - lkProcess *LiveKitProcess - registry *HandlerRegistry - permChecker *permissions.Checker + clients map[int64]*Client + mu syncutil.RWMutex + db *db.DB + limiter *auth.RateLimiter + broadcast chan broadcastMsg + register chan *Client + unregister chan *Client + stop chan struct{} + stopOnce sync.Once + gracefulOnce sync.Once + livekit *LiveKitClient + lkProcess *LiveKitProcess + registry *HandlerRegistry + permChecker *permissions.Checker seq uint64 // atomic monotonic sequence counter replayBuf *EventRingBuffer // recent broadcast events for reconnection replay @@ -211,27 +212,30 @@ func (h *Hub) Stop() { } // GracefulStop stops the LiveKit process (if managed) and then stops the hub. +// Safe to call multiple times concurrently. func (h *Hub) GracefulStop() { - // Broadcast restart notice to all connected clients. - h.BroadcastServerRestart("shutdown", 5) + h.gracefulOnce.Do(func() { + // Broadcast restart notice to all connected clients. + h.BroadcastServerRestart("shutdown", 5) - // Stop LiveKit process. - if h.lkProcess != nil { - h.lkProcess.Stop() - } + // Stop LiveKit process. + if h.lkProcess != nil { + h.lkProcess.Stop() + } - // Give clients 5 seconds to disconnect gracefully. - time.Sleep(5 * time.Second) + // Give clients 5 seconds to disconnect gracefully. + time.Sleep(5 * time.Second) - // Close all remaining client connections. - h.mu.Lock() - for _, c := range h.clients { - c.closeSend() - } - h.mu.Unlock() + // Close all remaining client connections. + h.mu.Lock() + for _, c := range h.clients { + c.closeSend() + } + h.mu.Unlock() - // Stop the hub dispatch loop. - h.stopOnce.Do(func() { close(h.stop) }) + // Stop the hub dispatch loop. + h.stopOnce.Do(func() { close(h.stop) }) + }) } // CleanupVoiceForChannel removes all voice participants from the given channel. diff --git a/Server/ws/hub_test.go b/Server/ws/hub_test.go index be04974a..846f9f3a 100644 --- a/Server/ws/hub_test.go +++ b/Server/ws/hub_test.go @@ -504,6 +504,33 @@ func TestHub_GracefulStop_NoPanic(t *testing.T) { hub.GracefulStop() } +func TestHub_GracefulStop_Idempotent(t *testing.T) { + // BUG-087: GracefulStop must be safe to call concurrently/twice. + // Without sync.Once protection, double lkProcess.Stop() can panic. + hub, _ := newTestHub(t) + go hub.Run() + + var wg sync.WaitGroup + wg.Add(2) + for range 2 { + go func() { + defer wg.Done() + hub.GracefulStop() + }() + } + done := make(chan struct{}) + go func() { + wg.Wait() + close(done) + }() + select { + case <-done: + // Success: no panic from concurrent GracefulStop. + case <-time.After(15 * time.Second): + t.Fatal("concurrent GracefulStop calls deadlocked") + } +} + // ─── CleanupVoiceForChannel ─────────────────────────────────────────────────── func TestHub_CleanupVoiceForChannel_NoVoiceState_NoPanic(t *testing.T) { diff --git a/Server/ws/ringbuffer.go b/Server/ws/ringbuffer.go index 3bf993c3..c35589ca 100644 --- a/Server/ws/ringbuffer.go +++ b/Server/ws/ringbuffer.go @@ -50,8 +50,9 @@ func (rb *EventRingBuffer) EventsSince(afterSeq uint64) [][]byte { oldestIdx := (rb.pos - rb.count + rb.size) % rb.size oldestSeq := rb.entries[oldestIdx].seq - // If the requested seq is older than our oldest, we can't replay. - if afterSeq < oldestSeq { + // If the requested seq is at or older than our oldest, we can't guarantee + // full coverage — return nil to trigger a full ready payload. + if afterSeq <= oldestSeq { return nil } diff --git a/Server/ws/ringbuffer_test.go b/Server/ws/ringbuffer_test.go index de79b33a..fb703c1e 100644 --- a/Server/ws/ringbuffer_test.go +++ b/Server/ws/ringbuffer_test.go @@ -36,13 +36,18 @@ func TestPush_MultipleInOrder(t *testing.T) { rb.Push(i, []byte(fmt.Sprintf("msg-%d", i))) } - // Request events after the oldest (seq 1) — should return seq 2..5. - got := rb.EventsSince(1) - if len(got) != 4 { - t.Fatalf("expected 4 events after seq 1, got %d", len(got)) + // afterSeq == oldestSeq (1) → nil (BUG-085: conservative boundary). + if got := rb.EventsSince(1); got != nil { + t.Fatalf("expected nil when afterSeq == oldestSeq, got %d events", len(got)) + } + + // afterSeq one past oldest → replay seq 3..5. + got := rb.EventsSince(2) + if len(got) != 3 { + t.Fatalf("expected 3 events after seq 2, got %d", len(got)) } for i, ev := range got { - want := fmt.Sprintf("msg-%d", i+2) + want := fmt.Sprintf("msg-%d", i+3) if string(ev) != want { t.Errorf("event[%d]: expected %q, got %q", i, want, string(ev)) } @@ -70,12 +75,18 @@ func TestPush_WrapsAround(t *testing.T) { t.Fatalf("expected nil (afterSeq 2 still evicted), got %d events", len(got)) } - // Ask for events after seq 3 — should get seq 4, 5, 6. + // afterSeq == oldestSeq (3) → nil (BUG-085). got = rb.EventsSince(3) - if len(got) != 3 { - t.Fatalf("expected 3 events after seq 3, got %d", len(got)) + if got != nil { + t.Fatalf("expected nil when afterSeq == oldestSeq (3), got %d events", len(got)) } - for i, want := range []string{"e4", "e5", "e6"} { + + // afterSeq one past oldest → replay seq 5, 6. + got = rb.EventsSince(4) + if len(got) != 2 { + t.Fatalf("expected 2 events after seq 4, got %d", len(got)) + } + for i, want := range []string{"e5", "e6"} { if string(got[i]) != want { t.Errorf("event[%d]: expected %q, got %q", i, want, string(got[i])) } @@ -100,12 +111,18 @@ func TestPush_OverwritesOldest(t *testing.T) { t.Fatalf("expected oldest seq 2 after overwrite, got %d", oldest) } - got := rb.EventsSince(2) - if len(got) != 2 { - t.Fatalf("expected 2 events, got %d", len(got)) + // afterSeq == oldestSeq (2) → nil (BUG-085). + if got := rb.EventsSince(2); got != nil { + t.Fatalf("expected nil when afterSeq == oldestSeq, got %d events", len(got)) } - if string(got[0]) != "c" || string(got[1]) != "d" { - t.Errorf("expected [c, d], got [%s, %s]", got[0], got[1]) + + // afterSeq one past oldest → replay seq 4 only. + got := rb.EventsSince(3) + if len(got) != 1 { + t.Fatalf("expected 1 event, got %d", len(got)) + } + if string(got[0]) != "d" { + t.Errorf("expected [d], got [%s]", got[0]) } } @@ -171,12 +188,17 @@ func TestEventsSince_WraparoundOrder(t *testing.T) { rb.Push(i, []byte(fmt.Sprintf("v%d", i))) } - // Oldest is seq 4. Get everything from seq 4 onward. - got := rb.EventsSince(4) - if len(got) != 3 { - t.Fatalf("expected 3 events, got %d", len(got)) + // afterSeq == oldestSeq (4) → nil (BUG-085). + if got := rb.EventsSince(4); got != nil { + t.Fatalf("expected nil when afterSeq == oldestSeq, got %d events", len(got)) } - for i, want := range []string{"v5", "v6", "v7"} { + + // afterSeq one past oldest → replay seq 6, 7. + got := rb.EventsSince(5) + if len(got) != 2 { + t.Fatalf("expected 2 events, got %d", len(got)) + } + for i, want := range []string{"v6", "v7"} { if string(got[i]) != want { t.Errorf("event[%d]: expected %q, got %q", i, want, string(got[i])) } @@ -197,19 +219,54 @@ func TestEventsSince_AfterSeqZero_ReturnsBehavior(t *testing.T) { t.Fatalf("expected nil for afterSeq=0 (before oldest), got %d events", len(got)) } - // If we start seqs from 0, then afterSeq=0 equals oldest, and we get events > 0. + // If we start seqs from 0, afterSeq=0 equals oldest → nil (BUG-085). rb2 := ws.NewEventRingBuffer(8) rb2.Push(0, []byte("z0")) rb2.Push(1, []byte("z1")) rb2.Push(2, []byte("z2")) got = rb2.EventsSince(0) - if len(got) != 2 { - t.Fatalf("expected 2 events after seq 0, got %d", len(got)) + if got != nil { + t.Fatalf("expected nil when afterSeq == oldestSeq (0), got %d events", len(got)) } - if string(got[0]) != "z1" || string(got[1]) != "z2" { - t.Errorf("expected [z1, z2], got [%s, %s]", got[0], got[1]) + + // afterSeq one past oldest → replay seq 2 only. + got = rb2.EventsSince(1) + if len(got) != 1 { + t.Fatalf("expected 1 event after seq 1, got %d", len(got)) } + if string(got[0]) != "z2" { + t.Errorf("expected [z2], got [%s]", got[0]) + } +} + +func TestEventsSince_AfterSeqEqualsOldest_ReturnsNil(t *testing.T) { + // BUG-085: When afterSeq == oldestSeq, the client's last event is the + // oldest in the buffer. We can't guarantee nothing was missed between + // the evicted event before oldest and oldest itself, so EventsSince + // must return nil to trigger a full ready payload. + const cap = 4 + rb := ws.NewEventRingBuffer(cap) + + // Push 6 events: buffer holds seq 3,4,5,6. Oldest = 3. + for i := uint64(1); i <= 6; i++ { + rb.Push(i, []byte(fmt.Sprintf("e%d", i))) + } + + if oldest := rb.OldestSeq(); oldest != 3 { + t.Fatalf("expected oldest seq 3, got %d", oldest) + } + + // afterSeq == oldestSeq (3): must return nil, not empty slice. + got := rb.EventsSince(3) + if got == nil { + // This is the CORRECT behavior after the fix. + return + } + // Before the fix, this returns a non-nil slice [e4, e5, e6]. + // That's wrong because the client at seq 3 might have missed events + // between the evicted seq 2 and seq 3. + t.Fatalf("expected nil when afterSeq == oldestSeq, got %d events", len(got)) } // ─── OldestSeq ─────────────────────────────────────────────────────────────── @@ -310,12 +367,19 @@ func TestEventsSince_CapacityBoundaries(t *testing.T) { wantLen: -1, }, { - name: "exactly at capacity, from oldest", + name: "exactly at capacity, afterSeq == oldest → nil", + cap: 4, + pushes: 4, + afterSeq: 1, // oldest=1, BUG-085: == returns nil + wantLen: -1, + }, + { + name: "exactly at capacity, afterSeq one past oldest", cap: 4, pushes: 4, - afterSeq: 1, - wantLen: 3, - wantFirst: "e2", + afterSeq: 2, + wantLen: 2, + wantFirst: "e3", }, { name: "one past capacity", @@ -325,34 +389,41 @@ func TestEventsSince_CapacityBoundaries(t *testing.T) { wantLen: -1, }, { - name: "one past capacity, valid afterSeq", + name: "one past capacity, afterSeq == oldest → nil", + cap: 4, + pushes: 5, + afterSeq: 2, // oldest=2, BUG-085: == returns nil + wantLen: -1, + }, + { + name: "one past capacity, afterSeq one past oldest", cap: 4, pushes: 5, - afterSeq: 2, - wantLen: 3, - wantFirst: "e3", + afterSeq: 3, + wantLen: 2, + wantFirst: "e4", }, { - name: "double capacity", + name: "double capacity, afterSeq == oldest → nil", + cap: 4, + pushes: 8, + afterSeq: 5, // oldest=5, BUG-085: == returns nil + wantLen: -1, + }, + { + name: "double capacity, afterSeq one past oldest", cap: 4, pushes: 8, - afterSeq: 5, - wantLen: 3, - wantFirst: "e6", + afterSeq: 6, + wantLen: 2, + wantFirst: "e7", }, { - name: "capacity 1", + name: "capacity 1, afterSeq == oldest → nil", cap: 1, pushes: 3, - afterSeq: 3, - wantLen: 0, - }, - { - name: "capacity 1, afterSeq matches oldest", - cap: 1, - pushes: 3, - afterSeq: 3, - wantLen: 0, + afterSeq: 3, // oldest=3, BUG-085: == returns nil + wantLen: -1, }, { name: "capacity 1, afterSeq too old", diff --git a/Server/ws/ws_integration_test.go b/Server/ws/ws_integration_test.go index 2d3b3aa5..80e6fb19 100644 --- a/Server/ws/ws_integration_test.go +++ b/Server/ws/ws_integration_test.go @@ -614,8 +614,11 @@ func TestServeWS_Reconnect_PreservesVoiceState(t *testing.T) { } ws.SetClientVoiceStateForTest(originalClient, chID, vsBeforeReconnect.JoinedAt) - // Second connection: reconnect (lastSeq > 0) — voice state should transfer - conn2 := dialAndAuth(1) + // Second connection: reconnect (lastSeq > oldestSeq) — voice state should transfer. + // Use lastSeq=2 because conn1's join produces at least 2 broadcasts + // (member_join seq=1, presence seq=2), and afterSeq must be > oldestSeq + // for EventsSince to return a replay instead of nil (BUG-085). + conn2 := dialAndAuth(2) defer func() { _ = conn2.Close(websocket.StatusNormalClosure, "") }() var replacementClient *ws.Client