From c7e6702c8c9a03a6631fc45ddb59ba52d2bbff87 Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 19 Jul 2026 15:43:46 +0000 Subject: [PATCH] refactor(server/db): delegate voice, dm, channels, admin to dbgen (D2) voice: all reads (GetVoiceState, GetChannelVoiceStates, GetAllVoiceStates) and writes (join/leave/mute/deafen/camera/screenshare/clear, capacity + camera-limit atomic guards); CountChannelVoiceUsers stays raw (no query). dm: OpenDM, CloseDM, IsDMParticipant, GetDMParticipantIDs; GetOrCreateDMChannel (serializable tx), GetUserDMChannels (aggregate), GetDMRecipient stay raw. channels: List/Get (shared channelFromFields mapper), Create/Update/Delete, slow-mode/max-users setters, permission overrides get/list-for-role/upsert/ delete; ListChannelRoleOverrides + GetChannelTypes (variable IN) stay raw. admin: UserCount, GetServerStats counts (PRAGMA stays raw), ListAllUsers, UpdateUserRole, ForceLogoutUser, GetUserSessions, AdminUpdateChannel, AdminDeleteChannel, LogAudit, GetAuditLog, GetSetting, SetSetting, GetAllSettings, CountUsersWithoutTOTP; AdminCreateChannel + Backup stay raw. Added b2i64 and strToNullPtr mapper helpers; retired the obsolete scanChannel/nullableString. Behavior and public signatures unchanged. Verified: go build ./...; go test ./db ./service ./permissions ./admin; sqlc-verify. Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_01UA17KPvqGBX3XbXYnMf1rA --- Server/db/admin_queries.go | 228 +++++++++++----------------- Server/db/channel_queries.go | 212 +++++++++++--------------- Server/db/dm_queries.go | 47 ++---- Server/db/mappers.go | 18 +++ Server/db/voice_queries.go | 279 ++++++++++++----------------------- 5 files changed, 299 insertions(+), 485 deletions(-) diff --git a/Server/db/admin_queries.go b/Server/db/admin_queries.go index 67537ffd..fcf47c47 100644 --- a/Server/db/admin_queries.go +++ b/Server/db/admin_queries.go @@ -6,14 +6,16 @@ import ( "fmt" "path/filepath" "strings" + + "github.com/owncord/server/db/dbgen" ) // ─── Setup ─────────────────────────────────────────────────────────────────── // UserCount returns the total number of registered users. func (d *DB) UserCount() (int64, error) { - var count int64 - if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM users`).Scan(&count); err != nil { + count, err := d.q.UserCount(dbCtx()) + if err != nil { return 0, fmt.Errorf("UserCount: %w", err) } return count, nil @@ -26,21 +28,23 @@ func (d *DB) UserCount() (int64, error) { // a meaningful value only for file-backed databases). func (d *DB) GetServerStats() (*ServerStats, error) { stats := &ServerStats{} + var err error - if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM users`).Scan(&stats.UserCount); err != nil { + if stats.UserCount, err = d.q.CountUsers(dbCtx()); err != nil { return nil, fmt.Errorf("GetServerStats users: %w", err) } - if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM messages WHERE deleted = 0`).Scan(&stats.MessageCount); err != nil { + if stats.MessageCount, err = d.q.CountActiveMessages(dbCtx()); err != nil { return nil, fmt.Errorf("GetServerStats messages: %w", err) } - if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM channels`).Scan(&stats.ChannelCount); err != nil { + if stats.ChannelCount, err = d.q.CountChannels(dbCtx()); err != nil { return nil, fmt.Errorf("GetServerStats channels: %w", err) } - if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM invites WHERE revoked = 0`).Scan(&stats.InviteCount); err != nil { + if stats.InviteCount, err = d.q.CountActiveInvites(dbCtx()); err != nil { return nil, fmt.Errorf("GetServerStats invites: %w", err) } - // page_count * page_size gives the database size in bytes. + // page_count * page_size gives the database size in bytes. PRAGMAs are not + // expressible as sqlc queries, so they stay on the raw connection. // For :memory: databases this still works (returns the in-memory size). var pageCount, pageSize int64 if err := d.sqlDB.QueryRow(`PRAGMA page_count`).Scan(&pageCount); err != nil { @@ -59,53 +63,40 @@ func (d *DB) GetServerStats() (*ServerStats, error) { // ListAllUsers returns users joined with their role name, ordered by ID. // limit=0 returns no rows. func (d *DB) ListAllUsers(limit, offset int) ([]UserWithRole, error) { - rows, err := d.sqlDB.Query( - `SELECT u.id, u.username, u.avatar, u.role_id, - u.status, u.created_at, u.last_seen, u.banned, u.ban_reason, u.ban_expires, - COALESCE(r.name, '') AS role_name - FROM users u - LEFT JOIN roles r ON r.id = u.role_id - ORDER BY u.id ASC - LIMIT ? OFFSET ?`, - limit, offset, - ) + rows, err := d.q.ListAllUsers(dbCtx(), dbgen.ListAllUsersParams{ + Limit: int64(limit), + Offset: int64(offset), + }) if err != nil { return nil, fmt.Errorf("ListAllUsers: %w", err) } - defer rows.Close() //nolint:errcheck - - var result []UserWithRole - for rows.Next() { - var uwr UserWithRole - var banned int - err := rows.Scan( - &uwr.ID, &uwr.Username, &uwr.Avatar, &uwr.RoleID, - &uwr.Status, &uwr.CreatedAt, &uwr.LastSeen, - &banned, &uwr.BanReason, &uwr.BanExpires, - &uwr.RoleName, - ) - if err != nil { - return nil, fmt.Errorf("ListAllUsers scan: %w", err) - } - uwr.Banned = banned != 0 - result = append(result, uwr) - } - if rows.Err() != nil { - return nil, fmt.Errorf("ListAllUsers rows: %w", rows.Err()) - } - if result == nil { - result = []UserWithRole{} + result := make([]UserWithRole, 0, len(rows)) + for _, r := range rows { + result = append(result, UserWithRole{ + User: User{ + ID: r.ID, + Username: r.Username, + Avatar: r.Avatar, + RoleID: r.RoleID, + Status: r.Status, + CreatedAt: r.CreatedAt, + LastSeen: r.LastSeen, + Banned: r.Banned != 0, + BanReason: r.BanReason, + BanExpires: r.BanExpires, + }, + RoleName: r.RoleName, + }) } return result, nil } // UpdateUserRole changes the role_id of a user. func (d *DB) UpdateUserRole(userID, roleID int64) error { - _, err := d.sqlDB.Exec( - `UPDATE users SET role_id = ? WHERE id = ?`, - roleID, userID, - ) - if err != nil { + if err := d.q.UpdateUserRole(dbCtx(), dbgen.UpdateUserRoleParams{ + RoleID: roleID, + ID: userID, + }); err != nil { return fmt.Errorf("UpdateUserRole: %w", err) } return nil @@ -113,8 +104,7 @@ func (d *DB) UpdateUserRole(userID, roleID int64) error { // ForceLogoutUser deletes all sessions for the given user ID. func (d *DB) ForceLogoutUser(userID int64) error { - _, err := d.sqlDB.Exec(`DELETE FROM sessions WHERE user_id = ?`, userID) - if err != nil { + if err := d.q.ForceLogoutUser(dbCtx(), userID); err != nil { return fmt.Errorf("ForceLogoutUser: %w", err) } return nil @@ -122,33 +112,13 @@ func (d *DB) ForceLogoutUser(userID int64) error { // GetUserSessions returns all active sessions for the given user ID. func (d *DB) GetUserSessions(userID int64) ([]Session, error) { - rows, err := d.sqlDB.Query( - `SELECT id, user_id, token, device, ip_address, created_at, last_used, expires_at - FROM sessions WHERE user_id = ? ORDER BY created_at DESC`, - userID, - ) + rows, err := d.q.GetUserSessions(dbCtx(), userID) if err != nil { return nil, fmt.Errorf("GetUserSessions: %w", err) } - defer rows.Close() //nolint:errcheck - - var sessions []Session - for rows.Next() { - var s Session - err := rows.Scan( - &s.ID, &s.UserID, &s.TokenHash, &s.Device, &s.IP, - &s.CreatedAt, &s.LastUsed, &s.ExpiresAt, - ) - if err != nil { - return nil, fmt.Errorf("GetUserSessions scan: %w", err) - } - sessions = append(sessions, s) - } - if rows.Err() != nil { - return nil, fmt.Errorf("GetUserSessions rows: %w", rows.Err()) - } - if sessions == nil { - sessions = []Session{} + sessions := make([]Session, 0, len(rows)) + for _, s := range rows { + sessions = append(sessions, sessionFromGen(s)) } return sessions, nil } @@ -156,11 +126,12 @@ func (d *DB) GetUserSessions(userID int64) ([]Session, error) { // ─── Channel Management (admin) ─────────────────────────────────────────────── // AdminCreateChannel creates a channel with full field control including position. +// No sqlc query covers this exact INSERT shape, so it stays on raw SQL. func (d *DB) AdminCreateChannel(name, chanType, category, topic string, position int) (int64, error) { res, err := d.sqlDB.Exec( `INSERT INTO channels (name, type, category, topic, position) VALUES (?, ?, ?, ?, ?)`, - name, chanType, nullableString(category), nullableString(topic), position, + name, chanType, strToNullPtr(category), strToNullPtr(topic), position, ) if err != nil { return 0, fmt.Errorf("AdminCreateChannel: %w", err) @@ -170,17 +141,14 @@ func (d *DB) AdminCreateChannel(name, chanType, category, topic string, position // AdminUpdateChannel updates all mutable channel fields. func (d *DB) AdminUpdateChannel(id int64, name, topic string, slowMode, position int, archived bool) error { - archivedInt := 0 - if archived { - archivedInt = 1 - } - _, err := d.sqlDB.Exec( - `UPDATE channels - SET name = ?, topic = ?, slow_mode = ?, position = ?, archived = ? - WHERE id = ?`, - name, nullableString(topic), slowMode, position, archivedInt, id, - ) - if err != nil { + if err := d.q.AdminUpdateChannel(dbCtx(), dbgen.AdminUpdateChannelParams{ + Name: name, + Topic: strToNullPtr(topic), + SlowMode: int64(slowMode), + Position: int64(position), + Archived: b2i64(archived), + ID: id, + }); err != nil { return fmt.Errorf("AdminUpdateChannel: %w", err) } return nil @@ -188,8 +156,7 @@ func (d *DB) AdminUpdateChannel(id int64, name, topic string, slowMode, position // AdminDeleteChannel removes a channel by ID (cascades to messages, etc.). func (d *DB) AdminDeleteChannel(id int64) error { - _, err := d.sqlDB.Exec(`DELETE FROM channels WHERE id = ?`, id) - if err != nil { + if err := d.q.DeleteChannel(dbCtx(), id); err != nil { return fmt.Errorf("AdminDeleteChannel: %w", err) } return nil @@ -199,12 +166,13 @@ func (d *DB) AdminDeleteChannel(id int64) error { // LogAudit inserts an audit log entry. func (d *DB) LogAudit(actorID int64, action, targetType string, targetID int64, detail string) error { - _, err := d.sqlDB.Exec( - `INSERT INTO audit_log (actor_id, action, target_type, target_id, detail) - VALUES (?, ?, ?, ?, ?)`, - actorID, action, targetType, targetID, detail, - ) - if err != nil { + if err := d.q.LogAudit(dbCtx(), dbgen.LogAuditParams{ + ActorID: actorID, + Action: action, + TargetType: targetType, + TargetID: targetID, + Detail: detail, + }); err != nil { return fmt.Errorf("LogAudit: %w", err) } return nil @@ -212,36 +180,25 @@ func (d *DB) LogAudit(actorID int64, action, targetType string, targetID int64, // GetAuditLog returns audit log entries ordered newest-first with pagination. func (d *DB) GetAuditLog(limit, offset int) ([]AuditEntry, error) { - rows, err := d.sqlDB.Query( - `SELECT a.id, a.actor_id, COALESCE(u.username, ''), a.action, - a.target_type, a.target_id, a.detail, a.created_at - FROM audit_log a - LEFT JOIN users u ON u.id = a.actor_id - ORDER BY a.id DESC - LIMIT ? OFFSET ?`, - limit, offset, - ) + rows, err := d.q.GetAuditLog(dbCtx(), dbgen.GetAuditLogParams{ + Limit: int64(limit), + Offset: int64(offset), + }) if err != nil { return nil, fmt.Errorf("GetAuditLog: %w", err) } - defer rows.Close() //nolint:errcheck - - var entries []AuditEntry - for rows.Next() { - var e AuditEntry - if err := rows.Scan( - &e.ID, &e.ActorID, &e.ActorName, &e.Action, - &e.TargetType, &e.TargetID, &e.Detail, &e.CreatedAt, - ); err != nil { - return nil, fmt.Errorf("GetAuditLog scan: %w", err) - } - entries = append(entries, e) - } - if rows.Err() != nil { - return nil, fmt.Errorf("GetAuditLog rows: %w", rows.Err()) - } - if entries == nil { - entries = []AuditEntry{} + entries := make([]AuditEntry, 0, len(rows)) + for _, r := range rows { + entries = append(entries, AuditEntry{ + ID: r.ID, + ActorID: r.ActorID, + ActorName: r.ActorName, + Action: r.Action, + TargetType: r.TargetType, + TargetID: r.TargetID, + Detail: r.Detail, + CreatedAt: r.CreatedAt, + }) } return entries, nil } @@ -251,8 +208,7 @@ func (d *DB) GetAuditLog(limit, offset int) ([]AuditEntry, error) { // GetSetting returns the value for the given settings key. // Returns an error (wrapping sql.ErrNoRows) when the key does not exist. func (d *DB) GetSetting(key string) (string, error) { - var value string - err := d.sqlDB.QueryRow(`SELECT value FROM settings WHERE key = ?`, key).Scan(&value) + value, err := d.q.GetSetting(dbCtx(), key) if errors.Is(err, sql.ErrNoRows) { return "", fmt.Errorf("GetSetting: key %q: %w", key, ErrNotFound) } @@ -264,12 +220,10 @@ func (d *DB) GetSetting(key string) (string, error) { // SetSetting upserts a setting value for the given key. func (d *DB) SetSetting(key, value string) error { - _, err := d.sqlDB.Exec( - `INSERT INTO settings (key, value) VALUES (?, ?) - ON CONFLICT(key) DO UPDATE SET value = excluded.value`, - key, value, - ) - if err != nil { + if err := d.q.SetSetting(dbCtx(), dbgen.SetSettingParams{ + Key: key, + Value: value, + }); err != nil { return fmt.Errorf("SetSetting: %w", err) } return nil @@ -277,22 +231,13 @@ func (d *DB) SetSetting(key, value string) error { // GetAllSettings returns all settings as a key→value map. func (d *DB) GetAllSettings() (map[string]string, error) { - rows, err := d.sqlDB.Query(`SELECT key, value FROM settings`) + rows, err := d.q.GetAllSettings(dbCtx()) if err != nil { return nil, fmt.Errorf("GetAllSettings: %w", err) } - defer rows.Close() //nolint:errcheck - - result := make(map[string]string) - for rows.Next() { - var k, v string - if err := rows.Scan(&k, &v); err != nil { - return nil, fmt.Errorf("GetAllSettings scan: %w", err) - } - result[k] = v - } - if rows.Err() != nil { - return nil, fmt.Errorf("GetAllSettings rows: %w", rows.Err()) + result := make(map[string]string, len(rows)) + for _, s := range rows { + result[s.Key] = s.Value } return result, nil } @@ -300,12 +245,11 @@ func (d *DB) GetAllSettings() (map[string]string, error) { // CountUsersWithoutTOTP returns the number of non-banned users that do not // currently have a confirmed TOTP secret. func (d *DB) CountUsersWithoutTOTP() (int, error) { - var count int - err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM users WHERE banned = 0 AND totp_secret IS NULL`).Scan(&count) + count, err := d.q.CountUsersWithoutTOTP(dbCtx()) if err != nil { return 0, fmt.Errorf("CountUsersWithoutTOTP: %w", err) } - return count, nil + return int(count), nil } // ─── Backup ─────────────────────────────────────────────────────────────────── diff --git a/Server/db/channel_queries.go b/Server/db/channel_queries.go index eae17e32..1350b227 100644 --- a/Server/db/channel_queries.go +++ b/Server/db/channel_queries.go @@ -5,76 +5,82 @@ import ( "errors" "fmt" "strings" + + "github.com/owncord/server/db/dbgen" ) +// channelFields carries the 13 columns shared by GetChannelRow and +// ListChannelsRow; both generated row types are structurally identical, so a +// single mapper narrows either to the domain Channel model. +type channelFields struct { + ID int64 + Name string + Type string + Category string + Topic string + Position int64 + SlowMode int64 + Archived int64 + CreatedAt string + VoiceMaxUsers int64 + VoiceQuality *string + MixingThreshold *int64 + VoiceMaxVideo int64 +} + +func channelFromFields(f channelFields) Channel { + return Channel{ + ID: f.ID, + Name: f.Name, + Type: f.Type, + Category: f.Category, + Topic: f.Topic, + Position: int(f.Position), + SlowMode: int(f.SlowMode), + Archived: f.Archived != 0, + CreatedAt: f.CreatedAt, + VoiceMaxUsers: int(f.VoiceMaxUsers), + VoiceQuality: f.VoiceQuality, + MixingThreshold: ptrI64toI(f.MixingThreshold), + VoiceMaxVideo: int(f.VoiceMaxVideo), + } +} + // ListChannels returns all channels ordered by position. func (d *DB) ListChannels() ([]Channel, error) { - rows, err := d.sqlDB.Query( - `SELECT id, name, type, COALESCE(category,''), COALESCE(topic,''), - position, slow_mode, archived, created_at, - COALESCE(voice_max_users, 0), - voice_quality, - mixing_threshold, - COALESCE(voice_max_video, 0) - FROM channels ORDER BY position ASC, id ASC`, - ) + rows, err := d.q.ListChannels(dbCtx()) if err != nil { return nil, fmt.Errorf("ListChannels: %w", err) } - defer rows.Close() //nolint:errcheck - - var channels []Channel - for rows.Next() { - ch, scanErr := scanChannel(rows) - if scanErr != nil { - return nil, fmt.Errorf("ListChannels scan: %w", scanErr) - } - channels = append(channels, ch) - } - if rows.Err() != nil { - return nil, fmt.Errorf("ListChannels rows: %w", rows.Err()) - } - if channels == nil { - channels = []Channel{} + channels := make([]Channel, 0, len(rows)) + for _, r := range rows { + channels = append(channels, channelFromFields(channelFields(r))) } return channels, nil } // GetChannel returns the channel with the given id, or nil if not found. func (d *DB) GetChannel(id int64) (*Channel, error) { - row := d.sqlDB.QueryRow( - `SELECT id, name, type, COALESCE(category,''), COALESCE(topic,''), - position, slow_mode, archived, created_at, - COALESCE(voice_max_users, 0), - voice_quality, - mixing_threshold, - COALESCE(voice_max_video, 0) - FROM channels WHERE id = ?`, - id, - ) - ch := &Channel{} - var archived int - err := row.Scan( - &ch.ID, &ch.Name, &ch.Type, &ch.Category, &ch.Topic, - &ch.Position, &ch.SlowMode, &archived, &ch.CreatedAt, - &ch.VoiceMaxUsers, &ch.VoiceQuality, &ch.MixingThreshold, &ch.VoiceMaxVideo, - ) + r, err := d.q.GetChannel(dbCtx(), id) if errors.Is(err, sql.ErrNoRows) { return nil, nil } if err != nil { return nil, fmt.Errorf("GetChannel: %w", err) } - ch.Archived = archived != 0 - return ch, nil + ch := channelFromFields(channelFields(r)) + return &ch, nil } // CreateChannel inserts a new channel and returns the assigned ID. func (d *DB) CreateChannel(name, chanType, category, topic string, position int) (int64, error) { - res, err := d.sqlDB.Exec( - `INSERT INTO channels (name, type, category, topic, position) VALUES (?, ?, ?, ?, ?)`, - name, chanType, nullableString(category), nullableString(topic), position, - ) + res, err := d.q.CreateChannel(dbCtx(), dbgen.CreateChannelParams{ + Name: name, + Type: chanType, + Category: strToNullPtr(category), + Topic: strToNullPtr(topic), + Position: int64(position), + }) if err != nil { return 0, fmt.Errorf("CreateChannel: %w", err) } @@ -83,11 +89,12 @@ func (d *DB) CreateChannel(name, chanType, category, topic string, position int) // UpdateChannel modifies name, topic, and slow_mode for the given channel. func (d *DB) UpdateChannel(id int64, name, topic string, slowMode int) error { - _, err := d.sqlDB.Exec( - `UPDATE channels SET name = ?, topic = ?, slow_mode = ? WHERE id = ?`, - name, nullableString(topic), slowMode, id, - ) - if err != nil { + if err := d.q.UpdateChannel(dbCtx(), dbgen.UpdateChannelParams{ + Name: name, + Topic: strToNullPtr(topic), + SlowMode: int64(slowMode), + ID: id, + }); err != nil { return fmt.Errorf("UpdateChannel: %w", err) } return nil @@ -95,11 +102,10 @@ func (d *DB) UpdateChannel(id int64, name, topic string, slowMode int) error { // SetChannelSlowMode updates only the slow_mode field for the given channel. func (d *DB) SetChannelSlowMode(id int64, slowMode int) error { - _, err := d.sqlDB.Exec( - `UPDATE channels SET slow_mode = ? WHERE id = ?`, - slowMode, id, - ) - if err != nil { + if err := d.q.SetChannelSlowMode(dbCtx(), dbgen.SetChannelSlowModeParams{ + SlowMode: int64(slowMode), + ID: id, + }); err != nil { return fmt.Errorf("SetChannelSlowMode: %w", err) } return nil @@ -107,8 +113,10 @@ func (d *DB) SetChannelSlowMode(id int64, slowMode int) error { // SetChannelVoiceMaxUsers updates the voice_max_users field for the given channel. func (d *DB) SetChannelVoiceMaxUsers(id int64, maxUsers int) error { - _, err := d.sqlDB.Exec(`UPDATE channels SET voice_max_users = ? WHERE id = ?`, maxUsers, id) - if err != nil { + if err := d.q.SetChannelVoiceMaxUsers(dbCtx(), dbgen.SetChannelVoiceMaxUsersParams{ + VoiceMaxUsers: int64(maxUsers), + ID: id, + }); err != nil { return fmt.Errorf("SetChannelVoiceMaxUsers: %w", err) } return nil @@ -116,8 +124,7 @@ func (d *DB) SetChannelVoiceMaxUsers(id int64, maxUsers int) error { // DeleteChannel removes the channel row (cascades to messages, overrides, etc.). func (d *DB) DeleteChannel(id int64) error { - _, err := d.sqlDB.Exec(`DELETE FROM channels WHERE id = ?`, id) - if err != nil { + if err := d.q.DeleteChannel(dbCtx(), id); err != nil { return fmt.Errorf("DeleteChannel: %w", err) } return nil @@ -126,18 +133,17 @@ func (d *DB) DeleteChannel(id int64) error { // GetChannelPermissions returns the allow/deny override bits for a role on a // channel. Returns (0, 0, nil) when no override exists. func (d *DB) GetChannelPermissions(channelID, roleID int64) (allow, deny int64, err error) { - row := d.sqlDB.QueryRow( - `SELECT allow, deny FROM channel_overrides WHERE channel_id = ? AND role_id = ?`, - channelID, roleID, - ) - scanErr := row.Scan(&allow, &deny) + r, scanErr := d.q.GetChannelPermission(dbCtx(), dbgen.GetChannelPermissionParams{ + ChannelID: channelID, + RoleID: roleID, + }) if errors.Is(scanErr, sql.ErrNoRows) { return 0, 0, nil } if scanErr != nil { return 0, 0, fmt.Errorf("GetChannelPermissions: %w", scanErr) } - return allow, deny, nil + return r.Allow, r.Deny, nil } // ChannelOverride holds the allow/deny permission bits for a single channel. @@ -150,26 +156,13 @@ type ChannelOverride struct { // a role in a single query, keyed by channel ID. Eliminates N+1 queries when // filtering channels by permission. func (d *DB) GetAllChannelPermissionsForRole(roleID int64) (map[int64]ChannelOverride, error) { - rows, err := d.sqlDB.Query( - `SELECT channel_id, allow, deny FROM channel_overrides WHERE role_id = ?`, - roleID, - ) + rows, err := d.q.GetRoleChannelPermissions(dbCtx(), roleID) if err != nil { return nil, fmt.Errorf("GetAllChannelPermissionsForRole: %w", err) } - defer rows.Close() //nolint:errcheck - - result := make(map[int64]ChannelOverride) - for rows.Next() { - var chID int64 - var o ChannelOverride - if scanErr := rows.Scan(&chID, &o.Allow, &o.Deny); scanErr != nil { - return nil, fmt.Errorf("GetAllChannelPermissionsForRole scan: %w", scanErr) - } - result[chID] = o - } - if rows.Err() != nil { - return nil, fmt.Errorf("GetAllChannelPermissionsForRole rows: %w", rows.Err()) + result := make(map[int64]ChannelOverride, len(rows)) + for _, r := range rows { + result[r.ChannelID] = ChannelOverride{Allow: r.Allow, Deny: r.Deny} } return result, nil } @@ -177,14 +170,12 @@ func (d *DB) GetAllChannelPermissionsForRole(roleID int64) (map[int64]ChannelOve // UpsertChannelOverride inserts or updates the allow/deny permission override // for a role on a channel. func (d *DB) UpsertChannelOverride(channelID, roleID, allow, deny int64) error { - _, err := d.sqlDB.Exec( - `INSERT INTO channel_overrides (channel_id, role_id, allow, deny) - VALUES (?, ?, ?, ?) - ON CONFLICT(channel_id, role_id) - DO UPDATE SET allow = excluded.allow, deny = excluded.deny`, - channelID, roleID, allow, deny, - ) - if err != nil { + if err := d.q.UpsertChannelPermission(dbCtx(), dbgen.UpsertChannelPermissionParams{ + ChannelID: channelID, + RoleID: roleID, + Allow: allow, + Deny: deny, + }); err != nil { return fmt.Errorf("UpsertChannelOverride: %w", err) } return nil @@ -193,11 +184,10 @@ func (d *DB) UpsertChannelOverride(channelID, roleID, allow, deny int64) error { // DeleteChannelOverride removes the permission override for a role on a // channel. Deleting a non-existent override is a no-op. func (d *DB) DeleteChannelOverride(channelID, roleID int64) error { - _, err := d.sqlDB.Exec( - `DELETE FROM channel_overrides WHERE channel_id = ? AND role_id = ?`, - channelID, roleID, - ) - if err != nil { + if err := d.q.DeleteChannelPermission(dbCtx(), dbgen.DeleteChannelPermissionParams{ + ChannelID: channelID, + RoleID: roleID, + }); err != nil { return fmt.Errorf("DeleteChannelOverride: %w", err) } return nil @@ -251,34 +241,6 @@ func (d *DB) ListChannelRoleOverrides(channelID int64) ([]ChannelRoleOverride, e // ─── helpers ────────────────────────────────────────────────────────────────── -// scanChannel scans a single channel row from *sql.Rows. -// The query must select the 13 columns: id, name, type, category, topic, -// position, slow_mode, archived, created_at, voice_max_users, -// voice_quality, mixing_threshold, voice_max_video. -func scanChannel(rows *sql.Rows) (Channel, error) { - var ch Channel - var archived int - err := rows.Scan( - &ch.ID, &ch.Name, &ch.Type, &ch.Category, &ch.Topic, - &ch.Position, &ch.SlowMode, &archived, &ch.CreatedAt, - &ch.VoiceMaxUsers, &ch.VoiceQuality, &ch.MixingThreshold, &ch.VoiceMaxVideo, - ) - if err != nil { - return Channel{}, err - } - ch.Archived = archived != 0 - return ch, nil -} - -// nullableString returns nil when s is empty, otherwise a pointer to s. -// Used so empty strings are stored as NULL in optional TEXT columns. -func nullableString(s string) any { - if s == "" { - return nil - } - return s -} - // GetChannelTypes returns a map of channel ID → type string for the given IDs // in a single query, avoiding N+1 lookups. func (d *DB) GetChannelTypes(ids []int64) (map[int64]string, error) { diff --git a/Server/db/dm_queries.go b/Server/db/dm_queries.go index f2b5ae18..f1700f7d 100644 --- a/Server/db/dm_queries.go +++ b/Server/db/dm_queries.go @@ -5,6 +5,8 @@ import ( "database/sql" "errors" "fmt" + + "github.com/owncord/server/db/dbgen" ) // ─── DM Models ────────────────────────────────────────────────────────────── @@ -202,11 +204,10 @@ func (d *DB) GetUserDMChannels(userID int64) ([]DMChannelInfo, error) { // OpenDM adds a DM channel to a user's open list (idempotent). func (d *DB) OpenDM(userID, channelID int64) error { - _, err := d.sqlDB.Exec( - `INSERT OR IGNORE INTO dm_open_state (user_id, channel_id) VALUES (?, ?)`, - userID, channelID, - ) - if err != nil { + if err := d.q.OpenDM(dbCtx(), dbgen.OpenDMParams{ + UserID: userID, + ChannelID: channelID, + }); err != nil { return fmt.Errorf("OpenDM: %w", err) } return nil @@ -214,11 +215,10 @@ func (d *DB) OpenDM(userID, channelID int64) error { // CloseDM removes a DM channel from a user's open list. func (d *DB) CloseDM(userID, channelID int64) error { - _, err := d.sqlDB.Exec( - `DELETE FROM dm_open_state WHERE user_id = ? AND channel_id = ?`, - userID, channelID, - ) - if err != nil { + if err := d.q.CloseDM(dbCtx(), dbgen.CloseDMParams{ + UserID: userID, + ChannelID: channelID, + }); err != nil { return fmt.Errorf("CloseDM: %w", err) } return nil @@ -228,11 +228,10 @@ func (d *DB) CloseDM(userID, channelID int64) error { // IsDMParticipant checks if a user is a participant in a DM channel. func (d *DB) IsDMParticipant(userID, channelID int64) (bool, error) { - var id int64 - err := d.sqlDB.QueryRow( - `SELECT user_id FROM dm_participants WHERE user_id = ? AND channel_id = ?`, - userID, channelID, - ).Scan(&id) + _, err := d.q.IsDMParticipant(dbCtx(), dbgen.IsDMParticipantParams{ + UserID: userID, + ChannelID: channelID, + }) if errors.Is(err, sql.ErrNoRows) { return false, nil } @@ -244,26 +243,10 @@ func (d *DB) IsDMParticipant(userID, channelID int64) (bool, error) { // GetDMParticipantIDs returns all participant user IDs for a DM channel. func (d *DB) GetDMParticipantIDs(channelID int64) ([]int64, error) { - rows, err := d.sqlDB.Query( - `SELECT user_id FROM dm_participants WHERE channel_id = ?`, - channelID, - ) + ids, err := d.q.GetDMParticipantIDs(dbCtx(), channelID) if err != nil { return nil, fmt.Errorf("GetDMParticipantIDs: %w", err) } - defer rows.Close() //nolint:errcheck - - var ids []int64 - for rows.Next() { - var id int64 - if scanErr := rows.Scan(&id); scanErr != nil { - return nil, fmt.Errorf("GetDMParticipantIDs scan: %w", scanErr) - } - ids = append(ids, id) - } - if rows.Err() != nil { - return nil, fmt.Errorf("GetDMParticipantIDs rows: %w", rows.Err()) - } return ids, nil } diff --git a/Server/db/mappers.go b/Server/db/mappers.go index 5a3edb55..cff9329a 100644 --- a/Server/db/mappers.go +++ b/Server/db/mappers.go @@ -38,6 +38,24 @@ func ptrItoI64(p *int) *int64 { return &v } +// b2i64 converts a bool to sqlc's int64 representation of a SQLite boolean. +func b2i64(b bool) int64 { + if b { + return 1 + } + return 0 +} + +// strToNullPtr returns nil for an empty string, else a pointer to it — so +// empty strings are written as NULL in optional TEXT columns (matches the +// former nullableString helper). +func strToNullPtr(s string) *string { + if s == "" { + return nil + } + return &s +} + // userFromGen maps a generated user row to the domain User model. func userFromGen(u dbgen.User) *User { return &User{ diff --git a/Server/db/voice_queries.go b/Server/db/voice_queries.go index b4a94f2b..7054cec8 100644 --- a/Server/db/voice_queries.go +++ b/Server/db/voice_queries.go @@ -6,6 +6,8 @@ import ( "fmt" "sync/atomic" "time" + + "github.com/owncord/server/db/dbgen" ) // ErrChannelFull is returned when a voice channel is at capacity. @@ -26,21 +28,11 @@ func newVoiceJoinToken() string { // target one specific voice session even if the user later rejoins the same // channel. func (d *DB) JoinVoiceChannel(userID, channelID int64) error { - joinToken := newVoiceJoinToken() - _, err := d.sqlDB.Exec( - `INSERT INTO voice_states (user_id, channel_id, muted, deafened, speaking, camera, screenshare, joined_at) - VALUES (?, ?, 0, 0, 0, 0, 0, ?) - ON CONFLICT(user_id) DO UPDATE SET - channel_id = excluded.channel_id, - muted = 0, - deafened = 0, - speaking = 0, - camera = 0, - screenshare = 0, - joined_at = excluded.joined_at`, - userID, channelID, joinToken, - ) - if err != nil { + if err := d.q.JoinVoiceChannel(dbCtx(), dbgen.JoinVoiceChannelParams{ + UserID: userID, + ChannelID: channelID, + JoinedAt: newVoiceJoinToken(), + }); err != nil { return fmt.Errorf("JoinVoiceChannel: %w", err) } return nil @@ -51,21 +43,13 @@ func (d *DB) JoinVoiceChannel(userID, channelID int64) error { // the channel is at capacity. This prevents the TOCTOU race where two // concurrent joins both observe capacity and both succeed. func (d *DB) JoinVoiceChannelIfCapacity(userID, channelID int64, maxUsers int) error { - joinToken := newVoiceJoinToken() - res, err := d.sqlDB.Exec( - `INSERT INTO voice_states (user_id, channel_id, muted, deafened, speaking, camera, screenshare, joined_at) - SELECT ?, ?, 0, 0, 0, 0, 0, ? - WHERE (SELECT COUNT(*) FROM voice_states WHERE channel_id = ?) < ? - ON CONFLICT(user_id) DO UPDATE SET - channel_id = excluded.channel_id, - muted = 0, - deafened = 0, - speaking = 0, - camera = 0, - screenshare = 0, - joined_at = excluded.joined_at`, - userID, channelID, joinToken, channelID, maxUsers, - ) + res, err := d.q.JoinVoiceChannelIfCapacity(dbCtx(), dbgen.JoinVoiceChannelIfCapacityParams{ + UserID: userID, + ChannelID: channelID, + JoinedAt: newVoiceJoinToken(), + ChannelID_2: channelID, + ChannelID_3: int64(maxUsers), + }) if err != nil { return fmt.Errorf("JoinVoiceChannelIfCapacity: %w", err) } @@ -79,8 +63,7 @@ func (d *DB) JoinVoiceChannelIfCapacity(userID, channelID int64, maxUsers int) e // LeaveVoiceChannel removes the user's voice state entirely. // It is safe to call when the user is not in any voice channel. func (d *DB) LeaveVoiceChannel(userID int64) error { - _, err := d.sqlDB.Exec(`DELETE FROM voice_states WHERE user_id = ?`, userID) - if err != nil { + if err := d.q.LeaveVoiceChannel(dbCtx(), userID); err != nil { return fmt.Errorf("LeaveVoiceChannel: %w", err) } return nil @@ -90,10 +73,11 @@ func (d *DB) LeaveVoiceChannel(userID int64) error { // still points at expectedChannelID and matches the expected join token. // Returns true if a row was deleted. func (d *DB) LeaveVoiceChannelIfMatch(userID, expectedChannelID int64, expectedJoinedAt string) (bool, error) { - result, err := d.sqlDB.Exec( - `DELETE FROM voice_states WHERE user_id = ? AND channel_id = ? AND joined_at = ?`, - userID, expectedChannelID, expectedJoinedAt, - ) + result, err := d.q.LeaveVoiceChannelIfMatch(dbCtx(), dbgen.LeaveVoiceChannelIfMatchParams{ + UserID: userID, + ChannelID: expectedChannelID, + JoinedAt: expectedJoinedAt, + }) if err != nil { return false, fmt.Errorf("LeaveVoiceChannelIfMatch: %w", err) } @@ -104,49 +88,47 @@ func (d *DB) LeaveVoiceChannelIfMatch(userID, expectedChannelID int64, expectedJ // GetVoiceState returns the current voice state for the given user, // or nil if the user is not in any voice channel. func (d *DB) GetVoiceState(userID int64) (*VoiceState, error) { - row := d.sqlDB.QueryRow( - `SELECT vs.user_id, vs.channel_id, u.username, - vs.muted, vs.deafened, vs.speaking, - vs.camera, vs.screenshare, vs.joined_at - FROM voice_states vs - JOIN users u ON u.id = vs.user_id - WHERE vs.user_id = ?`, - userID, - ) - return scanVoiceState(row) + r, err := d.q.GetUserVoiceState(dbCtx(), userID) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("GetVoiceState: %w", err) + } + vs := VoiceState{ + UserID: r.UserID, + ChannelID: r.ChannelID, + Username: r.Username, + Muted: r.Muted != 0, + Deafened: r.Deafened != 0, + Speaking: r.Speaking != 0, + Camera: r.Camera != 0, + Screenshare: r.Screenshare != 0, + JoinedAt: r.JoinedAt, + } + return &vs, nil } // GetChannelVoiceStates returns all voice states for users currently in the // given voice channel. func (d *DB) GetChannelVoiceStates(channelID int64) ([]VoiceState, error) { - rows, err := d.sqlDB.Query( - `SELECT vs.user_id, vs.channel_id, u.username, - vs.muted, vs.deafened, vs.speaking, - vs.camera, vs.screenshare, vs.joined_at - FROM voice_states vs - JOIN users u ON u.id = vs.user_id - WHERE vs.channel_id = ? - ORDER BY vs.joined_at ASC`, - channelID, - ) + rows, err := d.q.GetChannelVoiceStates(dbCtx(), channelID) if err != nil { return nil, fmt.Errorf("GetChannelVoiceStates: %w", err) } - defer rows.Close() //nolint:errcheck - - var states []VoiceState - for rows.Next() { - vs, scanErr := scanVoiceStateRow(rows) - if scanErr != nil { - return nil, fmt.Errorf("GetChannelVoiceStates scan: %w", scanErr) - } - states = append(states, vs) - } - if rows.Err() != nil { - return nil, fmt.Errorf("GetChannelVoiceStates rows: %w", rows.Err()) - } - if states == nil { - states = []VoiceState{} + states := make([]VoiceState, 0, len(rows)) + for _, r := range rows { + states = append(states, VoiceState{ + UserID: r.UserID, + ChannelID: r.ChannelID, + Username: r.Username, + Muted: r.Muted != 0, + Deafened: r.Deafened != 0, + Speaking: r.Speaking != 0, + Camera: r.Camera != 0, + Screenshare: r.Screenshare != 0, + JoinedAt: r.JoinedAt, + }) } return states, nil } @@ -154,32 +136,23 @@ func (d *DB) GetChannelVoiceStates(channelID int64) ([]VoiceState, error) { // GetAllVoiceStates returns voice states across all voice channels in a single // query. Used at startup to build the ready payload without N+1 per-channel queries. func (d *DB) GetAllVoiceStates() ([]VoiceState, error) { - rows, err := d.sqlDB.Query( - `SELECT vs.user_id, vs.channel_id, u.username, - vs.muted, vs.deafened, vs.speaking, - vs.camera, vs.screenshare, vs.joined_at - FROM voice_states vs - JOIN users u ON u.id = vs.user_id - ORDER BY vs.channel_id, vs.joined_at ASC`, - ) + rows, err := d.q.GetAllVoiceStates(dbCtx()) if err != nil { return nil, fmt.Errorf("GetAllVoiceStates: %w", err) } - defer rows.Close() //nolint:errcheck - - var states []VoiceState - for rows.Next() { - vs, scanErr := scanVoiceStateRow(rows) - if scanErr != nil { - return nil, fmt.Errorf("GetAllVoiceStates scan: %w", scanErr) - } - states = append(states, vs) - } - if rows.Err() != nil { - return nil, fmt.Errorf("GetAllVoiceStates rows: %w", rows.Err()) - } - if states == nil { - states = []VoiceState{} + states := make([]VoiceState, 0, len(rows)) + for _, r := range rows { + states = append(states, VoiceState{ + UserID: r.UserID, + ChannelID: r.ChannelID, + Username: r.Username, + Muted: r.Muted != 0, + Deafened: r.Deafened != 0, + Speaking: r.Speaking != 0, + Camera: r.Camera != 0, + Screenshare: r.Screenshare != 0, + JoinedAt: r.JoinedAt, + }) } return states, nil } @@ -187,12 +160,10 @@ func (d *DB) GetAllVoiceStates() ([]VoiceState, error) { // UpdateVoiceMute sets the muted field for the given user's voice state. // It is safe to call when the user is not in any channel (no-op). func (d *DB) UpdateVoiceMute(userID int64, muted bool) error { - muteInt := boolToInt(muted) - _, err := d.sqlDB.Exec( - `UPDATE voice_states SET muted = ? WHERE user_id = ?`, - muteInt, userID, - ) - if err != nil { + if err := d.q.UpdateVoiceMute(dbCtx(), dbgen.UpdateVoiceMuteParams{ + Muted: b2i64(muted), + UserID: userID, + }); err != nil { return fmt.Errorf("UpdateVoiceMute: %w", err) } return nil @@ -201,12 +172,10 @@ func (d *DB) UpdateVoiceMute(userID int64, muted bool) error { // UpdateVoiceDeafen sets the deafened field for the given user's voice state. // It is safe to call when the user is not in any channel (no-op). func (d *DB) UpdateVoiceDeafen(userID int64, deafened bool) error { - deafenInt := boolToInt(deafened) - _, err := d.sqlDB.Exec( - `UPDATE voice_states SET deafened = ? WHERE user_id = ?`, - deafenInt, userID, - ) - if err != nil { + if err := d.q.UpdateVoiceDeafen(dbCtx(), dbgen.UpdateVoiceDeafenParams{ + Deafened: b2i64(deafened), + UserID: userID, + }); err != nil { return fmt.Errorf("UpdateVoiceDeafen: %w", err) } return nil @@ -215,8 +184,7 @@ func (d *DB) UpdateVoiceDeafen(userID int64, deafened bool) error { // ClearVoiceState removes a user's voice state on disconnect. // Equivalent to LeaveVoiceChannel but named to clarify the disconnect use case. func (d *DB) ClearVoiceState(userID int64) error { - _, err := d.sqlDB.Exec(`DELETE FROM voice_states WHERE user_id = ?`, userID) - if err != nil { + if err := d.q.ClearVoiceState(dbCtx(), userID); err != nil { return fmt.Errorf("ClearVoiceState: %w", err) } return nil @@ -225,8 +193,7 @@ func (d *DB) ClearVoiceState(userID int64) error { // ClearAllVoiceStates removes all voice state rows. Called on server startup // to clear stale state from a previous run. func (d *DB) ClearAllVoiceStates() error { - _, err := d.sqlDB.Exec(`DELETE FROM voice_states`) - if err != nil { + if err := d.q.ClearAllVoiceStates(dbCtx()); err != nil { return fmt.Errorf("ClearAllVoiceStates: %w", err) } return nil @@ -236,24 +203,19 @@ func (d *DB) ClearAllVoiceStates() error { // given voice channel. Uses the DB as source of truth (race-free via SQLite // serialization) rather than querying LiveKit. func (d *DB) CountActiveCameras(channelID int64) (int, error) { - var count int - err := d.sqlDB.QueryRow( - `SELECT COUNT(*) FROM voice_states WHERE channel_id = ? AND camera = 1`, - channelID, - ).Scan(&count) + count, err := d.q.CountActiveCameras(dbCtx(), channelID) if err != nil { return 0, fmt.Errorf("CountActiveCameras: %w", err) } - return count, nil + return int(count), nil } // UpdateVoiceCamera sets the camera field for the given user's voice state. func (d *DB) UpdateVoiceCamera(userID int64, camera bool) error { - _, err := d.sqlDB.Exec( - `UPDATE voice_states SET camera = ? WHERE user_id = ?`, - boolToInt(camera), userID, - ) - if err != nil { + if err := d.q.UpdateVoiceCamera(dbCtx(), dbgen.UpdateVoiceCameraParams{ + Camera: b2i64(camera), + UserID: userID, + }); err != nil { return fmt.Errorf("UpdateVoiceCamera: %w", err) } return nil @@ -263,12 +225,12 @@ func (d *DB) UpdateVoiceCamera(userID int64, camera bool) error { // channel has not yet reached maxVideo active cameras. Returns true if the // camera was enabled, false if the limit was already reached. func (d *DB) EnableCameraIfUnderLimit(userID, channelID int64, maxVideo int) (bool, error) { - res, err := d.sqlDB.Exec( - `UPDATE voice_states SET camera = 1 - WHERE user_id = ? AND channel_id = ? - AND (SELECT COUNT(*) FROM voice_states WHERE channel_id = ? AND camera = 1) < ?`, - userID, channelID, channelID, maxVideo, - ) + res, err := d.q.EnableCameraIfUnderLimit(dbCtx(), dbgen.EnableCameraIfUnderLimitParams{ + UserID: userID, + ChannelID: channelID, + ChannelID_2: channelID, + ChannelID_3: int64(maxVideo), + }) if err != nil { return false, fmt.Errorf("EnableCameraIfUnderLimit: %w", err) } @@ -281,11 +243,10 @@ func (d *DB) EnableCameraIfUnderLimit(userID, channelID int64, maxVideo int) (bo // UpdateVoiceScreenshare sets the screenshare field for the given user's voice state. func (d *DB) UpdateVoiceScreenshare(userID int64, screenshare bool) error { - _, err := d.sqlDB.Exec( - `UPDATE voice_states SET screenshare = ? WHERE user_id = ?`, - boolToInt(screenshare), userID, - ) - if err != nil { + if err := d.q.UpdateVoiceScreenshare(dbCtx(), dbgen.UpdateVoiceScreenshareParams{ + Screenshare: b2i64(screenshare), + UserID: userID, + }); err != nil { return fmt.Errorf("UpdateVoiceScreenshare: %w", err) } return nil @@ -304,57 +265,3 @@ func (d *DB) CountChannelVoiceUsers(channelID int64) (int, error) { } return count, nil } - -// ─── helpers ────────────────────────────────────────────────────────────────── - -// scanVoiceState scans a single *sql.Row into a VoiceState. -// Returns nil (not an error) when the row is not found. -func scanVoiceState(row *sql.Row) (*VoiceState, error) { - vs := &VoiceState{} - var muted, deafened, speaking, camera, screenshare int - err := row.Scan( - &vs.UserID, &vs.ChannelID, &vs.Username, - &muted, &deafened, &speaking, - &camera, &screenshare, &vs.JoinedAt, - ) - if errors.Is(err, sql.ErrNoRows) { - return nil, nil - } - if err != nil { - return nil, fmt.Errorf("scanVoiceState: %w", err) - } - vs.Muted = muted != 0 - vs.Deafened = deafened != 0 - vs.Speaking = speaking != 0 - vs.Camera = camera != 0 - vs.Screenshare = screenshare != 0 - return vs, nil -} - -// scanVoiceStateRow scans a single row from *sql.Rows into a VoiceState. -func scanVoiceStateRow(rows *sql.Rows) (VoiceState, error) { - vs := VoiceState{} - var muted, deafened, speaking, camera, screenshare int - err := rows.Scan( - &vs.UserID, &vs.ChannelID, &vs.Username, - &muted, &deafened, &speaking, - &camera, &screenshare, &vs.JoinedAt, - ) - if err != nil { - return vs, fmt.Errorf("scanVoiceStateRow: %w", err) - } - vs.Muted = muted != 0 - vs.Deafened = deafened != 0 - vs.Speaking = speaking != 0 - vs.Camera = camera != 0 - vs.Screenshare = screenshare != 0 - return vs, nil -} - -// boolToInt converts a bool to 0/1 for SQLite storage. -func boolToInt(b bool) int { - if b { - return 1 - } - return 0 -}