From a7df9c2b3ce9eaf16ad5b6334953bbf56fba4823 Mon Sep 17 00:00:00 2001 From: jevb Date: Thu, 19 Mar 2026 04:19:03 +0100 Subject: [PATCH] fix: resolve 5 remaining medium/low issues from third-pass go-review - NEW-1: Add rows.Err() check in ListMembers to catch cursor errors - NEW-2: Add minVal parameter to queryInt so offset=0 is not rejected - NEW-3: Fix copyFile double-close by removing defer, using explicit close on both success and error paths - NEW-4: Add GetAllChannelPermissionsForRole batch query, eliminating N+1 GetChannelPermissions calls in channel list and search handlers - NEW-5: Cap fetchBody with io.LimitReader(1 MiB) to prevent memory exhaustion from malformed release assets --- Server/admin/api.go | 6 +++-- Server/admin/handlers_backup.go | 2 +- Server/admin/handlers_channels.go | 4 ++-- Server/admin/handlers_users.go | 4 ++-- Server/api/channel_handler.go | 38 ++++++++++++++++++++++++++++--- Server/db/auth_queries.go | 3 +++ Server/db/channel_queries.go | 34 +++++++++++++++++++++++++++ Server/updater/updater.go | 4 +++- 8 files changed, 84 insertions(+), 11 deletions(-) diff --git a/Server/admin/api.go b/Server/admin/api.go index b5603984..70dfa0e9 100644 --- a/Server/admin/api.go +++ b/Server/admin/api.go @@ -253,13 +253,15 @@ func pathInt64(r *http.Request, param string) (int64, error) { return strconv.ParseInt(raw, 10, 64) } -func queryInt(r *http.Request, key string, defaultVal int) int { +// queryInt parses an integer query parameter with a minimum and maximum bound. +// Use minVal=1 for limit parameters, minVal=0 for offset parameters. +func queryInt(r *http.Request, key string, defaultVal, minVal int) int { raw := r.URL.Query().Get(key) if raw == "" { return defaultVal } n, err := strconv.Atoi(raw) - if err != nil || n < 1 { + if err != nil || n < minVal { return defaultVal } // Cap to prevent unbounded result sets exhausting memory. diff --git a/Server/admin/handlers_backup.go b/Server/admin/handlers_backup.go index 4f598cb0..234391b6 100644 --- a/Server/admin/handlers_backup.go +++ b/Server/admin/handlers_backup.go @@ -172,9 +172,9 @@ func copyFile(src, dst string) error { if err != nil { return fmt.Errorf("create destination: %w", err) } - defer out.Close() //nolint:errcheck if _, err := io.Copy(out, in); err != nil { + _ = out.Close() return fmt.Errorf("copy: %w", err) } return out.Close() diff --git a/Server/admin/handlers_channels.go b/Server/admin/handlers_channels.go index 92ef1a2b..9c679386 100644 --- a/Server/admin/handlers_channels.go +++ b/Server/admin/handlers_channels.go @@ -209,8 +209,8 @@ func handleDeleteChannel(database *db.DB, hub HubBroadcaster) http.HandlerFunc { func handleGetAuditLog(database *db.DB) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - limit := queryInt(r, "limit", 50) - offset := queryInt(r, "offset", 0) + limit := queryInt(r, "limit", 50, 1) + offset := queryInt(r, "offset", 0, 0) entries, err := database.GetAuditLog(limit, offset) if err != nil { diff --git a/Server/admin/handlers_users.go b/Server/admin/handlers_users.go index 9b3b27e9..299f6afe 100644 --- a/Server/admin/handlers_users.go +++ b/Server/admin/handlers_users.go @@ -27,8 +27,8 @@ func handleGetStats(database *db.DB, hub HubBroadcaster) http.HandlerFunc { func handleListUsers(database *db.DB) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - limit := queryInt(r, "limit", 50) - offset := queryInt(r, "offset", 0) + limit := queryInt(r, "limit", 50, 1) + offset := queryInt(r, "offset", 0, 0) users, err := database.ListAllUsers(limit, offset) if err != nil { diff --git a/Server/api/channel_handler.go b/Server/api/channel_handler.go index e599527d..c34852e6 100644 --- a/Server/api/channel_handler.go +++ b/Server/api/channel_handler.go @@ -43,6 +43,20 @@ func hasChannelPermREST(database *db.DB, role *db.Role, channelID, perm int64) b return effective&perm == perm } +// hasChannelPermBatch checks permission using a pre-fetched overrides map, +// eliminating N+1 queries when filtering multiple channels. +func hasChannelPermBatch(role *db.Role, overrides map[int64]db.ChannelOverride, channelID, perm int64) bool { + if role == nil { + return false + } + if permissions.HasAdmin(role.Permissions) { + return true + } + o := overrides[channelID] // zero-value (0,0) when no override exists + effective := permissions.EffectivePerms(role.Permissions, o.Allow, o.Deny) + return effective&perm == perm +} + // handleListChannels returns all channels the authenticated user can see. func handleListChannels(database *db.DB) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { @@ -58,10 +72,20 @@ func handleListChannels(database *db.DB) http.HandlerFunc { return } + // Batch-fetch all channel permission overrides for this role in one query. + overrides := map[int64]db.ChannelOverride{} + if role != nil && !permissions.HasAdmin(role.Permissions) { + var oErr error + overrides, oErr = database.GetAllChannelPermissionsForRole(role.ID) + if oErr != nil { + slog.Error("handleListChannels GetAllChannelPermissionsForRole", "err", oErr) + } + } + // Filter channels by READ_MESSAGES permission. var visible []db.Channel for _, ch := range channels { - if hasChannelPermREST(database, role, ch.ID, permissions.ReadMessages) { + if hasChannelPermBatch(role, overrides, ch.ID, permissions.ReadMessages) { visible = append(visible, ch) } } @@ -221,11 +245,19 @@ func handleSearch(database *db.DB) http.HandlerFunc { return } - // Post-filter results by READ_MESSAGES permission on each channel. + // Batch-fetch overrides and post-filter results by READ_MESSAGES. role, _ := r.Context().Value(RoleKey).(*db.Role) + overrides := map[int64]db.ChannelOverride{} + if role != nil && !permissions.HasAdmin(role.Permissions) { + var oErr error + overrides, oErr = database.GetAllChannelPermissionsForRole(role.ID) + if oErr != nil { + slog.Error("handleSearch GetAllChannelPermissionsForRole", "err", oErr) + } + } var filtered []db.MessageSearchResult for _, res := range results { - if hasChannelPermREST(database, role, res.ChannelID, permissions.ReadMessages) { + if hasChannelPermBatch(role, overrides, res.ChannelID, permissions.ReadMessages) { filtered = append(filtered, res) } } diff --git a/Server/db/auth_queries.go b/Server/db/auth_queries.go index e105a88f..76109dc7 100644 --- a/Server/db/auth_queries.go +++ b/Server/db/auth_queries.go @@ -356,6 +356,9 @@ func (d *DB) ListMembers() ([]MemberSummary, error) { } members = append(members, m) } + if rows.Err() != nil { + return nil, fmt.Errorf("ListMembers rows: %w", rows.Err()) + } if members == nil { members = []MemberSummary{} } diff --git a/Server/db/channel_queries.go b/Server/db/channel_queries.go index b2f1b443..99a31cda 100644 --- a/Server/db/channel_queries.go +++ b/Server/db/channel_queries.go @@ -135,6 +135,40 @@ func (d *DB) GetChannelPermissions(channelID, roleID int64) (allow, deny int64, return allow, deny, nil } +// ChannelOverride holds the allow/deny permission bits for a single channel. +type ChannelOverride struct { + Allow int64 + Deny int64 +} + +// GetAllChannelPermissionsForRole returns all channel permission overrides for +// 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, + ) + 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()) + } + return result, nil +} + // ─── helpers ────────────────────────────────────────────────────────────────── // scanChannel scans a single channel row from *sql.Rows. diff --git a/Server/updater/updater.go b/Server/updater/updater.go index 7d724ef9..34fd8f81 100644 --- a/Server/updater/updater.go +++ b/Server/updater/updater.go @@ -296,7 +296,9 @@ func (u *Updater) fetchBody(ctx context.Context, url string) ([]byte, error) { return nil, fmt.Errorf("HTTP %d fetching %s", resp.StatusCode, url) } - return io.ReadAll(resp.Body) + // Cap reads at 1 MiB — checksum and signature files are tiny text; + // this prevents a malicious or corrupted release asset from exhausting memory. + return io.ReadAll(io.LimitReader(resp.Body, 1<<20)) } // FindClientAssets scans the cached release assets for the Tauri NSIS