From 4bb78bde0e744c5228f90dfcbf396bcf5d44fa79 Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:27:09 -0400 Subject: [PATCH 01/11] fix(server): persist bans and temp kicks across hub restart --- CHANGELOG.md | 1 + cmd/server/main.go | 5 +- server/admin_panel_test.go | 2 +- server/admin_web_test.go | 2 +- server/client_empty_content_test.go | 10 +- server/client_file_limit_test.go | 2 +- server/client_nullbyte_test.go | 2 +- server/client_sender_spoof_test.go | 2 +- server/client_test.go | 2 +- server/handlers.go | 34 ++++- server/health_test.go | 2 +- server/hub.go | 58 ++++++++- server/hub_moderation_persist_test.go | 173 ++++++++++++++++++++++++++ server/hub_test.go | 59 +++++---- server/integration_test.go | 8 +- server/loadverify_bench_test.go | 5 +- 16 files changed, 317 insertions(+), 50 deletions(-) create mode 100644 server/hub_moderation_persist_test.go diff --git a/CHANGELOG.md b/CHANGELOG.md index a76a8f0..cb97853 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,6 +11,7 @@ On **`main`** only; not part of the latest tagged release until you tag and publ - **Server**: **Fix:** reject empty or whitespace-only plaintext on `text`, `dm`, and `edit` when `encrypted` is false (System reply, no persist/broadcast); encrypted opaque `content` is not treated as empty ([#117](https://github.com/Cod-e-Codes/marchat/issues/117)). - **Server**: **Fix:** SQLite `InitDB` applies `busy_timeout` / WAL / related pragmas via the DSN on every connection and sets `MaxOpenConns(1)` / `MaxIdleConns(1)`, so concurrent inserts no longer fail with `SQLITE_BUSY` from one-shot `PRAGMA` + the default `database/sql` pool ([#118](https://github.com/Cod-e-Codes/marchat/issues/118)). - **Dependencies**: **modernc.org/sqlite** v1.56.0 (journal-rollback corruption fix; **modernc.org/libc** v1.74.4); **github.com/lucasb-eyer/go-colorful** v1.4.1. +- **Server**: **Fix:** active permanent bans and unexpired temp kicks load from `ban_history` on hub start (`expires_at` NULL = permanent; non-NULL = kick expiry), so moderation survives process restart. Pre-upgrade open rows without `expires_at` load as permanent. ## v1.3.4 diff --git a/cmd/server/main.go b/cmd/server/main.go index 50a0766..ec30494 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -305,7 +305,10 @@ func main() { }) } - hub := server.NewHub(pluginDir, dataDir, registryURL, db) + hub, err := server.NewHub(pluginDir, dataDir, registryURL, db) + if err != nil { + log.Fatalf("Failed to create hub: %v", err) + } go hub.Run() // Log server startup diff --git a/server/admin_panel_test.go b/server/admin_panel_test.go index 1c9c431..7781f5c 100644 --- a/server/admin_panel_test.go +++ b/server/admin_panel_test.go @@ -19,7 +19,7 @@ func setupPanelEnv(t *testing.T) (*AdminPanel, func()) { CreateSchema(db) pluginDir := filepath.Join(tdir, "plugins") dataDir := filepath.Join(tdir, "data") - hub := NewHub(pluginDir, dataDir, "", db) + hub := mustNewHub(t, pluginDir, dataDir, "", db) cfg := &appcfg.Config{Port: 8080, AdminKey: "k", Admins: []string{"a"}, DBPath: dbPath, ConfigDir: tdir} panel := NewAdminPanel(hub, db, hub.GetPluginManager(), cfg) return panel, func() { _ = db.Close() } diff --git a/server/admin_web_test.go b/server/admin_web_test.go index 1985ee2..214aeaf 100644 --- a/server/admin_web_test.go +++ b/server/admin_web_test.go @@ -31,7 +31,7 @@ func setupTestServerEnv(t *testing.T) (*sql.DB, *Hub, *appcfg.Config, func()) { _ = os.MkdirAll(pluginDir, 0o755) _ = os.MkdirAll(dataDir, 0o755) - hub := NewHub(pluginDir, dataDir, "", db) + hub := mustNewHub(t, pluginDir, dataDir, "", db) go func() { // run hub in background hub.Run() }() diff --git a/server/client_empty_content_test.go b/server/client_empty_content_test.go index d4cae40..4d17aa2 100644 --- a/server/client_empty_content_test.go +++ b/server/client_empty_content_test.go @@ -38,7 +38,7 @@ func TestIntegrationEmptyTextRejectedNoBroadcast(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, 10<<20, dbPath) @@ -107,7 +107,7 @@ func TestIntegrationEmptyDMRejected(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, 10<<20, dbPath) @@ -175,7 +175,7 @@ func TestIntegrationEmptyEditRejected(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, 10<<20, dbPath) @@ -247,7 +247,7 @@ func TestIntegrationEncryptedOpaqueTextAccepted(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, 10<<20, dbPath) @@ -314,7 +314,7 @@ func TestIntegrationCommandPathStillWorksWithEmptyCheck(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, 10<<20, dbPath) diff --git a/server/client_file_limit_test.go b/server/client_file_limit_test.go index 153ce26..39c35fd 100644 --- a/server/client_file_limit_test.go +++ b/server/client_file_limit_test.go @@ -119,7 +119,7 @@ func setupFileLimitHub(t *testing.T, maxFileBytes int64) (string, func()) { } CreateSchema(db) - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, maxFileBytes, dbPath) diff --git a/server/client_nullbyte_test.go b/server/client_nullbyte_test.go index 4c29f37..951be59 100644 --- a/server/client_nullbyte_test.go +++ b/server/client_nullbyte_test.go @@ -69,7 +69,7 @@ func TestIntegrationNullByteRejectedNoBroadcast(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, 10<<20, dbPath) diff --git a/server/client_sender_spoof_test.go b/server/client_sender_spoof_test.go index 7bea15f..22e8e81 100644 --- a/server/client_sender_spoof_test.go +++ b/server/client_sender_spoof_test.go @@ -24,7 +24,7 @@ func setupSpoofTestHub(t *testing.T) (*sql.DB, string, func()) { } CreateSchema(db) - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, 10<<20, dbPath) diff --git a/server/client_test.go b/server/client_test.go index 8425886..ef7d1fa 100644 --- a/server/client_test.go +++ b/server/client_test.go @@ -29,7 +29,7 @@ func setupTestClient(t *testing.T) (*Client, *Hub, *sql.DB, func()) { CreateSchema(db) // Create hub with correct parameters - hub := NewHub(tdir, tdir, "http://localhost:8080", db) + hub := mustNewHub(t, tdir, tdir, "http://localhost:8080", db) go hub.Run() // Create mock websocket connection diff --git a/server/handlers.go b/server/handlers.go index 6a34e17..c5bde47 100644 --- a/server/handlers.go +++ b/server/handlers.go @@ -194,7 +194,8 @@ func CreateSchema(db *sql.DB) { username ` + keyedTextType + ` NOT NULL, banned_at ` + dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP, unbanned_at ` + dateTimeType + `, - banned_by ` + keyedTextType + ` NOT NULL + banned_by ` + keyedTextType + ` NOT NULL, + expires_at ` + dateTimeType + ` );` _, err = dbExec(db, banHistorySchema) @@ -202,6 +203,30 @@ func CreateSchema(db *sql.DB) { log.Printf("Warning: failed to create ban_history table: %v", err) } + // Add expires_at for ban_history created before this column existed. + // NULL expires_at = permanent ban; non-NULL = temporary kick expiry. + { + var expiresExists int + switch dialect { + case DialectPostgres: + err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = current_schema() AND table_name = 'ban_history' AND column_name = ?`, "expires_at").Scan(&expiresExists) + case DialectMySQL: + err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = 'ban_history' AND column_name = ?`, "expires_at").Scan(&expiresExists) + default: + err = dbQueryRow(db, `SELECT COUNT(*) FROM pragma_table_info('ban_history') WHERE name=?`, "expires_at").Scan(&expiresExists) + } + if err != nil { + log.Printf("Warning: failed to check for ban_history.expires_at column: %v", err) + } else if expiresExists == 0 { + _, err = dbExec(db, `ALTER TABLE ban_history ADD COLUMN expires_at `+dateTimeType) + if err != nil { + log.Printf("Warning: failed to add ban_history.expires_at column: %v", err) + } else { + log.Printf("Added expires_at column to ban_history table") + } + } + } + // Create indexes for performance (MySQL needs a prefix length when indexing LONGTEXT) recipientIdx := `CREATE INDEX IF NOT EXISTS idx_messages_recipient ON messages(recipient)` if dialect == DialectMySQL { @@ -516,9 +541,10 @@ func clearUserMessageState(db *sql.DB, username string) error { return err } -// recordBanEvent records a ban event in the ban_history table -func recordBanEvent(db *sql.DB, username, bannedBy string) error { - _, err := dbExec(db, `INSERT INTO ban_history (username, banned_by) VALUES (?, ?)`, username, bannedBy) +// recordBanEvent records a ban or kick in ban_history. +// expiresAt nil means a permanent ban; non-nil is the temporary kick expiry. +func recordBanEvent(db *sql.DB, username, bannedBy string, expiresAt *time.Time) error { + _, err := dbExec(db, `INSERT INTO ban_history (username, banned_by, expires_at) VALUES (?, ?, ?)`, username, bannedBy, expiresAt) if err != nil { log.Printf("Warning: failed to record ban event for user %s: %v", username, err) } diff --git a/server/health_test.go b/server/health_test.go index 06a7611..695ead1 100644 --- a/server/health_test.go +++ b/server/health_test.go @@ -25,7 +25,7 @@ func setupTestHealthChecker(t *testing.T) (*HealthChecker, *sql.DB, func()) { CreateSchema(db) // Create hub with correct parameters - hub := NewHub(tdir, tdir, "http://localhost:8080", db) + hub := mustNewHub(t, tdir, tdir, "http://localhost:8080", db) go hub.Run() // Create health checker diff --git a/server/hub.go b/server/hub.go index 9deef77..910ef1e 100644 --- a/server/hub.go +++ b/server/hub.go @@ -3,6 +3,7 @@ package server import ( "database/sql" "errors" + "fmt" "log" "strings" "sync" @@ -52,11 +53,11 @@ type Hub struct { channelMutex sync.RWMutex } -func NewHub(pluginDir, dataDir, registryURL string, db *sql.DB) *Hub { +func NewHub(pluginDir, dataDir, registryURL string, db *sql.DB) (*Hub, error) { pluginManager := manager.NewPluginManager(pluginDir, dataDir, registryURL) pluginCommandHandler := NewPluginCommandHandler(pluginManager) - return &Hub{ + h := &Hub{ clients: make(map[*Client]bool), usernames: make(map[string]struct{}), broadcast: make(chan interface{}), @@ -69,6 +70,50 @@ func NewHub(pluginDir, dataDir, registryURL string, db *sql.DB) *Hub { db: db, channels: make(map[string]map[*Client]bool), } + if db != nil { + if err := h.loadModerationState(); err != nil { + return nil, err + } + } + return h, nil +} + +// loadModerationState restores active bans and unexpired temp kicks from ban_history. +// Open rows (unbanned_at IS NULL) with NULL expires_at are permanent bans. +// Open rows with expires_at in the future are temp kicks; expired rows are skipped. +func (h *Hub) loadModerationState() error { + rows, err := dbQuery(h.db, ` + SELECT username, expires_at FROM ban_history + WHERE unbanned_at IS NULL`) + if err != nil { + return fmt.Errorf("load moderation state: %w", err) + } + defer rows.Close() + + now := time.Now() + h.banMutex.Lock() + defer h.banMutex.Unlock() + + for rows.Next() { + var username string + var expiresAt sql.NullTime + if err := rows.Scan(&username, &expiresAt); err != nil { + return fmt.Errorf("scan moderation row: %w", err) + } + lower := strings.ToLower(username) + if !expiresAt.Valid { + // Permanent ban (or pre-upgrade open row treated as permanent). + h.bans[lower] = now.Add(100 * 365 * 24 * time.Hour) + continue + } + if now.Before(expiresAt.Time) { + h.tempKicks[lower] = expiresAt.Time + } + } + if err := rows.Err(); err != nil { + return fmt.Errorf("iterate moderation rows: %w", err) + } + return nil } func (h *Hub) TryReserveUsername(username string) bool { @@ -118,9 +163,9 @@ func (h *Hub) BanUser(username string, adminUsername string) error { "admin": adminUsername, }) - // Record ban event in database + // Record ban event in database (NULL expires_at = permanent) if h.getDB() != nil { - err := recordBanEvent(h.getDB(), lowerUsername, adminUsername) + err := recordBanEvent(h.getDB(), lowerUsername, adminUsername, nil) if err != nil { log.Printf("Warning: failed to record ban event for user %s: %v", username, err) } @@ -294,9 +339,10 @@ func (h *Hub) KickUser(username string, adminUsername string) error { "until": kickExpiry.Format("2006-01-02 15:04:05"), }) - // Record kick event in database (reuse ban event structure) + // Record kick event in database with 24h expiry if h.getDB() != nil { - err := recordBanEvent(h.getDB(), lowerUsername, adminUsername) + exp := kickExpiry + err := recordBanEvent(h.getDB(), lowerUsername, adminUsername, &exp) if err != nil { log.Printf("Warning: failed to record kick event for user %s: %v", username, err) } diff --git a/server/hub_moderation_persist_test.go b/server/hub_moderation_persist_test.go new file mode 100644 index 0000000..db963b5 --- /dev/null +++ b/server/hub_moderation_persist_test.go @@ -0,0 +1,173 @@ +package server + +import ( + "net/http/httptest" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/Cod-e-Codes/marchat/shared" + "github.com/gorilla/websocket" +) + +func TestPermanentBanPersistsAcrossHubRestart(t *testing.T) { + tdir := t.TempDir() + dbPath := filepath.Join(tdir, "moderation.db") + + db, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB: %v", err) + } + CreateSchema(db) + + hub1 := mustNewHub(t, tdir, tdir, "", db) + if err := hub1.BanUser("alice", "admin"); err != nil { + t.Fatalf("BanUser: %v", err) + } + if !hub1.IsUserBanned("alice") { + t.Fatal("alice should be banned before restart") + } + if err := db.Close(); err != nil { + t.Fatalf("close db: %v", err) + } + + db2, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB reopen: %v", err) + } + defer db2.Close() + CreateSchema(db2) + + hub2 := mustNewHub(t, tdir, tdir, "", db2) + if !hub2.IsUserBanned("alice") { + t.Fatal("alice should still be banned after hub restart") + } + if !hub2.IsUserBanned("ALICE") { + t.Fatal("persisted ban should remain case-insensitive") + } + + assertHandshakeRejectedAsBanned(t, hub2, db2, tdir, "alice") +} + +func TestTempKickPersistsAcrossHubRestart(t *testing.T) { + tdir := t.TempDir() + dbPath := filepath.Join(tdir, "moderation.db") + + db, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB: %v", err) + } + CreateSchema(db) + + hub1 := mustNewHub(t, tdir, tdir, "", db) + registerTestClient(hub1, "bob") + if err := hub1.KickUser("bob", "admin"); err != nil { + t.Fatalf("KickUser: %v", err) + } + if !hub1.IsUserBanned("bob") { + t.Fatal("bob should be kicked before restart") + } + if err := db.Close(); err != nil { + t.Fatalf("close db: %v", err) + } + + db2, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB reopen: %v", err) + } + defer db2.Close() + CreateSchema(db2) + + hub2 := mustNewHub(t, tdir, tdir, "", db2) + if !hub2.IsUserBanned("bob") { + t.Fatal("bob should still be kicked after hub restart") + } + + hub2.banMutex.RLock() + _, permanent := hub2.bans["bob"] + kickExpiry, kicked := hub2.tempKicks["bob"] + hub2.banMutex.RUnlock() + if permanent { + t.Fatal("temp kick must not load as permanent ban") + } + if !kicked || !time.Now().Before(kickExpiry) { + t.Fatalf("expected unexpired tempKick, got kicked=%v expiry=%v", kicked, kickExpiry) + } + + assertHandshakeRejectedAsBanned(t, hub2, db2, tdir, "bob") +} + +func TestExpiredTempKickNotLoadedOnRestart(t *testing.T) { + tdir := t.TempDir() + dbPath := filepath.Join(tdir, "moderation.db") + + db, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB: %v", err) + } + CreateSchema(db) + + past := time.Now().Add(-time.Hour) + if err := recordBanEvent(db, "carol", "admin", &past); err != nil { + t.Fatalf("recordBanEvent: %v", err) + } + if err := db.Close(); err != nil { + t.Fatalf("close db: %v", err) + } + + db2, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB reopen: %v", err) + } + defer db2.Close() + CreateSchema(db2) + + hub := mustNewHub(t, tdir, tdir, "", db2) + if hub.IsUserBanned("carol") { + t.Fatal("expired temp kick must not reject after restart") + } +} + +func assertHandshakeRejectedAsBanned(t *testing.T, hub *Hub, db interface { + Close() error +}, tdir, username string) { + t.Helper() + _ = db + sqlDB := hub.getDB() + go hub.Run() + + handler := ServeWs(hub, sqlDB, nil, "admin-key", false, 10<<20, filepath.Join(tdir, "ws.db")) + srv := httptest.NewServer(handler) + defer srv.Close() + + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer conn.Close() + + if err := conn.WriteJSON(shared.Handshake{Username: username}); err != nil { + t.Fatalf("write handshake: %v", err) + } + + _ = conn.SetReadDeadline(time.Now().Add(3 * time.Second)) + _, _, err = conn.ReadMessage() + if err == nil { + t.Fatal("expected close after ban reject, got a message") + } + closeErr, ok := err.(*websocket.CloseError) + if !ok { + if !strings.Contains(strings.ToLower(err.Error()), "banned") { + t.Fatalf("expected ban close, got %T: %v", err, err) + } + return + } + if closeErr.Code != websocket.ClosePolicyViolation { + t.Fatalf("expected ClosePolicyViolation, got %d (%s)", closeErr.Code, closeErr.Text) + } + if !strings.Contains(strings.ToLower(closeErr.Text), "banned") { + t.Fatalf("expected ban close text, got %q", closeErr.Text) + } +} diff --git a/server/hub_test.go b/server/hub_test.go index 760692c..707510c 100644 --- a/server/hub_test.go +++ b/server/hub_test.go @@ -23,6 +23,15 @@ func registerTestClient(hub *Hub, username string) *Client { return client } +func mustNewHub(t *testing.T, pluginDir, dataDir, registryURL string, db *sql.DB) *Hub { + t.Helper() + h, err := NewHub(pluginDir, dataDir, registryURL, db) + if err != nil { + t.Fatalf("NewHub: %v", err) + } + return h +} + func countBanHistoryRows(db *sql.DB, username string) (int, error) { var count int err := db.QueryRow(`SELECT COUNT(*) FROM ban_history WHERE username = ?`, strings.ToLower(username)).Scan(&count) @@ -37,7 +46,9 @@ func TestNewHub(t *testing.T) { } defer db.Close() - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + CreateSchema(db) + + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) if hub == nil { t.Fatal("NewHub returned nil") @@ -91,7 +102,7 @@ func TestHubBanUser(t *testing.T) { // Create schema for database operations CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "testuser" adminUsername := "admin" @@ -128,7 +139,7 @@ func TestHubUnbanUser(t *testing.T) { // Create schema for database operations CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "testuser" adminUsername := "admin" @@ -168,7 +179,7 @@ func TestHubKickUser(t *testing.T) { CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "testuser" adminUsername := "admin" @@ -214,7 +225,7 @@ func TestHubKickUserOfflineReturnsError(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "offlineuser" adminUsername := "admin" @@ -251,7 +262,7 @@ func TestHubKickUserCaseInsensitiveOnlineMatch(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) registerTestClient(hub, "TestUser") adminUsername := "admin" @@ -271,7 +282,7 @@ func TestHubRejectsSelfKickAndBan(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) adminUsername := "Alice" cases := []string{"Alice", "alice", "ALICE", "aLiCe"} @@ -322,7 +333,7 @@ func TestHubKickPermanentlyBannedReturnsError(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "banneduser" adminUsername := "admin" registerTestClient(hub, username) @@ -358,7 +369,7 @@ func TestHubAllowUser(t *testing.T) { // Create schema for database operations CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "testuser" adminUsername := "admin" @@ -400,7 +411,7 @@ func TestHubBanOverridesKick(t *testing.T) { // Create schema for database operations CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "testuser" adminUsername := "admin" @@ -442,7 +453,7 @@ func TestHubCleanupExpiredBans(t *testing.T) { // Create schema for database operations CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "testuser" adminUsername := "admin" @@ -477,7 +488,9 @@ func TestHubForceDisconnectUser(t *testing.T) { } defer db.Close() - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + CreateSchema(db) + + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "testuser" adminUsername := "admin" @@ -499,7 +512,9 @@ func TestHubGetPluginManager(t *testing.T) { } defer db.Close() - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + CreateSchema(db) + + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) pluginManager := hub.GetPluginManager() if pluginManager == nil { @@ -521,7 +536,7 @@ func TestHubBanCaseInsensitive(t *testing.T) { // Create schema for database operations CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "TestUser" adminUsername := "admin" @@ -556,7 +571,7 @@ func TestHubMultipleBansAndKicks(t *testing.T) { // Create schema for database operations CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) adminUsername := "admin" users := []string{"user1", "user2", "user3"} @@ -604,7 +619,7 @@ func TestHubConcurrentBanOperations(t *testing.T) { // Create schema for database operations CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "testuser" adminUsername := "admin" @@ -649,7 +664,7 @@ func TestHubConcurrentBanOperations(t *testing.T) { } func TestBroadcastUserListNonBlocking(t *testing.T) { - hub := NewHub("", "", "", nil) + hub := mustNewHub(t, "", "", "", nil) // Create a client with a tiny send buffer that we intentionally fill. stalled := &Client{username: "stalled", send: make(chan interface{}, 1)} @@ -689,7 +704,7 @@ func TestKickUserNonBlocking(t *testing.T) { defer db.Close() CreateSchema(db) - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) // Create a client with a full send buffer. client := &Client{ @@ -719,7 +734,7 @@ func TestKickUserNonBlocking(t *testing.T) { } func TestChannelManagement(t *testing.T) { - hub := NewHub("", "", "", nil) + hub := mustNewHub(t, "", "", "", nil) client := &Client{username: "testuser", send: make(chan interface{}, 10)} @@ -814,7 +829,7 @@ func TestChannelManagement(t *testing.T) { } func TestBroadcastDM(t *testing.T) { - hub := NewHub("", "", "", nil) + hub := mustNewHub(t, "", "", "", nil) sender := &Client{username: "alice", send: make(chan interface{}, 10)} recipient := &Client{username: "bob", send: make(chan interface{}, 10)} @@ -847,7 +862,7 @@ func TestBroadcastDM(t *testing.T) { } func TestBroadcastDMCaseInsensitive(t *testing.T) { - hub := NewHub("", "", "", nil) + hub := mustNewHub(t, "", "", "", nil) sender := &Client{username: "Alice", send: make(chan interface{}, 10)} recipient := &Client{username: "BOB", send: make(chan interface{}, 10)} @@ -875,7 +890,7 @@ func TestBroadcastDMCaseInsensitive(t *testing.T) { } func TestConcurrentChannelOperations(t *testing.T) { - hub := NewHub("", "", "", nil) + hub := mustNewHub(t, "", "", "", nil) done := make(chan bool, 4) clients := make([]*Client, 10) diff --git a/server/integration_test.go b/server/integration_test.go index 6bef536..7969c12 100644 --- a/server/integration_test.go +++ b/server/integration_test.go @@ -29,7 +29,7 @@ func TestIntegrationMessageFlow(t *testing.T) { CreateSchema(db) // Create hub (for future use in tests) - _ = NewHub("./plugins", "./data", "http://registry.example.com", db) + _ = mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) // Test message insertion and retrieval now := time.Now() @@ -80,7 +80,7 @@ func TestIntegrationUserBanFlow(t *testing.T) { CreateSchema(db) // Create hub - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) username := "troublemaker" adminUsername := "admin" @@ -282,7 +282,7 @@ func TestIntegrationWebSocketHandshakeReplayOnReconnect(t *testing.T) { } tdir := t.TempDir() - hub := NewHub(tdir, tdir, "", db) + hub := mustNewHub(t, tdir, tdir, "", db) go hub.Run() handler := ServeWs(hub, db, nil, "admin-key", false, 10<<20, filepath.Join(tdir, "test.db")) @@ -355,7 +355,7 @@ func TestIntegrationConcurrentOperations(t *testing.T) { CreateSchema(db) // Create hub - hub := NewHub("./plugins", "./data", "http://registry.example.com", db) + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) // Test concurrent message insertions with proper synchronization var wg sync.WaitGroup diff --git a/server/loadverify_bench_test.go b/server/loadverify_bench_test.go index 29e7b0f..32c61c3 100644 --- a/server/loadverify_bench_test.go +++ b/server/loadverify_bench_test.go @@ -57,7 +57,10 @@ func setupLoadverifyHub(b *testing.B, total, inChannel int) *Hub { } b.Cleanup(func() { db.Close() }) CreateSchema(db) - hub := NewHub("", "", "", db) + hub, err := NewHub("", "", "", db) + if err != nil { + b.Fatalf("NewHub: %v", err) + } go hub.Run() time.Sleep(20 * time.Millisecond) From a2812b5c0d9800d940b90af5d5b3f1d4d18a1338 Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:32:11 -0400 Subject: [PATCH 02/11] refactor(server): split readPump into typed inbound handlers --- server/client.go | 300 +-------------------------- server/client_dispatch.go | 369 +++++++++++++++++++++++++++++++++ server/client_dispatch_test.go | 255 +++++++++++++++++++++++ 3 files changed, 625 insertions(+), 299 deletions(-) create mode 100644 server/client_dispatch.go create mode 100644 server/client_dispatch_test.go diff --git a/server/client.go b/server/client.go index dafc44c..19353b5 100644 --- a/server/client.go +++ b/server/client.go @@ -150,305 +150,7 @@ func (c *Client) readPump() { } msgTimestamps = append(msgTimestamps, now) - if msg.Type == shared.FileMessageType && msg.File != nil { - maxBytes := c.maxFileBytes - if maxBytes <= 0 { - maxBytes = 1024 * 1024 - } - if msg.File.Size > maxBytes || int64(len(msg.File.Data)) > maxBytes { - log.Printf("Rejected file from %s: too large (declared %d bytes, payload %d bytes)", c.username, msg.File.Size, len(msg.File.Data)) - c.send <- fileTooLargeSystemMessage(maxBytes) - continue - } - c.stampSenderTimedOutbound(&msg) - c.hub.broadcast <- msg - continue - } - - if msg.Type == shared.EditMessageType && msg.MessageID > 0 { - if contentContainsNUL(msg.Content) { - c.send <- shared.Message{ - Sender: "System", - Content: "Message not sent: invalid character in content", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - if plaintextContentEmpty(msg.Content, msg.Encrypted) { - c.send <- shared.Message{ - Sender: "System", - Content: "Message not sent: empty content", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - if err := EditMessage(c.db, msg.MessageID, c.username, msg.Content, msg.Encrypted); err != nil { - c.send <- shared.Message{ - Sender: "System", - Content: "Edit failed: " + err.Error(), - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - } else { - msg.Edited = true - c.stampSenderTimedOutbound(&msg) - c.hub.broadcast <- msg - } - continue - } - - if msg.Type == shared.DeleteMessage && msg.MessageID > 0 { - if err := DeleteMessage(c.db, msg.MessageID, c.username, c.isAdmin); err != nil { - c.send <- shared.Message{ - Sender: "System", - Content: "Delete failed: " + err.Error(), - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - } else { - c.stampSenderTimedOutbound(&msg) - c.hub.broadcast <- msg - } - continue - } - - if msg.Type == shared.TypingMessage { - msg.Sender = c.username - if strings.TrimSpace(msg.Recipient) != "" { - c.hub.broadcastDM(msg) - } else { - c.stampClientChannel(&msg) - c.hub.broadcast <- msg - } - continue - } - - if msg.Type == shared.ReactionMessage && msg.Reaction != nil { - c.stampSenderTimedOutbound(&msg) - PersistReaction(c.db, msg) - c.hub.broadcast <- msg - continue - } - - if msg.Type == shared.DirectMessage && msg.Recipient != "" { - if contentContainsNUL(msg.Content) { - c.send <- shared.Message{ - Sender: "System", - Content: "Message not sent: invalid character in content", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - if plaintextContentEmpty(msg.Content, msg.Encrypted) { - c.send <- shared.Message{ - Sender: "System", - Content: "Message not sent: empty content", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - c.stampSenderTimedOutbound(&msg) - msgID, err := InsertMessage(c.db, msg) - if err != nil { - log.Printf("Failed to persist DM from %s to %s: %v", c.username, msg.Recipient, err) - c.send <- shared.Message{ - Sender: "System", - Content: "Message not sent: could not save message", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - msg.MessageID = msgID - c.hub.broadcastDM(msg) - continue - } - - if msg.Type == shared.SearchMessage { - results := SearchMessages(c.db, msg.Content, 20) - var sb strings.Builder - if len(results) == 0 { - fmt.Fprintf(&sb, "No results found for: %s", msg.Content) - } else { - sb.WriteString(fmt.Sprintf("Search results for '%s' (%d found):\n", msg.Content, len(results))) - for _, r := range results { - sb.WriteString(fmt.Sprintf(" [%s] %s: %s\n", r.CreatedAt.Format("01/02 15:04"), r.Sender, r.Content)) - } - } - c.send <- shared.Message{ - Sender: "System", - Content: sb.String(), - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - - if msg.Type == shared.PinMessage { - if msg.MessageID == 0 { - pinned := GetPinnedMessages(c.db) - var sb strings.Builder - if len(pinned) == 0 { - sb.WriteString("No pinned messages") - } else { - sb.WriteString(fmt.Sprintf("Pinned messages (%d):\n", len(pinned))) - for _, p := range pinned { - sb.WriteString(fmt.Sprintf(" #%d [%s] %s: %s\n", p.MessageID, p.CreatedAt.Format("01/02 15:04"), p.Sender, p.Content)) - } - } - c.send <- shared.Message{ - Sender: "System", - Content: sb.String(), - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - if !c.isAdmin { - c.send <- shared.Message{ - Sender: "System", - Content: "Only admins can pin messages", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - pinned, err := TogglePinMessage(c.db, msg.MessageID) - if err != nil { - c.send <- shared.Message{ - Sender: "System", - Content: "Pin failed: " + err.Error(), - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - } else { - action := "pinned" - if !pinned { - action = "unpinned" - } - c.hub.broadcast <- shared.Message{ - Sender: "System", - Content: fmt.Sprintf("Message %d %s by %s", msg.MessageID, action, c.username), - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - } - continue - } - - if msg.Type == shared.ReadReceiptType { - msg.Sender = c.username - if msg.MessageID > 0 { - PersistReadReceipt(c.db, c.username, msg.MessageID) - } - c.stampClientChannel(&msg) - c.hub.broadcast <- msg - continue - } - - if msg.Type == shared.JoinChannelType && msg.Channel != "" { - msg.Channel = strings.ToLower(strings.TrimSpace(msg.Channel)) - if msg.Channel == "" { - continue - } - old := c.hub.getClientChannel(c) - if old != msg.Channel { - c.hub.leaveChannel(c, old) - } - c.hub.joinChannel(c, msg.Channel) - PersistUserChannel(c.db, c.username, msg.Channel) - c.send <- shared.Message{ - Sender: "System", - Content: "Joined channel #" + msg.Channel, - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - if msg.Type == shared.LeaveChannelType { - current := c.hub.getClientChannel(c) - if current != "general" { - c.hub.leaveChannel(c, current) - c.hub.joinChannel(c, "general") - PersistUserChannel(c.db, c.username, "general") - c.send <- shared.Message{ - Sender: "System", - Content: "Left #" + current + ", back to #general", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - } - continue - } - if msg.Type == shared.ListChannelsType { - channels := c.hub.listChannels() - current := c.hub.getClientChannel(c) - var lines []string - for _, ch := range channels { - n := len(c.hub.getChannelUsers(ch)) - marker := "" - if ch == current { - marker = " (current)" - } - lines = append(lines, fmt.Sprintf(" #%s - %d user(s)%s", ch, n, marker)) - } - c.send <- shared.Message{ - Sender: "System", - Content: "Active channels:\n" + strings.Join(lines, "\n"), - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - - if strings.HasPrefix(msg.Content, ":") || msg.Type == shared.AdminCommandType { - AdminLogger.Info("Command received", map[string]interface{}{ - "user": c.username, - "command": msg.Content, - "admin": c.isAdmin, - "type": msg.Type, - }) - c.handleCommand(msg.Content) - continue - } - if contentContainsNUL(msg.Content) { - c.send <- shared.Message{ - Sender: "System", - Content: "Message not sent: invalid character in content", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - if plaintextContentEmpty(msg.Content, msg.Encrypted) { - c.send <- shared.Message{ - Sender: "System", - Content: "Message not sent: empty content", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - c.stampSenderTimedOutbound(&msg) - if msg.Type == "" || msg.Type == shared.TextMessage { - msgID, err := InsertMessage(c.db, msg) - if err != nil { - log.Printf("Failed to persist message from %s: %v", c.username, err) - c.send <- shared.Message{ - Sender: "System", - Content: "Message not sent: could not save message", - CreatedAt: time.Now(), - Type: shared.TextMessage, - } - continue - } - msg.MessageID = msgID - } - c.hub.broadcast <- msg + c.dispatchInbound(&msg) } } diff --git a/server/client_dispatch.go b/server/client_dispatch.go new file mode 100644 index 0000000..a4f79dd --- /dev/null +++ b/server/client_dispatch.go @@ -0,0 +1,369 @@ +package server + +import ( + "fmt" + "log" + "strings" + "time" + + "github.com/Cod-e-Codes/marchat/shared" +) + +func (c *Client) dispatchInbound(msg *shared.Message) { + if msg.Type == shared.FileMessageType && msg.File != nil { + c.handleInboundFile(msg) + return + } + + if msg.Type == shared.EditMessageType && msg.MessageID > 0 { + c.handleInboundEdit(msg) + return + } + + if msg.Type == shared.DeleteMessage && msg.MessageID > 0 { + c.handleInboundDelete(msg) + return + } + + if msg.Type == shared.TypingMessage { + c.handleInboundTyping(msg) + return + } + + if msg.Type == shared.ReactionMessage && msg.Reaction != nil { + c.handleInboundReaction(msg) + return + } + + if msg.Type == shared.DirectMessage && msg.Recipient != "" { + c.handleInboundDM(msg) + return + } + + if msg.Type == shared.SearchMessage { + c.handleInboundSearch(msg) + return + } + + if msg.Type == shared.PinMessage { + c.handleInboundPin(msg) + return + } + + if msg.Type == shared.ReadReceiptType { + c.handleInboundReadReceipt(msg) + return + } + + if msg.Type == shared.JoinChannelType && msg.Channel != "" { + c.handleInboundJoinChannel(msg) + return + } + if msg.Type == shared.LeaveChannelType { + c.handleInboundLeaveChannel(msg) + return + } + if msg.Type == shared.ListChannelsType { + c.handleInboundListChannels(msg) + return + } + + if strings.HasPrefix(msg.Content, ":") || msg.Type == shared.AdminCommandType { + c.handleInboundCommand(msg) + return + } + + c.handleInboundText(msg) +} + +func (c *Client) handleInboundFile(msg *shared.Message) { + maxBytes := c.maxFileBytes + if maxBytes <= 0 { + maxBytes = 1024 * 1024 + } + if msg.File.Size > maxBytes || int64(len(msg.File.Data)) > maxBytes { + log.Printf("Rejected file from %s: too large (declared %d bytes, payload %d bytes)", c.username, msg.File.Size, len(msg.File.Data)) + c.send <- fileTooLargeSystemMessage(maxBytes) + return + } + c.stampSenderTimedOutbound(msg) + c.hub.broadcast <- *msg +} + +func (c *Client) handleInboundEdit(msg *shared.Message) { + if contentContainsNUL(msg.Content) { + c.send <- shared.Message{ + Sender: "System", + Content: "Message not sent: invalid character in content", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + if plaintextContentEmpty(msg.Content, msg.Encrypted) { + c.send <- shared.Message{ + Sender: "System", + Content: "Message not sent: empty content", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + if err := EditMessage(c.db, msg.MessageID, c.username, msg.Content, msg.Encrypted); err != nil { + c.send <- shared.Message{ + Sender: "System", + Content: "Edit failed: " + err.Error(), + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + } else { + msg.Edited = true + c.stampSenderTimedOutbound(msg) + c.hub.broadcast <- *msg + } +} + +func (c *Client) handleInboundDelete(msg *shared.Message) { + if err := DeleteMessage(c.db, msg.MessageID, c.username, c.isAdmin); err != nil { + c.send <- shared.Message{ + Sender: "System", + Content: "Delete failed: " + err.Error(), + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + } else { + c.stampSenderTimedOutbound(msg) + c.hub.broadcast <- *msg + } +} + +func (c *Client) handleInboundTyping(msg *shared.Message) { + msg.Sender = c.username + if strings.TrimSpace(msg.Recipient) != "" { + c.hub.broadcastDM(*msg) + } else { + c.stampClientChannel(msg) + c.hub.broadcast <- *msg + } +} + +func (c *Client) handleInboundReaction(msg *shared.Message) { + c.stampSenderTimedOutbound(msg) + PersistReaction(c.db, *msg) + c.hub.broadcast <- *msg +} + +func (c *Client) handleInboundDM(msg *shared.Message) { + if contentContainsNUL(msg.Content) { + c.send <- shared.Message{ + Sender: "System", + Content: "Message not sent: invalid character in content", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + if plaintextContentEmpty(msg.Content, msg.Encrypted) { + c.send <- shared.Message{ + Sender: "System", + Content: "Message not sent: empty content", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + c.stampSenderTimedOutbound(msg) + msgID, err := InsertMessage(c.db, *msg) + if err != nil { + log.Printf("Failed to persist DM from %s to %s: %v", c.username, msg.Recipient, err) + c.send <- shared.Message{ + Sender: "System", + Content: "Message not sent: could not save message", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + msg.MessageID = msgID + c.hub.broadcastDM(*msg) +} + +func (c *Client) handleInboundSearch(msg *shared.Message) { + results := SearchMessages(c.db, msg.Content, 20) + var sb strings.Builder + if len(results) == 0 { + fmt.Fprintf(&sb, "No results found for: %s", msg.Content) + } else { + sb.WriteString(fmt.Sprintf("Search results for '%s' (%d found):\n", msg.Content, len(results))) + for _, r := range results { + sb.WriteString(fmt.Sprintf(" [%s] %s: %s\n", r.CreatedAt.Format("01/02 15:04"), r.Sender, r.Content)) + } + } + c.send <- shared.Message{ + Sender: "System", + Content: sb.String(), + CreatedAt: time.Now(), + Type: shared.TextMessage, + } +} + +func (c *Client) handleInboundPin(msg *shared.Message) { + if msg.MessageID == 0 { + pinned := GetPinnedMessages(c.db) + var sb strings.Builder + if len(pinned) == 0 { + sb.WriteString("No pinned messages") + } else { + sb.WriteString(fmt.Sprintf("Pinned messages (%d):\n", len(pinned))) + for _, p := range pinned { + sb.WriteString(fmt.Sprintf(" #%d [%s] %s: %s\n", p.MessageID, p.CreatedAt.Format("01/02 15:04"), p.Sender, p.Content)) + } + } + c.send <- shared.Message{ + Sender: "System", + Content: sb.String(), + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + if !c.isAdmin { + c.send <- shared.Message{ + Sender: "System", + Content: "Only admins can pin messages", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + pinned, err := TogglePinMessage(c.db, msg.MessageID) + if err != nil { + c.send <- shared.Message{ + Sender: "System", + Content: "Pin failed: " + err.Error(), + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + } else { + action := "pinned" + if !pinned { + action = "unpinned" + } + c.hub.broadcast <- shared.Message{ + Sender: "System", + Content: fmt.Sprintf("Message %d %s by %s", msg.MessageID, action, c.username), + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + } +} + +func (c *Client) handleInboundReadReceipt(msg *shared.Message) { + msg.Sender = c.username + if msg.MessageID > 0 { + PersistReadReceipt(c.db, c.username, msg.MessageID) + } + c.stampClientChannel(msg) + c.hub.broadcast <- *msg +} + +func (c *Client) handleInboundJoinChannel(msg *shared.Message) { + msg.Channel = strings.ToLower(strings.TrimSpace(msg.Channel)) + if msg.Channel == "" { + return + } + old := c.hub.getClientChannel(c) + if old != msg.Channel { + c.hub.leaveChannel(c, old) + } + c.hub.joinChannel(c, msg.Channel) + PersistUserChannel(c.db, c.username, msg.Channel) + c.send <- shared.Message{ + Sender: "System", + Content: "Joined channel #" + msg.Channel, + CreatedAt: time.Now(), + Type: shared.TextMessage, + } +} + +func (c *Client) handleInboundLeaveChannel(msg *shared.Message) { + current := c.hub.getClientChannel(c) + if current != "general" { + c.hub.leaveChannel(c, current) + c.hub.joinChannel(c, "general") + PersistUserChannel(c.db, c.username, "general") + c.send <- shared.Message{ + Sender: "System", + Content: "Left #" + current + ", back to #general", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + } +} + +func (c *Client) handleInboundListChannels(msg *shared.Message) { + channels := c.hub.listChannels() + current := c.hub.getClientChannel(c) + var lines []string + for _, ch := range channels { + n := len(c.hub.getChannelUsers(ch)) + marker := "" + if ch == current { + marker = " (current)" + } + lines = append(lines, fmt.Sprintf(" #%s - %d user(s)%s", ch, n, marker)) + } + c.send <- shared.Message{ + Sender: "System", + Content: "Active channels:\n" + strings.Join(lines, "\n"), + CreatedAt: time.Now(), + Type: shared.TextMessage, + } +} + +func (c *Client) handleInboundCommand(msg *shared.Message) { + AdminLogger.Info("Command received", map[string]interface{}{ + "user": c.username, + "command": msg.Content, + "admin": c.isAdmin, + "type": msg.Type, + }) + c.handleCommand(msg.Content) +} + +func (c *Client) handleInboundText(msg *shared.Message) { + if contentContainsNUL(msg.Content) { + c.send <- shared.Message{ + Sender: "System", + Content: "Message not sent: invalid character in content", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + if plaintextContentEmpty(msg.Content, msg.Encrypted) { + c.send <- shared.Message{ + Sender: "System", + Content: "Message not sent: empty content", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + c.stampSenderTimedOutbound(msg) + if msg.Type == "" || msg.Type == shared.TextMessage { + msgID, err := InsertMessage(c.db, *msg) + if err != nil { + log.Printf("Failed to persist message from %s: %v", c.username, err) + c.send <- shared.Message{ + Sender: "System", + Content: "Message not sent: could not save message", + CreatedAt: time.Now(), + Type: shared.TextMessage, + } + return + } + msg.MessageID = msgID + } + c.hub.broadcast <- *msg +} diff --git a/server/client_dispatch_test.go b/server/client_dispatch_test.go new file mode 100644 index 0000000..ddadf31 --- /dev/null +++ b/server/client_dispatch_test.go @@ -0,0 +1,255 @@ +package server + +import ( + "path/filepath" + "strings" + "testing" + "time" + + "github.com/Cod-e-Codes/marchat/shared" +) + +func setupDispatchTestClient(t *testing.T) (*Client, *Hub) { + t.Helper() + + tdir := t.TempDir() + dbPath := filepath.Join(tdir, "test.db") + db, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB: %v", err) + } + t.Cleanup(func() { db.Close() }) + CreateSchema(db) + + hub := mustNewHub(t, tdir, tdir, "", db) + go hub.Run() + + client := &Client{ + hub: hub, + send: make(chan interface{}, 64), + db: db, + username: "dispatchuser", + maxFileBytes: 1024, + pluginCommandHandler: hub.pluginCommandHandler, + } + hub.clientsMutex.Lock() + hub.clients[client] = true + hub.clientsMutex.Unlock() + hub.joinChannel(client, "general") + + return client, hub +} + +func drainSend(ch chan interface{}) { + for { + select { + case <-ch: + default: + return + } + } +} + +func waitSystemSend(t *testing.T, ch chan interface{}, substr string, timeout time.Duration) bool { + t.Helper() + deadline := time.After(timeout) + for { + select { + case <-deadline: + return false + case v := <-ch: + msg, ok := v.(shared.Message) + if !ok { + continue + } + if msg.Sender == "System" && strings.Contains(msg.Content, substr) { + return true + } + } + } +} + +func TestDispatchInboundRoutesNoPanic(t *testing.T) { + client, _ := setupDispatchTestClient(t) + + msgID, err := InsertMessage(client.db, shared.Message{ + Sender: client.username, + Content: "seed for edit/delete", + CreatedAt: time.Now(), + Type: shared.TextMessage, + }) + if err != nil { + t.Fatalf("InsertMessage: %v", err) + } + + cases := []struct { + name string + msg shared.Message + }{ + {name: "file", msg: shared.Message{ + Type: shared.FileMessageType, + File: &shared.FileMeta{Filename: "a.txt", Size: 3, Data: []byte("abc")}, + }}, + {name: "file_too_large", msg: shared.Message{ + Type: shared.FileMessageType, + File: &shared.FileMeta{Filename: "big.bin", Size: 2048, Data: make([]byte, 2048)}, + }}, + {name: "edit", msg: shared.Message{ + Type: shared.EditMessageType, MessageID: msgID, Content: "edited body", + }}, + {name: "delete", msg: shared.Message{ + Type: shared.DeleteMessage, MessageID: msgID, + }}, + {name: "typing_channel", msg: shared.Message{Type: shared.TypingMessage}}, + {name: "typing_dm", msg: shared.Message{ + Type: shared.TypingMessage, Recipient: "dmpeer", + }}, + {name: "reaction", msg: shared.Message{ + Type: shared.ReactionMessage, + Reaction: &shared.ReactionMeta{ + Emoji: "thumbsup", TargetID: msgID, + }, + }}, + {name: "dm", msg: shared.Message{ + Type: shared.DirectMessage, Recipient: "dmpeer", Content: "hello dm", + }}, + {name: "dm_empty", msg: shared.Message{ + Type: shared.DirectMessage, Recipient: "dmpeer", Content: " ", + }}, + {name: "search", msg: shared.Message{ + Type: shared.SearchMessage, Content: "seed", + }}, + {name: "pin_list", msg: shared.Message{Type: shared.PinMessage}}, + {name: "pin_toggle_non_admin", msg: shared.Message{ + Type: shared.PinMessage, MessageID: msgID, + }}, + {name: "read_receipt", msg: shared.Message{ + Type: shared.ReadReceiptType, MessageID: msgID, + }}, + {name: "join_channel", msg: shared.Message{ + Type: shared.JoinChannelType, Channel: "random", + }}, + {name: "leave_channel", msg: shared.Message{Type: shared.LeaveChannelType}}, + {name: "list_channels", msg: shared.Message{Type: shared.ListChannelsType}}, + {name: "command", msg: shared.Message{Content: ":stats"}}, + {name: "admin_command_type", msg: shared.Message{ + Type: shared.AdminCommandType, Content: ":stats", + }}, + {name: "text", msg: shared.Message{ + Type: shared.TextMessage, Content: "hello world", + }}, + {name: "text_empty", msg: shared.Message{ + Type: shared.TextMessage, Content: "", + }}, + {name: "edit_skipped_zero_id", msg: shared.Message{ + Type: shared.EditMessageType, Content: "orphan edit", + }}, + {name: "file_nil_meta", msg: shared.Message{Type: shared.FileMessageType}}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + drainSend(client.send) + msg := tc.msg + defer func() { + if r := recover(); r != nil { + t.Fatalf("dispatchInbound panicked: %v", r) + } + }() + client.dispatchInbound(&msg) + }) + } +} + +func TestDispatchInboundSystemMessages(t *testing.T) { + client, _ := setupDispatchTestClient(t) + + t.Run("empty_text", func(t *testing.T) { + drainSend(client.send) + msg := shared.Message{Type: shared.TextMessage, Content: " "} + client.dispatchInbound(&msg) + if !waitSystemSend(t, client.send, "empty content", time.Second) { + t.Fatal("expected System reply about empty content") + } + }) + + t.Run("empty_dm", func(t *testing.T) { + drainSend(client.send) + msg := shared.Message{ + Type: shared.DirectMessage, Recipient: "someone", Content: "", + } + client.dispatchInbound(&msg) + if !waitSystemSend(t, client.send, "empty content", time.Second) { + t.Fatal("expected System reply about empty DM content") + } + }) + + t.Run("file_too_large", func(t *testing.T) { + drainSend(client.send) + msg := shared.Message{ + Type: shared.FileMessageType, + File: &shared.FileMeta{ + Filename: "huge.bin", + Size: 2048, + Data: make([]byte, 2048), + }, + } + client.dispatchInbound(&msg) + if !waitSystemSend(t, client.send, "exceeds maximum size limit", time.Second) { + t.Fatal("expected System reply about file size") + } + }) + + t.Run("pin_non_admin", func(t *testing.T) { + msgID, err := InsertMessage(client.db, shared.Message{ + Sender: client.username, Content: "pin me", CreatedAt: time.Now(), + }) + if err != nil { + t.Fatalf("InsertMessage: %v", err) + } + drainSend(client.send) + msg := shared.Message{Type: shared.PinMessage, MessageID: msgID} + client.dispatchInbound(&msg) + if !waitSystemSend(t, client.send, "Only admins can pin messages", time.Second) { + t.Fatal("expected System reply denying pin") + } + }) + + t.Run("join_channel", func(t *testing.T) { + drainSend(client.send) + msg := shared.Message{Type: shared.JoinChannelType, Channel: "ops"} + client.dispatchInbound(&msg) + if !waitSystemSend(t, client.send, "Joined channel #ops", time.Second) { + t.Fatal("expected System join confirmation") + } + }) + + t.Run("search_no_results", func(t *testing.T) { + drainSend(client.send) + msg := shared.Message{ + Type: shared.SearchMessage, Content: "zzznomatchzzzz", + } + client.dispatchInbound(&msg) + if !waitSystemSend(t, client.send, "No results found for:", time.Second) { + t.Fatal("expected System search empty reply") + } + }) + + t.Run("pin_list_empty", func(t *testing.T) { + drainSend(client.send) + msg := shared.Message{Type: shared.PinMessage} + client.dispatchInbound(&msg) + if !waitSystemSend(t, client.send, "No pinned messages", time.Second) { + t.Fatal("expected System pinned list empty reply") + } + }) + + t.Run("list_channels", func(t *testing.T) { + drainSend(client.send) + msg := shared.Message{Type: shared.ListChannelsType} + client.dispatchInbound(&msg) + if !waitSystemSend(t, client.send, "Active channels:", time.Second) { + t.Fatal("expected System channel list reply") + } + }) +} From 5e1b30ca3e7ac124f064148c5d20bb04990f8b91 Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:37:00 -0400 Subject: [PATCH 03/11] feat(server): add versioned schema migrator with hard-fail verification --- .cursor/skills/database-marchat/SKILL.md | 11 +- CHANGELOG.md | 1 + cmd/server/main.go | 4 +- server/handlers.go | 226 --------------- server/migrate.go | 346 +++++++++++++++++++++++ server/migrate_test.go | 101 +++++++ 6 files changed, 458 insertions(+), 231 deletions(-) create mode 100644 server/migrate.go create mode 100644 server/migrate_test.go diff --git a/.cursor/skills/database-marchat/SKILL.md b/.cursor/skills/database-marchat/SKILL.md index 032160b..4f71edb 100644 --- a/.cursor/skills/database-marchat/SKILL.md +++ b/.cursor/skills/database-marchat/SKILL.md @@ -7,7 +7,9 @@ description: >- paths: - "server/db.go" - "server/db_dialect.go" + - "server/migrate.go" - "server/db_*_test.go" + - "server/migrate_test.go" - "server/handlers.go" - "server/message_state.go" --- @@ -46,10 +48,11 @@ Locally, CI smoke tests skip without env vars. See `testing-marchat` skill. ## Schema change workflow -1. Update `CreateSchema` / migrations in `db.go` with dialect branches. -2. Add or extend `db_dialect_test.go` for new SQL fragments. -3. Run `go test ./server/...`. -4. Document env or migration notes in `ARCHITECTURE.md` / `CHANGELOG.md` if user-visible. +1. Add a new migration step in `server/migrate.go` (`applyMigrationV2`, etc.) and bump `currentSchemaVersion`; extend `verifySchema` when new required tables or columns ship. +2. `MigrateSchema` runs ordered migrations, records `schema_version`, and verifies required tables (including `ban_history.expires_at`). `CreateSchema` in the same file is a thin `log.Fatal` wrapper for tests. +3. Add or extend `db_dialect_test.go` for new SQL fragments. +4. Run `go test ./server/...`. +5. Document env or migration notes in `ARCHITECTURE.md` / `CHANGELOG.md` if user-visible. ## Backup diff --git a/CHANGELOG.md b/CHANGELOG.md index cb97853..1ba4a34 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,7 @@ On **`main`** only; not part of the latest tagged release until you tag and publ - **Server**: **Fix:** SQLite `InitDB` applies `busy_timeout` / WAL / related pragmas via the DSN on every connection and sets `MaxOpenConns(1)` / `MaxIdleConns(1)`, so concurrent inserts no longer fail with `SQLITE_BUSY` from one-shot `PRAGMA` + the default `database/sql` pool ([#118](https://github.com/Cod-e-Codes/marchat/issues/118)). - **Dependencies**: **modernc.org/sqlite** v1.56.0 (journal-rollback corruption fix; **modernc.org/libc** v1.74.4); **github.com/lucasb-eyer/go-colorful** v1.4.1. - **Server**: **Fix:** active permanent bans and unexpired temp kicks load from `ban_history` on hub start (`expires_at` NULL = permanent; non-NULL = kick expiry), so moderation survives process restart. Pre-upgrade open rows without `expires_at` load as permanent. +- **Server**: Database schema uses versioned migrations (`MigrateSchema`, `schema_version` table) and **hard-fails** startup when required tables or columns are missing instead of logging warnings and continuing. ## v1.3.4 diff --git a/cmd/server/main.go b/cmd/server/main.go index ec30494..e23518b 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -282,7 +282,9 @@ func main() { if err != nil { log.Fatalf("Failed to initialize database: %v", err) } - server.CreateSchema(db) + if err := server.MigrateSchema(db); err != nil { + log.Fatalf("Failed to migrate database schema: %v", err) + } // Set up plugin directories pluginDir := cfg.ConfigDir + "/plugins" diff --git a/server/handlers.go b/server/handlers.go index c5bde47..255db88 100644 --- a/server/handlers.go +++ b/server/handlers.go @@ -79,232 +79,6 @@ type UserList struct { Users []string `json:"users"` } -func CreateSchema(db *sql.DB) { - dialect := getDBDialect(db) - idColumn := "id INTEGER PRIMARY KEY AUTOINCREMENT" - boolDefault := "BOOLEAN DEFAULT 0" - dateTimeType := "DATETIME" - blobType := "BLOB" - textType := "TEXT" - keyedTextType := "TEXT" - channelColumnType := "TEXT" - switch dialect { - case DialectPostgres: - idColumn = "id BIGSERIAL PRIMARY KEY" - boolDefault = "BOOLEAN DEFAULT FALSE" - dateTimeType = "TIMESTAMPTZ" - blobType = "BYTEA" - case DialectMySQL: - idColumn = "id BIGINT PRIMARY KEY AUTO_INCREMENT" - boolDefault = "BOOLEAN DEFAULT FALSE" - dateTimeType = "DATETIME" - blobType = "LONGBLOB" - textType = "LONGTEXT" - keyedTextType = "VARCHAR(191)" - channelColumnType = keyedTextType - } - - // First, create the basic messages table if it doesn't exist - basicSchema := fmt.Sprintf(` - CREATE TABLE IF NOT EXISTS messages ( - %s, - sender %s, - content %s, - created_at %s, - is_encrypted %s, - message_id INTEGER NOT NULL DEFAULT 0, - edited %s, - deleted %s, - pinned %s, - encrypted_data %s, - nonce %s, - recipient %s, - channel %s NOT NULL DEFAULT 'general' - );`, idColumn, textType, textType, dateTimeType, boolDefault, boolDefault, boolDefault, boolDefault, blobType, blobType, textType, channelColumnType) - - _, err := dbExec(db, basicSchema) - if err != nil { - log.Fatal("failed to create basic schema:", err) - } - - // Migrations: add columns if they don't exist - boolMigrationDefault := "BOOLEAN DEFAULT 0" - if dialect == DialectPostgres || dialect == DialectMySQL { - boolMigrationDefault = "BOOLEAN DEFAULT FALSE" - } - migrations := []struct { - column string - ddl string - }{ - {"message_id", `ALTER TABLE messages ADD COLUMN message_id INTEGER DEFAULT 0`}, - {"edited", fmt.Sprintf("ALTER TABLE messages ADD COLUMN edited %s", boolMigrationDefault)}, - {"deleted", fmt.Sprintf("ALTER TABLE messages ADD COLUMN deleted %s", boolMigrationDefault)}, - {"pinned", fmt.Sprintf("ALTER TABLE messages ADD COLUMN pinned %s", boolMigrationDefault)}, - {"channel", fmt.Sprintf("ALTER TABLE messages ADD COLUMN channel %s NOT NULL DEFAULT 'general'", channelColumnType)}, - } - - for _, m := range migrations { - var exists int - switch dialect { - case DialectPostgres: - err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = current_schema() AND table_name = 'messages' AND column_name = ?`, m.column).Scan(&exists) - case DialectMySQL: - err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = 'messages' AND column_name = ?`, m.column).Scan(&exists) - default: - err = dbQueryRow(db, `SELECT COUNT(*) FROM pragma_table_info('messages') WHERE name=?`, m.column).Scan(&exists) - } - if err != nil { - log.Printf("Warning: failed to check for %s column: %v", m.column, err) - continue - } - if exists == 0 { - _, err = dbExec(db, m.ddl) - if err != nil { - log.Printf("Warning: failed to add %s column: %v", m.column, err) - } else { - log.Printf("Added %s column to messages table", m.column) - } - } - } - - // Create user_message_state table - userStateSchema := ` - CREATE TABLE IF NOT EXISTS user_message_state ( - username ` + keyedTextType + ` PRIMARY KEY, - last_message_id INTEGER NOT NULL DEFAULT 0, - last_seen ` + dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP - );` - - _, err = dbExec(db, userStateSchema) - if err != nil { - log.Fatal("failed to create user_message_state table:", err) - } - - // Create ban_history table - banHistoryID := "id INTEGER PRIMARY KEY AUTOINCREMENT" - switch dialect { - case DialectPostgres: - banHistoryID = "id BIGSERIAL PRIMARY KEY" - case DialectMySQL: - banHistoryID = "id BIGINT PRIMARY KEY AUTO_INCREMENT" - } - banHistorySchema := ` - CREATE TABLE IF NOT EXISTS ban_history ( - ` + banHistoryID + `, - username ` + keyedTextType + ` NOT NULL, - banned_at ` + dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP, - unbanned_at ` + dateTimeType + `, - banned_by ` + keyedTextType + ` NOT NULL, - expires_at ` + dateTimeType + ` - );` - - _, err = dbExec(db, banHistorySchema) - if err != nil { - log.Printf("Warning: failed to create ban_history table: %v", err) - } - - // Add expires_at for ban_history created before this column existed. - // NULL expires_at = permanent ban; non-NULL = temporary kick expiry. - { - var expiresExists int - switch dialect { - case DialectPostgres: - err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = current_schema() AND table_name = 'ban_history' AND column_name = ?`, "expires_at").Scan(&expiresExists) - case DialectMySQL: - err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = 'ban_history' AND column_name = ?`, "expires_at").Scan(&expiresExists) - default: - err = dbQueryRow(db, `SELECT COUNT(*) FROM pragma_table_info('ban_history') WHERE name=?`, "expires_at").Scan(&expiresExists) - } - if err != nil { - log.Printf("Warning: failed to check for ban_history.expires_at column: %v", err) - } else if expiresExists == 0 { - _, err = dbExec(db, `ALTER TABLE ban_history ADD COLUMN expires_at `+dateTimeType) - if err != nil { - log.Printf("Warning: failed to add ban_history.expires_at column: %v", err) - } else { - log.Printf("Added expires_at column to ban_history table") - } - } - } - - // Create indexes for performance (MySQL needs a prefix length when indexing LONGTEXT) - recipientIdx := `CREATE INDEX IF NOT EXISTS idx_messages_recipient ON messages(recipient)` - if dialect == DialectMySQL { - recipientIdx = `CREATE INDEX IF NOT EXISTS idx_messages_recipient ON messages(recipient(191))` - } - indexes := []string{ - `CREATE INDEX IF NOT EXISTS idx_messages_message_id ON messages(message_id)`, - `CREATE INDEX IF NOT EXISTS idx_messages_created_at ON messages(created_at)`, - recipientIdx, - `CREATE INDEX IF NOT EXISTS idx_messages_deleted_created_at ON messages(deleted, created_at)`, - `CREATE INDEX IF NOT EXISTS idx_user_message_state_username ON user_message_state(username)`, - `CREATE INDEX IF NOT EXISTS idx_ban_history_username ON ban_history(username)`, - `CREATE INDEX IF NOT EXISTS idx_ban_history_banned_at ON ban_history(banned_at)`, - `CREATE INDEX IF NOT EXISTS idx_ban_history_unbanned_at ON ban_history(unbanned_at)`, - } - - for _, index := range indexes { - q := index - if dialect == DialectMySQL { - // MySQL does not support "CREATE INDEX IF NOT EXISTS ..." (syntax error). - q = strings.Replace(index, "IF NOT EXISTS ", "", 1) - } - _, err = dbExec(db, q) - if err != nil { - if dialect == DialectMySQL && isMySQLDuplicateKeyName(err) { - continue - } - log.Printf("Warning: failed to create index: %v", err) - } - } - - // Migration: Update existing messages to have message_id = id - _, err = dbExec(db, `UPDATE messages SET message_id = id WHERE message_id = 0 OR message_id IS NULL`) - if err != nil { - log.Printf("Warning: failed to migrate existing messages: %v", err) - } else { - log.Printf("Successfully migrated existing messages") - } - - // Reactions table (durable reactions across reconnects) - _, err = dbExec(db, ` - CREATE TABLE IF NOT EXISTS message_reactions ( - `+idColumn+`, - message_id INTEGER NOT NULL, - username `+keyedTextType+` NOT NULL, - emoji `+keyedTextType+` NOT NULL, - created_at `+dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, - UNIQUE(message_id, username, emoji) - );`) - if err != nil { - log.Printf("Warning: failed to create message_reactions table: %v", err) - } - - // Channel memberships table (durable memberships across reconnects) - _, err = dbExec(db, ` - CREATE TABLE IF NOT EXISTS user_channels ( - username `+keyedTextType+` NOT NULL, - channel `+keyedTextType+` NOT NULL, - updated_at `+dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, - PRIMARY KEY (username) - );`) - if err != nil { - log.Printf("Warning: failed to create user_channels table: %v", err) - } - - // Read receipt state tracking - _, err = dbExec(db, ` - CREATE TABLE IF NOT EXISTS read_receipts ( - username `+keyedTextType+` NOT NULL, - message_id INTEGER NOT NULL, - read_at `+dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, - PRIMARY KEY (username, message_id) - );`) - if err != nil { - log.Printf("Warning: failed to create read_receipts table: %v", err) - } -} - func InsertMessage(db *sql.DB, msg shared.Message) (int64, error) { var id int64 channel := strings.ToLower(strings.TrimSpace(msg.Channel)) diff --git a/server/migrate.go b/server/migrate.go new file mode 100644 index 0000000..8ff7841 --- /dev/null +++ b/server/migrate.go @@ -0,0 +1,346 @@ +package server + +import ( + "database/sql" + "fmt" + "log" + "strings" +) + +const currentSchemaVersion = 1 + +type schemaTypes struct { + idColumn string + boolDefault string + boolMigrationDef string + dateTimeType string + blobType string + textType string + keyedTextType string + channelColumnType string + banHistoryID string +} + +func schemaTypesForDialect(dialect DBDialect) schemaTypes { + st := schemaTypes{ + idColumn: "id INTEGER PRIMARY KEY AUTOINCREMENT", + boolDefault: "BOOLEAN DEFAULT 0", + boolMigrationDef: "BOOLEAN DEFAULT 0", + dateTimeType: "DATETIME", + blobType: "BLOB", + textType: "TEXT", + keyedTextType: "TEXT", + channelColumnType: "TEXT", + banHistoryID: "id INTEGER PRIMARY KEY AUTOINCREMENT", + } + switch dialect { + case DialectPostgres: + st.idColumn = "id BIGSERIAL PRIMARY KEY" + st.boolDefault = "BOOLEAN DEFAULT FALSE" + st.boolMigrationDef = "BOOLEAN DEFAULT FALSE" + st.dateTimeType = "TIMESTAMPTZ" + st.blobType = "BYTEA" + st.banHistoryID = "id BIGSERIAL PRIMARY KEY" + case DialectMySQL: + st.idColumn = "id BIGINT PRIMARY KEY AUTO_INCREMENT" + st.boolDefault = "BOOLEAN DEFAULT FALSE" + st.boolMigrationDef = "BOOLEAN DEFAULT FALSE" + st.dateTimeType = "DATETIME" + st.blobType = "LONGBLOB" + st.textType = "LONGTEXT" + st.keyedTextType = "VARCHAR(191)" + st.channelColumnType = st.keyedTextType + st.banHistoryID = "id BIGINT PRIMARY KEY AUTO_INCREMENT" + } + return st +} + +// MigrateSchema applies ordered schema migrations and verifies required tables exist. +// Existing databases without a schema_version row run the v1 baseline idempotently, then record version 1. +func MigrateSchema(db *sql.DB) error { + if err := ensureSchemaVersionTable(db); err != nil { + return fmt.Errorf("schema_version table: %w", err) + } + + version, err := readSchemaVersion(db) + if err != nil { + return fmt.Errorf("read schema version: %w", err) + } + + if version < 1 { + if err := applyMigrationV1(db); err != nil { + return fmt.Errorf("migration v1: %w", err) + } + if err := setSchemaVersion(db, 1); err != nil { + return fmt.Errorf("record schema version 1: %w", err) + } + } + + return verifySchema(db) +} + +func ensureSchemaVersionTable(db *sql.DB) error { + st := schemaTypesForDialect(getDBDialect(db)) + _, err := dbExec(db, fmt.Sprintf(` + CREATE TABLE IF NOT EXISTS schema_version ( + version INTEGER PRIMARY KEY, + applied_at %s NOT NULL DEFAULT CURRENT_TIMESTAMP + );`, st.dateTimeType)) + return err +} + +func readSchemaVersion(db *sql.DB) (int, error) { + var version sql.NullInt64 + err := dbQueryRow(db, `SELECT MAX(version) FROM schema_version`).Scan(&version) + if err != nil { + return 0, err + } + if !version.Valid { + return 0, nil + } + return int(version.Int64), nil +} + +func setSchemaVersion(db *sql.DB, version int) error { + switch getDBDialect(db) { + case DialectPostgres: + _, err := dbExec(db, `INSERT INTO schema_version (version) VALUES (?) ON CONFLICT (version) DO NOTHING`, version) + return err + case DialectMySQL: + _, err := dbExec(db, `INSERT IGNORE INTO schema_version (version) VALUES (?)`, version) + return err + default: + _, err := dbExec(db, `INSERT OR IGNORE INTO schema_version (version) VALUES (?)`, version) + return err + } +} + +func applyMigrationV1(db *sql.DB) error { + dialect := getDBDialect(db) + st := schemaTypesForDialect(dialect) + + basicSchema := fmt.Sprintf(` + CREATE TABLE IF NOT EXISTS messages ( + %s, + sender %s, + content %s, + created_at %s, + is_encrypted %s, + message_id INTEGER NOT NULL DEFAULT 0, + edited %s, + deleted %s, + pinned %s, + encrypted_data %s, + nonce %s, + recipient %s, + channel %s NOT NULL DEFAULT 'general' + );`, st.idColumn, st.textType, st.textType, st.dateTimeType, st.boolDefault, + st.boolDefault, st.boolDefault, st.boolDefault, st.blobType, st.blobType, st.textType, st.channelColumnType) + + if _, err := dbExec(db, basicSchema); err != nil { + return fmt.Errorf("create messages table: %w", err) + } + + migrations := []struct { + column string + ddl string + }{ + {"message_id", `ALTER TABLE messages ADD COLUMN message_id INTEGER DEFAULT 0`}, + {"edited", fmt.Sprintf("ALTER TABLE messages ADD COLUMN edited %s", st.boolMigrationDef)}, + {"deleted", fmt.Sprintf("ALTER TABLE messages ADD COLUMN deleted %s", st.boolMigrationDef)}, + {"pinned", fmt.Sprintf("ALTER TABLE messages ADD COLUMN pinned %s", st.boolMigrationDef)}, + {"channel", fmt.Sprintf("ALTER TABLE messages ADD COLUMN channel %s NOT NULL DEFAULT 'general'", st.channelColumnType)}, + } + + for _, m := range migrations { + exists, err := columnExists(db, "messages", m.column) + if err != nil { + return fmt.Errorf("check messages.%s column: %w", m.column, err) + } + if !exists { + if _, err := dbExec(db, m.ddl); err != nil { + return fmt.Errorf("add messages.%s column: %w", m.column, err) + } + } + } + + userStateSchema := ` + CREATE TABLE IF NOT EXISTS user_message_state ( + username ` + st.keyedTextType + ` PRIMARY KEY, + last_message_id INTEGER NOT NULL DEFAULT 0, + last_seen ` + st.dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP + );` + if _, err := dbExec(db, userStateSchema); err != nil { + return fmt.Errorf("create user_message_state table: %w", err) + } + + banHistorySchema := ` + CREATE TABLE IF NOT EXISTS ban_history ( + ` + st.banHistoryID + `, + username ` + st.keyedTextType + ` NOT NULL, + banned_at ` + st.dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP, + unbanned_at ` + st.dateTimeType + `, + banned_by ` + st.keyedTextType + ` NOT NULL, + expires_at ` + st.dateTimeType + ` + );` + if _, err := dbExec(db, banHistorySchema); err != nil { + return fmt.Errorf("create ban_history table: %w", err) + } + + expiresExists, err := columnExists(db, "ban_history", "expires_at") + if err != nil { + return fmt.Errorf("check ban_history.expires_at column: %w", err) + } + if !expiresExists { + if _, err := dbExec(db, `ALTER TABLE ban_history ADD COLUMN expires_at `+st.dateTimeType); err != nil { + return fmt.Errorf("add ban_history.expires_at column: %w", err) + } + } + + recipientIdx := `CREATE INDEX IF NOT EXISTS idx_messages_recipient ON messages(recipient)` + if dialect == DialectMySQL { + recipientIdx = `CREATE INDEX IF NOT EXISTS idx_messages_recipient ON messages(recipient(191))` + } + indexes := []string{ + `CREATE INDEX IF NOT EXISTS idx_messages_message_id ON messages(message_id)`, + `CREATE INDEX IF NOT EXISTS idx_messages_created_at ON messages(created_at)`, + recipientIdx, + `CREATE INDEX IF NOT EXISTS idx_messages_deleted_created_at ON messages(deleted, created_at)`, + `CREATE INDEX IF NOT EXISTS idx_user_message_state_username ON user_message_state(username)`, + `CREATE INDEX IF NOT EXISTS idx_ban_history_username ON ban_history(username)`, + `CREATE INDEX IF NOT EXISTS idx_ban_history_banned_at ON ban_history(banned_at)`, + `CREATE INDEX IF NOT EXISTS idx_ban_history_unbanned_at ON ban_history(unbanned_at)`, + } + + for _, index := range indexes { + q := index + if dialect == DialectMySQL { + q = strings.Replace(index, "IF NOT EXISTS ", "", 1) + } + if _, err := dbExec(db, q); err != nil { + if dialect == DialectMySQL && isMySQLDuplicateKeyName(err) { + continue + } + return fmt.Errorf("create index %q: %w", index, err) + } + } + + if _, err := dbExec(db, `UPDATE messages SET message_id = id WHERE message_id = 0 OR message_id IS NULL`); err != nil { + return fmt.Errorf("backfill messages.message_id: %w", err) + } + + if _, err := dbExec(db, ` + CREATE TABLE IF NOT EXISTS message_reactions ( + `+st.idColumn+`, + message_id INTEGER NOT NULL, + username `+st.keyedTextType+` NOT NULL, + emoji `+st.keyedTextType+` NOT NULL, + created_at `+st.dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, + UNIQUE(message_id, username, emoji) + );`); err != nil { + return fmt.Errorf("create message_reactions table: %w", err) + } + + if _, err := dbExec(db, ` + CREATE TABLE IF NOT EXISTS user_channels ( + username `+st.keyedTextType+` NOT NULL, + channel `+st.keyedTextType+` NOT NULL, + updated_at `+st.dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (username) + );`); err != nil { + return fmt.Errorf("create user_channels table: %w", err) + } + + if _, err := dbExec(db, ` + CREATE TABLE IF NOT EXISTS read_receipts ( + username `+st.keyedTextType+` NOT NULL, + message_id INTEGER NOT NULL, + read_at `+st.dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (username, message_id) + );`); err != nil { + return fmt.Errorf("create read_receipts table: %w", err) + } + + return nil +} + +func columnExists(db *sql.DB, table, column string) (bool, error) { + var exists int + var err error + switch getDBDialect(db) { + case DialectPostgres: + err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = current_schema() AND table_name = ? AND column_name = ?`, table, column).Scan(&exists) + case DialectMySQL: + err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = ? AND column_name = ?`, table, column).Scan(&exists) + default: + err = dbQueryRow(db, fmt.Sprintf(`SELECT COUNT(*) FROM pragma_table_info(%q) WHERE name=?`, table), column).Scan(&exists) + } + if err != nil { + return false, err + } + return exists > 0, nil +} + +func tableExists(db *sql.DB, table string) (bool, error) { + var exists int + var err error + switch getDBDialect(db) { + case DialectPostgres: + err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = current_schema() AND table_name = ?`, table).Scan(&exists) + case DialectMySQL: + err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = DATABASE() AND table_name = ?`, table).Scan(&exists) + default: + err = dbQueryRow(db, `SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name=?`, table).Scan(&exists) + } + if err != nil { + return false, err + } + return exists > 0, nil +} + +func verifySchema(db *sql.DB) error { + requiredTables := []string{ + "messages", + "user_message_state", + "ban_history", + "message_reactions", + "user_channels", + "read_receipts", + "schema_version", + } + for _, name := range requiredTables { + ok, err := tableExists(db, name) + if err != nil { + return fmt.Errorf("verify table %q: %w", name, err) + } + if !ok { + return fmt.Errorf("required table %q is missing", name) + } + } + + hasExpires, err := columnExists(db, "ban_history", "expires_at") + if err != nil { + return fmt.Errorf("verify ban_history.expires_at: %w", err) + } + if !hasExpires { + return fmt.Errorf("required column ban_history.expires_at is missing") + } + + version, err := readSchemaVersion(db) + if err != nil { + return fmt.Errorf("verify schema version: %w", err) + } + if version < currentSchemaVersion { + return fmt.Errorf("schema version %d is below required %d", version, currentSchemaVersion) + } + + return nil +} + +// CreateSchema applies schema migrations and terminates the process on failure. +// Tests and legacy callers use this wrapper; production startup should call MigrateSchema directly. +func CreateSchema(db *sql.DB) { + if err := MigrateSchema(db); err != nil { + log.Fatal(err) + } +} diff --git a/server/migrate_test.go b/server/migrate_test.go new file mode 100644 index 0000000..3d02fda --- /dev/null +++ b/server/migrate_test.go @@ -0,0 +1,101 @@ +package server + +import ( + "database/sql" + "path/filepath" + "strings" + "testing" +) + +func TestMigrateSchemaFromEmpty(t *testing.T) { + t.Run("memory", func(t *testing.T) { + db, err := InitDB(":memory:") + if err != nil { + t.Fatalf("InitDB: %v", err) + } + defer db.Close() + assertMigrateSchemaFromEmpty(t, db) + }) + + t.Run("file", func(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "migrate.db") + db, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB: %v", err) + } + defer db.Close() + assertMigrateSchemaFromEmpty(t, db) + }) +} + +func assertMigrateSchemaFromEmpty(t *testing.T, db *sql.DB) { + t.Helper() + + if err := MigrateSchema(db); err != nil { + t.Fatalf("first MigrateSchema: %v", err) + } + + version, err := readSchemaVersion(db) + if err != nil { + t.Fatalf("readSchemaVersion: %v", err) + } + if version != 1 { + t.Fatalf("schema version = %d, want 1", version) + } + + for _, table := range []string{ + "messages", "user_message_state", "ban_history", "message_reactions", + "user_channels", "read_receipts", "schema_version", + } { + ok, err := tableExists(db, table) + if err != nil { + t.Fatalf("tableExists(%q): %v", table, err) + } + if !ok { + t.Fatalf("table %q missing after migration", table) + } + } + + hasExpires, err := columnExists(db, "ban_history", "expires_at") + if err != nil { + t.Fatalf("columnExists ban_history.expires_at: %v", err) + } + if !hasExpires { + t.Fatal("ban_history.expires_at column missing") + } + + if err := MigrateSchema(db); err != nil { + t.Fatalf("second MigrateSchema (idempotent): %v", err) + } +} + +func TestMigrateSchemaPartialFailsVerification(t *testing.T) { + db, err := InitDB(":memory:") + if err != nil { + t.Fatalf("InitDB: %v", err) + } + defer db.Close() + + if err := ensureSchemaVersionTable(db); err != nil { + t.Fatalf("ensureSchemaVersionTable: %v", err) + } + if _, err := dbExec(db, `INSERT INTO schema_version (version) VALUES (1)`); err != nil { + t.Fatalf("seed schema_version: %v", err) + } + + // Partial schema: version recorded but ban_history never created. + if err := applyMigrationV1(db); err != nil { + t.Fatalf("applyMigrationV1: %v", err) + } + if _, err := dbExec(db, `DROP TABLE ban_history`); err != nil { + t.Fatalf("drop ban_history: %v", err) + } + + err = MigrateSchema(db) + if err == nil { + t.Fatal("expected MigrateSchema to fail when required tables are missing") + } + if !strings.Contains(err.Error(), "ban_history") { + t.Fatalf("error = %v, want mention of missing ban_history", err) + } +} From 345edba00967b94ed172edf513db69e800c81b9f Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:37:37 -0400 Subject: [PATCH 04/11] fix(server): store permanent bans without far-future expiry sentinel --- CHANGELOG.md | 3 ++- server/hub.go | 25 ++++++++----------------- server/hub_test.go | 31 +++++++++++++++++++++++++++++++ 3 files changed, 41 insertions(+), 18 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 1ba4a34..2c8fa1f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,7 +11,8 @@ On **`main`** only; not part of the latest tagged release until you tag and publ - **Server**: **Fix:** reject empty or whitespace-only plaintext on `text`, `dm`, and `edit` when `encrypted` is false (System reply, no persist/broadcast); encrypted opaque `content` is not treated as empty ([#117](https://github.com/Cod-e-Codes/marchat/issues/117)). - **Server**: **Fix:** SQLite `InitDB` applies `busy_timeout` / WAL / related pragmas via the DSN on every connection and sets `MaxOpenConns(1)` / `MaxIdleConns(1)`, so concurrent inserts no longer fail with `SQLITE_BUSY` from one-shot `PRAGMA` + the default `database/sql` pool ([#118](https://github.com/Cod-e-Codes/marchat/issues/118)). - **Dependencies**: **modernc.org/sqlite** v1.56.0 (journal-rollback corruption fix; **modernc.org/libc** v1.74.4); **github.com/lucasb-eyer/go-colorful** v1.4.1. -- **Server**: **Fix:** active permanent bans and unexpired temp kicks load from `ban_history` on hub start (`expires_at` NULL = permanent; non-NULL = kick expiry), so moderation survives process restart. Pre-upgrade open rows without `expires_at` load as permanent. +- **Server**: **Fix:** active permanent bans and unexpired temp kicks load from `ban_history` on hub start (`expires_at` NULL = permanent; non-NULL = kick expiry), so moderation survives process restart. Pre-upgrade open rows without `expires_at` load as permanent. Permanent bans are presence-only in memory (no 100-year sentinel expiry). +- **Server**: **Fix:** schema bootstrap uses versioned `MigrateSchema` (`schema_version`) and hard-fails when required tables or `ban_history.expires_at` are missing instead of warning and continuing. - **Server**: Database schema uses versioned migrations (`MigrateSchema`, `schema_version` table) and **hard-fails** startup when required tables or columns are missing instead of logging warnings and continuing. ## v1.3.4 diff --git a/server/hub.go b/server/hub.go index 910ef1e..401267b 100644 --- a/server/hub.go +++ b/server/hub.go @@ -33,7 +33,7 @@ type Hub struct { unregister chan *Client // Ban management - bans map[string]time.Time // username -> expiry time (permanent bans use far future time) + bans map[string]struct{} // username -> permanently banned (no expiry) tempKicks map[string]time.Time // username -> kick expiry time (24h temporary) banMutex sync.RWMutex @@ -63,7 +63,7 @@ func NewHub(pluginDir, dataDir, registryURL string, db *sql.DB) (*Hub, error) { broadcast: make(chan interface{}), register: make(chan *Client), unregister: make(chan *Client), - bans: make(map[string]time.Time), + bans: make(map[string]struct{}), tempKicks: make(map[string]time.Time), pluginManager: pluginManager, pluginCommandHandler: pluginCommandHandler, @@ -103,7 +103,7 @@ func (h *Hub) loadModerationState() error { lower := strings.ToLower(username) if !expiresAt.Valid { // Permanent ban (or pre-upgrade open row treated as permanent). - h.bans[lower] = now.Add(100 * 365 * 24 * time.Hour) + h.bans[lower] = struct{}{} continue } if now.Before(expiresAt.Time) { @@ -155,9 +155,8 @@ func (h *Hub) BanUser(username string, adminUsername string) error { // Remove from temporary kicks if present delete(h.tempKicks, lowerUsername) - // Add to permanent bans (using far future time to indicate permanent) - permanentBanTime := time.Now().Add(100 * 365 * 24 * time.Hour) // 100 years in the future - h.bans[lowerUsername] = permanentBanTime + // Add to permanent bans (no expiry) + h.bans[lowerUsername] = struct{}{} AdminLogger.Info("User permanently banned", map[string]interface{}{ "banned_user": username, "admin": adminUsername, @@ -227,7 +226,7 @@ func (h *Hub) IsUserBanned(username string) bool { lowerUsername := strings.ToLower(username) - // Check permanent bans (these don't expire automatically) + // Check permanent bans (no expiry) if _, exists := h.bans[lowerUsername]; exists { return true } @@ -389,22 +388,14 @@ func (h *Hub) AllowUser(username string, adminUsername string) bool { return false } -// CleanupExpiredBans removes expired bans and kicks from the lists +// CleanupExpiredBans removes expired temporary kicks from the lists. +// Permanent bans have no expiry and are never cleared here. func (h *Hub) CleanupExpiredBans() { h.banMutex.Lock() defer h.banMutex.Unlock() now := time.Now() - // Clean up expired permanent bans (shouldn't happen with our 100-year approach, but just in case) - for username, banTime := range h.bans { - if now.After(banTime) { - delete(h.bans, username) - log.Printf("[SYSTEM] Expired permanent ban removed for user: %s", username) - } - } - - // Clean up expired temporary kicks for username, kickTime := range h.tempKicks { if now.After(kickTime) { delete(h.tempKicks, username) diff --git a/server/hub_test.go b/server/hub_test.go index 707510c..3c370cc 100644 --- a/server/hub_test.go +++ b/server/hub_test.go @@ -481,6 +481,37 @@ func TestHubCleanupExpiredBans(t *testing.T) { } } +func TestPermanentBanHasNoExpiry(t *testing.T) { + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + t.Fatalf("Failed to open test database: %v", err) + } + defer db.Close() + CreateSchema(db) + + hub := mustNewHub(t, "./plugins", "./data", "http://registry.example.com", db) + if err := hub.BanUser("dave", "admin"); err != nil { + t.Fatalf("BanUser: %v", err) + } + + hub.banMutex.RLock() + _, permanent := hub.bans["dave"] + _, inTemp := hub.tempKicks["dave"] + hub.banMutex.RUnlock() + + if !permanent { + t.Fatal("permanent ban must be stored in bans map") + } + if inTemp { + t.Fatal("permanent ban must not appear in tempKicks") + } + + hub.CleanupExpiredBans() + if !hub.IsUserBanned("dave") { + t.Fatal("CleanupExpiredBans must not clear permanent bans") + } +} + func TestHubForceDisconnectUser(t *testing.T) { db, err := sql.Open("sqlite", ":memory:") if err != nil { From f11a9dbc05d5deee7dd3ebc0ff0f5e8743d0abb4 Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:38:08 -0400 Subject: [PATCH 05/11] ci: pin golangci-lint v2.12.2 and govulncheck v1.6.0 --- .github/workflows/go.yml | 28 ++++++++++++---------------- CHANGELOG.md | 2 +- 2 files changed, 13 insertions(+), 17 deletions(-) diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 4cfdfbe..a2b4600 100644 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -35,25 +35,28 @@ jobs: - name: Test run: go test -race ./... - - name: Lint (golangci-lint if available, else go vet) + - name: Lint (golangci-lint) run: | - if go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest 2>/dev/null; then - $(go env GOPATH)/bin/golangci-lint run ./... - else - go vet ./... - fi + # Bump from https://github.com/golangci/golangci-lint/releases + go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.12.2 + $(go env GOPATH)/bin/golangci-lint run ./... - name: Govulncheck run: | - go install golang.org/x/vuln/cmd/govulncheck@latest + # Bump from https://pkg.go.dev/golang.org/x/vuln?tab=versions + go install golang.org/x/vuln/cmd/govulncheck@v1.6.0 "$(go env GOPATH)/bin/govulncheck" -show verbose ./... # Nested modules are not part of root `go test ./...` / `go build ./...` (separate go.mod). - name: Nested Go modules (plugin/sdk, plugin/examples/echo) run: | set -euo pipefail - go install golang.org/x/vuln/cmd/govulncheck@latest + # Bump from https://pkg.go.dev/golang.org/x/vuln?tab=versions + go install golang.org/x/vuln/cmd/govulncheck@v1.6.0 VULN="$(go env GOPATH)/bin/govulncheck" + # Bump from https://github.com/golangci/golangci-lint/releases + go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.12.2 + LINT="$(go env GOPATH)/bin/golangci-lint" for dir in plugin/sdk plugin/examples/echo; do echo "::group::$dir" ( @@ -65,17 +68,10 @@ jobs: go test -race ./... go vet ./... "$VULN" -show verbose ./... + "$LINT" run ./... ) echo "::endgroup::" done - LINT="$(go env GOPATH)/bin/golangci-lint" - if [ -x "$LINT" ]; then - for dir in plugin/sdk plugin/examples/echo; do - echo "::group::lint $dir" - (cd "$dir" && "$LINT" run ./...) - echo "::endgroup::" - done - fi database-smoke: runs-on: ubuntu-latest diff --git a/CHANGELOG.md b/CHANGELOG.md index 2c8fa1f..6ec0c08 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,7 +13,7 @@ On **`main`** only; not part of the latest tagged release until you tag and publ - **Dependencies**: **modernc.org/sqlite** v1.56.0 (journal-rollback corruption fix; **modernc.org/libc** v1.74.4); **github.com/lucasb-eyer/go-colorful** v1.4.1. - **Server**: **Fix:** active permanent bans and unexpired temp kicks load from `ban_history` on hub start (`expires_at` NULL = permanent; non-NULL = kick expiry), so moderation survives process restart. Pre-upgrade open rows without `expires_at` load as permanent. Permanent bans are presence-only in memory (no 100-year sentinel expiry). - **Server**: **Fix:** schema bootstrap uses versioned `MigrateSchema` (`schema_version`) and hard-fails when required tables or `ban_history.expires_at` are missing instead of warning and continuing. -- **Server**: Database schema uses versioned migrations (`MigrateSchema`, `schema_version` table) and **hard-fails** startup when required tables or columns are missing instead of logging warnings and continuing. +- **CI**: Pin **golangci-lint** v2.12.2 and **govulncheck** v1.6.0 (no `@latest`) in the main and nested-module jobs. ## v1.3.4 From 093b718378d0923a2ac8a06da29f384ccc527021 Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:40:26 -0400 Subject: [PATCH 06/11] perf(server): O(1) username to client lookup via clientsByUsername --- CHANGELOG.md | 1 + server/client_dispatch_test.go | 1 + server/hub.go | 94 +++++++++++++++++++-------------- server/hub_test.go | 27 ++++++++++ server/loadverify_bench_test.go | 2 + 5 files changed, 85 insertions(+), 40 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 6ec0c08..6efc30e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -14,6 +14,7 @@ On **`main`** only; not part of the latest tagged release until you tag and publ - **Server**: **Fix:** active permanent bans and unexpired temp kicks load from `ban_history` on hub start (`expires_at` NULL = permanent; non-NULL = kick expiry), so moderation survives process restart. Pre-upgrade open rows without `expires_at` load as permanent. Permanent bans are presence-only in memory (no 100-year sentinel expiry). - **Server**: **Fix:** schema bootstrap uses versioned `MigrateSchema` (`schema_version`) and hard-fails when required tables or `ban_history.expires_at` are missing instead of warning and continuing. - **CI**: Pin **golangci-lint** v2.12.2 and **govulncheck** v1.6.0 (no `@latest`) in the main and nested-module jobs. +- **Server**: Connected-user lookups (`KickUser`, `kickUser`, `ForceDisconnectUser`, `broadcastDM`) use an O(1) `clientsByUsername` map under `clientsMutex`. ## v1.3.4 diff --git a/server/client_dispatch_test.go b/server/client_dispatch_test.go index ddadf31..72fcc2b 100644 --- a/server/client_dispatch_test.go +++ b/server/client_dispatch_test.go @@ -34,6 +34,7 @@ func setupDispatchTestClient(t *testing.T) (*Client, *Hub) { } hub.clientsMutex.Lock() hub.clients[client] = true + hub.clientsByUsername[strings.ToLower(client.username)] = client hub.clientsMutex.Unlock() hub.joinChannel(client, "general") diff --git a/server/hub.go b/server/hub.go index 401267b..285ede1 100644 --- a/server/hub.go +++ b/server/hub.go @@ -25,12 +25,13 @@ var ( ) type Hub struct { - clients map[*Client]bool - usernames map[string]struct{} - clientsMutex sync.RWMutex - broadcast chan interface{} - register chan *Client - unregister chan *Client + clients map[*Client]bool + usernames map[string]struct{} // reserved names (handshake), lowercased + clientsByUsername map[string]*Client // connected clients by lowercased username + clientsMutex sync.RWMutex + broadcast chan interface{} + register chan *Client + unregister chan *Client // Ban management bans map[string]struct{} // username -> permanently banned (no expiry) @@ -60,6 +61,7 @@ func NewHub(pluginDir, dataDir, registryURL string, db *sql.DB) (*Hub, error) { h := &Hub{ clients: make(map[*Client]bool), usernames: make(map[string]struct{}), + clientsByUsername: make(map[string]*Client), broadcast: make(chan interface{}), register: make(chan *Client), unregister: make(chan *Client), @@ -270,16 +272,28 @@ func (h *Hub) disconnectClient(target *Client, reason string) { } } +// clientByUsernameLocked returns the connected client for a lowercased username. +// Caller must hold clientsMutex (read or write). +func (h *Hub) clientByUsernameLocked(lowerUsername string) *Client { + return h.clientsByUsername[lowerUsername] +} + +// removeClientLocked removes a client from clients and clientsByUsername. +// Caller must hold clientsMutex for write. +func (h *Hub) removeClientLocked(client *Client) { + delete(h.clients, client) + if client != nil && client.username != "" { + lower := strings.ToLower(client.username) + if h.clientsByUsername[lower] == client { + delete(h.clientsByUsername, lower) + } + } +} + // kickUser forcibly disconnects a user by username. func (h *Hub) kickUser(username string, reason string) { h.clientsMutex.RLock() - var target *Client - for client := range h.clients { - if strings.EqualFold(client.username, username) { - target = client - break - } - } + target := h.clientByUsernameLocked(strings.ToLower(username)) h.clientsMutex.RUnlock() if target == nil { @@ -305,13 +319,7 @@ func (h *Hub) KickUser(username string, adminUsername string) error { } h.clientsMutex.RLock() - var target *Client - for client := range h.clients { - if strings.EqualFold(client.username, username) { - target = client - break - } - } + target := h.clientByUsernameLocked(strings.ToLower(username)) h.clientsMutex.RUnlock() if target == nil { @@ -428,7 +436,7 @@ func (h *Hub) CleanupStaleConnections() { for _, client := range staleClients { if _, exists := h.clients[client]; exists { log.Printf("[CLEANUP] Removing stale connection for user '%s' (IP: %s)", client.username, client.ipAddr) - delete(h.clients, client) + h.removeClientLocked(client) delete(h.usernames, strings.ToLower(client.username)) client.conn.Close() } @@ -442,13 +450,7 @@ func (h *Hub) CleanupStaleConnections() { // ForceDisconnectUser forcibly removes a user from the clients map (admin command for stale connections) func (h *Hub) ForceDisconnectUser(username string, adminUsername string) bool { h.clientsMutex.Lock() - var target *Client - for client := range h.clients { - if strings.EqualFold(client.username, username) { - target = client - break - } - } + target := h.clientByUsernameLocked(strings.ToLower(username)) if target == nil { h.clientsMutex.Unlock() log.Printf("[ADMIN] Force disconnect attempt for '%s' by '%s' - user not found", username, adminUsername) @@ -457,8 +459,7 @@ func (h *Hub) ForceDisconnectUser(username string, adminUsername string) bool { log.Printf("[ADMIN] Force disconnecting user '%s' (IP: %s) by admin '%s'", username, target.ipAddr, adminUsername) - // Remove from clients map - delete(h.clients, target) + h.removeClientLocked(target) delete(h.usernames, strings.ToLower(target.username)) h.clientsMutex.Unlock() @@ -504,6 +505,9 @@ func (h *Hub) Run() { case client := <-h.register: h.clientsMutex.Lock() h.clients[client] = true + if client.username != "" { + h.clientsByUsername[strings.ToLower(client.username)] = client + } h.clientsMutex.Unlock() HubLogger.Info("Client registered", map[string]interface{}{ "username": client.username, @@ -523,7 +527,7 @@ func (h *Hub) Run() { case client := <-h.unregister: h.clientsMutex.Lock() if _, ok := h.clients[client]; ok { - delete(h.clients, client) + h.removeClientLocked(client) delete(h.usernames, strings.ToLower(client.username)) // Intentionally do not close client.send here. // Closing send while readPump is still processing can trigger send-on-closed-channel panics. @@ -562,7 +566,7 @@ func (h *Hub) Run() { case client.send <- message: default: log.Printf("Dropping client %s due to full send channel\n", client.username) - delete(h.clients, client) + h.removeClientLocked(client) delete(h.usernames, strings.ToLower(client.username)) // Fail-fast backpressure handling: drop slow client and close socket. client.conn.Close() @@ -576,7 +580,7 @@ func (h *Hub) Run() { case client.send <- message: default: log.Printf("Dropping client %s due to full send channel\n", client.username) - delete(h.clients, client) + h.removeClientLocked(client) delete(h.usernames, strings.ToLower(client.username)) // Fail-fast backpressure handling: drop slow client and close socket. client.conn.Close() @@ -599,13 +603,23 @@ func (h *Hub) Run() { func (h *Hub) broadcastDM(msg shared.Message) { h.clientsMutex.RLock() defer h.clientsMutex.RUnlock() - for client := range h.clients { - if strings.EqualFold(client.username, msg.Sender) || strings.EqualFold(client.username, msg.Recipient) { - select { - case client.send <- msg: - default: - log.Printf("Dropping DM for client %s due to full send channel", client.username) - } + seen := make(map[*Client]struct{}, 2) + for _, name := range []string{msg.Sender, msg.Recipient} { + if name == "" { + continue + } + client := h.clientByUsernameLocked(strings.ToLower(name)) + if client == nil { + continue + } + if _, ok := seen[client]; ok { + continue + } + seen[client] = struct{}{} + select { + case client.send <- msg: + default: + log.Printf("Dropping DM for client %s due to full send channel", client.username) } } } diff --git a/server/hub_test.go b/server/hub_test.go index 3c370cc..8e028ee 100644 --- a/server/hub_test.go +++ b/server/hub_test.go @@ -19,6 +19,7 @@ func registerTestClient(hub *Hub, username string) *Client { } hub.clientsMutex.Lock() hub.clients[client] = true + hub.clientsByUsername[strings.ToLower(username)] = client hub.clientsMutex.Unlock() return client } @@ -747,6 +748,7 @@ func TestKickUserNonBlocking(t *testing.T) { hub.clientsMutex.Lock() hub.clients[client] = true + hub.clientsByUsername[strings.ToLower(client.username)] = client hub.clientsMutex.Unlock() // kickUser must not block even though the buffer is full, and must not @@ -859,6 +861,26 @@ func TestChannelManagement(t *testing.T) { }) } +func TestClientByUsernameLookup(t *testing.T) { + hub := mustNewHub(t, "", "", "", nil) + alice := registerTestClient(hub, "Alice") + registerTestClient(hub, "bob") + + hub.clientsMutex.RLock() + got := hub.clientByUsernameLocked("alice") + hub.clientsMutex.RUnlock() + if got != alice { + t.Fatalf("expected alice client, got %#v", got) + } + + if err := hub.KickUser("ALICE", "admin"); err != nil { + t.Fatalf("KickUser via O(1) lookup: %v", err) + } + if !hub.IsUserBanned("alice") { + t.Fatal("kick via case-insensitive O(1) lookup should temp-ban") + } +} + func TestBroadcastDM(t *testing.T) { hub := mustNewHub(t, "", "", "", nil) @@ -870,6 +892,9 @@ func TestBroadcastDM(t *testing.T) { hub.clients[sender] = true hub.clients[recipient] = true hub.clients[bystander] = true + hub.clientsByUsername[strings.ToLower(sender.username)] = sender + hub.clientsByUsername[strings.ToLower(recipient.username)] = recipient + hub.clientsByUsername[strings.ToLower(bystander.username)] = bystander hub.clientsMutex.Unlock() msg := shared.Message{ @@ -901,6 +926,8 @@ func TestBroadcastDMCaseInsensitive(t *testing.T) { hub.clientsMutex.Lock() hub.clients[sender] = true hub.clients[recipient] = true + hub.clientsByUsername[strings.ToLower(sender.username)] = sender + hub.clientsByUsername[strings.ToLower(recipient.username)] = recipient hub.clientsMutex.Unlock() msg := shared.Message{ diff --git a/server/loadverify_bench_test.go b/server/loadverify_bench_test.go index 32c61c3..42d398c 100644 --- a/server/loadverify_bench_test.go +++ b/server/loadverify_bench_test.go @@ -30,6 +30,7 @@ import ( "database/sql" "encoding/json" "fmt" + "strings" "testing" "time" @@ -75,6 +76,7 @@ func setupLoadverifyHub(b *testing.B, total, inChannel int) *Hub { loadverifyDrain(c.send) hub.clientsMutex.Lock() hub.clients[c] = true + hub.clientsByUsername[strings.ToLower(c.username)] = c hub.clientsMutex.Unlock() if i < inChannel { hub.joinChannel(c, "bench") From 704be5b8d2e833ac86f25300407bdb43bd4bf4dd Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:40:57 -0400 Subject: [PATCH 07/11] docs(server): document Hub mutex lock ordering --- .cursor/skills/server-marchat/SKILL.md | 2 ++ server/hub.go | 12 ++++++++++++ 2 files changed, 14 insertions(+) diff --git a/.cursor/skills/server-marchat/SKILL.md b/.cursor/skills/server-marchat/SKILL.md index 483b40d..bbb031a 100644 --- a/.cursor/skills/server-marchat/SKILL.md +++ b/.cursor/skills/server-marchat/SKILL.md @@ -22,6 +22,8 @@ App entry: `cmd/server/main.go`. Library: `server/` (hub, client, handlers, db, ## Hub and WebSocket - Per-channel routing, DMs, typing, read receipts, reactions. +- Moderation: permanent bans and unexpired temp kicks load from `ban_history` on hub start; in-memory maps are the hot path. See lock ordering comment on `Hub` in `hub.go`. +- Inbound WebSocket messages: `readPump` rate-limits then `dispatchInbound` / typed handlers in `client_dispatch.go`. - Outbound client messages are channel-stamped from hub membership (`stampClientChannel`); client-supplied `channel` values are ignored for routing. - All outbound/persist paths stamp `sender` from the authenticated session (`stampSenderTimedOutbound`); NUL bytes in persistable `content` are rejected before insert; empty or whitespace-only plaintext on `text` / `dm` / `edit` is rejected when `encrypted` is false (encrypted opaque ciphertext is never treated as empty). - Reserved usernames during handshake (no double-book before registration). diff --git a/server/hub.go b/server/hub.go index 285ede1..27cc486 100644 --- a/server/hub.go +++ b/server/hub.go @@ -24,6 +24,18 @@ var ( ErrKickNotConnected = errors.New("user is not connected") ) +// Hub coordinates WebSocket clients, channels, and moderation state. +// +// Lock ordering (when more than one mutex is needed in the same call path): +// 1. Never hold banMutex across disconnectClient or any potentially blocking +// send on client.send (BanUser / KickUser release banMutex first). +// 2. clientsMutex then banMutex is OK only as separate critical sections +// (lookup under clientsMutex, unlock, then mutate bans under banMutex). +// Do not nest banMutex inside clientsMutex for long operations. +// 3. Prefer clientsMutex before channelMutex when both are required. +// Do not hold banMutex together with channelMutex. +// 4. metricsMutex may be taken briefly while already holding clientsMutex on +// register/unregister paths. Never take clientsMutex while holding metricsMutex. type Hub struct { clients map[*Client]bool usernames map[string]struct{} // reserved names (handshake), lowercased From 80adbe1be13d4569e604c19e899ecb6a0dcc6269 Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:46:43 -0400 Subject: [PATCH 08/11] ci: add golangci-lint v2 config for pinned CI lint --- .golangci.yml | 23 +++++++++++++++++++++++ CHANGELOG.md | 2 +- 2 files changed, 24 insertions(+), 1 deletion(-) create mode 100644 .golangci.yml diff --git a/.golangci.yml b/.golangci.yml new file mode 100644 index 0000000..a40fab7 --- /dev/null +++ b/.golangci.yml @@ -0,0 +1,23 @@ +# golangci-lint v2 config for marchat. +# Pin bumps: https://github.com/golangci/golangci-lint/releases +# Schema: https://golangci-lint.run/docs/product/migration-guide/ +version: "2" + +linters: + # Keep CI focused: govet-equivalent plus staticcheck SA* (bugs), without + # default errcheck-on-Close or ST/QF style nits across the whole tree. + default: none + enable: + - govet + - ineffassign + - staticcheck + settings: + staticcheck: + checks: + - all + - "-ST*" + - "-QF*" + +issues: + max-issues-per-linter: 0 + max-same-issues: 0 diff --git a/CHANGELOG.md b/CHANGELOG.md index 6efc30e..9183f49 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,7 +13,7 @@ On **`main`** only; not part of the latest tagged release until you tag and publ - **Dependencies**: **modernc.org/sqlite** v1.56.0 (journal-rollback corruption fix; **modernc.org/libc** v1.74.4); **github.com/lucasb-eyer/go-colorful** v1.4.1. - **Server**: **Fix:** active permanent bans and unexpired temp kicks load from `ban_history` on hub start (`expires_at` NULL = permanent; non-NULL = kick expiry), so moderation survives process restart. Pre-upgrade open rows without `expires_at` load as permanent. Permanent bans are presence-only in memory (no 100-year sentinel expiry). - **Server**: **Fix:** schema bootstrap uses versioned `MigrateSchema` (`schema_version`) and hard-fails when required tables or `ban_history.expires_at` are missing instead of warning and continuing. -- **CI**: Pin **golangci-lint** v2.12.2 and **govulncheck** v1.6.0 (no `@latest`) in the main and nested-module jobs. +- **CI**: Pin **golangci-lint** v2.12.2 and **govulncheck** v1.6.0 (no `@latest`) in the main and nested-module jobs; add `.golangci.yml` (v2) enabling govet/ineffassign/staticcheck SA* checks. - **Server**: Connected-user lookups (`KickUser`, `kickUser`, `ForceDisconnectUser`, `broadcastDM`) use an O(1) `clientsByUsername` map under `clientsMutex`. ## v1.3.4 From 2af8de535f2da67d8f5a5bde487fc22560837c40 Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 07:56:07 -0400 Subject: [PATCH 09/11] docs: align ARCHITECTURE and TESTING with MigrateSchema and CI lint pins ARCHITECTURE still pointed at CreateSchema in handlers.go and omitted ban_history.expires_at / schema_version; TESTING still suggested @latest. --- ARCHITECTURE.md | 15 +++++++++++++-- CHANGELOG.md | 1 + TESTING.md | 2 +- 3 files changed, 15 insertions(+), 3 deletions(-) diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md index bb2929b..73ba091 100644 --- a/ARCHITECTURE.md +++ b/ARCHITECTURE.md @@ -291,7 +291,7 @@ See [PROTOCOL.md](PROTOCOL.md) for the full message format specification. ## Database Schema -DDL below uses the **SQLite** dialect for readability. PostgreSQL and MySQL variants differ in type names (`BIGSERIAL`/`BIGINT AUTO_INCREMENT` for IDs, `VARCHAR(191)` for indexed text on MySQL, `LONGTEXT`/`LONGBLOB` for large fields) and are generated by `CreateSchema` in `server/handlers.go`. +DDL below uses the **SQLite** dialect for readability. PostgreSQL and MySQL variants differ in type names (`BIGSERIAL`/`BIGINT AUTO_INCREMENT` for IDs, `VARCHAR(191)` for indexed text on MySQL, `LONGTEXT`/`LONGBLOB` for large fields) and are applied by versioned `MigrateSchema` in `server/migrate.go` (`schema_version`; `CreateSchema` is a thin fatal wrapper for tests). ### Tables @@ -333,9 +333,20 @@ CREATE TABLE ban_history ( username TEXT NOT NULL, banned_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, unbanned_at DATETIME, - banned_by TEXT NOT NULL + banned_by TEXT NOT NULL, + expires_at DATETIME ); ``` +Open rows (`unbanned_at` NULL) are the source of truth for active moderation across restart: `expires_at` NULL means permanent ban; non-NULL is a temporary kick expiry. Hub start loads those into in-memory maps. + +#### `schema_version` +```sql +CREATE TABLE schema_version ( + version INTEGER PRIMARY KEY, + applied_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP +); +``` +Records the highest applied migration version from `MigrateSchema`. #### `message_reactions` ```sql diff --git a/CHANGELOG.md b/CHANGELOG.md index 9183f49..60bbad5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -15,6 +15,7 @@ On **`main`** only; not part of the latest tagged release until you tag and publ - **Server**: **Fix:** schema bootstrap uses versioned `MigrateSchema` (`schema_version`) and hard-fails when required tables or `ban_history.expires_at` are missing instead of warning and continuing. - **CI**: Pin **golangci-lint** v2.12.2 and **govulncheck** v1.6.0 (no `@latest`) in the main and nested-module jobs; add `.golangci.yml` (v2) enabling govet/ineffassign/staticcheck SA* checks. - **Server**: Connected-user lookups (`KickUser`, `kickUser`, `ForceDisconnectUser`, `broadcastDM`) use an O(1) `clientsByUsername` map under `clientsMutex`. +- **Docs**: **ARCHITECTURE** documents `MigrateSchema` / `schema_version` / `ban_history.expires_at`; **TESTING** local lint install pins match CI (no `@latest`). ## v1.3.4 diff --git a/TESTING.md b/TESTING.md index 8e15e9d..d7716be 100644 --- a/TESTING.md +++ b/TESTING.md @@ -470,7 +470,7 @@ go test -v -race ./... ### Local validation (PowerShell) -From the repo root, with `golangci-lint` and `govulncheck` on `PATH` (install with `go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest` and `go install golang.org/x/vuln/cmd/govulncheck@latest`): +From the repo root, with `golangci-lint` and `govulncheck` on `PATH` (match CI pins in `.github/workflows/go.yml`: `go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.12.2` and `go install golang.org/x/vuln/cmd/govulncheck@v1.6.0`): ```powershell $env:GOTOOLCHAIN = "auto" From 0f6f461be02ad9fda7d0d797d0b7db03f10054c3 Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 13:07:04 -0400 Subject: [PATCH 10/11] fix(server): enforce single open ban_history row and DB-first moderation Close open rows before insert, persist before updating in-memory maps, and load the latest open row by id so kick-then-ban and legacy duplicates restart cleanly. --- server/client_dispatch_test.go | 5 ++ server/handlers.go | 26 ++++-- server/hub.go | 110 ++++++++++++-------------- server/hub_moderation_persist_test.go | 100 +++++++++++++++++++++++ 4 files changed, 173 insertions(+), 68 deletions(-) diff --git a/server/client_dispatch_test.go b/server/client_dispatch_test.go index 72fcc2b..874494d 100644 --- a/server/client_dispatch_test.go +++ b/server/client_dispatch_test.go @@ -146,6 +146,11 @@ func TestDispatchInboundRoutesNoPanic(t *testing.T) { Type: shared.EditMessageType, Content: "orphan edit", }}, {name: "file_nil_meta", msg: shared.Message{Type: shared.FileMessageType}}, + {name: "reaction_nil_meta", msg: shared.Message{Type: shared.ReactionMessage}}, + {name: "dm_empty_recipient", msg: shared.Message{ + Type: shared.DirectMessage, Content: "orphan dm", + }}, + {name: "join_empty_channel", msg: shared.Message{Type: shared.JoinChannelType}}, } for _, tc := range cases { diff --git a/server/handlers.go b/server/handlers.go index 255db88..765cfbb 100644 --- a/server/handlers.go +++ b/server/handlers.go @@ -317,20 +317,32 @@ func clearUserMessageState(db *sql.DB, username string) error { // recordBanEvent records a ban or kick in ban_history. // expiresAt nil means a permanent ban; non-nil is the temporary kick expiry. +// Any previously open rows for the username are closed first so at most one +// open (unbanned_at IS NULL) row remains - the current moderation state. func recordBanEvent(db *sql.DB, username, bannedBy string, expiresAt *time.Time) error { - _, err := dbExec(db, `INSERT INTO ban_history (username, banned_by, expires_at) VALUES (?, ?, ?)`, username, bannedBy, expiresAt) + tx, err := db.Begin() if err != nil { - log.Printf("Warning: failed to record ban event for user %s: %v", username, err) + return fmt.Errorf("begin ban_history tx: %w", err) } - return err + defer func() { _ = tx.Rollback() }() + + closeQ := rebindQuery(db, `UPDATE ban_history SET unbanned_at = CURRENT_TIMESTAMP WHERE username = ? AND unbanned_at IS NULL`) + if _, err := tx.Exec(closeQ, username); err != nil { + return fmt.Errorf("close open ban_history rows: %w", err) + } + insertQ := rebindQuery(db, `INSERT INTO ban_history (username, banned_by, expires_at) VALUES (?, ?, ?)`) + if _, err := tx.Exec(insertQ, username, bannedBy, expiresAt); err != nil { + return fmt.Errorf("insert ban_history row: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit ban_history tx: %w", err) + } + return nil } -// recordUnbanEvent records an unban event in the ban_history table +// recordUnbanEvent closes all open ban_history rows for username. func recordUnbanEvent(db *sql.DB, username string) error { _, err := dbExec(db, `UPDATE ban_history SET unbanned_at = CURRENT_TIMESTAMP WHERE username = ? AND unbanned_at IS NULL`, username) - if err != nil { - log.Printf("Warning: failed to record unban event for user %s: %v", username, err) - } return err } diff --git a/server/hub.go b/server/hub.go index 27cc486..b4012bc 100644 --- a/server/hub.go +++ b/server/hub.go @@ -26,7 +26,7 @@ var ( // Hub coordinates WebSocket clients, channels, and moderation state. // -// Lock ordering (when more than one mutex is needed in the same call path): +// Hub mutex and blocking-operation rules (not a strict total lock order): // 1. Never hold banMutex across disconnectClient or any potentially blocking // send on client.send (BanUser / KickUser release banMutex first). // 2. clientsMutex then banMutex is OK only as separate critical sections @@ -93,18 +93,21 @@ func NewHub(pluginDir, dataDir, registryURL string, db *sql.DB) (*Hub, error) { } // loadModerationState restores active bans and unexpired temp kicks from ban_history. -// Open rows (unbanned_at IS NULL) with NULL expires_at are permanent bans. -// Open rows with expires_at in the future are temp kicks; expired rows are skipped. +// Writers keep at most one open row per username (close-then-insert). For legacy +// duplicates, the latest open row by id wins (ORDER BY id DESC, first seen). +// NULL expires_at = permanent ban; future expires_at = temp kick; past = skip. func (h *Hub) loadModerationState() error { rows, err := dbQuery(h.db, ` SELECT username, expires_at FROM ban_history - WHERE unbanned_at IS NULL`) + WHERE unbanned_at IS NULL + ORDER BY id DESC`) if err != nil { return fmt.Errorf("load moderation state: %w", err) } defer rows.Close() now := time.Now() + seen := make(map[string]struct{}) h.banMutex.Lock() defer h.banMutex.Unlock() @@ -115,13 +118,19 @@ func (h *Hub) loadModerationState() error { return fmt.Errorf("scan moderation row: %w", err) } lower := strings.ToLower(username) + if _, already := seen[lower]; already { + continue + } + seen[lower] = struct{}{} if !expiresAt.Valid { // Permanent ban (or pre-upgrade open row treated as permanent). h.bans[lower] = struct{}{} + delete(h.tempKicks, lower) continue } if now.Before(expiresAt.Time) { h.tempKicks[lower] = expiresAt.Time + delete(h.bans, lower) } } if err := rows.Err(); err != nil { @@ -166,32 +175,24 @@ func (h *Hub) BanUser(username string, adminUsername string) error { lowerUsername := strings.ToLower(username) - // Remove from temporary kicks if present - delete(h.tempKicks, lowerUsername) + // Persist first when a DB is configured so a successful BanUser survives restart. + if h.getDB() != nil { + if err := recordBanEvent(h.getDB(), lowerUsername, adminUsername, nil); err != nil { + h.banMutex.Unlock() + return fmt.Errorf("persist ban: %w", err) + } + if err := clearUserMessageState(h.getDB(), lowerUsername); err != nil { + log.Printf("Warning: failed to clear message state for banned user %s: %v", username, err) + } + } - // Add to permanent bans (no expiry) + delete(h.tempKicks, lowerUsername) h.bans[lowerUsername] = struct{}{} AdminLogger.Info("User permanently banned", map[string]interface{}{ "banned_user": username, "admin": adminUsername, }) - // Record ban event in database (NULL expires_at = permanent) - if h.getDB() != nil { - err := recordBanEvent(h.getDB(), lowerUsername, adminUsername, nil) - if err != nil { - log.Printf("Warning: failed to record ban event for user %s: %v", username, err) - } - } - - // Clear per-user last_seen bookkeeping on ban (ban gaps use ban_history, not this table). - if h.getDB() != nil { - err := clearUserMessageState(h.getDB(), lowerUsername) - if err != nil { - log.Printf("Warning: failed to clear message state for banned user %s: %v", username, err) - } - } - h.banMutex.Unlock() h.kickUser(username, "You have been permanently banned by an administrator") @@ -205,28 +206,22 @@ func (h *Hub) UnbanUser(username string, adminUsername string) bool { lowerUsername := strings.ToLower(username) if _, exists := h.bans[lowerUsername]; exists { - delete(h.bans, lowerUsername) - AdminLogger.Info("User unbanned", map[string]interface{}{ - "unbanned_user": username, - "admin": adminUsername, - }) - - // Record unban event in database if h.getDB() != nil { - err := recordUnbanEvent(h.getDB(), lowerUsername) - if err != nil { + if err := recordUnbanEvent(h.getDB(), lowerUsername); err != nil { log.Printf("Warning: failed to record unban event for user %s: %v", username, err) + return false } - } - - // Clear per-user last_seen bookkeeping on unban. - if h.getDB() != nil { - err := clearUserMessageState(h.getDB(), lowerUsername) - if err != nil { + if err := clearUserMessageState(h.getDB(), lowerUsername); err != nil { log.Printf("Warning: failed to clear message state for unbanned user %s: %v", username, err) } } + delete(h.bans, lowerUsername) + AdminLogger.Info("User unbanned", map[string]interface{}{ + "unbanned_user": username, + "admin": adminUsername, + }) + return true } log.Printf("[ADMIN] Unban attempt for '%s' by '%s' - user not found in ban list", username, adminUsername) @@ -349,32 +344,27 @@ func (h *Hub) KickUser(username string, adminUsername string) error { return ErrKickPermanentlyBanned } - // Add to temporary kicks for 24 hours kickExpiry := time.Now().Add(24 * time.Hour) - h.tempKicks[lowerUsername] = kickExpiry - AdminLogger.Info("User kicked", map[string]interface{}{ - "kicked_user": username, - "admin": adminUsername, - "until": kickExpiry.Format("2006-01-02 15:04:05"), - }) - // Record kick event in database with 24h expiry + // Persist first when a DB is configured so a successful KickUser survives restart. if h.getDB() != nil { exp := kickExpiry - err := recordBanEvent(h.getDB(), lowerUsername, adminUsername, &exp) - if err != nil { - log.Printf("Warning: failed to record kick event for user %s: %v", username, err) + if err := recordBanEvent(h.getDB(), lowerUsername, adminUsername, &exp); err != nil { + h.banMutex.Unlock() + return fmt.Errorf("persist kick: %w", err) } - } - - // Clear per-user last_seen bookkeeping on kick (ban gaps use ban_history, not this table). - if h.getDB() != nil { - err := clearUserMessageState(h.getDB(), lowerUsername) - if err != nil { + if err := clearUserMessageState(h.getDB(), lowerUsername); err != nil { log.Printf("Warning: failed to clear message state for kicked user %s: %v", username, err) } } + h.tempKicks[lowerUsername] = kickExpiry + AdminLogger.Info("User kicked", map[string]interface{}{ + "kicked_user": username, + "admin": adminUsername, + "until": kickExpiry.Format("2006-01-02 15:04:05"), + }) + h.banMutex.Unlock() h.disconnectClient(target, "You have been kicked by an administrator (24 hour temporary ban)") @@ -390,17 +380,15 @@ func (h *Hub) AllowUser(username string, adminUsername string) bool { // Check if user is in temporary kick list if _, exists := h.tempKicks[lowerUsername]; exists { - delete(h.tempKicks, lowerUsername) - log.Printf("[ADMIN] User '%s' allowed back by '%s' (kick override)", username, adminUsername) - - // Record unban event in database if h.getDB() != nil { - err := recordUnbanEvent(h.getDB(), lowerUsername) - if err != nil { + if err := recordUnbanEvent(h.getDB(), lowerUsername); err != nil { log.Printf("Warning: failed to record allow event for user %s: %v", username, err) + return false } } + delete(h.tempKicks, lowerUsername) + log.Printf("[ADMIN] User '%s' allowed back by '%s' (kick override)", username, adminUsername) return true } diff --git a/server/hub_moderation_persist_test.go b/server/hub_moderation_persist_test.go index db963b5..f9bdbb4 100644 --- a/server/hub_moderation_persist_test.go +++ b/server/hub_moderation_persist_test.go @@ -1,6 +1,7 @@ package server import ( + "database/sql" "net/http/httptest" "path/filepath" "strings" @@ -129,6 +130,105 @@ func TestExpiredTempKickNotLoadedOnRestart(t *testing.T) { } } +func countOpenBanRows(t *testing.T, db *sql.DB, username string) int { + t.Helper() + var n int + if err := dbQueryRow(db, `SELECT COUNT(*) FROM ban_history WHERE username = ? AND unbanned_at IS NULL`, strings.ToLower(username)).Scan(&n); err != nil { + t.Fatalf("count open ban rows: %v", err) + } + return n +} + +func TestKickThenBanLeavesSingleOpenRow(t *testing.T) { + tdir := t.TempDir() + dbPath := filepath.Join(tdir, "moderation.db") + db, err := InitDB(dbPath) + if err != nil { + t.Fatalf("InitDB: %v", err) + } + CreateSchema(db) + + hub := mustNewHub(t, tdir, tdir, "", db) + registerTestClient(hub, "bob") + if err := hub.KickUser("bob", "admin"); err != nil { + t.Fatalf("KickUser: %v", err) + } + if n := countOpenBanRows(t, db, "bob"); n != 1 { + t.Fatalf("after kick open rows = %d, want 1", n) + } + + if err := hub.BanUser("bob", "admin"); err != nil { + t.Fatalf("BanUser: %v", err) + } + if n := countOpenBanRows(t, db, "bob"); n != 1 { + t.Fatalf("after kick+ban open rows = %d, want 1", n) + } + + var expires sql.NullTime + if err := dbQueryRow(db, `SELECT expires_at FROM ban_history WHERE username = ? AND unbanned_at IS NULL`, "bob").Scan(&expires); err != nil { + t.Fatalf("scan open row: %v", err) + } + if expires.Valid { + t.Fatal("open row after permanent ban should have NULL expires_at") + } + + if err := db.Close(); err != nil { + t.Fatalf("close: %v", err) + } + db2, err := InitDB(dbPath) + if err != nil { + t.Fatalf("reopen: %v", err) + } + defer db2.Close() + CreateSchema(db2) + hub2 := mustNewHub(t, tdir, tdir, "", db2) + + if !hub2.IsUserBanned("bob") { + t.Fatal("bob should be permanently banned after restart") + } + hub2.banMutex.RLock() + _, permanent := hub2.bans["bob"] + _, kicked := hub2.tempKicks["bob"] + hub2.banMutex.RUnlock() + if !permanent { + t.Fatal("expected bans map entry after kick-then-ban restart") + } + if kicked { + t.Fatal("tempKicks must not be set when latest open row is permanent") + } +} + +func TestLoadModerationStateLatestOpenRowWins(t *testing.T) { + tdir := t.TempDir() + db, err := InitDB(filepath.Join(tdir, "legacy.db")) + if err != nil { + t.Fatalf("InitDB: %v", err) + } + defer db.Close() + CreateSchema(db) + + kickExpiry := time.Now().Add(12 * time.Hour) + // Legacy duplicate open rows: older permanent, newer kick (without close-before-insert). + if _, err := dbExec(db, `INSERT INTO ban_history (username, banned_by, expires_at) VALUES (?, ?, ?)`, "legacy", "admin", nil); err != nil { + t.Fatalf("insert permanent: %v", err) + } + if _, err := dbExec(db, `INSERT INTO ban_history (username, banned_by, expires_at) VALUES (?, ?, ?)`, "legacy", "admin", kickExpiry); err != nil { + t.Fatalf("insert kick: %v", err) + } + + hub := mustNewHub(t, tdir, tdir, "", db) + hub.banMutex.RLock() + _, permanent := hub.bans["legacy"] + loadedExpiry, kicked := hub.tempKicks["legacy"] + hub.banMutex.RUnlock() + if permanent { + t.Fatal("latest open row is temp kick; bans must be empty") + } + if !kicked || !time.Now().Before(loadedExpiry) { + t.Fatalf("expected temp kick from latest open row, kicked=%v expiry=%v", kicked, loadedExpiry) + } +} + func assertHandshakeRejectedAsBanned(t *testing.T, hub *Hub, db interface { Close() error }, tdir, username string) { From ecc712a85d6bd2ec4aecd89fb04eef5f1dd2cbbb Mon Sep 17 00:00:00 2001 From: Cod-e-Codes Date: Mon, 10 Aug 2026 13:07:04 -0400 Subject: [PATCH 11/11] fix(server): wrap schema migrations in transactions where DDL allows SQLite and Postgres apply each version inside a transaction with rollback on mid-migration failure; MySQL DDL implicitly commits so it stays non-wrapping with schema_version recorded only after a successful apply. --- .cursor/skills/database-marchat/SKILL.md | 4 +- .cursor/skills/server-marchat/SKILL.md | 2 +- ARCHITECTURE.md | 4 +- CHANGELOG.md | 6 +- server/migrate.go | 152 +++++++++++++++++------ server/migrate_test.go | 56 ++++++++- 6 files changed, 175 insertions(+), 49 deletions(-) diff --git a/.cursor/skills/database-marchat/SKILL.md b/.cursor/skills/database-marchat/SKILL.md index 4f71edb..efd5844 100644 --- a/.cursor/skills/database-marchat/SKILL.md +++ b/.cursor/skills/database-marchat/SKILL.md @@ -48,8 +48,8 @@ Locally, CI smoke tests skip without env vars. See `testing-marchat` skill. ## Schema change workflow -1. Add a new migration step in `server/migrate.go` (`applyMigrationV2`, etc.) and bump `currentSchemaVersion`; extend `verifySchema` when new required tables or columns ship. -2. `MigrateSchema` runs ordered migrations, records `schema_version`, and verifies required tables (including `ban_history.expires_at`). `CreateSchema` in the same file is a thin `log.Fatal` wrapper for tests. +1. Add a new migration step in `server/migrate.go` (`applyMigrationV2`, etc.) and bump `currentSchemaVersion`; extend `verifySchema` when new required tables or columns ship. Prefer deterministic DDL for versions after the v1 baseline (avoid inspect-and-reconcile). +2. `MigrateSchema` runs ordered migrations, records `schema_version`, and verifies required tables (including `ban_history.expires_at`). SQLite/Postgres wrap each version in a transaction; MySQL cannot (DDL implicit commit) - document that when changing migrator behavior. `CreateSchema` in the same file is a thin `log.Fatal` wrapper for tests. 3. Add or extend `db_dialect_test.go` for new SQL fragments. 4. Run `go test ./server/...`. 5. Document env or migration notes in `ARCHITECTURE.md` / `CHANGELOG.md` if user-visible. diff --git a/.cursor/skills/server-marchat/SKILL.md b/.cursor/skills/server-marchat/SKILL.md index bbb031a..1727cb4 100644 --- a/.cursor/skills/server-marchat/SKILL.md +++ b/.cursor/skills/server-marchat/SKILL.md @@ -22,7 +22,7 @@ App entry: `cmd/server/main.go`. Library: `server/` (hub, client, handlers, db, ## Hub and WebSocket - Per-channel routing, DMs, typing, read receipts, reactions. -- Moderation: permanent bans and unexpired temp kicks load from `ban_history` on hub start; in-memory maps are the hot path. See lock ordering comment on `Hub` in `hub.go`. +- Moderation: permanent bans and unexpired temp kicks load from `ban_history` on hub start (latest open row per user); writers close open rows before insert and persist before updating in-memory maps. See Hub mutex rules comment on `Hub` in `hub.go`. - Inbound WebSocket messages: `readPump` rate-limits then `dispatchInbound` / typed handlers in `client_dispatch.go`. - Outbound client messages are channel-stamped from hub membership (`stampClientChannel`); client-supplied `channel` values are ignored for routing. - All outbound/persist paths stamp `sender` from the authenticated session (`stampSenderTimedOutbound`); NUL bytes in persistable `content` are rejected before insert; empty or whitespace-only plaintext on `text` / `dm` / `edit` is rejected when `encrypted` is false (encrypted opaque ciphertext is never treated as empty). diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md index 73ba091..d4a4297 100644 --- a/ARCHITECTURE.md +++ b/ARCHITECTURE.md @@ -291,7 +291,7 @@ See [PROTOCOL.md](PROTOCOL.md) for the full message format specification. ## Database Schema -DDL below uses the **SQLite** dialect for readability. PostgreSQL and MySQL variants differ in type names (`BIGSERIAL`/`BIGINT AUTO_INCREMENT` for IDs, `VARCHAR(191)` for indexed text on MySQL, `LONGTEXT`/`LONGBLOB` for large fields) and are applied by versioned `MigrateSchema` in `server/migrate.go` (`schema_version`; `CreateSchema` is a thin fatal wrapper for tests). +DDL below uses the **SQLite** dialect for readability. PostgreSQL and MySQL variants differ in type names (`BIGSERIAL`/`BIGINT AUTO_INCREMENT` for IDs, `VARCHAR(191)` for indexed text on MySQL, `LONGTEXT`/`LONGBLOB` for large fields) and are applied by versioned `MigrateSchema` in `server/migrate.go` (`schema_version`; `CreateSchema` is a thin fatal wrapper for tests). SQLite and PostgreSQL wrap each versioned migration in a transaction. MySQL DDL implicitly commits, so MySQL runs migration steps without a multi-statement transaction; `schema_version` is still recorded only after a successful apply. Future versions after the v1 baseline should be deterministic DDL steps, not inspect-and-reconcile passes. ### Tables @@ -337,7 +337,7 @@ CREATE TABLE ban_history ( expires_at DATETIME ); ``` -Open rows (`unbanned_at` NULL) are the source of truth for active moderation across restart: `expires_at` NULL means permanent ban; non-NULL is a temporary kick expiry. Hub start loads those into in-memory maps. +Open rows (`unbanned_at` NULL) are the source of truth for active moderation across restart: writers close any open row for the username before inserting a new one (at most one open row). `expires_at` NULL means permanent ban; non-NULL is a temporary kick expiry. Hub start loads the latest open row per user (`ORDER BY id DESC`). #### `schema_version` ```sql diff --git a/CHANGELOG.md b/CHANGELOG.md index 60bbad5..e62222a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,11 +12,11 @@ On **`main`** only; not part of the latest tagged release until you tag and publ - **Server**: **Fix:** SQLite `InitDB` applies `busy_timeout` / WAL / related pragmas via the DSN on every connection and sets `MaxOpenConns(1)` / `MaxIdleConns(1)`, so concurrent inserts no longer fail with `SQLITE_BUSY` from one-shot `PRAGMA` + the default `database/sql` pool ([#118](https://github.com/Cod-e-Codes/marchat/issues/118)). - **Dependencies**: **modernc.org/sqlite** v1.56.0 (journal-rollback corruption fix; **modernc.org/libc** v1.74.4); **github.com/lucasb-eyer/go-colorful** v1.4.1. - **Server**: **Fix:** active permanent bans and unexpired temp kicks load from `ban_history` on hub start (`expires_at` NULL = permanent; non-NULL = kick expiry), so moderation survives process restart. Pre-upgrade open rows without `expires_at` load as permanent. Permanent bans are presence-only in memory (no 100-year sentinel expiry). -- **Server**: **Fix:** schema bootstrap uses versioned `MigrateSchema` (`schema_version`) and hard-fails when required tables or `ban_history.expires_at` are missing instead of warning and continuing. -- **CI**: Pin **golangci-lint** v2.12.2 and **govulncheck** v1.6.0 (no `@latest`) in the main and nested-module jobs; add `.golangci.yml` (v2) enabling govet/ineffassign/staticcheck SA* checks. +- **Server**: **Fix:** schema bootstrap uses versioned `MigrateSchema` (`schema_version`) and hard-fails when required tables or `ban_history.expires_at` are missing instead of warning and continuing. SQLite/Postgres apply each version inside a transaction (mid-migration failure rolls back); MySQL DDL cannot participate in multi-statement transactions (implicit commit), so steps run without a wrapping transaction and `schema_version` is recorded only after a successful apply. +- **CI**: Pin **golangci-lint** v2.12.2 and **govulncheck** v1.6.0 (no `@latest`) in the main and nested-module jobs; add `.golangci.yml` (v2) enabling govet/ineffassign/staticcheck with `all` minus `ST*`/`QF*` (bug-focused checks, not SA*-only). - **Server**: Connected-user lookups (`KickUser`, `kickUser`, `ForceDisconnectUser`, `broadcastDM`) use an O(1) `clientsByUsername` map under `clientsMutex`. - **Docs**: **ARCHITECTURE** documents `MigrateSchema` / `schema_version` / `ban_history.expires_at`; **TESTING** local lint install pins match CI (no `@latest`). - +- **Server**: **Fix:** `:ban` / `:kick` close any open `ban_history` row before inserting (at most one open row per user), persist to the DB before updating in-memory enforcement state, and load the latest open row by `id` on hub start so kick-then-ban (and legacy duplicates) restart cleanly. ## v1.3.4 **Released 2026-08-03.** Since **[v1.3.3](https://github.com/Cod-e-Codes/marchat/releases/tag/v1.3.3)**; compare [`v1.3.3...v1.3.4`](https://github.com/Cod-e-Codes/marchat/compare/v1.3.3...v1.3.4). Commits: **`git log v1.3.3..v1.3.4 --oneline`**. diff --git a/server/migrate.go b/server/migrate.go index 8ff7841..b82f065 100644 --- a/server/migrate.go +++ b/server/migrate.go @@ -9,6 +9,10 @@ import ( const currentSchemaVersion = 1 +// migrationFailAfterStep is a test-only hook. When non-empty, applyMigrationV1 +// returns an error after completing the named step (used to verify rollback). +var migrationFailAfterStep string + type schemaTypes struct { idColumn string boolDefault string @@ -55,8 +59,43 @@ func schemaTypesForDialect(dialect DBDialect) schemaTypes { return st } +// migrationConn executes SQL against either *sql.DB or *sql.Tx while keeping +// dialect rebinding keyed off the parent *sql.DB handle. +type migrationConn struct { + db *sql.DB + tx *sql.Tx +} + +func (c migrationConn) Exec(query string, args ...interface{}) (sql.Result, error) { + q := rebindQuery(c.db, query) + if c.tx != nil { + return c.tx.Exec(q, args...) + } + return c.db.Exec(q, args...) +} + +func (c migrationConn) QueryRow(query string, args ...interface{}) *sql.Row { + q := rebindQuery(c.db, query) + if c.tx != nil { + return c.tx.QueryRow(q, args...) + } + return c.db.QueryRow(q, args...) +} + +func (c migrationConn) maybeFail(step string) error { + if migrationFailAfterStep != "" && migrationFailAfterStep == step { + return fmt.Errorf("injected migration failure after %s", step) + } + return nil +} + // MigrateSchema applies ordered schema migrations and verifies required tables exist. // Existing databases without a schema_version row run the v1 baseline idempotently, then record version 1. +// +// SQLite and PostgreSQL apply each versioned migration (DDL + version row) inside a single +// transaction so a mid-migration failure rolls back. MySQL DDL implicitly commits, so MySQL +// runs steps without a multi-statement transaction; each statement is still statement-atomic +// on InnoDB, and schema_version is recorded only after applyMigrationV1 returns nil. func MigrateSchema(db *sql.DB) error { if err := ensureSchemaVersionTable(db); err != nil { return fmt.Errorf("schema_version table: %w", err) @@ -68,17 +107,43 @@ func MigrateSchema(db *sql.DB) error { } if version < 1 { - if err := applyMigrationV1(db); err != nil { + if err := applyVersionedMigration(db, 1, applyMigrationV1); err != nil { return fmt.Errorf("migration v1: %w", err) } - if err := setSchemaVersion(db, 1); err != nil { - return fmt.Errorf("record schema version 1: %w", err) - } } return verifySchema(db) } +func applyVersionedMigration(db *sql.DB, version int, apply func(migrationConn) error) error { + dialect := getDBDialect(db) + if dialect == DialectMySQL { + // MySQL DDL ends any open transaction (implicit commit). Do not wrap. + if err := apply(migrationConn{db: db}); err != nil { + return err + } + return setSchemaVersion(db, nil, version) + } + + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("begin migration transaction: %w", err) + } + defer func() { _ = tx.Rollback() }() + + conn := migrationConn{db: db, tx: tx} + if err := apply(conn); err != nil { + return err + } + if err := setSchemaVersion(db, tx, version); err != nil { + return fmt.Errorf("record schema version %d: %w", version, err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit migration transaction: %w", err) + } + return nil +} + func ensureSchemaVersionTable(db *sql.DB) error { st := schemaTypesForDialect(getDBDialect(db)) _, err := dbExec(db, fmt.Sprintf(` @@ -101,22 +166,23 @@ func readSchemaVersion(db *sql.DB) (int, error) { return int(version.Int64), nil } -func setSchemaVersion(db *sql.DB, version int) error { +func setSchemaVersion(db *sql.DB, tx *sql.Tx, version int) error { + conn := migrationConn{db: db, tx: tx} + var q string switch getDBDialect(db) { case DialectPostgres: - _, err := dbExec(db, `INSERT INTO schema_version (version) VALUES (?) ON CONFLICT (version) DO NOTHING`, version) - return err + q = `INSERT INTO schema_version (version) VALUES (?) ON CONFLICT (version) DO NOTHING` case DialectMySQL: - _, err := dbExec(db, `INSERT IGNORE INTO schema_version (version) VALUES (?)`, version) - return err + q = `INSERT IGNORE INTO schema_version (version) VALUES (?)` default: - _, err := dbExec(db, `INSERT OR IGNORE INTO schema_version (version) VALUES (?)`, version) - return err + q = `INSERT OR IGNORE INTO schema_version (version) VALUES (?)` } + _, err := conn.Exec(q, version) + return err } -func applyMigrationV1(db *sql.DB) error { - dialect := getDBDialect(db) +func applyMigrationV1(conn migrationConn) error { + dialect := getDBDialect(conn.db) st := schemaTypesForDialect(dialect) basicSchema := fmt.Sprintf(` @@ -137,9 +203,12 @@ func applyMigrationV1(db *sql.DB) error { );`, st.idColumn, st.textType, st.textType, st.dateTimeType, st.boolDefault, st.boolDefault, st.boolDefault, st.boolDefault, st.blobType, st.blobType, st.textType, st.channelColumnType) - if _, err := dbExec(db, basicSchema); err != nil { + if _, err := conn.Exec(basicSchema); err != nil { return fmt.Errorf("create messages table: %w", err) } + if err := conn.maybeFail("create messages"); err != nil { + return err + } migrations := []struct { column string @@ -153,12 +222,12 @@ func applyMigrationV1(db *sql.DB) error { } for _, m := range migrations { - exists, err := columnExists(db, "messages", m.column) + exists, err := columnExistsConn(conn, "messages", m.column) if err != nil { return fmt.Errorf("check messages.%s column: %w", m.column, err) } if !exists { - if _, err := dbExec(db, m.ddl); err != nil { + if _, err := conn.Exec(m.ddl); err != nil { return fmt.Errorf("add messages.%s column: %w", m.column, err) } } @@ -170,9 +239,12 @@ func applyMigrationV1(db *sql.DB) error { last_message_id INTEGER NOT NULL DEFAULT 0, last_seen ` + st.dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP );` - if _, err := dbExec(db, userStateSchema); err != nil { + if _, err := conn.Exec(userStateSchema); err != nil { return fmt.Errorf("create user_message_state table: %w", err) } + if err := conn.maybeFail("create user_message_state"); err != nil { + return err + } banHistorySchema := ` CREATE TABLE IF NOT EXISTS ban_history ( @@ -183,16 +255,16 @@ func applyMigrationV1(db *sql.DB) error { banned_by ` + st.keyedTextType + ` NOT NULL, expires_at ` + st.dateTimeType + ` );` - if _, err := dbExec(db, banHistorySchema); err != nil { + if _, err := conn.Exec(banHistorySchema); err != nil { return fmt.Errorf("create ban_history table: %w", err) } - expiresExists, err := columnExists(db, "ban_history", "expires_at") + expiresExists, err := columnExistsConn(conn, "ban_history", "expires_at") if err != nil { return fmt.Errorf("check ban_history.expires_at column: %w", err) } if !expiresExists { - if _, err := dbExec(db, `ALTER TABLE ban_history ADD COLUMN expires_at `+st.dateTimeType); err != nil { + if _, err := conn.Exec(`ALTER TABLE ban_history ADD COLUMN expires_at ` + st.dateTimeType); err != nil { return fmt.Errorf("add ban_history.expires_at column: %w", err) } } @@ -217,7 +289,7 @@ func applyMigrationV1(db *sql.DB) error { if dialect == DialectMySQL { q = strings.Replace(index, "IF NOT EXISTS ", "", 1) } - if _, err := dbExec(db, q); err != nil { + if _, err := conn.Exec(q); err != nil { if dialect == DialectMySQL && isMySQLDuplicateKeyName(err) { continue } @@ -225,37 +297,37 @@ func applyMigrationV1(db *sql.DB) error { } } - if _, err := dbExec(db, `UPDATE messages SET message_id = id WHERE message_id = 0 OR message_id IS NULL`); err != nil { + if _, err := conn.Exec(`UPDATE messages SET message_id = id WHERE message_id = 0 OR message_id IS NULL`); err != nil { return fmt.Errorf("backfill messages.message_id: %w", err) } - if _, err := dbExec(db, ` + if _, err := conn.Exec(` CREATE TABLE IF NOT EXISTS message_reactions ( - `+st.idColumn+`, + ` + st.idColumn + `, message_id INTEGER NOT NULL, - username `+st.keyedTextType+` NOT NULL, - emoji `+st.keyedTextType+` NOT NULL, - created_at `+st.dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, + username ` + st.keyedTextType + ` NOT NULL, + emoji ` + st.keyedTextType + ` NOT NULL, + created_at ` + st.dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP, UNIQUE(message_id, username, emoji) );`); err != nil { return fmt.Errorf("create message_reactions table: %w", err) } - if _, err := dbExec(db, ` + if _, err := conn.Exec(` CREATE TABLE IF NOT EXISTS user_channels ( - username `+st.keyedTextType+` NOT NULL, - channel `+st.keyedTextType+` NOT NULL, - updated_at `+st.dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, + username ` + st.keyedTextType + ` NOT NULL, + channel ` + st.keyedTextType + ` NOT NULL, + updated_at ` + st.dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP, PRIMARY KEY (username) );`); err != nil { return fmt.Errorf("create user_channels table: %w", err) } - if _, err := dbExec(db, ` + if _, err := conn.Exec(` CREATE TABLE IF NOT EXISTS read_receipts ( - username `+st.keyedTextType+` NOT NULL, + username ` + st.keyedTextType + ` NOT NULL, message_id INTEGER NOT NULL, - read_at `+st.dateTimeType+` NOT NULL DEFAULT CURRENT_TIMESTAMP, + read_at ` + st.dateTimeType + ` NOT NULL DEFAULT CURRENT_TIMESTAMP, PRIMARY KEY (username, message_id) );`); err != nil { return fmt.Errorf("create read_receipts table: %w", err) @@ -265,15 +337,19 @@ func applyMigrationV1(db *sql.DB) error { } func columnExists(db *sql.DB, table, column string) (bool, error) { + return columnExistsConn(migrationConn{db: db}, table, column) +} + +func columnExistsConn(conn migrationConn, table, column string) (bool, error) { var exists int var err error - switch getDBDialect(db) { + switch getDBDialect(conn.db) { case DialectPostgres: - err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = current_schema() AND table_name = ? AND column_name = ?`, table, column).Scan(&exists) + err = conn.QueryRow(`SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = current_schema() AND table_name = ? AND column_name = ?`, table, column).Scan(&exists) case DialectMySQL: - err = dbQueryRow(db, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = ? AND column_name = ?`, table, column).Scan(&exists) + err = conn.QueryRow(`SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = ? AND column_name = ?`, table, column).Scan(&exists) default: - err = dbQueryRow(db, fmt.Sprintf(`SELECT COUNT(*) FROM pragma_table_info(%q) WHERE name=?`, table), column).Scan(&exists) + err = conn.QueryRow(fmt.Sprintf(`SELECT COUNT(*) FROM pragma_table_info(%q) WHERE name=?`, table), column).Scan(&exists) } if err != nil { return false, err diff --git a/server/migrate_test.go b/server/migrate_test.go index 3d02fda..23403b7 100644 --- a/server/migrate_test.go +++ b/server/migrate_test.go @@ -69,7 +69,7 @@ func assertMigrateSchemaFromEmpty(t *testing.T, db *sql.DB) { } } -func TestMigrateSchemaPartialFailsVerification(t *testing.T) { +func TestMigrateSchemaFailsWhenRecordedSchemaIsIncomplete(t *testing.T) { db, err := InitDB(":memory:") if err != nil { t.Fatalf("InitDB: %v", err) @@ -83,8 +83,8 @@ func TestMigrateSchemaPartialFailsVerification(t *testing.T) { t.Fatalf("seed schema_version: %v", err) } - // Partial schema: version recorded but ban_history never created. - if err := applyMigrationV1(db); err != nil { + // Version claims v1 but required table is missing (corruption / incomplete install). + if err := applyMigrationV1(migrationConn{db: db}); err != nil { t.Fatalf("applyMigrationV1: %v", err) } if _, err := dbExec(db, `DROP TABLE ban_history`); err != nil { @@ -99,3 +99,53 @@ func TestMigrateSchemaPartialFailsVerification(t *testing.T) { t.Fatalf("error = %v, want mention of missing ban_history", err) } } + +func TestMigrateSchemaV1RollsBackOnInjectedFailure(t *testing.T) { + db, err := InitDB(":memory:") + if err != nil { + t.Fatalf("InitDB: %v", err) + } + defer db.Close() + + migrationFailAfterStep = "create user_message_state" + t.Cleanup(func() { migrationFailAfterStep = "" }) + + err = MigrateSchema(db) + if err == nil { + t.Fatal("expected MigrateSchema to fail on injected mid-migration error") + } + if !strings.Contains(err.Error(), "injected migration failure") { + t.Fatalf("error = %v, want injected failure", err) + } + + version, err := readSchemaVersion(db) + if err != nil { + t.Fatalf("readSchemaVersion: %v", err) + } + if version != 0 { + t.Fatalf("schema version = %d after failed migration, want 0", version) + } + + // messages + user_message_state were created inside the rolled-back transaction. + for _, table := range []string{"messages", "user_message_state", "ban_history"} { + ok, err := tableExists(db, table) + if err != nil { + t.Fatalf("tableExists(%q): %v", table, err) + } + if ok { + t.Fatalf("table %q should not exist after migration rollback", table) + } + } + + migrationFailAfterStep = "" + if err := MigrateSchema(db); err != nil { + t.Fatalf("retry MigrateSchema after rollback: %v", err) + } + version, err = readSchemaVersion(db) + if err != nil { + t.Fatalf("readSchemaVersion after retry: %v", err) + } + if version != 1 { + t.Fatalf("schema version = %d after retry, want 1", version) + } +}