mirror of
https://github.com/J3vb/OwnCord.git
synced 2026-09-03 03:50:00 +03:00
refactor(server): work off the complexity backlog — 62 findings to 0 (#1389)
* refactor(ws): split handleVoiceJoin into cohesive join-stage helpers handleVoiceJoin was 130 statements / cyclomatic 59 / nestif 11, breaking all three complexity budgets at once. Split along the stage boundaries the doc comment already described: precheck, leave-current, persist, restore moderator flags, grant token, complete. The publish-permission derivation becomes its own helper because it is the one branch-heavy block inside the token grant. Pure move: every statement is preserved verbatim. The only edits are bare `return`s becoming the typed returns of their new helper, `c.userID` becoming the `userID` parameter inside voiceJoinPublishPerms, and voiceJoinComplete re-reading `ch.VoiceMaxUsers` instead of receiving it — `ch` is never mutated, so the value is identical. Verified by normalising both revisions of the region to sorted, comment- and whitespace-stripped statements and diffing: the only deltas are the ones listed above. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor: collapse the three duplicated sibling pairs dupl flagged three pairs of adjacent near-identical functions. Each pair is now one parameterised implementation plus two thin, still-greppable wrappers. - ws/voice_controls.go: handleVoiceMuteV2 / handleVoiceDeafenV2 share voiceSelfToggleV2; handleVoiceCameraV2 / handleVoiceScreenshareV2 share voiceStreamToggleV2. Camera and screenshare drawing from one voice_max_video budget (OC-0023) was a bug caused by exactly this duplication drifting, so one body is the point, not a side effect. - db/mention_queries.go: ListMentionTargetsByRoles / ListMentionTargetsByUserIDs share listMentionTargets. The matched column is a closed named type (mentionTargetColumn) rather than a bare string, so the value interpolated into the SELECT cannot become caller-supplied. Behaviour is unchanged: every rate-limit key, error code, error string, slog message and slog key is preserved verbatim, including the two "failed to update <kind> state" messages, which are now assembled the same way enableVideoSlot already assembled them. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(api): extract readEmojiUpload from handleCreateEmoji handleCreateEmoji was 101 lines against a 100-line budget. The upload-bytes stage — pull the file out of the parsed form, cap its size, sniff its MIME type and sniff its dimensions — is the one self-contained block in it, and it already wrote its own refusals, so it moves out whole as readEmojiUpload. The permission-before-parse ordering the doc comment calls out is unchanged; so is every error string. file.Close() now runs when the helper returns rather than when the handler does, which is strictly earlier and unobservable: the bytes are already copied into raw and nothing else touches the handle. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor: extract one cohesive block from three single-budget offenders Each of these was over exactly one budget, so each gets exactly one extraction rather than a restructure: - api/totp_handler.go handleVerifyTOTP (102 lines / 100): the block that resolves the user behind the partial-auth challenge and decrypts their TOTP secret becomes totpChallengeSecret. The ban-inside-the-partial-window check moves with it. - service/message_reactions.go handleReaction (cyclop 21 / 20): the whole authorisation chain — channel lookup, archived gate, DM participant and block checks, non-DM permission check — becomes reactionAudience, which also returns the DM fan-out audience it already resolved. Check order is unchanged and load-bearing. - db/admin_queries.go BackupToSafe (cyclop 21 / 20): the character allowlist loop and the SQL-comment rejection become validateBackupPathChars. That loop alone was most of the branch count. No error string, no check and no ordering changed. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(plugin): split InstallFromZip into staged install helpers 104 statements / cyclomatic 44 / nestif 12. Split along the stages the code already had: installZipExtract (the per-entry write loop, with installZipEntryDest holding the mode/symlink/zip-slip guard chain and installZipWriteEntry the size-capped copy), installZipStagedManifest, installZipPromote, and installZipReactivate for the :399 nested block. Every zip-slip, symlink, entry-mode and uncompressed-size check is preserved in the same order relative to the writes it guards. The 19 inline `cleanup(); return` sites collapse to 4 in the orchestrator, one per stage, because each helper now returns an error instead of unwinding itself — the staging directory is still removed on exactly the same set of failures. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(api): split newWAFMiddleware into engine build and per-phase helpers 184 lines / cyclomatic 38, and the request-body block at :382 was the worst nested site in the tree at nestif 17. Engine construction moves out of the closure (wafInlineEngine, wafCRSEngine — the Coraza directive string is lifted verbatim), and each request phase becomes its own helper: wafInlineRequestHeaders, wafCRSRequestHeaders (including the Host/Transfer-Encoding re-add for CRS 920280), wafFeedCRSBody and wafInspectRequestBody, which is the old :382 block. The three `handleWAFInterruption(w, it); return` sites inside the body block become one: the helper now returns the interruption and the orchestrator handles it. No statement runs between the two points on either side, so the verdict is honoured identically — in particular a CRS body interruption still returns without replacing r.Body. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(service): split SendMessage and lift EditMessage's access check SendMessage was 79 statements / cyclomatic 35 with an 11-deep nested attachment block at :101; EditMessage was one point over cyclop. SendMessage becomes sendMessagePrecheck (permission and DM-block gates, content sanitisation), sendMessageLinkAttachments (the :101 block: attachment ownership, claim and link) and sendMessageDMSideEffects. EditMessage gets editMessageCheckAccess and nothing else — one budget over earns one extraction. The sanitizeContent fixpoint and the attachment ownership check are unchanged, as is the order of every gate. The DM side effects run behind `isDM && !s.sendMessageDMSideEffects(...)`, so a non-DM never enters them; inside, only the GetDMParticipantIDs failure returns false, matching the one error the original early-returned on. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(admin): split handlePatchUser into per-field apply helpers 106 lines / cyclomatic 29, with the ban block at :154 nested 9 deep. Each optional field of the partial edit becomes its own helper — patchUserPrecheck, patchUserAuthorizeRole, patchUserApplyBan (the :154 block, including the session disconnect and the broadcast) and patchUserApplyRole. Each returns a bool meaning "keep going"; none of them writes a success response, so the single response site in the orchestrator is unchanged. Field application order, the permission-cache invalidation on a role change and the disconnect-and-broadcast on a ban are all preserved, as are the three fail-closed `mod == nil` guards, which now sit at the top of their own helper and still fire on exactly the same conditions. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(admin): split handleSetup into first-run setup stages 143 lines / cyclomatic 30, with the optional-wizard block at :219 sitting exactly on the nestif threshold. Split into the stages the endpoint already had: request gating (rate limit and origin check, which run before any auth exists on a fresh server), owner account creation, and the wizard application that was the :219 block. Every gate in front of the handler is a security control on an unauthenticated endpoint; none moved relative to the work it protects. setup_wizard.go is untouched. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor: split run() into named bootstrap and shutdown steps 131 statements / cyclomatic 57, with the executable-path fallback at :126 nested 9 deep. The five anonymous `defer func(){...}()` blocks become named functions — telemetryStop, runClosePlugins, runStopEventPersistence, runStopAuditWriter, maintenanceStop — and the bootstrap stages move out likewise. Every defer is still registered in run() itself, at the same point in the sequence, so the LIFO teardown order is unchanged; that order is documented in the surrounding comments and is load-bearing (the audit-writer stop must follow database.Close's registration, the event-persistence stop must precede it). runStopEventPersistence is now registered unconditionally with a nil persister meaning "disabled", where the old code registered its defer inside the enabled branch — a no-op occupying that slot cannot change the relative order of the others. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(ws): split handleReconnect into resume stages 77 statements / cyclomatic 41, plus the replay block at :199 and, in handleFreshConnect, the voice-state restore at :622. handleReconnect becomes reconnectPrecheck, reconnectSelectReplay (with reconnectVetColdTail for the cold-tier gap check), reconnectRegister and reconnectWriteReplay. handleFreshConnect's stale-voice cleanup moves to its own helper, where the `if h.livekit != nil` wrapper becomes a guard clause — that block was the tail of its scope, so returning early and falling off the end are the same. The parts that carry the invariants are moved verbatim: reconnectRegister still takes h.seqMu, still calls registerNow inside that same critical section (BUG-123 / OC-0206), still unlocks on every exit, and still emits the "full" tier counter and telemetry on each of its three re-check failures. handleReconnect's two-boolean contract is unchanged — the collapsed `return false, false` sites are all fall-through-to-full-ready, and the single `return true, false` is still the handshake-write-failure path whose teardown already ran (OC-0051). Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * docs(server): fold in the adversarial review of the complexity refactors Eleven skeptic passes over the refactor commits on this branch found no blocker and no major — behaviour is preserved throughout. They did find comment and accuracy defects worth correcting: - db/mention_queries.go: the mentionTargetColumn rationale claimed the named type made the interpolated column "only ever one of the two constants". A Go named type is not closed, so that is a convention the type makes visible, not one it enforces. Reworded, gosec justification included. - ws/voice_controls.go: the dupl collapse generalised away three specifics — that a server deafen is the moderator's to lift (now on the serverDeafen field), the concrete voice_states.camera / voice_states.screenshare column names, and the half of the OC-0023 rationale about neither stream kind hiding from the other's count. All three restored. - ws/voice_join.go: `maxUsers := ch.VoiceMaxUsers` had been hoisted to the top of voiceJoinComplete, moving a read across the tail supersession guard. The read is inert, but it was the one statement in that commit whose position relative to a security guard changed; it now sits at its use, as before. - ws/*_test.go: three test comments cited voice_join.go line numbers that the split invalidated. They now cite the helper by name instead. - service/message_reactions.go: reactionAudience's doc claimed to enforce "every gate on reacting"; it enforces the channel-scoped ones, and the doc now says which gates stay with the caller. - api/emoji_handler.go: the readEmojiUpload call reused the outer `ok` from the auth check by assignment; it gets its own readOK. - admin/setup_handler.go: a moved comment kept a "the response above" deictic that no longer had a response above it. No behaviour change. Build, vet, full tests and -race on five packages green. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(ws): clear the remaining complexity budgets across the hub Eight files, thirteen findings. Each function is split at the stages it already had; no branch is reordered, merged or inverted. - handlers.go handleMessage (cyclop 28, 88 stmts): session re-check, frame decode and result application become handleMessageSessionRecheck, handleMessageDecode and handleMessageApply. The V2 constructor lookup -> DispatchV2 -> Result resolution order is untouched. - serve_ready.go buildReady (cyclop 26, 61 stmts): the per-section fetches split out, readyChannelPayloads among them. Every visibility predicate is preserved verbatim — this is the payload that decides what a client may see. - serve_pumps.go writePump (cyclop 31): writePumpWrite, writePumpDeliver, writePumpDrainChannel and writePumpDrainAndClose. Every channel receive stays in the same select statement, so scheduling is unchanged. - hub_sweep.go sweepStaleVoiceStates (cyclop 22, 56 stmts): the staleness predicate, the hub-lock ordering and the position of the race hook are all as they were — handleVoiceJoin's BUG-088 ordering depends on them. - hub_broadcast.go channelReadAudienceImpl and RefreshChannelVisibility (cyclop 22 each, 57 stmts): channelReadAudienceDM and refreshChannelVisibilityCanSend. The audience predicate is the OC-0090 group-DM leak surface, so it is extracted, never simplified. - livekit_webhook.go (nestif 13 and 14): webhookJoinedEnforceVoiceState, webhookLeftCleanupClient and webhookLeftFinishLeave. DB delete still precedes broadcast on every path. - livekit_download.go EnsureLiveKitBinary (52 stmts): one extraction, ensureLiveKitStageBinary, keeping every archive path check intact. - voice_moderation.go (nestif 8): voiceModDeafenRollback. The persisted server_muted flag remains the authority. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor(api): clear the remaining complexity budgets across the HTTP layer - router.go NewRouter (cyclop 28, 84 stmts): split by wiring concern into routerTOTPKey, routerHealthDeps, routerMiddleware, routerUploadRoutes, routerPluginWiring, routerVoiceRoutes and routerMetricsRoutes. Middleware ORDER is a security property (auth before handler, WAF before body parse, rate limit before work) and is unchanged; the returned cleanup func still closes over and releases everything it did before. - auth_handler.go handleRegister (133 lines) and handleLogin (cyclop 21, 152 lines): registerPolicyGate, registerReadRequest, loginReadRequest and loginAuthenticate. The always-compare posture, every rate-limit key, every counter reset and the ban-check-versus-password-compare order are all preserved — including loginUserFailureThreshold staying unscaled by scaledAuthLimit, which is deliberate and commented. - upload_handler.go handleServeFile (cyclop 31, 128 lines): serveFileResolve and serveFileAuthorize. Every header this sets — Content-Disposition included, which is what stops a stored file being served as active content — is still set with the same value in the same circumstances. - profile_handler.go handleUploadAvatar (120 lines): avatarUploadReadImage, mirroring readEmojiUpload in shape but with the avatar caps and MIME set. The two deliberately do not share a helper. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> * refactor: clear the last complexity budgets in db and admin - db/account.go DeleteAccount (cyclop 28, 55 stmts): grouped by subsystem into deleteAccountAdminGuard, deleteAccountDMChannels and deleteAccountCloseDMChannels, each taking the same transaction. The transaction boundary, the delete ORDER (which foreign keys depend on) and the rollback path are unchanged. - admin/logstream.go handleLogStream (cyclop 24): logStreamAuthorize. Flush cadence, heartbeat and disconnect detection untouched. - admin/setup_wizard.go validateWizard (cyclop 23): grouped by section into wizardValidateIdentity, wizardValidateNetwork and wizardValidateMedia. Every message and bound is unchanged — this is the first input-validation boundary on a fresh server, before any auth exists. With this the tree is at zero: golangci-lint run reports 0 issues against the budgets set in #1384 (funlen 100/50, cyclop 20, nestif 8, dupl 150), with no //nolint and no exclusion added anywhere. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
+168
-118
@@ -90,36 +90,170 @@ func writeModerationErr(w http.ResponseWriter, err error) {
|
||||
}
|
||||
}
|
||||
|
||||
// patchUserPrecheck resolves and validates the target of a
|
||||
// PATCH /admin/api/users/{id} before any mutation is attempted. It reports
|
||||
// whether the handler may continue; on false it has already written the error
|
||||
// response.
|
||||
func patchUserPrecheck(w http.ResponseWriter, r *http.Request, database *db.DB) (int64, patchUserRequest, int64, bool) {
|
||||
var req patchUserRequest
|
||||
|
||||
id, err := pathInt64(r, "id")
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid user id")
|
||||
return 0, req, 0, false
|
||||
}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid request body")
|
||||
return 0, req, 0, false
|
||||
}
|
||||
|
||||
user, err := database.GetUserByID(r.Context(), id)
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch user")
|
||||
return 0, req, 0, false
|
||||
}
|
||||
if user == nil {
|
||||
writeErr(w, http.StatusNotFound, "NOT_FOUND", "user not found")
|
||||
return 0, req, 0, false
|
||||
}
|
||||
|
||||
actor := actorFromContext(r)
|
||||
|
||||
// Prevent admins from modifying their own role or ban status, which
|
||||
// could lock them out of the admin panel with no recovery path.
|
||||
if id == actor {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "cannot modify your own account via admin panel")
|
||||
return 0, req, 0, false
|
||||
}
|
||||
|
||||
return id, req, actor, true
|
||||
}
|
||||
|
||||
// patchUserAuthorizeRole runs every ChangeUserRole precondition for a PATCH
|
||||
// carrying role_id without committing anything; a request without role_id is
|
||||
// a no-op. It reports whether the handler may continue; on false it has
|
||||
// already written the error response.
|
||||
func patchUserAuthorizeRole(w http.ResponseWriter, r *http.Request, mod *service.ModerationService, actor, id int64, req patchUserRequest) bool {
|
||||
if req.RoleID == nil {
|
||||
return true
|
||||
}
|
||||
if mod == nil {
|
||||
// Fail closed rather than fall back to an unchecked UPDATE.
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "moderation service unavailable")
|
||||
return false
|
||||
}
|
||||
if _, _, _, err := mod.AuthorizeRoleChange(r.Context(), actor, id, *req.RoleID); err != nil {
|
||||
writeModerationErr(w, err)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// patchUserApplyBan commits the ban/unban half of the PATCH and fans the
|
||||
// result out to connected clients; a request without banned is a no-op. It
|
||||
// reports whether the handler may continue; on false it has already written
|
||||
// the error response.
|
||||
func patchUserApplyBan(w http.ResponseWriter, r *http.Request, hub HubBroadcaster, mod *service.ModerationService, actor, id int64, req patchUserRequest) bool {
|
||||
if req.Banned == nil {
|
||||
return true
|
||||
}
|
||||
if mod == nil {
|
||||
// Fail closed rather than fall back to an unchecked UPDATE.
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "moderation service unavailable")
|
||||
return false
|
||||
}
|
||||
banReason := ""
|
||||
if req.BanReason != nil {
|
||||
banReason = *req.BanReason
|
||||
}
|
||||
var banExpires *time.Time
|
||||
if req.BanDurationHours != nil && *req.BanDurationHours != 0 {
|
||||
hours := *req.BanDurationHours
|
||||
if hours < 0 || hours > maxBanDurationHours {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "ban_duration_hours must be between 1 and 8760")
|
||||
return false
|
||||
}
|
||||
t := time.Now().Add(time.Duration(hours) * time.Hour)
|
||||
banExpires = &t
|
||||
}
|
||||
var actionErr error
|
||||
if *req.Banned {
|
||||
actionErr = mod.BanUser(r.Context(), actor, id, banReason, banExpires)
|
||||
} else {
|
||||
actionErr = mod.UnbanUser(r.Context(), actor, id)
|
||||
}
|
||||
if actionErr != nil {
|
||||
writeModerationErr(w, actionErr)
|
||||
return false
|
||||
}
|
||||
switch {
|
||||
case *req.Banned && hub != nil:
|
||||
hub.BroadcastMemberBan(id)
|
||||
case !*req.Banned && hub != nil:
|
||||
// Ban had no WS event on the way out (member_ban hard-deletes
|
||||
// the row client-side); unban needs one on the way back in, or
|
||||
// every already-connected client keeps the user missing from
|
||||
// its member store while a freshly connecting client sees them.
|
||||
if mub, ok := hub.(memberUnbanBroadcaster); ok {
|
||||
mub.BroadcastMemberUnban(id)
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// patchUserApplyRole commits the role half of the PATCH and fans the result
|
||||
// out to connected clients; a request without role_id is a no-op. It reports
|
||||
// whether the handler may continue; on false it has already written the error
|
||||
// response.
|
||||
func patchUserApplyRole(w http.ResponseWriter, r *http.Request, hub HubBroadcaster, permInvalidator PermissionInvalidator, mod *service.ModerationService, actor, id int64, req patchUserRequest) bool {
|
||||
if req.RoleID == nil {
|
||||
return true
|
||||
}
|
||||
// Routed through ModerationService, which re-runs the same
|
||||
// MANAGE_ROLES, actor-outranks-target, and assign-below-own-rank
|
||||
// checks the AuthorizeRoleChange pre-flight above already passed
|
||||
// (a second pass, not a redundant one: it catches anything that
|
||||
// changed in the window between the pre-flight and here, e.g. a
|
||||
// concurrent role delete), then commits and writes the audit row.
|
||||
if mod == nil {
|
||||
// Fail closed rather than fall back to an unchecked UPDATE.
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "moderation service unavailable")
|
||||
return false
|
||||
}
|
||||
newRole, err := mod.ChangeUserRole(r.Context(), actor, id, *req.RoleID)
|
||||
if err != nil {
|
||||
writeModerationErr(w, err)
|
||||
return false
|
||||
}
|
||||
if permInvalidator != nil {
|
||||
permInvalidator.InvalidateUser(id)
|
||||
}
|
||||
// Use the role ChangeUserRole already loaded and validated rather
|
||||
// than re-reading it: a re-read can race a concurrent role delete
|
||||
// (or a transient read error) and silently skip this whole
|
||||
// fan-out, leaving the demoted user's socket subscribed to
|
||||
// channels it can no longer read (OC-0045). The role change
|
||||
// itself already committed, so the fan-out must not be
|
||||
// conditional on anything past that point.
|
||||
if hub != nil {
|
||||
hub.BroadcastMemberUpdate(id, newRole.Name)
|
||||
// BroadcastMemberUpdate only revokes subscriptions the new
|
||||
// role can no longer read (hub_broadcast.go's
|
||||
// revokeUnreadableChannels); it never grants the ones the
|
||||
// new role newly gained READ_MESSAGES on. Without this,
|
||||
// a promoted user's sidebar is missing channels until
|
||||
// their next reconnect, unlike a role permission edit or
|
||||
// a role delete, which both re-derive visibility fully.
|
||||
hub.RefreshAllChannelVisibility()
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func handlePatchUser(database *db.DB, hub HubBroadcaster, permInvalidator PermissionInvalidator, mod *service.ModerationService) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := pathInt64(r, "id")
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid user id")
|
||||
return
|
||||
}
|
||||
|
||||
var req patchUserRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid request body")
|
||||
return
|
||||
}
|
||||
|
||||
user, err := database.GetUserByID(r.Context(), id)
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch user")
|
||||
return
|
||||
}
|
||||
if user == nil {
|
||||
writeErr(w, http.StatusNotFound, "NOT_FOUND", "user not found")
|
||||
return
|
||||
}
|
||||
|
||||
actor := actorFromContext(r)
|
||||
|
||||
// Prevent admins from modifying their own role or ban status, which
|
||||
// could lock them out of the admin panel with no recovery path.
|
||||
if id == actor {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "cannot modify your own account via admin panel")
|
||||
id, req, actor, ok := patchUserPrecheck(w, r, database)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -133,16 +267,8 @@ func handlePatchUser(database *db.DB, hub HubBroadcaster, permInvalidator Permis
|
||||
// (OC-0215). Running every ChangeUserRole precondition up front,
|
||||
// before either mutation lands, keeps the PATCH all-or-nothing from
|
||||
// the caller's perspective.
|
||||
if req.RoleID != nil {
|
||||
if mod == nil {
|
||||
// Fail closed rather than fall back to an unchecked UPDATE.
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "moderation service unavailable")
|
||||
return
|
||||
}
|
||||
if _, _, _, err := mod.AuthorizeRoleChange(r.Context(), actor, id, *req.RoleID); err != nil {
|
||||
writeModerationErr(w, err)
|
||||
return
|
||||
}
|
||||
if !patchUserAuthorizeRole(w, r, mod, actor, id, req) {
|
||||
return
|
||||
}
|
||||
|
||||
// Ban/unban first: it routes through ModerationService, which enforces
|
||||
@@ -151,88 +277,12 @@ func handlePatchUser(database *db.DB, hub HubBroadcaster, permInvalidator Permis
|
||||
// role change, if requested, was already authorized above, so a ban
|
||||
// committing here cannot be followed by a refused role change leaving
|
||||
// a half-applied PATCH behind.
|
||||
if req.Banned != nil {
|
||||
if mod == nil {
|
||||
// Fail closed rather than fall back to an unchecked UPDATE.
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "moderation service unavailable")
|
||||
return
|
||||
}
|
||||
banReason := ""
|
||||
if req.BanReason != nil {
|
||||
banReason = *req.BanReason
|
||||
}
|
||||
var banExpires *time.Time
|
||||
if req.BanDurationHours != nil && *req.BanDurationHours != 0 {
|
||||
hours := *req.BanDurationHours
|
||||
if hours < 0 || hours > maxBanDurationHours {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "ban_duration_hours must be between 1 and 8760")
|
||||
return
|
||||
}
|
||||
t := time.Now().Add(time.Duration(hours) * time.Hour)
|
||||
banExpires = &t
|
||||
}
|
||||
var actionErr error
|
||||
if *req.Banned {
|
||||
actionErr = mod.BanUser(r.Context(), actor, id, banReason, banExpires)
|
||||
} else {
|
||||
actionErr = mod.UnbanUser(r.Context(), actor, id)
|
||||
}
|
||||
if actionErr != nil {
|
||||
writeModerationErr(w, actionErr)
|
||||
return
|
||||
}
|
||||
switch {
|
||||
case *req.Banned && hub != nil:
|
||||
hub.BroadcastMemberBan(id)
|
||||
case !*req.Banned && hub != nil:
|
||||
// Ban had no WS event on the way out (member_ban hard-deletes
|
||||
// the row client-side); unban needs one on the way back in, or
|
||||
// every already-connected client keeps the user missing from
|
||||
// its member store while a freshly connecting client sees them.
|
||||
if mub, ok := hub.(memberUnbanBroadcaster); ok {
|
||||
mub.BroadcastMemberUnban(id)
|
||||
}
|
||||
}
|
||||
if !patchUserApplyBan(w, r, hub, mod, actor, id, req) {
|
||||
return
|
||||
}
|
||||
|
||||
if req.RoleID != nil {
|
||||
// Routed through ModerationService, which re-runs the same
|
||||
// MANAGE_ROLES, actor-outranks-target, and assign-below-own-rank
|
||||
// checks the AuthorizeRoleChange pre-flight above already passed
|
||||
// (a second pass, not a redundant one: it catches anything that
|
||||
// changed in the window between the pre-flight and here, e.g. a
|
||||
// concurrent role delete), then commits and writes the audit row.
|
||||
if mod == nil {
|
||||
// Fail closed rather than fall back to an unchecked UPDATE.
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "moderation service unavailable")
|
||||
return
|
||||
}
|
||||
newRole, err := mod.ChangeUserRole(r.Context(), actor, id, *req.RoleID)
|
||||
if err != nil {
|
||||
writeModerationErr(w, err)
|
||||
return
|
||||
}
|
||||
if permInvalidator != nil {
|
||||
permInvalidator.InvalidateUser(id)
|
||||
}
|
||||
// Use the role ChangeUserRole already loaded and validated rather
|
||||
// than re-reading it: a re-read can race a concurrent role delete
|
||||
// (or a transient read error) and silently skip this whole
|
||||
// fan-out, leaving the demoted user's socket subscribed to
|
||||
// channels it can no longer read (OC-0045). The role change
|
||||
// itself already committed, so the fan-out must not be
|
||||
// conditional on anything past that point.
|
||||
if hub != nil {
|
||||
hub.BroadcastMemberUpdate(id, newRole.Name)
|
||||
// BroadcastMemberUpdate only revokes subscriptions the new
|
||||
// role can no longer read (hub_broadcast.go's
|
||||
// revokeUnreadableChannels); it never grants the ones the
|
||||
// new role newly gained READ_MESSAGES on. Without this,
|
||||
// a promoted user's sidebar is missing channels until
|
||||
// their next reconnect, unlike a role permission edit or
|
||||
// a role delete, which both re-derive visibility fully.
|
||||
hub.RefreshAllChannelVisibility()
|
||||
}
|
||||
if !patchUserApplyRole(w, r, hub, permInvalidator, mod, actor, id, req) {
|
||||
return
|
||||
}
|
||||
|
||||
updated, err := database.GetUserByID(r.Context(), id)
|
||||
|
||||
+56
-43
@@ -408,57 +408,70 @@ func categorizeSource(r slog.Record) string {
|
||||
}
|
||||
}
|
||||
|
||||
// logStreamAuthorize runs the log stream's authentication prologue: it redeems
|
||||
// the single-use ticket, resolves the principal behind it, and returns the
|
||||
// re-check closure the stream must call before every write. It writes the error
|
||||
// response itself and reports false when the caller must stop.
|
||||
func logStreamAuthorize(w http.ResponseWriter, r *http.Request, database *db.DB) (func() bool, bool) {
|
||||
// Authenticate via single-use ticket.
|
||||
ticket := r.URL.Query().Get("ticket")
|
||||
entry, ok := logTickets.redeem(ticket)
|
||||
if ticket == "" || !ok {
|
||||
errResp, _ := json.Marshal(map[string]string{
|
||||
"error": "UNAUTHORIZED",
|
||||
"message": "invalid or expired ticket",
|
||||
})
|
||||
http.Error(w, string(errResp), http.StatusUnauthorized)
|
||||
return nil, false
|
||||
}
|
||||
// Stream lifetime == request lifetime, so all principal re-checks below
|
||||
// use the stream request's context. The ticket's hash is resolved the
|
||||
// same way adminAuthMiddleware resolves a bearer credential — login
|
||||
// session first, then API token — so revoking either kind mid-stream
|
||||
// cuts the stream.
|
||||
ctx := r.Context()
|
||||
user, role, _, err := auth.ResolveTokenHash(ctx, database, entry.tokenHash)
|
||||
if err != nil || user == nil || role == nil {
|
||||
errResp, _ := json.Marshal(map[string]string{
|
||||
"error": "UNAUTHORIZED",
|
||||
"message": "invalid or expired session",
|
||||
})
|
||||
http.Error(w, string(errResp), http.StatusUnauthorized)
|
||||
return nil, false
|
||||
}
|
||||
principalStillAuthorized := func() bool {
|
||||
current, currentRole, _, resolveErr := auth.ResolveTokenHash(ctx, database, entry.tokenHash)
|
||||
if resolveErr != nil || current == nil || currentRole == nil {
|
||||
return false
|
||||
}
|
||||
// A ban mid-stream must cut the stream, same as adminAuthMiddleware
|
||||
// rejects a banned user on the request path.
|
||||
if auth.IsEffectivelyBanned(current) {
|
||||
return false
|
||||
}
|
||||
return permissions.HasAdmin(currentRole.Permissions)
|
||||
}
|
||||
if !principalStillAuthorized() {
|
||||
errResp, _ := json.Marshal(map[string]string{
|
||||
"error": "FORBIDDEN",
|
||||
"message": "administrator permission required",
|
||||
})
|
||||
http.Error(w, string(errResp), http.StatusForbidden)
|
||||
return nil, false
|
||||
}
|
||||
return principalStillAuthorized, true
|
||||
}
|
||||
|
||||
// handleLogStream serves an SSE endpoint that streams log entries in real-time.
|
||||
// Auth is via query param ?ticket= — a short-lived single-use ticket obtained
|
||||
// from POST /admin/api/logs/ticket (which requires normal admin auth).
|
||||
func handleLogStream(database *db.DB, ringBuf *RingBuffer) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
// Authenticate via single-use ticket.
|
||||
ticket := r.URL.Query().Get("ticket")
|
||||
entry, ok := logTickets.redeem(ticket)
|
||||
if ticket == "" || !ok {
|
||||
errResp, _ := json.Marshal(map[string]string{
|
||||
"error": "UNAUTHORIZED",
|
||||
"message": "invalid or expired ticket",
|
||||
})
|
||||
http.Error(w, string(errResp), http.StatusUnauthorized)
|
||||
principalStillAuthorized, ok := logStreamAuthorize(w, r, database)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
// Stream lifetime == request lifetime, so all principal re-checks below
|
||||
// use the stream request's context. The ticket's hash is resolved the
|
||||
// same way adminAuthMiddleware resolves a bearer credential — login
|
||||
// session first, then API token — so revoking either kind mid-stream
|
||||
// cuts the stream.
|
||||
ctx := r.Context()
|
||||
user, role, _, err := auth.ResolveTokenHash(ctx, database, entry.tokenHash)
|
||||
if err != nil || user == nil || role == nil {
|
||||
errResp, _ := json.Marshal(map[string]string{
|
||||
"error": "UNAUTHORIZED",
|
||||
"message": "invalid or expired session",
|
||||
})
|
||||
http.Error(w, string(errResp), http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
principalStillAuthorized := func() bool {
|
||||
current, currentRole, _, resolveErr := auth.ResolveTokenHash(ctx, database, entry.tokenHash)
|
||||
if resolveErr != nil || current == nil || currentRole == nil {
|
||||
return false
|
||||
}
|
||||
// A ban mid-stream must cut the stream, same as adminAuthMiddleware
|
||||
// rejects a banned user on the request path.
|
||||
if auth.IsEffectivelyBanned(current) {
|
||||
return false
|
||||
}
|
||||
return permissions.HasAdmin(currentRole.Permissions)
|
||||
}
|
||||
if !principalStillAuthorized() {
|
||||
errResp, _ := json.Marshal(map[string]string{
|
||||
"error": "FORBIDDEN",
|
||||
"message": "administrator permission required",
|
||||
})
|
||||
http.Error(w, string(errResp), http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
// Check that we can flush (required for SSE).
|
||||
flusher, ok := w.(http.Flusher)
|
||||
|
||||
+194
-146
@@ -104,139 +104,17 @@ func handleSetupStatus(database *db.DB, opts SetupOptions) http.HandlerFunc {
|
||||
// users exist in the database, preventing abuse after initial setup.
|
||||
func handleSetup(database *db.DB, limiter *auth.RateLimiter, allowedOrigins []string, hub HubBroadcaster, opts SetupOptions) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
// CSRF protection: reject cross-origin requests (BUG-097).
|
||||
// A request is accepted when it is same-origin, or when its Origin is
|
||||
// explicitly allowlisted. Absent Origin = non-browser client (allow).
|
||||
if origin := r.Header.Get("Origin"); origin != "" {
|
||||
if !isSameOrigin(origin, r.Host) && !isSetupOriginAllowed(origin, allowedOrigins) {
|
||||
writeErr(w, http.StatusForbidden, "FORBIDDEN", "cross-origin setup request blocked")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Rate limit: 5 attempts per minute per IP.
|
||||
// Strip the port so that different source ports from the same IP
|
||||
// are correctly grouped under a single rate-limit bucket.
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
host = r.RemoteAddr
|
||||
}
|
||||
setupKey := "setup:" + host
|
||||
if !limiter.Allow(setupKey, 5, time.Minute) {
|
||||
writeErr(w, http.StatusTooManyRequests, "RATE_LIMITED", "too many setup attempts, try again later")
|
||||
req, host, ok := setupPrecheck(w, r, limiter, allowedOrigins)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
var req setupRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid request body")
|
||||
uid, token, inviteCode, ok := setupCreateOwner(w, r, database, req, host)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
req.Username = strings.TrimSpace(setupSanitizer.Sanitize(req.Username))
|
||||
if req.Username == "" || req.Password == "" {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "username and password are required")
|
||||
return
|
||||
}
|
||||
|
||||
// Validate username format (length, no control/invisible chars).
|
||||
if err := auth.ValidateUsername(req.Username); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := auth.ValidatePasswordStrength(req.Password); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// Validate the whole wizard payload BEFORE creating the account so a
|
||||
// bad value rejects the request instead of leaving a half-configured
|
||||
// server behind an already-created owner.
|
||||
if req.Wizard != nil {
|
||||
if err := validateWizard(req.Wizard); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Hash the password.
|
||||
hash, err := auth.HashPassword(req.Password)
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to hash password")
|
||||
return
|
||||
}
|
||||
|
||||
// Atomically check no users exist and create the owner (BUG-119).
|
||||
// This closes the TOCTOU race between UserCount() and CreateUser().
|
||||
uid, err := database.CreateOwnerIfEmpty(r.Context(), req.Username, hash, ownerRoleID)
|
||||
if errors.Is(err, db.ErrConflict) {
|
||||
writeErr(w, http.StatusForbidden, "FORBIDDEN", "setup has already been completed")
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to create user")
|
||||
return
|
||||
}
|
||||
|
||||
// Issue a session token so the user is immediately logged in.
|
||||
token, err := auth.GenerateToken()
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to generate session token")
|
||||
return
|
||||
}
|
||||
|
||||
device := r.Header.Get("User-Agent")
|
||||
const maxDeviceLen = 512
|
||||
if len(device) > maxDeviceLen {
|
||||
device = device[:maxDeviceLen]
|
||||
}
|
||||
if _, err := database.CreateSession(r.Context(), uid, auth.HashToken(token), device, host); err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to create session")
|
||||
return
|
||||
}
|
||||
|
||||
// Create default channels under canonical categories.
|
||||
_, _ = database.CreateChannel(r.Context(), "general", "text", "Text Channels", "Welcome to the server!", 0)
|
||||
_, _ = database.CreateChannel(r.Context(), "General", "voice", "Voice Channels", "", 0)
|
||||
|
||||
// Generate a bootstrap invite code so the owner can invite others.
|
||||
// Bound it (5 uses / 24h) rather than minting an unlimited, non-expiring
|
||||
// invite — the owner can create fresh invites once logged in.
|
||||
bootstrapInviteExpiry := time.Now().Add(24 * time.Hour)
|
||||
inviteCode, err := database.CreateInvite(r.Context(), uid, 5, &bootstrapInviteExpiry)
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to generate invite code")
|
||||
return
|
||||
}
|
||||
|
||||
// Apply the wizard payload. The account exists from here on, so any
|
||||
// failure downgrades to a warning — never a 5xx that would orphan the
|
||||
// owner behind an opaque error.
|
||||
var warnings []string
|
||||
restartRequired := false
|
||||
restartURL := ""
|
||||
if req.Wizard != nil {
|
||||
if err := applyWizardSettings(r.Context(), database, req.Wizard); err != nil {
|
||||
slog.Error("setup wizard: saving settings failed", "error", err)
|
||||
warnings = append(warnings,
|
||||
"could not save server settings: "+err.Error()+" — adjust them later in the admin panel's Settings page")
|
||||
}
|
||||
if opts.ConfigPath != "" {
|
||||
if err := config.Save(opts.ConfigPath, buildConfigPatch(req.Wizard, opts.RunningCfg)); err != nil {
|
||||
slog.Error("setup wizard: writing config failed", "path", opts.ConfigPath, "error", err)
|
||||
warnings = append(warnings,
|
||||
"could not write "+opts.ConfigPath+": "+err.Error()+" — your account was created; edit the file manually to apply these settings")
|
||||
} else {
|
||||
db.WriteAudit(context.WithoutCancel(r.Context()), database, uid, "config_write", "server", 0,
|
||||
"setup wizard wrote "+opts.ConfigPath+" ("+patchedConfigKeys(req.Wizard)+")")
|
||||
if opts.RunningCfg != nil && wizardChangesRunningConfig(req.Wizard, opts.RunningCfg) {
|
||||
restartRequired = true
|
||||
restartURL = computeRestartURL(r.Host, req.Wizard, opts.RunningCfg)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
warnings, restartRequired, restartURL := setupApplyWizard(r.Context(), database, req.Wizard, uid, r.Host, opts)
|
||||
|
||||
slog.Info("server setup completed", "owner", req.Username, "user_id", uid, "wizard", req.Wizard != nil, "restart", restartRequired)
|
||||
db.WriteAudit(context.WithoutCancel(r.Context()), database, uid, "server_setup", "server", 0,
|
||||
@@ -252,30 +130,200 @@ func handleSetup(database *db.DB, limiter *auth.RateLimiter, allowedOrigins []st
|
||||
Warnings: warnings,
|
||||
})
|
||||
|
||||
// Restart after the response is written so the browser receives the
|
||||
// token and the reconnect URL. Mirrors handleRestoreBackup /
|
||||
// handleApplyUpdate: broadcast, then request the restart in a
|
||||
// goroutine — main.go drains the server and performs the handoff.
|
||||
// tryDirectRestartPending loses only to an already in-flight update
|
||||
// or restore, which will itself restart the process; skipping is
|
||||
// correct then (the response above is already written either way).
|
||||
if restartRequired {
|
||||
if !tryDirectRestartPending() {
|
||||
slog.Warn("setup restart skipped: another restart-sensitive operation is already in progress")
|
||||
return
|
||||
}
|
||||
if hub != nil {
|
||||
hub.BroadcastServerRestart("setup", restartBroadcastDelaySeconds)
|
||||
}
|
||||
restartFn := opts.Restart
|
||||
if restartFn == nil {
|
||||
restartFn = requestRestart
|
||||
}
|
||||
go restartFn("setup_wizard")
|
||||
setupRestartAfterResponse(hub, opts)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// setupPrecheck runs every gate in front of the first-run setup endpoint —
|
||||
// origin check, rate limit, body decode, credential and wizard validation —
|
||||
// before any state is created. It writes the error response itself; ok=false
|
||||
// means the caller must return immediately. The returned host is the
|
||||
// rate-limit bucket key, reused as the session IP.
|
||||
func setupPrecheck(w http.ResponseWriter, r *http.Request, limiter *auth.RateLimiter, allowedOrigins []string) (setupRequest, string, bool) {
|
||||
var req setupRequest
|
||||
|
||||
// CSRF protection: reject cross-origin requests (BUG-097).
|
||||
// A request is accepted when it is same-origin, or when its Origin is
|
||||
// explicitly allowlisted. Absent Origin = non-browser client (allow).
|
||||
if origin := r.Header.Get("Origin"); origin != "" {
|
||||
if !isSameOrigin(origin, r.Host) && !isSetupOriginAllowed(origin, allowedOrigins) {
|
||||
writeErr(w, http.StatusForbidden, "FORBIDDEN", "cross-origin setup request blocked")
|
||||
return req, "", false
|
||||
}
|
||||
}
|
||||
|
||||
// Rate limit: 5 attempts per minute per IP.
|
||||
// Strip the port so that different source ports from the same IP
|
||||
// are correctly grouped under a single rate-limit bucket.
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
host = r.RemoteAddr
|
||||
}
|
||||
setupKey := "setup:" + host
|
||||
if !limiter.Allow(setupKey, 5, time.Minute) {
|
||||
writeErr(w, http.StatusTooManyRequests, "RATE_LIMITED", "too many setup attempts, try again later")
|
||||
return req, "", false
|
||||
}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid request body")
|
||||
return req, "", false
|
||||
}
|
||||
|
||||
req.Username = strings.TrimSpace(setupSanitizer.Sanitize(req.Username))
|
||||
if req.Username == "" || req.Password == "" {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "username and password are required")
|
||||
return req, "", false
|
||||
}
|
||||
|
||||
// Validate username format (length, no control/invisible chars).
|
||||
if err := auth.ValidateUsername(req.Username); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", err.Error())
|
||||
return req, "", false
|
||||
}
|
||||
|
||||
if err := auth.ValidatePasswordStrength(req.Password); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", err.Error())
|
||||
return req, "", false
|
||||
}
|
||||
|
||||
// Validate the whole wizard payload BEFORE creating the account so a
|
||||
// bad value rejects the request instead of leaving a half-configured
|
||||
// server behind an already-created owner.
|
||||
if req.Wizard != nil {
|
||||
if err := validateWizard(req.Wizard); err != nil {
|
||||
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", err.Error())
|
||||
return req, "", false
|
||||
}
|
||||
}
|
||||
|
||||
return req, host, true
|
||||
}
|
||||
|
||||
// setupCreateOwner creates the owner account and everything that ships with
|
||||
// it: the session token, the default channels and the bootstrap invite. It
|
||||
// writes the error response itself; ok=false means the caller must return
|
||||
// immediately.
|
||||
func setupCreateOwner(w http.ResponseWriter, r *http.Request, database *db.DB, req setupRequest, host string) (int64, string, string, bool) {
|
||||
// Hash the password.
|
||||
hash, err := auth.HashPassword(req.Password)
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to hash password")
|
||||
return 0, "", "", false
|
||||
}
|
||||
|
||||
// Atomically check no users exist and create the owner (BUG-119).
|
||||
// This closes the TOCTOU race between UserCount() and CreateUser().
|
||||
uid, err := database.CreateOwnerIfEmpty(r.Context(), req.Username, hash, ownerRoleID)
|
||||
if errors.Is(err, db.ErrConflict) {
|
||||
writeErr(w, http.StatusForbidden, "FORBIDDEN", "setup has already been completed")
|
||||
return 0, "", "", false
|
||||
}
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to create user")
|
||||
return 0, "", "", false
|
||||
}
|
||||
|
||||
// Issue a session token so the user is immediately logged in.
|
||||
token, err := auth.GenerateToken()
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to generate session token")
|
||||
return 0, "", "", false
|
||||
}
|
||||
|
||||
device := r.Header.Get("User-Agent")
|
||||
const maxDeviceLen = 512
|
||||
if len(device) > maxDeviceLen {
|
||||
device = device[:maxDeviceLen]
|
||||
}
|
||||
if _, err := database.CreateSession(r.Context(), uid, auth.HashToken(token), device, host); err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to create session")
|
||||
return 0, "", "", false
|
||||
}
|
||||
|
||||
// Create default channels under canonical categories.
|
||||
_, _ = database.CreateChannel(r.Context(), "general", "text", "Text Channels", "Welcome to the server!", 0)
|
||||
_, _ = database.CreateChannel(r.Context(), "General", "voice", "Voice Channels", "", 0)
|
||||
|
||||
// Generate a bootstrap invite code so the owner can invite others.
|
||||
// Bound it (5 uses / 24h) rather than minting an unlimited, non-expiring
|
||||
// invite — the owner can create fresh invites once logged in.
|
||||
bootstrapInviteExpiry := time.Now().Add(24 * time.Hour)
|
||||
inviteCode, err := database.CreateInvite(r.Context(), uid, 5, &bootstrapInviteExpiry)
|
||||
if err != nil {
|
||||
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to generate invite code")
|
||||
return 0, "", "", false
|
||||
}
|
||||
|
||||
return uid, token, inviteCode, true
|
||||
}
|
||||
|
||||
// setupApplyWizard applies the wizard payload. wr == nil is the legacy
|
||||
// request shape (create the owner account only) and applies nothing. reqHost
|
||||
// is the request's Host header, from which the post-restart admin-panel URL
|
||||
// is derived.
|
||||
//
|
||||
// The account exists from here on, so any
|
||||
// failure downgrades to a warning — never a 5xx that would orphan the
|
||||
// owner behind an opaque error.
|
||||
func setupApplyWizard(ctx context.Context, database *db.DB, wr *setupWizardRequest, uid int64, reqHost string, opts SetupOptions) ([]string, bool, string) {
|
||||
var warnings []string
|
||||
restartRequired := false
|
||||
restartURL := ""
|
||||
if wr == nil {
|
||||
return warnings, restartRequired, restartURL
|
||||
}
|
||||
|
||||
if err := applyWizardSettings(ctx, database, wr); err != nil {
|
||||
slog.Error("setup wizard: saving settings failed", "error", err)
|
||||
warnings = append(warnings,
|
||||
"could not save server settings: "+err.Error()+" — adjust them later in the admin panel's Settings page")
|
||||
}
|
||||
if opts.ConfigPath != "" {
|
||||
if err := config.Save(opts.ConfigPath, buildConfigPatch(wr, opts.RunningCfg)); err != nil {
|
||||
slog.Error("setup wizard: writing config failed", "path", opts.ConfigPath, "error", err)
|
||||
warnings = append(warnings,
|
||||
"could not write "+opts.ConfigPath+": "+err.Error()+" — your account was created; edit the file manually to apply these settings")
|
||||
} else {
|
||||
db.WriteAudit(context.WithoutCancel(ctx), database, uid, "config_write", "server", 0,
|
||||
"setup wizard wrote "+opts.ConfigPath+" ("+patchedConfigKeys(wr)+")")
|
||||
if opts.RunningCfg != nil && wizardChangesRunningConfig(wr, opts.RunningCfg) {
|
||||
restartRequired = true
|
||||
restartURL = computeRestartURL(reqHost, wr, opts.RunningCfg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return warnings, restartRequired, restartURL
|
||||
}
|
||||
|
||||
// setupRestartAfterResponse hands the setup-wizard restart off to main.go.
|
||||
//
|
||||
// Called by handleSetup only after it has written its response, so the
|
||||
// browser receives the token and the reconnect URL before the process
|
||||
// goes away. Mirrors handleRestoreBackup / handleApplyUpdate: broadcast,
|
||||
// then request the restart in a goroutine — main.go drains the server and
|
||||
// performs the handoff. tryDirectRestartPending loses only to an already
|
||||
// in-flight update or restore, which will itself restart the process;
|
||||
// skipping is correct then, since the caller's response is written either
|
||||
// way.
|
||||
func setupRestartAfterResponse(hub HubBroadcaster, opts SetupOptions) {
|
||||
if !tryDirectRestartPending() {
|
||||
slog.Warn("setup restart skipped: another restart-sensitive operation is already in progress")
|
||||
return
|
||||
}
|
||||
if hub != nil {
|
||||
hub.BroadcastServerRestart("setup", restartBroadcastDelaySeconds)
|
||||
}
|
||||
restartFn := opts.Restart
|
||||
if restartFn == nil {
|
||||
restartFn = requestRestart
|
||||
}
|
||||
go restartFn("setup_wizard")
|
||||
}
|
||||
|
||||
// restartBroadcastDelaySeconds is the countdown clients are told before the
|
||||
// setup-wizard restart. There are normally no chat clients connected during
|
||||
// first-run setup, so this is informational.
|
||||
|
||||
@@ -85,6 +85,21 @@ var validVoiceQualities = map[string]struct{}{
|
||||
// be called BEFORE the owner account is created so a bad payload rejects the
|
||||
// whole request instead of leaving a half-configured server.
|
||||
func validateWizard(wr *setupWizardRequest) error {
|
||||
if err := wizardValidateIdentity(wr); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := wizardValidateNetwork(wr); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := wizardValidateMedia(wr); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// wizardValidateIdentity checks and normalises the settings-table fields the
|
||||
// server reads live: the display name and the message of the day.
|
||||
func wizardValidateIdentity(wr *setupWizardRequest) error {
|
||||
if wr.ServerName != nil {
|
||||
name := strings.TrimSpace(setupSanitizer.Sanitize(*wr.ServerName))
|
||||
if name == "" {
|
||||
@@ -102,6 +117,12 @@ func validateWizard(wr *setupWizardRequest) error {
|
||||
}
|
||||
*wr.Motd = motd
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// wizardValidateNetwork checks and normalises the listener and TLS fields,
|
||||
// including the cross-field rule that ACME issuance needs a domain.
|
||||
func wizardValidateNetwork(wr *setupWizardRequest) error {
|
||||
if wr.Port != nil && (*wr.Port < 1 || *wr.Port > 65535) {
|
||||
return fmt.Errorf("port must be between 1 and 65535")
|
||||
}
|
||||
@@ -125,6 +146,12 @@ func validateWizard(wr *setupWizardRequest) error {
|
||||
(wr.TLSDomain == nil || *wr.TLSDomain == "") {
|
||||
return fmt.Errorf("tls_domain is required when tls_mode is acme")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// wizardValidateMedia checks and normalises the upload-size cap and the voice
|
||||
// quality preset.
|
||||
func wizardValidateMedia(wr *setupWizardRequest) error {
|
||||
if wr.UploadMaxSizeMB != nil && (*wr.UploadMaxSizeMB < 1 || *wr.UploadMaxSizeMB > maxUploadSizeMB) {
|
||||
return fmt.Errorf("upload_max_size_mb must be between 1 and %d", maxUploadSizeMB)
|
||||
}
|
||||
|
||||
+239
-199
@@ -145,79 +145,12 @@ func MountAuthRoutes(r chi.Router, database *db.DB, limiter *auth.RateLimiter, t
|
||||
func handleRegister(database *db.DB, trustedProxies []string) http.HandlerFunc {
|
||||
proxyNets := parseCIDRList(trustedProxies) // W3-3a: parse once at construction
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
registrationOpen, err := isRegistrationOpen(r.Context(), database)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "failed to load registration policy",
|
||||
})
|
||||
return
|
||||
}
|
||||
if !registrationOpen {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "registration is currently closed",
|
||||
})
|
||||
if !registerPolicyGate(w, r, database) {
|
||||
return
|
||||
}
|
||||
|
||||
require2FA, err := isRequire2FAEnabled(r.Context(), database)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "failed to load registration policy",
|
||||
})
|
||||
return
|
||||
}
|
||||
if require2FA {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "registration is unavailable while two-factor authentication is required",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
var req registerRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "malformed request body",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// F: use the fixpoint sanitizer (service.SanitizeText), not the bare
|
||||
// sanitizer.Sanitize below — Sanitize's output is always HTML-escaped
|
||||
// (' -> ', & -> &, " -> "), so a plain call here would store
|
||||
// a different string than what handleLogin looks up (which only
|
||||
// trims), permanently locking out any username containing one of
|
||||
// those characters. See service.SanitizeText's doc comment.
|
||||
req.Username = strings.TrimSpace(service.SanitizeText(req.Username))
|
||||
req.InviteCode = strings.TrimSpace(req.InviteCode)
|
||||
|
||||
if req.Username == "" || req.Password == "" || req.InviteCode == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "username, password, and invite_code are required",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Validate username format (length, no control/invisible chars).
|
||||
if err := auth.ValidateUsername(req.Username); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Validate password strength before anything else.
|
||||
if err := auth.ValidatePasswordStrength(req.Password); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: err.Error(),
|
||||
})
|
||||
req, ok := registerReadRequest(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -294,147 +227,108 @@ func handleRegister(database *db.DB, trustedProxies []string) http.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// registerPolicyGate reports whether registration is currently permitted,
|
||||
// writing the refusal response itself when it is not.
|
||||
func registerPolicyGate(w http.ResponseWriter, r *http.Request, database *db.DB) bool {
|
||||
registrationOpen, err := isRegistrationOpen(r.Context(), database)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "failed to load registration policy",
|
||||
})
|
||||
return false
|
||||
}
|
||||
if !registrationOpen {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "registration is currently closed",
|
||||
})
|
||||
return false
|
||||
}
|
||||
|
||||
require2FA, err := isRequire2FAEnabled(r.Context(), database)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "failed to load registration policy",
|
||||
})
|
||||
return false
|
||||
}
|
||||
if require2FA {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "registration is unavailable while two-factor authentication is required",
|
||||
})
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// registerReadRequest decodes and validates the registration body, writing the
|
||||
// rejection response itself when the input cannot be used.
|
||||
func registerReadRequest(w http.ResponseWriter, r *http.Request) (registerRequest, bool) {
|
||||
var req registerRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "malformed request body",
|
||||
})
|
||||
return req, false
|
||||
}
|
||||
|
||||
// F: use the fixpoint sanitizer (service.SanitizeText), not the bare
|
||||
// sanitizer.Sanitize below — Sanitize's output is always HTML-escaped
|
||||
// (' -> ', & -> &, " -> "), so a plain call here would store
|
||||
// a different string than what handleLogin looks up (which only
|
||||
// trims), permanently locking out any username containing one of
|
||||
// those characters. See service.SanitizeText's doc comment.
|
||||
req.Username = strings.TrimSpace(service.SanitizeText(req.Username))
|
||||
req.InviteCode = strings.TrimSpace(req.InviteCode)
|
||||
|
||||
if req.Username == "" || req.Password == "" || req.InviteCode == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "username, password, and invite_code are required",
|
||||
})
|
||||
return req, false
|
||||
}
|
||||
|
||||
// Validate username format (length, no control/invisible chars).
|
||||
if err := auth.ValidateUsername(req.Username); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: err.Error(),
|
||||
})
|
||||
return req, false
|
||||
}
|
||||
|
||||
// Validate password strength before anything else.
|
||||
if err := auth.ValidatePasswordStrength(req.Password); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: err.Error(),
|
||||
})
|
||||
return req, false
|
||||
}
|
||||
return req, true
|
||||
}
|
||||
|
||||
// handleLogin processes POST /api/v1/auth/login.
|
||||
func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.PartialAuthStore, trustedProxies []string) http.HandlerFunc {
|
||||
proxyNets := parseCIDRList(trustedProxies) // W3-3a: parse once at construction
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
var req loginRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "malformed request body",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
// Do NOT trim req.Password — passwords may intentionally contain
|
||||
// leading/trailing whitespace. Bcrypt handles arbitrary bytes.
|
||||
|
||||
if req.Username == "" || req.Password == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "username and password are required",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// F: reject an over-long username before it is ever used to build a
|
||||
// RateLimiter map key below (unameKey, failKey, userFailKey, lockout
|
||||
// keys). Unlike registration, login has no account to validate
|
||||
// against yet, so nothing else bounds this value — an unauthenticated
|
||||
// caller could otherwise pin an arbitrarily large, body-sized string
|
||||
// as a retained key (Cleanup only evicts it after hours). Mirrors the
|
||||
// same 32-rune cap auth.ValidateUsername enforces at registration.
|
||||
if utf8.RuneCountInString(req.Username) > maxLoginUsernameLen {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "username is too long",
|
||||
})
|
||||
req, ok := loginReadRequest(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
ip := clientIPWithProxies(r, proxyNets)
|
||||
|
||||
// Check per-IP lockout first.
|
||||
lockKey := "login_lock:" + ip
|
||||
if limiter.IsLockedOut(lockKey) {
|
||||
writeJSON(w, http.StatusTooManyRequests, errorResponse{
|
||||
Error: "RATE_LIMITED",
|
||||
Message: "account temporarily locked due to too many failed attempts",
|
||||
})
|
||||
user, ok := loginAuthenticate(w, r, database, limiter, req, ip)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
// BUG-110: Also check per-username lockout to prevent distributed brute force.
|
||||
// F1: canonicalize the username the same way GetUserByUsername does (COLLATE
|
||||
// NOCASE) before keying the lockout, so case variants of one account
|
||||
// (admin/Admin/ADMIN) share a single bucket instead of each getting its own.
|
||||
unameKey := strings.ToLower(req.Username)
|
||||
userLockKey := "login_user_lock:" + unameKey
|
||||
if limiter.IsLockedOut(userLockKey) {
|
||||
writeJSON(w, http.StatusTooManyRequests, errorResponse{
|
||||
Error: "RATE_LIMITED",
|
||||
Message: "account temporarily locked due to too many failed attempts",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Constant-time lookup: always attempt bcrypt compare even when user
|
||||
// does not exist to prevent timing-based username enumeration.
|
||||
user, err := database.GetUserByUsername(r.Context(), req.Username)
|
||||
|
||||
// Distinguish DB errors from authentication failures. DB errors
|
||||
// should NOT increment the rate limiter — otherwise a transient
|
||||
// DB outage would lock out legitimate users.
|
||||
if err != nil && user == nil {
|
||||
// Could be a real DB error or simply "user not found".
|
||||
// GetUserByUsername returns (nil, nil) for not-found, so a
|
||||
// non-nil error here is a genuine DB failure.
|
||||
slog.Error("login: GetUserByUsername failed", "err", err, "ip", ip)
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "login temporarily unavailable",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
failKey := "login_fail:" + ip
|
||||
userFailKey := "login_user_fail:" + unameKey
|
||||
// F3: atomically reserve this attempt BEFORE the bcrypt compare. The
|
||||
// read-only IsLockedOut gates above are check-then-act: N concurrent
|
||||
// requests all pass them before any failure is recorded below, so the
|
||||
// per-username cap — the only cross-IP brute-force defence — bound
|
||||
// only sequential attackers. Allow records the attempt under the
|
||||
// limiter's lock, capping a concurrent burst at the same budget a
|
||||
// sequential attacker gets. Sized at threshold+1 so the sequential
|
||||
// accepted-input set is unchanged: failures 1–10 still land, the 10th
|
||||
// still trips the lockout (via the Check below), and a correct
|
||||
// password on attempt 10 still succeeds — successful logins reset
|
||||
// both counters. The reservation sits after the DB-error return above
|
||||
// so a transient DB outage still does not consume attempts.
|
||||
if !limiter.Allow(failKey, scaledAuthLimit(loginFailureThreshold)+1, loginFailureWindow) ||
|
||||
!limiter.Allow(userFailKey, loginUserFailureThreshold+1, loginUserFailureWindow) {
|
||||
writeJSON(w, http.StatusTooManyRequests, errorResponse{
|
||||
Error: "RATE_LIMITED",
|
||||
Message: "account temporarily locked due to too many failed attempts",
|
||||
})
|
||||
return
|
||||
}
|
||||
// Always run the password check — with an empty hash when the user does
|
||||
// not exist. auth.CheckPassword performs a dummy bcrypt comparison for an
|
||||
// empty hash, so bcrypt executes on every path and response time stays
|
||||
// constant, preventing timing-based username enumeration. (A `user == nil
|
||||
// || CheckPassword(...)` short-circuit would skip bcrypt entirely for
|
||||
// unknown usernames, reintroducing the timing side-channel.)
|
||||
storedHash := ""
|
||||
if user != nil {
|
||||
storedHash = user.PasswordHash
|
||||
}
|
||||
if !auth.CheckPassword(storedHash, req.Password) {
|
||||
// The attempt was already recorded atomically up-front (F3); here
|
||||
// only decide the lockouts, at the same boundary as before: the
|
||||
// 10th in-window failure locks the key. Check is read-only, so
|
||||
// the reservation is not double-counted.
|
||||
if !limiter.Check(failKey, scaledAuthLimit(loginFailureThreshold)+1, loginFailureWindow) {
|
||||
limiter.Lockout(r.Context(), lockKey, loginLockoutDuration)
|
||||
}
|
||||
// BUG-110: per-username lockout on threshold.
|
||||
if !limiter.Check(userFailKey, loginUserFailureThreshold+1, loginUserFailureWindow) {
|
||||
limiter.Lockout(r.Context(), userLockKey, loginUserLockoutDuration)
|
||||
}
|
||||
slog.Info("login failed", "ip", ip, "username_len", len(req.Username))
|
||||
writeJSON(w, http.StatusUnauthorized, errorResponse{
|
||||
Error: "UNAUTHORIZED",
|
||||
Message: "invalid credentials",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Reset failure counters on success.
|
||||
limiter.Reset(r.Context(), failKey)
|
||||
limiter.Reset(r.Context(), userFailKey)
|
||||
|
||||
if auth.IsEffectivelyBanned(user) {
|
||||
slog.Warn("banned user login attempt", "username", user.Username, "user_id", user.ID, "ip", ip)
|
||||
db.WriteAudit(context.WithoutCancel(r.Context()), database, user.ID, "login_blocked_banned", "user", user.ID,
|
||||
@@ -502,6 +396,152 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
|
||||
}
|
||||
}
|
||||
|
||||
// loginReadRequest decodes and validates the login body, writing the rejection
|
||||
// response itself when the input cannot be used.
|
||||
func loginReadRequest(w http.ResponseWriter, r *http.Request) (loginRequest, bool) {
|
||||
var req loginRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "malformed request body",
|
||||
})
|
||||
return req, false
|
||||
}
|
||||
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
// Do NOT trim req.Password — passwords may intentionally contain
|
||||
// leading/trailing whitespace. Bcrypt handles arbitrary bytes.
|
||||
|
||||
if req.Username == "" || req.Password == "" {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "username and password are required",
|
||||
})
|
||||
return req, false
|
||||
}
|
||||
|
||||
// F: reject an over-long username before it is ever used to build a
|
||||
// RateLimiter map key below (unameKey, failKey, userFailKey, lockout
|
||||
// keys). Unlike registration, login has no account to validate
|
||||
// against yet, so nothing else bounds this value — an unauthenticated
|
||||
// caller could otherwise pin an arbitrarily large, body-sized string
|
||||
// as a retained key (Cleanup only evicts it after hours). Mirrors the
|
||||
// same 32-rune cap auth.ValidateUsername enforces at registration.
|
||||
if utf8.RuneCountInString(req.Username) > maxLoginUsernameLen {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "INVALID_INPUT",
|
||||
Message: "username is too long",
|
||||
})
|
||||
return req, false
|
||||
}
|
||||
return req, true
|
||||
}
|
||||
|
||||
// loginAuthenticate runs the lockout gates, the constant-time password compare
|
||||
// and the failure accounting for one login attempt. It returns the
|
||||
// authenticated user, or false after writing the rejection response itself.
|
||||
func loginAuthenticate(w http.ResponseWriter, r *http.Request, database *db.DB, limiter *auth.RateLimiter, req loginRequest, ip string) (*db.User, bool) {
|
||||
// Check per-IP lockout first.
|
||||
lockKey := "login_lock:" + ip
|
||||
if limiter.IsLockedOut(lockKey) {
|
||||
writeJSON(w, http.StatusTooManyRequests, errorResponse{
|
||||
Error: "RATE_LIMITED",
|
||||
Message: "account temporarily locked due to too many failed attempts",
|
||||
})
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// BUG-110: Also check per-username lockout to prevent distributed brute force.
|
||||
// F1: canonicalize the username the same way GetUserByUsername does (COLLATE
|
||||
// NOCASE) before keying the lockout, so case variants of one account
|
||||
// (admin/Admin/ADMIN) share a single bucket instead of each getting its own.
|
||||
unameKey := strings.ToLower(req.Username)
|
||||
userLockKey := "login_user_lock:" + unameKey
|
||||
if limiter.IsLockedOut(userLockKey) {
|
||||
writeJSON(w, http.StatusTooManyRequests, errorResponse{
|
||||
Error: "RATE_LIMITED",
|
||||
Message: "account temporarily locked due to too many failed attempts",
|
||||
})
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// Constant-time lookup: always attempt bcrypt compare even when user
|
||||
// does not exist to prevent timing-based username enumeration.
|
||||
user, err := database.GetUserByUsername(r.Context(), req.Username)
|
||||
|
||||
// Distinguish DB errors from authentication failures. DB errors
|
||||
// should NOT increment the rate limiter — otherwise a transient
|
||||
// DB outage would lock out legitimate users.
|
||||
if err != nil && user == nil {
|
||||
// Could be a real DB error or simply "user not found".
|
||||
// GetUserByUsername returns (nil, nil) for not-found, so a
|
||||
// non-nil error here is a genuine DB failure.
|
||||
slog.Error("login: GetUserByUsername failed", "err", err, "ip", ip)
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "login temporarily unavailable",
|
||||
})
|
||||
return nil, false
|
||||
}
|
||||
|
||||
failKey := "login_fail:" + ip
|
||||
userFailKey := "login_user_fail:" + unameKey
|
||||
// F3: atomically reserve this attempt BEFORE the bcrypt compare. The
|
||||
// read-only IsLockedOut gates above are check-then-act: N concurrent
|
||||
// requests all pass them before any failure is recorded below, so the
|
||||
// per-username cap — the only cross-IP brute-force defence — bound
|
||||
// only sequential attackers. Allow records the attempt under the
|
||||
// limiter's lock, capping a concurrent burst at the same budget a
|
||||
// sequential attacker gets. Sized at threshold+1 so the sequential
|
||||
// accepted-input set is unchanged: failures 1–10 still land, the 10th
|
||||
// still trips the lockout (via the Check below), and a correct
|
||||
// password on attempt 10 still succeeds — successful logins reset
|
||||
// both counters. The reservation sits after the DB-error return above
|
||||
// so a transient DB outage still does not consume attempts.
|
||||
if !limiter.Allow(failKey, scaledAuthLimit(loginFailureThreshold)+1, loginFailureWindow) ||
|
||||
!limiter.Allow(userFailKey, loginUserFailureThreshold+1, loginUserFailureWindow) {
|
||||
writeJSON(w, http.StatusTooManyRequests, errorResponse{
|
||||
Error: "RATE_LIMITED",
|
||||
Message: "account temporarily locked due to too many failed attempts",
|
||||
})
|
||||
return nil, false
|
||||
}
|
||||
// Always run the password check — with an empty hash when the user does
|
||||
// not exist. auth.CheckPassword performs a dummy bcrypt comparison for an
|
||||
// empty hash, so bcrypt executes on every path and response time stays
|
||||
// constant, preventing timing-based username enumeration. (A `user == nil
|
||||
// || CheckPassword(...)` short-circuit would skip bcrypt entirely for
|
||||
// unknown usernames, reintroducing the timing side-channel.)
|
||||
storedHash := ""
|
||||
if user != nil {
|
||||
storedHash = user.PasswordHash
|
||||
}
|
||||
if !auth.CheckPassword(storedHash, req.Password) {
|
||||
// The attempt was already recorded atomically up-front (F3); here
|
||||
// only decide the lockouts, at the same boundary as before: the
|
||||
// 10th in-window failure locks the key. Check is read-only, so
|
||||
// the reservation is not double-counted.
|
||||
if !limiter.Check(failKey, scaledAuthLimit(loginFailureThreshold)+1, loginFailureWindow) {
|
||||
limiter.Lockout(r.Context(), lockKey, loginLockoutDuration)
|
||||
}
|
||||
// BUG-110: per-username lockout on threshold.
|
||||
if !limiter.Check(userFailKey, loginUserFailureThreshold+1, loginUserFailureWindow) {
|
||||
limiter.Lockout(r.Context(), userLockKey, loginUserLockoutDuration)
|
||||
}
|
||||
slog.Info("login failed", "ip", ip, "username_len", len(req.Username))
|
||||
writeJSON(w, http.StatusUnauthorized, errorResponse{
|
||||
Error: "UNAUTHORIZED",
|
||||
Message: "invalid credentials",
|
||||
})
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// Reset failure counters on success.
|
||||
limiter.Reset(r.Context(), failKey)
|
||||
limiter.Reset(r.Context(), userFailKey)
|
||||
return user, true
|
||||
}
|
||||
|
||||
// handleLogout processes POST /api/v1/auth/logout.
|
||||
func handleLogout(database *db.DB) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
+63
-50
@@ -141,56 +141,8 @@ func handleCreateEmoji(svc *service.Services, store FileStore, limiter *auth.Rat
|
||||
return
|
||||
}
|
||||
|
||||
file, _, err := r.FormFile("file")
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "missing file field",
|
||||
})
|
||||
return
|
||||
}
|
||||
defer file.Close() //nolint:errcheck
|
||||
|
||||
// Read at most one byte past the cap so "exactly at the limit" passes
|
||||
// and "one byte over" is caught, without buffering an unbounded body.
|
||||
raw, err := io.ReadAll(io.LimitReader(file, maxEmojiFileBytes+1))
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "failed to read uploaded file",
|
||||
})
|
||||
return
|
||||
}
|
||||
if int64(len(raw)) > maxEmojiFileBytes {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: fmt.Sprintf("emoji must be at most %d KB", maxEmojiFileBytes>>10),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
mimeType := http.DetectContentType(raw)
|
||||
if !allowedEmojiMIME[mimeType] {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: "emoji must be a PNG, JPEG, GIF or WebP image",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
width, height, err := imageDimensions(raw, mimeType)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "could not read image dimensions",
|
||||
})
|
||||
return
|
||||
}
|
||||
// Re-check the sniffed dimensions rather than trusting anything the
|
||||
// client said about the image: the cap is what keeps an "emoji" from
|
||||
// being a full-size picture inlined into every message that names it.
|
||||
if width <= 0 || height <= 0 || width > maxEmojiDimension || height > maxEmojiDimension {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: fmt.Sprintf("emoji must be at most %dx%d pixels (got %dx%d)", maxEmojiDimension, maxEmojiDimension, width, height),
|
||||
})
|
||||
raw, mimeType, readOK := readEmojiUpload(w, r)
|
||||
if !readOK {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -217,6 +169,67 @@ func handleCreateEmoji(svc *service.Services, store FileStore, limiter *auth.Rat
|
||||
}
|
||||
}
|
||||
|
||||
// readEmojiUpload pulls the uploaded file out of the already-parsed multipart
|
||||
// form and enforces every property of the bytes themselves: the size cap, the
|
||||
// sniffed MIME type and the sniffed pixel dimensions. It writes the refusal
|
||||
// itself, so a false third result means the response is already complete.
|
||||
func readEmojiUpload(w http.ResponseWriter, r *http.Request) ([]byte, string, bool) {
|
||||
file, _, err := r.FormFile("file")
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "missing file field",
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
defer file.Close() //nolint:errcheck
|
||||
|
||||
// Read at most one byte past the cap so "exactly at the limit" passes
|
||||
// and "one byte over" is caught, without buffering an unbounded body.
|
||||
raw, err := io.ReadAll(io.LimitReader(file, maxEmojiFileBytes+1))
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "failed to read uploaded file",
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
if int64(len(raw)) > maxEmojiFileBytes {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: fmt.Sprintf("emoji must be at most %d KB", maxEmojiFileBytes>>10),
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
|
||||
mimeType := http.DetectContentType(raw)
|
||||
if !allowedEmojiMIME[mimeType] {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: "emoji must be a PNG, JPEG, GIF or WebP image",
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
|
||||
width, height, err := imageDimensions(raw, mimeType)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "could not read image dimensions",
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
// Re-check the sniffed dimensions rather than trusting anything the
|
||||
// client said about the image: the cap is what keeps an "emoji" from
|
||||
// being a full-size picture inlined into every message that names it.
|
||||
if width <= 0 || height <= 0 || width > maxEmojiDimension || height > maxEmojiDimension {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: fmt.Sprintf("emoji must be at most %dx%d pixels (got %dx%d)", maxEmojiDimension, maxEmojiDimension, width, height),
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
|
||||
return raw, mimeType, true
|
||||
}
|
||||
|
||||
func handleDeleteEmoji(svc *service.Services, store FileStore, broadcaster EmojiBroadcaster) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
user, ok := r.Context().Value(UserKey).(*db.User)
|
||||
|
||||
@@ -508,49 +508,8 @@ func handleUploadAvatar(
|
||||
}
|
||||
defer file.Close() //nolint:errcheck
|
||||
|
||||
// Read one byte past the cap so "exactly at the limit" passes and "one
|
||||
// byte over" is caught, without buffering an unbounded body.
|
||||
raw, err := io.ReadAll(io.LimitReader(file, maxAvatarFileBytes+1))
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "failed to read uploaded file",
|
||||
})
|
||||
return
|
||||
}
|
||||
if int64(len(raw)) > maxAvatarFileBytes {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: fmt.Sprintf("avatar must be at most %d KB", maxAvatarFileBytes>>10),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Never trust the client's Content-Type — sniff the bytes.
|
||||
mimeType := http.DetectContentType(raw)
|
||||
if !allowedAvatarMIME[mimeType] {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "avatar must be a PNG, JPEG or WebP image",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
width, height, err := imageDimensions(raw, mimeType)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "could not read image dimensions",
|
||||
})
|
||||
return
|
||||
}
|
||||
// Measured from the sniffed image, not from anything the client said.
|
||||
// The client crops to a square before uploading; the server does not
|
||||
// re-encode (that would mean decoding and re-compressing every upload
|
||||
// to change nothing a CSS circle mask does not already do), it just
|
||||
// refuses a picture too big to be an avatar.
|
||||
if width <= 0 || height <= 0 || width > maxAvatarDimension || height > maxAvatarDimension {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: fmt.Sprintf("avatar must be at most %dx%d pixels (got %dx%d)", maxAvatarDimension, maxAvatarDimension, width, height),
|
||||
})
|
||||
raw, mimeType, width, height, ok := avatarUploadReadImage(w, file)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -612,3 +571,58 @@ func handleUploadAvatar(
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// avatarUploadReadImage is the bytes stage of handleUploadAvatar: read the
|
||||
// uploaded file under its cap, sniff its type and measure it. It writes its own
|
||||
// 400 and reports ok=false when the upload is not an acceptable avatar, so the
|
||||
// caller only has to return. Deliberately not shared with the emoji route: the
|
||||
// two carry different caps and a different allowed MIME set.
|
||||
func avatarUploadReadImage(w http.ResponseWriter, file io.Reader) (raw []byte, mimeType string, width, height int, ok bool) {
|
||||
// Read one byte past the cap so "exactly at the limit" passes and "one
|
||||
// byte over" is caught, without buffering an unbounded body.
|
||||
raw, err := io.ReadAll(io.LimitReader(file, maxAvatarFileBytes+1))
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "failed to read uploaded file",
|
||||
})
|
||||
return nil, "", 0, 0, false
|
||||
}
|
||||
if int64(len(raw)) > maxAvatarFileBytes {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: fmt.Sprintf("avatar must be at most %d KB", maxAvatarFileBytes>>10),
|
||||
})
|
||||
return nil, "", 0, 0, false
|
||||
}
|
||||
|
||||
// Never trust the client's Content-Type — sniff the bytes.
|
||||
mimeType = http.DetectContentType(raw)
|
||||
if !allowedAvatarMIME[mimeType] {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "avatar must be a PNG, JPEG or WebP image",
|
||||
})
|
||||
return nil, "", 0, 0, false
|
||||
}
|
||||
|
||||
width, height, err = imageDimensions(raw, mimeType)
|
||||
if err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST", Message: "could not read image dimensions",
|
||||
})
|
||||
return nil, "", 0, 0, false
|
||||
}
|
||||
// Measured from the sniffed image, not from anything the client said.
|
||||
// The client crops to a square before uploading; the server does not
|
||||
// re-encode (that would mean decoding and re-compressing every upload
|
||||
// to change nothing a CSS circle mask does not already do), it just
|
||||
// refuses a picture too big to be an avatar.
|
||||
if width <= 0 || height <= 0 || width > maxAvatarDimension || height > maxAvatarDimension {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "BAD_REQUEST",
|
||||
Message: fmt.Sprintf("avatar must be at most %dx%d pixels (got %dx%d)", maxAvatarDimension, maxAvatarDimension, width, height),
|
||||
})
|
||||
return nil, "", 0, 0, false
|
||||
}
|
||||
|
||||
return raw, mimeType, width, height, true
|
||||
}
|
||||
|
||||
+217
-166
@@ -44,52 +44,11 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri
|
||||
// (M1). Done first, before any other setup, so a fatal failure here
|
||||
// (below) doesn't leave background goroutines or partially-mounted
|
||||
// routes behind.
|
||||
totpKey, totpKeyErr := auth.LoadOrGenerateTOTPKey(cfg.Server.DataDir)
|
||||
if totpKeyErr != nil {
|
||||
if cfg.Server.DataDir != "" {
|
||||
// A configured data directory means this is a real deployment —
|
||||
// main.go creates cfg.Server.DataDir before calling NewRouter, so
|
||||
// by this point LoadOrGenerateTOTPKey only fails for a malformed
|
||||
// OWNCORD_TOTP_KEY or a corrupt/truncated totp.key file, never for
|
||||
// a missing directory. (The zero-value "" DataDir used by handler
|
||||
// tests that never touch TOTP crypto is exempted below so the
|
||||
// existing test suite keeps passing.)
|
||||
//
|
||||
// Continuing here would leave totpKey nil: every AES call in
|
||||
// auth.EncryptTOTPSecret/DecryptTOTPSecret then hits
|
||||
// aes.NewCipher(nil) and 500s, so every 2FA-enabled account
|
||||
// (including the owner) would be locked out of login and unable
|
||||
// to re-enroll, forever, while /health kept reporting OK. Refuse
|
||||
// to start instead.
|
||||
panic(fmt.Sprintf("api: failed to load TOTP encryption key: %v", totpKeyErr))
|
||||
}
|
||||
slog.Error("failed to load TOTP encryption key", "error", totpKeyErr)
|
||||
// Fall through — only reachable when DataDir is unset; TOTP handlers
|
||||
// cannot encrypt/decrypt until a data directory is configured.
|
||||
}
|
||||
totpKey := routerTOTPKey(cfg)
|
||||
|
||||
r := chi.NewRouter()
|
||||
|
||||
// Middleware stack.
|
||||
r.Use(boundRequestID) // must precede RequestID — it reads the header verbatim
|
||||
r.Use(middleware.RequestID)
|
||||
r.Use(setRequestIDHeader) // echo request ID into response header
|
||||
// NOTE: middleware.RealIP is intentionally omitted — trusting X-Real-IP from
|
||||
// any source allows IP spoofing for rate-limit bypass. IP header trust is now
|
||||
// handled explicitly in clientIPWithProxies using the trusted_proxies config.
|
||||
r.Use(recoverer) // slog-routing panic recovery (replaces chi's stderr-only Recoverer)
|
||||
r.Use(requestLogger) // structured request/response logging
|
||||
// Phase B Step 8 — OpenTelemetry HTTP tracing. No-op when telemetry is
|
||||
// disabled or the otel build tag is not set, so this is safe to mount
|
||||
// unconditionally.
|
||||
r.Use(telemetry.HTTPMiddleware())
|
||||
r.Use(SecurityHeadersWithTLS(cfg.TLS.Mode))
|
||||
r.Use(MaxBodySizeUnless(defaultMaxBodySize, bodyCapExemptPrefixes...))
|
||||
|
||||
// Coraza WAF — opt-in via config.
|
||||
if cfg.Server.WAFEnabled {
|
||||
r.Use(NewWAFMiddlewareCRS(cfg.Server.WAFParanoiaLevel, cfg.Server.WAFCRSMode))
|
||||
}
|
||||
routerMiddleware(r, cfg)
|
||||
|
||||
// Health check — unauthenticated, no versioning prefix.
|
||||
// The hub-backed callbacks are set after hub creation below (late-bound
|
||||
@@ -97,33 +56,7 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri
|
||||
// handler instance backs both /health mounts so they share the check cache.
|
||||
var getOnlineUsers func() int
|
||||
var hubAlive func() bool
|
||||
healthHandler := handleHealth(healthDeps{
|
||||
onlineUsers: func() int {
|
||||
if getOnlineUsers != nil {
|
||||
return getOnlineUsers()
|
||||
}
|
||||
return 0
|
||||
},
|
||||
dbPing: func(ctx context.Context) error {
|
||||
if database == nil {
|
||||
return nil
|
||||
}
|
||||
// Reader pool, not the writer: a scheduled backup's VACUUM INTO
|
||||
// holds the sole writer connection for its whole duration, and
|
||||
// the server keeps serving reads throughout — /health must not
|
||||
// call that outage (see db.PingRead).
|
||||
return database.PingRead(ctx)
|
||||
},
|
||||
dispatchAlive: func() bool {
|
||||
if hubAlive != nil {
|
||||
return hubAlive()
|
||||
}
|
||||
return true
|
||||
},
|
||||
freeDiskBytes: func() (uint64, error) {
|
||||
return diskutil.FreeBytes(cfg.Server.DataDir)
|
||||
},
|
||||
})
|
||||
healthHandler := handleHealth(routerHealthDeps(cfg, database, &getOnlineUsers, &hubAlive))
|
||||
r.Get("/health", healthHandler)
|
||||
|
||||
// Shared rate limiter for auth endpoints. Lockouts are persisted to the
|
||||
@@ -167,18 +100,7 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri
|
||||
// be passed as a DMBroadcaster for real-time close events.
|
||||
|
||||
// File upload and serving routes.
|
||||
// L12: verify config upload size fits within the HTTP body limit.
|
||||
if int64(cfg.Upload.MaxSizeMB)<<20 > uploadMaxBodySize {
|
||||
slog.Warn("upload.max_size_mb exceeds HTTP body limit, capping",
|
||||
"configured_mb", cfg.Upload.MaxSizeMB,
|
||||
"http_limit_bytes", uploadMaxBodySize)
|
||||
}
|
||||
store, storeErr := storage.New(cfg.Upload.StorageDir, cfg.Upload.MaxSizeMB)
|
||||
if storeErr != nil {
|
||||
slog.Error("failed to create file storage", "error", storeErr)
|
||||
} else {
|
||||
MountUploadRoutes(r, database, store, limiter, cfg.Server.AllowedOrigins, svc.Permissions)
|
||||
}
|
||||
store, storeErr := routerUploadRoutes(r, database, limiter, cfg, svc.Permissions)
|
||||
|
||||
// WebSocket hub — WS does its own in-band auth, so no AuthMiddleware here.
|
||||
hub := ws.NewHub(database, limiter, svc)
|
||||
@@ -194,6 +116,209 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri
|
||||
// anonymise-and-ban DB state.
|
||||
MountAuthRoutes(r, database, limiter, cfg.Server.TrustedProxies, totpKey, hub)
|
||||
|
||||
routerPluginWiring(hub, pluginRegistry)
|
||||
|
||||
// Voice: LiveKit client, optional companion process, webhook and proxy routes.
|
||||
routerVoiceRoutes(r, cfg, limiter, hub)
|
||||
|
||||
// Profile routes: update profile, change password, session management.
|
||||
// Mounted after hub creation so the hub can broadcast user_update events.
|
||||
// A storage failure leaves store unusable, so the avatar-upload route is
|
||||
// simply not registered; the rest of the profile surface is unaffected.
|
||||
// Built as a FileStore interface value from scratch — assigning the typed
|
||||
// nil pointer would produce a non-nil interface and defeat the mount-time
|
||||
// nil check.
|
||||
var profileStore FileStore
|
||||
if storeErr == nil {
|
||||
profileStore = store
|
||||
}
|
||||
MountProfileRoutes(r, database, svc, profileStore, limiter, cfg.Server.TrustedProxies, hub)
|
||||
|
||||
// DM (direct message) REST routes — mounted after hub creation so the
|
||||
// hub can send real-time dm_channel_close events to WebSocket clients.
|
||||
MountDMRoutes(r, database, svc, hub)
|
||||
|
||||
// Channel and message REST routes — mounted after hub creation so a
|
||||
// message purge can broadcast chat_bulk_deleted to the channel.
|
||||
MountChannelRoutes(r, database, svc, limiter, cfg.Server.TrustedProxies, hub)
|
||||
|
||||
// Custom emoji REST routes — mounted after hub creation so an upload or a
|
||||
// delete can fan the new set out as an emoji_update. Requires the same file
|
||||
// storage the attachment routes use; without it the emoji endpoints are not
|
||||
// mounted at all (a 404 the client reads as "this server has no emoji").
|
||||
if storeErr == nil {
|
||||
MountEmojiRoutes(r, database, svc, store, limiter, hub)
|
||||
}
|
||||
|
||||
// H-8: Connectivity diagnostics restricted to admin users only.
|
||||
// Exposes Go runtime version and LiveKit node IP which aid targeted attacks.
|
||||
r.With(AuthMiddleware(database),
|
||||
RequirePermission(permissions.Administrator),
|
||||
RateLimitMiddleware(limiter, "diag:", 5, time.Minute, cfg.Server.TrustedProxies)).
|
||||
Get("/api/v1/diagnostics/connectivity",
|
||||
handleDiagnosticsConnectivity(cfg, ver, hub))
|
||||
|
||||
go hub.Run()
|
||||
r.Get("/api/v1/ws", ws.ServeWS(hub, database, cfg.Server.AllowedOrigins, cfg.Server.MaxWSConnections))
|
||||
|
||||
routerMetricsRoutes(r, cfg, database, svc, hub)
|
||||
|
||||
// Admin panel: static files + REST API (Phase 6).
|
||||
// Restrict /admin to configured CIDRs (default: private networks only).
|
||||
u := updater.NewUpdater(ver, cfg.GitHub.Token, cfg.GitHub.Owner, cfg.GitHub.Repo)
|
||||
adminHandler := admin.NewHandler(database, ver, hub, u, logBuf, cfg.Server.AllowedOrigins, svc.Permissions, svc.Moderation, svc.Roles,
|
||||
admin.SetupOptions{ConfigPath: config.DefaultPath, RunningCfg: cfg})
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(AdminIPRestrict(cfg.Server.AdminAllowedCIDRs, cfg.Server.TrustedProxies))
|
||||
r.Mount("/admin", adminHandler)
|
||||
|
||||
// Phase C Step 9 — plugin admin REST surface. The IP gate above is
|
||||
// only the outer perimeter; plugin lifecycle endpoints additionally
|
||||
// require a valid admin Bearer token via admin.RequireAdminAuth so a
|
||||
// LAN attacker on the allowed CIDR cannot install/enable plugins
|
||||
// without a session. The handler is wired with the live registry
|
||||
// constructed in main.go (nil when plugin support is disabled, in
|
||||
// which case lifecycle calls return 503 and list returns []).
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(admin.RequireAdminAuth(database))
|
||||
r.Mount("/api/v1/admin/plugins", NewPluginAdminHandler(pluginRegistry, database))
|
||||
})
|
||||
})
|
||||
|
||||
// Client auto-update endpoint (unauthenticated). Per-IP rate limited to
|
||||
// bound abuse; the signature fetch is cached inside the updater (DoS fix).
|
||||
// Dedicated key prefix (mirroring "livekit_proxy:"): the empty-prefix
|
||||
// middleware would share per-IP buckets with verify-totp, password change,
|
||||
// and the other sensitive endpoints, so a client's 30/min auto-poll could
|
||||
// 429 its user's own 2FA or password change.
|
||||
MountClientUpdateRoute(
|
||||
r.With(rateLimitMiddlewareWithPrefix(limiter, "client_update:", clientUpdateRateLimitPerMinute, time.Minute, cfg.Server.TrustedProxies)),
|
||||
u,
|
||||
)
|
||||
|
||||
// Issue 15: Warn if AllowedOrigins contains wildcard.
|
||||
if slices.Contains(cfg.Server.AllowedOrigins, "*") {
|
||||
slog.Warn("AllowedOrigins contains wildcard '*' — consider restricting to specific origins for production use")
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
close(limiterStopCh)
|
||||
}
|
||||
|
||||
return r, hub, cleanup
|
||||
}
|
||||
|
||||
// routerTOTPKey loads (or auto-generates) the AES-256 key NewRouter hands to the
|
||||
// auth routes for TOTP secret encryption (M1).
|
||||
func routerTOTPKey(cfg *config.Config) []byte {
|
||||
totpKey, totpKeyErr := auth.LoadOrGenerateTOTPKey(cfg.Server.DataDir)
|
||||
if totpKeyErr != nil {
|
||||
if cfg.Server.DataDir != "" {
|
||||
// A configured data directory means this is a real deployment —
|
||||
// main.go creates cfg.Server.DataDir before calling NewRouter, so
|
||||
// by this point LoadOrGenerateTOTPKey only fails for a malformed
|
||||
// OWNCORD_TOTP_KEY or a corrupt/truncated totp.key file, never for
|
||||
// a missing directory. (The zero-value "" DataDir used by handler
|
||||
// tests that never touch TOTP crypto is exempted below so the
|
||||
// existing test suite keeps passing.)
|
||||
//
|
||||
// Continuing here would leave totpKey nil: every AES call in
|
||||
// auth.EncryptTOTPSecret/DecryptTOTPSecret then hits
|
||||
// aes.NewCipher(nil) and 500s, so every 2FA-enabled account
|
||||
// (including the owner) would be locked out of login and unable
|
||||
// to re-enroll, forever, while /health kept reporting OK. Refuse
|
||||
// to start instead.
|
||||
panic(fmt.Sprintf("api: failed to load TOTP encryption key: %v", totpKeyErr))
|
||||
}
|
||||
slog.Error("failed to load TOTP encryption key", "error", totpKeyErr)
|
||||
// Fall through — only reachable when DataDir is unset; TOTP handlers
|
||||
// cannot encrypt/decrypt until a data directory is configured.
|
||||
}
|
||||
return totpKey
|
||||
}
|
||||
|
||||
// routerHealthDeps builds the liveness probes behind the shared /health
|
||||
// handler. getOnlineUsers and hubAlive are taken as pointers because NewRouter
|
||||
// only assigns them once the hub exists, after this handler is already mounted;
|
||||
// the closures read whatever the variables hold at request time.
|
||||
func routerHealthDeps(cfg *config.Config, database *db.DB, getOnlineUsers *func() int, hubAlive *func() bool) healthDeps {
|
||||
return healthDeps{
|
||||
onlineUsers: func() int {
|
||||
if *getOnlineUsers != nil {
|
||||
return (*getOnlineUsers)()
|
||||
}
|
||||
return 0
|
||||
},
|
||||
dbPing: func(ctx context.Context) error {
|
||||
if database == nil {
|
||||
return nil
|
||||
}
|
||||
// Reader pool, not the writer: a scheduled backup's VACUUM INTO
|
||||
// holds the sole writer connection for its whole duration, and
|
||||
// the server keeps serving reads throughout — /health must not
|
||||
// call that outage (see db.PingRead).
|
||||
return database.PingRead(ctx)
|
||||
},
|
||||
dispatchAlive: func() bool {
|
||||
if *hubAlive != nil {
|
||||
return (*hubAlive)()
|
||||
}
|
||||
return true
|
||||
},
|
||||
freeDiskBytes: func() (uint64, error) {
|
||||
return diskutil.FreeBytes(cfg.Server.DataDir)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// routerMiddleware installs NewRouter's global middleware stack. The order is a
|
||||
// security property (request-id binding before the logger reads it, security
|
||||
// headers and the body cap before any handler runs) — keep it exactly as
|
||||
// written.
|
||||
func routerMiddleware(r chi.Router, cfg *config.Config) {
|
||||
// Middleware stack.
|
||||
r.Use(boundRequestID) // must precede RequestID — it reads the header verbatim
|
||||
r.Use(middleware.RequestID)
|
||||
r.Use(setRequestIDHeader) // echo request ID into response header
|
||||
// NOTE: middleware.RealIP is intentionally omitted — trusting X-Real-IP from
|
||||
// any source allows IP spoofing for rate-limit bypass. IP header trust is now
|
||||
// handled explicitly in clientIPWithProxies using the trusted_proxies config.
|
||||
r.Use(recoverer) // slog-routing panic recovery (replaces chi's stderr-only Recoverer)
|
||||
r.Use(requestLogger) // structured request/response logging
|
||||
// Phase B Step 8 — OpenTelemetry HTTP tracing. No-op when telemetry is
|
||||
// disabled or the otel build tag is not set, so this is safe to mount
|
||||
// unconditionally.
|
||||
r.Use(telemetry.HTTPMiddleware())
|
||||
r.Use(SecurityHeadersWithTLS(cfg.TLS.Mode))
|
||||
r.Use(MaxBodySizeUnless(defaultMaxBodySize, bodyCapExemptPrefixes...))
|
||||
|
||||
// Coraza WAF — opt-in via config.
|
||||
if cfg.Server.WAFEnabled {
|
||||
r.Use(NewWAFMiddlewareCRS(cfg.Server.WAFParanoiaLevel, cfg.Server.WAFCRSMode))
|
||||
}
|
||||
}
|
||||
|
||||
// routerUploadRoutes mounts the file upload and serving routes and returns the
|
||||
// shared file storage (and its construction error) for the profile-avatar and
|
||||
// emoji mounts, which reuse the same store.
|
||||
func routerUploadRoutes(r chi.Router, database *db.DB, limiter *auth.RateLimiter, cfg *config.Config, permSvc *service.PermissionService) (*storage.Storage, error) {
|
||||
// L12: verify config upload size fits within the HTTP body limit.
|
||||
if int64(cfg.Upload.MaxSizeMB)<<20 > uploadMaxBodySize {
|
||||
slog.Warn("upload.max_size_mb exceeds HTTP body limit, capping",
|
||||
"configured_mb", cfg.Upload.MaxSizeMB,
|
||||
"http_limit_bytes", uploadMaxBodySize)
|
||||
}
|
||||
store, storeErr := storage.New(cfg.Upload.StorageDir, cfg.Upload.MaxSizeMB)
|
||||
if storeErr != nil {
|
||||
slog.Error("failed to create file storage", "error", storeErr)
|
||||
} else {
|
||||
MountUploadRoutes(r, database, store, limiter, cfg.Server.AllowedOrigins, permSvc)
|
||||
}
|
||||
return store, storeErr
|
||||
}
|
||||
|
||||
// routerPluginWiring wires the plugin registry and its event sink into the hub.
|
||||
func routerPluginWiring(hub *ws.Hub, pluginRegistry *plugin.Registry) {
|
||||
// Phase C Step 9 — wire plugin registry and event sink into the hub.
|
||||
// nil pluginRegistry means plugins are disabled; the hub no-ops cleanly.
|
||||
if pluginRegistry != nil {
|
||||
@@ -202,7 +327,13 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri
|
||||
sink.SetBroadcaster(hub.BroadcastToChannel)
|
||||
hub.SetPluginEventSink(sink)
|
||||
}
|
||||
}
|
||||
|
||||
// routerVoiceRoutes creates the LiveKit client, optionally starts the companion
|
||||
// LiveKit process, and mounts the webhook, LiveKit health and signaling-proxy
|
||||
// routes. Voice is disabled — and none of those routes are mounted — when the
|
||||
// client fails to build.
|
||||
func routerVoiceRoutes(r chi.Router, cfg *config.Config, limiter *auth.RateLimiter, hub *ws.Hub) {
|
||||
// Create LiveKit client if voice config is present; voice is disabled on failure.
|
||||
lk, lkErr := ws.NewLiveKitClient(&cfg.Voice)
|
||||
if lkErr != nil {
|
||||
@@ -274,47 +405,11 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri
|
||||
r.With(rateLimitMiddlewareWithPrefix(limiter, "livekit_proxy:", livekitProxyRateLimitPerMinute, time.Minute, cfg.Server.TrustedProxies)).
|
||||
Handle("/livekit/*", http.StripPrefix("/livekit", NewLiveKitProxy(cfg.Voice.LiveKitURL, cfg.Server.AllowedOrigins)))
|
||||
}
|
||||
}
|
||||
|
||||
// Profile routes: update profile, change password, session management.
|
||||
// Mounted after hub creation so the hub can broadcast user_update events.
|
||||
// A storage failure leaves store unusable, so the avatar-upload route is
|
||||
// simply not registered; the rest of the profile surface is unaffected.
|
||||
// Built as a FileStore interface value from scratch — assigning the typed
|
||||
// nil pointer would produce a non-nil interface and defeat the mount-time
|
||||
// nil check.
|
||||
var profileStore FileStore
|
||||
if storeErr == nil {
|
||||
profileStore = store
|
||||
}
|
||||
MountProfileRoutes(r, database, svc, profileStore, limiter, cfg.Server.TrustedProxies, hub)
|
||||
|
||||
// DM (direct message) REST routes — mounted after hub creation so the
|
||||
// hub can send real-time dm_channel_close events to WebSocket clients.
|
||||
MountDMRoutes(r, database, svc, hub)
|
||||
|
||||
// Channel and message REST routes — mounted after hub creation so a
|
||||
// message purge can broadcast chat_bulk_deleted to the channel.
|
||||
MountChannelRoutes(r, database, svc, limiter, cfg.Server.TrustedProxies, hub)
|
||||
|
||||
// Custom emoji REST routes — mounted after hub creation so an upload or a
|
||||
// delete can fan the new set out as an emoji_update. Requires the same file
|
||||
// storage the attachment routes use; without it the emoji endpoints are not
|
||||
// mounted at all (a 404 the client reads as "this server has no emoji").
|
||||
if storeErr == nil {
|
||||
MountEmojiRoutes(r, database, svc, store, limiter, hub)
|
||||
}
|
||||
|
||||
// H-8: Connectivity diagnostics restricted to admin users only.
|
||||
// Exposes Go runtime version and LiveKit node IP which aid targeted attacks.
|
||||
r.With(AuthMiddleware(database),
|
||||
RequirePermission(permissions.Administrator),
|
||||
RateLimitMiddleware(limiter, "diag:", 5, time.Minute, cfg.Server.TrustedProxies)).
|
||||
Get("/api/v1/diagnostics/connectivity",
|
||||
handleDiagnosticsConnectivity(cfg, ver, hub))
|
||||
|
||||
go hub.Run()
|
||||
r.Get("/api/v1/ws", ws.ServeWS(hub, database, cfg.Server.AllowedOrigins, cfg.Server.MaxWSConnections))
|
||||
|
||||
// routerMetricsRoutes mounts the JSON metrics endpoint and, when an OTel
|
||||
// Prometheus exporter is wired, the Prometheus handler beside it.
|
||||
func routerMetricsRoutes(r chi.Router, cfg *config.Config, database *db.DB, svc *service.Services, hub *ws.Hub) {
|
||||
// Metrics endpoint — IP-restricted by metrics_allowed_cidrs (falls back to
|
||||
// admin_allowed_cidrs) so a central scraper can be admitted without
|
||||
// widening /admin. The shape is documented in docs/deployment.md — keep
|
||||
@@ -342,50 +437,6 @@ func NewRouter(cfg *config.Config, database *db.DB, ver string, logBuf *admin.Ri
|
||||
r.With(AdminIPRestrict(cfg.Server.MetricsCIDRs(), cfg.Server.TrustedProxies)).
|
||||
Mount("/metrics", promH)
|
||||
}
|
||||
|
||||
// Admin panel: static files + REST API (Phase 6).
|
||||
// Restrict /admin to configured CIDRs (default: private networks only).
|
||||
u := updater.NewUpdater(ver, cfg.GitHub.Token, cfg.GitHub.Owner, cfg.GitHub.Repo)
|
||||
adminHandler := admin.NewHandler(database, ver, hub, u, logBuf, cfg.Server.AllowedOrigins, svc.Permissions, svc.Moderation, svc.Roles,
|
||||
admin.SetupOptions{ConfigPath: config.DefaultPath, RunningCfg: cfg})
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(AdminIPRestrict(cfg.Server.AdminAllowedCIDRs, cfg.Server.TrustedProxies))
|
||||
r.Mount("/admin", adminHandler)
|
||||
|
||||
// Phase C Step 9 — plugin admin REST surface. The IP gate above is
|
||||
// only the outer perimeter; plugin lifecycle endpoints additionally
|
||||
// require a valid admin Bearer token via admin.RequireAdminAuth so a
|
||||
// LAN attacker on the allowed CIDR cannot install/enable plugins
|
||||
// without a session. The handler is wired with the live registry
|
||||
// constructed in main.go (nil when plugin support is disabled, in
|
||||
// which case lifecycle calls return 503 and list returns []).
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(admin.RequireAdminAuth(database))
|
||||
r.Mount("/api/v1/admin/plugins", NewPluginAdminHandler(pluginRegistry, database))
|
||||
})
|
||||
})
|
||||
|
||||
// Client auto-update endpoint (unauthenticated). Per-IP rate limited to
|
||||
// bound abuse; the signature fetch is cached inside the updater (DoS fix).
|
||||
// Dedicated key prefix (mirroring "livekit_proxy:"): the empty-prefix
|
||||
// middleware would share per-IP buckets with verify-totp, password change,
|
||||
// and the other sensitive endpoints, so a client's 30/min auto-poll could
|
||||
// 429 its user's own 2FA or password change.
|
||||
MountClientUpdateRoute(
|
||||
r.With(rateLimitMiddlewareWithPrefix(limiter, "client_update:", clientUpdateRateLimitPerMinute, time.Minute, cfg.Server.TrustedProxies)),
|
||||
u,
|
||||
)
|
||||
|
||||
// Issue 15: Warn if AllowedOrigins contains wildcard.
|
||||
if slices.Contains(cfg.Server.AllowedOrigins, "*") {
|
||||
slog.Warn("AllowedOrigins contains wildcard '*' — consider restricting to specific origins for production use")
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
close(limiterStopCh)
|
||||
}
|
||||
|
||||
return r, hub, cleanup
|
||||
}
|
||||
|
||||
// serverStartTime records when the process started; used for uptime in /health.
|
||||
|
||||
+39
-27
@@ -86,33 +86,8 @@ func handleVerifyTOTP(database *db.DB, partialStore *auth.PartialAuthStore, limi
|
||||
return
|
||||
}
|
||||
|
||||
user, err := database.GetUserByID(r.Context(), challenge.UserID)
|
||||
if err != nil || user == nil || user.TOTPSecret == nil {
|
||||
writeJSON(w, http.StatusUnauthorized, errorResponse{
|
||||
Error: "UNAUTHORIZED",
|
||||
Message: "invalid or expired two-factor challenge",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// A ban can land inside the partial-token window; the login path
|
||||
// refuses banned users right after the password compare, so the
|
||||
// second factor must refuse them too.
|
||||
if auth.IsEffectivelyBanned(user) {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "your account has been suspended",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
secret, decErr := auth.DecryptTOTPSecret(totpKey, *user.TOTPSecret)
|
||||
if decErr != nil {
|
||||
slog.Error("failed to decrypt TOTP secret", "user_id", user.ID, "error", decErr)
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "failed to verify two-factor code",
|
||||
})
|
||||
user, secret, ok := totpChallengeSecret(w, r, database, totpKey, challenge.UserID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -158,6 +133,43 @@ func handleVerifyTOTP(database *db.DB, partialStore *auth.PartialAuthStore, limi
|
||||
}
|
||||
}
|
||||
|
||||
// totpChallengeSecret resolves the user behind a partial-auth challenge and
|
||||
// returns their decrypted TOTP secret. It writes its own refusal, so a false
|
||||
// third result means the response is already complete.
|
||||
func totpChallengeSecret(w http.ResponseWriter, r *http.Request, database *db.DB, totpKey []byte, challengeUserID int64) (*db.User, string, bool) {
|
||||
user, err := database.GetUserByID(r.Context(), challengeUserID)
|
||||
if err != nil || user == nil || user.TOTPSecret == nil {
|
||||
writeJSON(w, http.StatusUnauthorized, errorResponse{
|
||||
Error: "UNAUTHORIZED",
|
||||
Message: "invalid or expired two-factor challenge",
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
|
||||
// A ban can land inside the partial-token window; the login path
|
||||
// refuses banned users right after the password compare, so the
|
||||
// second factor must refuse them too.
|
||||
if auth.IsEffectivelyBanned(user) {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "your account has been suspended",
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
|
||||
secret, decErr := auth.DecryptTOTPSecret(totpKey, *user.TOTPSecret)
|
||||
if decErr != nil {
|
||||
slog.Error("failed to decrypt TOTP secret", "user_id", user.ID, "error", decErr)
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "failed to verify two-factor code",
|
||||
})
|
||||
return nil, "", false
|
||||
}
|
||||
|
||||
return user, secret, true
|
||||
}
|
||||
|
||||
func handleEnableTOTP(pendingStore *auth.PendingTOTPStore, limiter *auth.RateLimiter) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
user, ok := r.Context().Value(UserKey).(*db.User)
|
||||
|
||||
+134
-112
@@ -286,122 +286,13 @@ func handleServeFile(database *db.DB, store FileStore, allowedOrigins []string,
|
||||
return
|
||||
}
|
||||
|
||||
user, _ := r.Context().Value(UserKey).(*db.User)
|
||||
role, _ := r.Context().Value(RoleKey).(*db.Role)
|
||||
|
||||
// Look up attachment metadata with channel context.
|
||||
aa, err := database.GetAttachmentWithChannel(r.Context(), fileID)
|
||||
if err != nil {
|
||||
slog.Error("failed to look up attachment", "id", fileID, "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "internal server error",
|
||||
})
|
||||
return
|
||||
}
|
||||
aa := serveFileResolve(w, r, database, fileID)
|
||||
if aa == nil {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
// A soft-deleted message's attachments must stop being servable the
|
||||
// moment the message is deleted — the client shows a tombstone, but
|
||||
// without this check the file stays reachable by URL forever (no
|
||||
// sweep can ever reclaim a linked row either, since the only reaper
|
||||
// requires message_id IS NULL). Checked before the ACL branch so it
|
||||
// also covers admins, matching the tombstone applying to everyone.
|
||||
//
|
||||
// Queried directly rather than through database.GetMessage: that
|
||||
// wrapper's SELECT list carries every message column, and the
|
||||
// `deleted` flag is the only one this check needs.
|
||||
if aa.MessageID != nil {
|
||||
var deleted bool
|
||||
deletedErr := database.QueryRowContext(r.Context(),
|
||||
`SELECT deleted FROM messages WHERE id = ?`, *aa.MessageID).Scan(&deleted)
|
||||
switch {
|
||||
case errors.Is(deletedErr, sql.ErrNoRows):
|
||||
// No message row — leave ACL to decide (unlinked-shaped by now).
|
||||
case deletedErr != nil:
|
||||
slog.Error("failed to look up message for attachment", "id", fileID, "error", deletedErr)
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "internal server error",
|
||||
})
|
||||
return
|
||||
case deleted:
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// ── Access control ──────────────────────────────────────────────
|
||||
isAdmin := role != nil && permissions.HasAdmin(role.Permissions)
|
||||
|
||||
// DM participation is required of everyone, including admins — this
|
||||
// matches every other DM read gate in the codebase (requireChannelRead,
|
||||
// PermissionService.RequireChannelAccess, checkSendPermission), none of
|
||||
// which have an admin bypass. Checked ahead of the `!isAdmin` block so
|
||||
// the admin bypass below cannot skip it.
|
||||
if aa.ChannelID != nil && aa.ChannelType == "dm" {
|
||||
if user == nil {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return
|
||||
}
|
||||
ok, dmErr := database.IsDMParticipant(r.Context(), user.ID, *aa.ChannelID)
|
||||
if dmErr != nil || !ok {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if !isAdmin {
|
||||
if aa.ChannelID == nil {
|
||||
// An unlinked attachment that some user's avatar points at is
|
||||
// readable by every authenticated user: an avatar has to be
|
||||
// visible to the people who see the messages it sits next to.
|
||||
// The check is by the exact URL the column stores, so the file
|
||||
// stops being public the instant the avatar is replaced.
|
||||
isAvatar, avatarErr := database.IsAvatarFileURL(r.Context(), service.AvatarFileURL(fileID))
|
||||
if avatarErr != nil {
|
||||
slog.Error("failed to check avatar file", "id", fileID, "error", avatarErr)
|
||||
}
|
||||
switch {
|
||||
case isAvatar:
|
||||
// Public while in use — fall through to serving.
|
||||
// Unlinked attachment — only the uploader may access.
|
||||
// M-2: Legacy rows (NULL uploader_id) are now denied rather than
|
||||
// served to any authenticated user.
|
||||
case aa.UploaderID == nil:
|
||||
slog.Warn("legacy attachment access denied (NULL uploader_id)", "id", fileID)
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return
|
||||
case user == nil || *aa.UploaderID != user.ID:
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return
|
||||
}
|
||||
} else if aa.ChannelType != "dm" {
|
||||
// Linked attachment in a guild channel — check channel
|
||||
// permissions. The DM case is handled unconditionally above.
|
||||
if user == nil || !permSvc.HasChannelPerm(r.Context(), user.ID, *aa.ChannelID, permissions.ReadMessages) {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
if !serveFileAuthorize(w, r, database, permSvc, aa, fileID) {
|
||||
return
|
||||
}
|
||||
|
||||
// Open file from storage.
|
||||
@@ -450,3 +341,134 @@ func handleServeFile(database *db.DB, store FileStore, allowedOrigins []string,
|
||||
http.ServeContent(w, r, aa.Filename, modTime, f)
|
||||
}
|
||||
}
|
||||
|
||||
// serveFileResolve looks up the attachment behind {id} and applies the checks
|
||||
// that make a file unservable regardless of who is asking. It returns nil once
|
||||
// it has written the response, so the caller only has to return.
|
||||
func serveFileResolve(w http.ResponseWriter, r *http.Request, database *db.DB, fileID string) *db.AttachmentAccess {
|
||||
// Look up attachment metadata with channel context.
|
||||
aa, err := database.GetAttachmentWithChannel(r.Context(), fileID)
|
||||
if err != nil {
|
||||
slog.Error("failed to look up attachment", "id", fileID, "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "internal server error",
|
||||
})
|
||||
return nil
|
||||
}
|
||||
if aa == nil {
|
||||
http.NotFound(w, r)
|
||||
return nil
|
||||
}
|
||||
|
||||
// A soft-deleted message's attachments must stop being servable the
|
||||
// moment the message is deleted — the client shows a tombstone, but
|
||||
// without this check the file stays reachable by URL forever (no
|
||||
// sweep can ever reclaim a linked row either, since the only reaper
|
||||
// requires message_id IS NULL). Checked before the ACL branch so it
|
||||
// also covers admins, matching the tombstone applying to everyone.
|
||||
//
|
||||
// Queried directly rather than through database.GetMessage: that
|
||||
// wrapper's SELECT list carries every message column, and the
|
||||
// `deleted` flag is the only one this check needs.
|
||||
if aa.MessageID != nil {
|
||||
var deleted bool
|
||||
deletedErr := database.QueryRowContext(r.Context(),
|
||||
`SELECT deleted FROM messages WHERE id = ?`, *aa.MessageID).Scan(&deleted)
|
||||
switch {
|
||||
case errors.Is(deletedErr, sql.ErrNoRows):
|
||||
// No message row — leave ACL to decide (unlinked-shaped by now).
|
||||
case deletedErr != nil:
|
||||
slog.Error("failed to look up message for attachment", "id", fileID, "error", deletedErr)
|
||||
writeJSON(w, http.StatusInternalServerError, errorResponse{
|
||||
Error: "INTERNAL_ERROR",
|
||||
Message: "internal server error",
|
||||
})
|
||||
return nil
|
||||
case deleted:
|
||||
http.NotFound(w, r)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return aa
|
||||
}
|
||||
|
||||
// serveFileAuthorize decides whether the caller may read aa. It returns false
|
||||
// once it has written the response, so the caller only has to return.
|
||||
func serveFileAuthorize(w http.ResponseWriter, r *http.Request, database *db.DB, permSvc *service.PermissionService, aa *db.AttachmentAccess, fileID string) bool {
|
||||
user, _ := r.Context().Value(UserKey).(*db.User)
|
||||
role, _ := r.Context().Value(RoleKey).(*db.Role)
|
||||
|
||||
// ── Access control ──────────────────────────────────────────────
|
||||
isAdmin := role != nil && permissions.HasAdmin(role.Permissions)
|
||||
|
||||
// DM participation is required of everyone, including admins — this
|
||||
// matches every other DM read gate in the codebase (requireChannelRead,
|
||||
// PermissionService.RequireChannelAccess, checkSendPermission), none of
|
||||
// which have an admin bypass. Checked ahead of the `!isAdmin` block so
|
||||
// the admin bypass below cannot skip it.
|
||||
if aa.ChannelID != nil && aa.ChannelType == "dm" {
|
||||
if user == nil {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return false
|
||||
}
|
||||
ok, dmErr := database.IsDMParticipant(r.Context(), user.ID, *aa.ChannelID)
|
||||
if dmErr != nil || !ok {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
if !isAdmin {
|
||||
if aa.ChannelID == nil {
|
||||
// An unlinked attachment that some user's avatar points at is
|
||||
// readable by every authenticated user: an avatar has to be
|
||||
// visible to the people who see the messages it sits next to.
|
||||
// The check is by the exact URL the column stores, so the file
|
||||
// stops being public the instant the avatar is replaced.
|
||||
isAvatar, avatarErr := database.IsAvatarFileURL(r.Context(), service.AvatarFileURL(fileID))
|
||||
if avatarErr != nil {
|
||||
slog.Error("failed to check avatar file", "id", fileID, "error", avatarErr)
|
||||
}
|
||||
switch {
|
||||
case isAvatar:
|
||||
// Public while in use — fall through to serving.
|
||||
// Unlinked attachment — only the uploader may access.
|
||||
// M-2: Legacy rows (NULL uploader_id) are now denied rather than
|
||||
// served to any authenticated user.
|
||||
case aa.UploaderID == nil:
|
||||
slog.Warn("legacy attachment access denied (NULL uploader_id)", "id", fileID)
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return false
|
||||
case user == nil || *aa.UploaderID != user.ID:
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return false
|
||||
}
|
||||
} else if aa.ChannelType != "dm" {
|
||||
// Linked attachment in a guild channel — check channel
|
||||
// permissions. The DM case is handled unconditionally above.
|
||||
if user == nil || !permSvc.HasChannelPerm(r.Context(), user.ID, *aa.ChannelID, permissions.ReadMessages) {
|
||||
writeJSON(w, http.StatusForbidden, errorResponse{
|
||||
Error: "FORBIDDEN",
|
||||
Message: "you do not have access to this file",
|
||||
})
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
+131
-84
@@ -205,16 +205,11 @@ func NewWAFMiddlewareCRS(paranoiaLevel int, crsMode string) func(http.Handler) h
|
||||
return newWAFMiddleware(paranoiaLevel, crsMode, nil)
|
||||
}
|
||||
|
||||
// newWAFMiddleware is the implementation behind NewWAFMiddlewareCRS.
|
||||
// onCRSMatch overrides the CRS match logger (used by tests to observe
|
||||
// detect-mode matches); nil means log via slog.
|
||||
func newWAFMiddleware(paranoiaLevel int, crsMode string, onCRSMatch func(types.MatchedRule)) func(http.Handler) http.Handler {
|
||||
if paranoiaLevel < 1 || paranoiaLevel > 4 {
|
||||
paranoiaLevel = 2
|
||||
}
|
||||
crsMode = normalizeCRSMode(crsMode)
|
||||
|
||||
waf, err := coraza.NewWAF(
|
||||
// wafInlineEngine builds the long-standing inline-rules Coraza engine used by
|
||||
// newWAFMiddleware. Its rules keep their exact, test-pinned blocking behavior
|
||||
// regardless of the CRS mode.
|
||||
func wafInlineEngine(paranoiaLevel int) (coraza.WAF, error) {
|
||||
return coraza.NewWAF(
|
||||
coraza.NewWAFConfig().
|
||||
WithDirectives(fmt.Sprintf(`
|
||||
SecRuleEngine On
|
||||
@@ -262,11 +257,12 @@ func newWAFMiddleware(paranoiaLevel int, crsMode string, onCRSMatch func(types.M
|
||||
SecRule REQUEST_URI "@beginsWith /api/v1/users/me/avatar" "id:900005,phase:1,pass,nolog,ctl:requestBodyAccess=Off"
|
||||
`, paranoiaLevel)),
|
||||
)
|
||||
if err != nil {
|
||||
slog.Error("waf: failed to create WAF engine, continuing without WAF", "error", err)
|
||||
return func(next http.Handler) http.Handler { return next }
|
||||
}
|
||||
}
|
||||
|
||||
// wafCRSEngine builds the OWASP CRS layer for newWAFMiddleware. It returns the
|
||||
// CRS engine (nil when the layer is off or failed to load) and whether
|
||||
// detect-mode match logging is aggregated per request.
|
||||
func wafCRSEngine(paranoiaLevel int, crsMode string, onCRSMatch func(types.MatchedRule)) (coraza.WAF, bool) {
|
||||
// OWASP CRS layer — a second engine so the inline rules above keep their
|
||||
// exact blocking behavior in every CRS mode. If the CRS fails to load the
|
||||
// server continues with the inline engine only (same failure philosophy
|
||||
@@ -300,6 +296,124 @@ func newWAFMiddleware(paranoiaLevel int, crsMode string, onCRSMatch func(types.M
|
||||
crsWAF = cw
|
||||
}
|
||||
}
|
||||
return crsWAF, aggregateCRSLog
|
||||
}
|
||||
|
||||
// wafInlineRequestHeaders feeds the connection, URI and request headers into
|
||||
// the inline engine and runs its phase 1. A non-nil interruption means the
|
||||
// request must be blocked.
|
||||
func wafInlineRequestHeaders(tx types.Transaction, r *http.Request) *types.Interruption {
|
||||
tx.ProcessConnection(r.RemoteAddr, 0, "", 0)
|
||||
tx.ProcessURI(r.URL.String(), r.Method, r.Proto)
|
||||
for name, values := range r.Header {
|
||||
for _, value := range values {
|
||||
tx.AddRequestHeader(name, value)
|
||||
}
|
||||
}
|
||||
return tx.ProcessRequestHeaders()
|
||||
}
|
||||
|
||||
// wafCRSRequestHeaders feeds the connection, URI and request headers into the
|
||||
// CRS engine and runs its phase 1. A non-nil interruption means the request
|
||||
// must be blocked.
|
||||
func wafCRSRequestHeaders(crsTx types.Transaction, r *http.Request) *types.Interruption {
|
||||
crsTx.ProcessConnection(r.RemoteAddr, 0, "", 0)
|
||||
crsTx.ProcessURI(r.URL.String(), r.Method, r.Proto)
|
||||
for name, values := range r.Header {
|
||||
for _, value := range values {
|
||||
crsTx.AddRequestHeader(name, value)
|
||||
}
|
||||
}
|
||||
// net/http promotes Host and Transfer-Encoding out of
|
||||
// r.Header; re-add them like the official coraza http
|
||||
// connector does, otherwise CRS rule 920280 ("Request
|
||||
// Missing a Host Header", anomaly score 5) fires on every
|
||||
// request. The inline engine is left as-is on purpose — its
|
||||
// rules never look at these headers and its behavior is
|
||||
// pinned by tests.
|
||||
if r.Host != "" {
|
||||
crsTx.AddRequestHeader("Host", r.Host)
|
||||
crsTx.SetServerName(r.Host)
|
||||
}
|
||||
for _, te := range r.TransferEncoding {
|
||||
crsTx.AddRequestHeader("Transfer-Encoding", te)
|
||||
}
|
||||
return crsTx.ProcessRequestHeaders()
|
||||
}
|
||||
|
||||
// wafFeedCRSBody mirrors the inline engine's buffered request body into the
|
||||
// CRS engine so the body is only read from the wire once. A non-nil
|
||||
// interruption means the request must be blocked.
|
||||
func wafFeedCRSBody(tx, crsTx types.Transaction) *types.Interruption {
|
||||
if reader, err := tx.RequestBodyReader(); err == nil && reader != nil {
|
||||
if it, _, err := crsTx.ReadRequestBodyFrom(reader); it != nil {
|
||||
return it
|
||||
} else if err != nil {
|
||||
slog.Debug("waf: error reading CRS request body", "error", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// wafInspectRequestBody buffers the request body through the inline engine,
|
||||
// runs its phase 2, mirrors the buffer into the CRS engine and hands the
|
||||
// buffered body to the downstream handler. A non-nil interruption means the
|
||||
// request must be blocked.
|
||||
func wafInspectRequestBody(r *http.Request, tx, crsTx types.Transaction) *types.Interruption {
|
||||
it, written, err := tx.ReadRequestBodyFrom(r.Body)
|
||||
if it != nil {
|
||||
return it
|
||||
} else if err != nil {
|
||||
slog.Debug("waf: error reading request body", "error", err)
|
||||
}
|
||||
|
||||
if it, err := tx.ProcessRequestBody(); it != nil {
|
||||
return it
|
||||
} else if err != nil {
|
||||
slog.Debug("waf: error processing request body", "error", err)
|
||||
}
|
||||
|
||||
// Feed the CRS engine from the inline engine's buffer so the
|
||||
// body is only read from the wire once. written == 0 means the
|
||||
// inline engine skipped buffering (requestBodyAccess turned
|
||||
// off for this route, e.g. uploads) — the CRS engine excludes
|
||||
// those routes too, so skip it as well and leave r.Body alone.
|
||||
if written > 0 {
|
||||
if crsTx != nil {
|
||||
if it := wafFeedCRSBody(tx, crsTx); it != nil {
|
||||
return it
|
||||
}
|
||||
}
|
||||
|
||||
// Replace body with buffered version so downstream handlers
|
||||
// can read it. Only done when the inline engine actually
|
||||
// buffered the body — replacing unconditionally would hand
|
||||
// routes with body inspection disabled (uploads) an empty
|
||||
// reader instead of the original stream.
|
||||
reader, err := tx.RequestBodyReader()
|
||||
if err == nil && reader != nil {
|
||||
r.Body = io.NopCloser(reader)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// newWAFMiddleware is the implementation behind NewWAFMiddlewareCRS.
|
||||
// onCRSMatch overrides the CRS match logger (used by tests to observe
|
||||
// detect-mode matches); nil means log via slog.
|
||||
func newWAFMiddleware(paranoiaLevel int, crsMode string, onCRSMatch func(types.MatchedRule)) func(http.Handler) http.Handler {
|
||||
if paranoiaLevel < 1 || paranoiaLevel > 4 {
|
||||
paranoiaLevel = 2
|
||||
}
|
||||
crsMode = normalizeCRSMode(crsMode)
|
||||
|
||||
waf, err := wafInlineEngine(paranoiaLevel)
|
||||
if err != nil {
|
||||
slog.Error("waf: failed to create WAF engine, continuing without WAF", "error", err)
|
||||
return func(next http.Handler) http.Handler { return next }
|
||||
}
|
||||
|
||||
crsWAF, aggregateCRSLog := wafCRSEngine(paranoiaLevel, crsMode, onCRSMatch)
|
||||
|
||||
slog.Info("waf: Coraza WAF enabled", "paranoia_level", paranoiaLevel, "crs_mode", crsMode)
|
||||
|
||||
@@ -331,15 +445,7 @@ func newWAFMiddleware(paranoiaLevel int, crsMode string, onCRSMatch func(types.M
|
||||
}
|
||||
|
||||
// Process request headers
|
||||
tx.ProcessConnection(r.RemoteAddr, 0, "", 0)
|
||||
tx.ProcessURI(r.URL.String(), r.Method, r.Proto)
|
||||
for name, values := range r.Header {
|
||||
for _, value := range values {
|
||||
tx.AddRequestHeader(name, value)
|
||||
}
|
||||
}
|
||||
|
||||
if it := tx.ProcessRequestHeaders(); it != nil {
|
||||
if it := wafInlineRequestHeaders(tx, r); it != nil {
|
||||
handleWAFInterruption(w, it)
|
||||
return
|
||||
}
|
||||
@@ -347,28 +453,7 @@ func newWAFMiddleware(paranoiaLevel int, crsMode string, onCRSMatch func(types.M
|
||||
// CRS phase 1. In detect mode the engine never interrupts, so the
|
||||
// returned interruption is only non-nil in block mode.
|
||||
if crsTx != nil {
|
||||
crsTx.ProcessConnection(r.RemoteAddr, 0, "", 0)
|
||||
crsTx.ProcessURI(r.URL.String(), r.Method, r.Proto)
|
||||
for name, values := range r.Header {
|
||||
for _, value := range values {
|
||||
crsTx.AddRequestHeader(name, value)
|
||||
}
|
||||
}
|
||||
// net/http promotes Host and Transfer-Encoding out of
|
||||
// r.Header; re-add them like the official coraza http
|
||||
// connector does, otherwise CRS rule 920280 ("Request
|
||||
// Missing a Host Header", anomaly score 5) fires on every
|
||||
// request. The inline engine is left as-is on purpose — its
|
||||
// rules never look at these headers and its behavior is
|
||||
// pinned by tests.
|
||||
if r.Host != "" {
|
||||
crsTx.AddRequestHeader("Host", r.Host)
|
||||
crsTx.SetServerName(r.Host)
|
||||
}
|
||||
for _, te := range r.TransferEncoding {
|
||||
crsTx.AddRequestHeader("Transfer-Encoding", te)
|
||||
}
|
||||
if it := crsTx.ProcessRequestHeaders(); it != nil {
|
||||
if it := wafCRSRequestHeaders(crsTx, r); it != nil {
|
||||
handleWAFInterruption(w, it)
|
||||
return
|
||||
}
|
||||
@@ -380,47 +465,9 @@ func newWAFMiddleware(paranoiaLevel int, crsMode string, onCRSMatch func(types.M
|
||||
// skipped for them. The read is bounded by SecRequestBodyLimit inside
|
||||
// Coraza. ContentLength == 0 (no body) still skips inspection.
|
||||
if r.Body != nil && r.ContentLength != 0 {
|
||||
it, written, err := tx.ReadRequestBodyFrom(r.Body)
|
||||
if it != nil {
|
||||
if it := wafInspectRequestBody(r, tx, crsTx); it != nil {
|
||||
handleWAFInterruption(w, it)
|
||||
return
|
||||
} else if err != nil {
|
||||
slog.Debug("waf: error reading request body", "error", err)
|
||||
}
|
||||
|
||||
if it, err := tx.ProcessRequestBody(); it != nil {
|
||||
handleWAFInterruption(w, it)
|
||||
return
|
||||
} else if err != nil {
|
||||
slog.Debug("waf: error processing request body", "error", err)
|
||||
}
|
||||
|
||||
// Feed the CRS engine from the inline engine's buffer so the
|
||||
// body is only read from the wire once. written == 0 means the
|
||||
// inline engine skipped buffering (requestBodyAccess turned
|
||||
// off for this route, e.g. uploads) — the CRS engine excludes
|
||||
// those routes too, so skip it as well and leave r.Body alone.
|
||||
if written > 0 {
|
||||
if crsTx != nil {
|
||||
if reader, err := tx.RequestBodyReader(); err == nil && reader != nil {
|
||||
if it, _, err := crsTx.ReadRequestBodyFrom(reader); it != nil {
|
||||
handleWAFInterruption(w, it)
|
||||
return
|
||||
} else if err != nil {
|
||||
slog.Debug("waf: error reading CRS request body", "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Replace body with buffered version so downstream handlers
|
||||
// can read it. Only done when the inline engine actually
|
||||
// buffered the body — replacing unconditionally would hand
|
||||
// routes with body inspection disabled (uploads) an empty
|
||||
// reader instead of the original stream.
|
||||
reader, err := tx.RequestBodyReader()
|
||||
if err == nil && reader != nil {
|
||||
r.Body = io.NopCloser(reader)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+171
-143
@@ -36,89 +36,13 @@ func (d *DB) DeleteAccount(ctx context.Context, userID int64) error {
|
||||
}
|
||||
defer tx.Rollback() //nolint:errcheck
|
||||
|
||||
// ── Guard: last admin/owner check ────────────────────────────────────
|
||||
// Resolve admin-class roles by the canonical criteria — the seeded
|
||||
// Owner/Admin role IDs plus any custom role holding the Administrator
|
||||
// bypass bit. Names are user-editable (the Owner can rename the seeded
|
||||
// Admin role), so a name lookup would silently disable the guard.
|
||||
adminRows, err := tx.QueryContext(ctx,
|
||||
`SELECT id FROM roles WHERE id IN (?, ?) OR (permissions & ?) != 0`,
|
||||
permissions.OwnerRoleID, permissions.AdminRoleID, permissions.Administrator,
|
||||
)
|
||||
if err := deleteAccountAdminGuard(ctx, tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
dmChannelIDs, err := deleteAccountDMChannels(ctx, tx, userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("DeleteAccount fetch admin roles: %w", err)
|
||||
}
|
||||
var adminRoleIDs []int64
|
||||
for adminRows.Next() {
|
||||
var rid int64
|
||||
if scanErr := adminRows.Scan(&rid); scanErr != nil {
|
||||
adminRows.Close() //nolint:errcheck
|
||||
return fmt.Errorf("DeleteAccount scan admin role: %w", scanErr)
|
||||
}
|
||||
adminRoleIDs = append(adminRoleIDs, rid)
|
||||
}
|
||||
adminRows.Close() //nolint:errcheck
|
||||
if adminRows.Err() != nil {
|
||||
return fmt.Errorf("DeleteAccount admin roles rows: %w", adminRows.Err())
|
||||
}
|
||||
|
||||
if len(adminRoleIDs) == 0 {
|
||||
// No admin-class roles defined; skip the guard.
|
||||
} else {
|
||||
var userRoleID int64
|
||||
if err := tx.QueryRowContext(ctx,
|
||||
`SELECT role_id FROM users WHERE id = ?`, userID,
|
||||
).Scan(&userRoleID); err != nil {
|
||||
return fmt.Errorf("DeleteAccount fetch role: %w", err)
|
||||
}
|
||||
|
||||
isAdminClass := slices.Contains(adminRoleIDs, userRoleID)
|
||||
|
||||
if isAdminClass {
|
||||
// Build IN clause dynamically for the admin role IDs.
|
||||
placeholders := make([]string, len(adminRoleIDs))
|
||||
args := make([]any, 0, len(adminRoleIDs)+1)
|
||||
for i, rid := range adminRoleIDs {
|
||||
placeholders[i] = "?"
|
||||
args = append(args, rid)
|
||||
}
|
||||
args = append(args, userID)
|
||||
|
||||
var adminCount int
|
||||
if err := tx.QueryRowContext(ctx,
|
||||
fmt.Sprintf(`SELECT COUNT(*) FROM users WHERE role_id IN (%s) AND id != ? AND banned = 0`,
|
||||
strings.Join(placeholders, ",")),
|
||||
args...,
|
||||
).Scan(&adminCount); err != nil {
|
||||
return fmt.Errorf("DeleteAccount count admins: %w", err)
|
||||
}
|
||||
if adminCount == 0 {
|
||||
return ErrLastAdmin
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot the user's DM channels before the participant rows go away,
|
||||
// so channels left with zero participants can be removed below —
|
||||
// LeaveGroupDM's invariant: a participant-less DM channel is an
|
||||
// unreachable, undeletable row.
|
||||
var dmChannelIDs []int64
|
||||
dmRows, err := tx.QueryContext(ctx,
|
||||
`SELECT channel_id FROM dm_participants WHERE user_id = ?`, userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("DeleteAccount list dm channels: %w", err)
|
||||
}
|
||||
for dmRows.Next() {
|
||||
var chID int64
|
||||
if scanErr := dmRows.Scan(&chID); scanErr != nil {
|
||||
dmRows.Close() //nolint:errcheck
|
||||
return fmt.Errorf("DeleteAccount scan dm channel: %w", scanErr)
|
||||
}
|
||||
dmChannelIDs = append(dmChannelIDs, chID)
|
||||
}
|
||||
dmRows.Close() //nolint:errcheck
|
||||
if dmRows.Err() != nil {
|
||||
return fmt.Errorf("DeleteAccount dm channels rows: %w", dmRows.Err())
|
||||
return err
|
||||
}
|
||||
|
||||
// ── Purge related data ───────────────────────────────────────────────
|
||||
@@ -141,67 +65,8 @@ func (d *DB) DeleteAccount(ctx context.Context, userID int64) error {
|
||||
}
|
||||
}
|
||||
|
||||
// Close and, where emptied, remove the deleted user's DM channels.
|
||||
for _, chID := range dmChannelIDs {
|
||||
var isGroup bool
|
||||
if err := tx.QueryRowContext(ctx,
|
||||
`SELECT is_group FROM channels WHERE id = ?`, chID,
|
||||
).Scan(&isGroup); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
continue // channel already gone
|
||||
}
|
||||
return fmt.Errorf("DeleteAccount dm channel is_group: %w", err)
|
||||
}
|
||||
|
||||
if !isGroup {
|
||||
// The purge above removed only this user's dm_participants row, so
|
||||
// a 1:1 DM with a live other side is untouched: its dm_participants
|
||||
// row (and the channel) survive, but the survivor's own
|
||||
// dm_open_state row does too. Left alone that renders as a
|
||||
// sidebar entry with a blank, unnamed recipient (GetDMParticipantsForUser
|
||||
// skips the viewer's own row and this user has none left to
|
||||
// return) that the survivor can still open and send into. Closing
|
||||
// it for them removes it from their sidebar, same as if they had
|
||||
// closed it themselves.
|
||||
if _, err := tx.ExecContext(ctx,
|
||||
`DELETE FROM dm_open_state WHERE channel_id = ? AND user_id != ?`,
|
||||
chID, userID,
|
||||
); err != nil {
|
||||
return fmt.Errorf("DeleteAccount close dm for survivor: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Hard-delete DM channels the deletion left with zero participants
|
||||
// (always true for the last member of a group DM; true for a 1:1 DM
|
||||
// only when the other side had already deleted their own account).
|
||||
//
|
||||
// Unlink attachments first: messages.channel_id and
|
||||
// attachments.message_id both cascade ON DELETE (migrations/001), so
|
||||
// deleting the channel row destroys the attachment rows too. Those
|
||||
// rows are the only handle DeleteOrphanedAttachments (the periodic
|
||||
// sweep in main.go) has on the uploaded files — once the cascade
|
||||
// removes them the files are stranded on disk forever. Setting
|
||||
// message_id to NULL first turns them into ordinary orphaned
|
||||
// attachments the sweep already knows how to reclaim.
|
||||
if _, err := tx.ExecContext(ctx,
|
||||
`UPDATE attachments SET message_id = NULL
|
||||
WHERE message_id IN (SELECT id FROM messages WHERE channel_id = ?)
|
||||
AND EXISTS (
|
||||
SELECT 1 FROM channels
|
||||
WHERE channels.id = ? AND channels.type = 'dm'
|
||||
AND NOT EXISTS (SELECT 1 FROM dm_participants WHERE channel_id = channels.id)
|
||||
)`,
|
||||
chID, chID,
|
||||
); err != nil {
|
||||
return fmt.Errorf("DeleteAccount unlink dm attachments: %w", err)
|
||||
}
|
||||
if _, err := tx.ExecContext(ctx,
|
||||
`DELETE FROM channels WHERE id = ? AND type = 'dm'
|
||||
AND NOT EXISTS (SELECT 1 FROM dm_participants WHERE channel_id = channels.id)`,
|
||||
chID,
|
||||
); err != nil {
|
||||
return fmt.Errorf("DeleteAccount empty dm channel: %w", err)
|
||||
}
|
||||
if err := deleteAccountCloseDMChannels(ctx, tx, userID, dmChannelIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Soft-delete messages: mark as deleted and clear content so the rows
|
||||
@@ -280,3 +145,166 @@ func anonymiseUser(ctx context.Context, tx *sql.Tx, userID int64) error {
|
||||
}
|
||||
return fmt.Errorf("DeleteAccount anonymise: %w", lastErr)
|
||||
}
|
||||
|
||||
// deleteAccountAdminGuard blocks the deletion when userID is the last
|
||||
// remaining admin-class account, returning ErrLastAdmin.
|
||||
func deleteAccountAdminGuard(ctx context.Context, tx *sql.Tx, userID int64) error {
|
||||
// ── Guard: last admin/owner check ────────────────────────────────────
|
||||
// Resolve admin-class roles by the canonical criteria — the seeded
|
||||
// Owner/Admin role IDs plus any custom role holding the Administrator
|
||||
// bypass bit. Names are user-editable (the Owner can rename the seeded
|
||||
// Admin role), so a name lookup would silently disable the guard.
|
||||
adminRows, err := tx.QueryContext(ctx,
|
||||
`SELECT id FROM roles WHERE id IN (?, ?) OR (permissions & ?) != 0`,
|
||||
permissions.OwnerRoleID, permissions.AdminRoleID, permissions.Administrator,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("DeleteAccount fetch admin roles: %w", err)
|
||||
}
|
||||
var adminRoleIDs []int64
|
||||
for adminRows.Next() {
|
||||
var rid int64
|
||||
if scanErr := adminRows.Scan(&rid); scanErr != nil {
|
||||
adminRows.Close() //nolint:errcheck
|
||||
return fmt.Errorf("DeleteAccount scan admin role: %w", scanErr)
|
||||
}
|
||||
adminRoleIDs = append(adminRoleIDs, rid)
|
||||
}
|
||||
adminRows.Close() //nolint:errcheck
|
||||
if adminRows.Err() != nil {
|
||||
return fmt.Errorf("DeleteAccount admin roles rows: %w", adminRows.Err())
|
||||
}
|
||||
|
||||
if len(adminRoleIDs) == 0 {
|
||||
// No admin-class roles defined; skip the guard.
|
||||
} else {
|
||||
var userRoleID int64
|
||||
if err := tx.QueryRowContext(ctx,
|
||||
`SELECT role_id FROM users WHERE id = ?`, userID,
|
||||
).Scan(&userRoleID); err != nil {
|
||||
return fmt.Errorf("DeleteAccount fetch role: %w", err)
|
||||
}
|
||||
|
||||
isAdminClass := slices.Contains(adminRoleIDs, userRoleID)
|
||||
|
||||
if isAdminClass {
|
||||
// Build IN clause dynamically for the admin role IDs.
|
||||
placeholders := make([]string, len(adminRoleIDs))
|
||||
args := make([]any, 0, len(adminRoleIDs)+1)
|
||||
for i, rid := range adminRoleIDs {
|
||||
placeholders[i] = "?"
|
||||
args = append(args, rid)
|
||||
}
|
||||
args = append(args, userID)
|
||||
|
||||
var adminCount int
|
||||
if err := tx.QueryRowContext(ctx,
|
||||
fmt.Sprintf(`SELECT COUNT(*) FROM users WHERE role_id IN (%s) AND id != ? AND banned = 0`,
|
||||
strings.Join(placeholders, ",")),
|
||||
args...,
|
||||
).Scan(&adminCount); err != nil {
|
||||
return fmt.Errorf("DeleteAccount count admins: %w", err)
|
||||
}
|
||||
if adminCount == 0 {
|
||||
return ErrLastAdmin
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// deleteAccountDMChannels snapshots the DM channels the user takes part in,
|
||||
// before the dm_participants purge removes the rows that name them.
|
||||
func deleteAccountDMChannels(ctx context.Context, tx *sql.Tx, userID int64) ([]int64, error) {
|
||||
// Snapshot the user's DM channels before the participant rows go away,
|
||||
// so channels left with zero participants can be removed below —
|
||||
// LeaveGroupDM's invariant: a participant-less DM channel is an
|
||||
// unreachable, undeletable row.
|
||||
var dmChannelIDs []int64
|
||||
dmRows, err := tx.QueryContext(ctx,
|
||||
`SELECT channel_id FROM dm_participants WHERE user_id = ?`, userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("DeleteAccount list dm channels: %w", err)
|
||||
}
|
||||
for dmRows.Next() {
|
||||
var chID int64
|
||||
if scanErr := dmRows.Scan(&chID); scanErr != nil {
|
||||
dmRows.Close() //nolint:errcheck
|
||||
return nil, fmt.Errorf("DeleteAccount scan dm channel: %w", scanErr)
|
||||
}
|
||||
dmChannelIDs = append(dmChannelIDs, chID)
|
||||
}
|
||||
dmRows.Close() //nolint:errcheck
|
||||
if dmRows.Err() != nil {
|
||||
return nil, fmt.Errorf("DeleteAccount dm channels rows: %w", dmRows.Err())
|
||||
}
|
||||
return dmChannelIDs, nil
|
||||
}
|
||||
|
||||
// deleteAccountCloseDMChannels closes, and where the purge emptied them,
|
||||
// removes the DM channels snapshotted by deleteAccountDMChannels.
|
||||
func deleteAccountCloseDMChannels(ctx context.Context, tx *sql.Tx, userID int64, dmChannelIDs []int64) error {
|
||||
// Close and, where emptied, remove the deleted user's DM channels.
|
||||
for _, chID := range dmChannelIDs {
|
||||
var isGroup bool
|
||||
if err := tx.QueryRowContext(ctx,
|
||||
`SELECT is_group FROM channels WHERE id = ?`, chID,
|
||||
).Scan(&isGroup); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
continue // channel already gone
|
||||
}
|
||||
return fmt.Errorf("DeleteAccount dm channel is_group: %w", err)
|
||||
}
|
||||
|
||||
if !isGroup {
|
||||
// The purge above removed only this user's dm_participants row, so
|
||||
// a 1:1 DM with a live other side is untouched: its dm_participants
|
||||
// row (and the channel) survive, but the survivor's own
|
||||
// dm_open_state row does too. Left alone that renders as a
|
||||
// sidebar entry with a blank, unnamed recipient (GetDMParticipantsForUser
|
||||
// skips the viewer's own row and this user has none left to
|
||||
// return) that the survivor can still open and send into. Closing
|
||||
// it for them removes it from their sidebar, same as if they had
|
||||
// closed it themselves.
|
||||
if _, err := tx.ExecContext(ctx,
|
||||
`DELETE FROM dm_open_state WHERE channel_id = ? AND user_id != ?`,
|
||||
chID, userID,
|
||||
); err != nil {
|
||||
return fmt.Errorf("DeleteAccount close dm for survivor: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Hard-delete DM channels the deletion left with zero participants
|
||||
// (always true for the last member of a group DM; true for a 1:1 DM
|
||||
// only when the other side had already deleted their own account).
|
||||
//
|
||||
// Unlink attachments first: messages.channel_id and
|
||||
// attachments.message_id both cascade ON DELETE (migrations/001), so
|
||||
// deleting the channel row destroys the attachment rows too. Those
|
||||
// rows are the only handle DeleteOrphanedAttachments (the periodic
|
||||
// sweep in main.go) has on the uploaded files — once the cascade
|
||||
// removes them the files are stranded on disk forever. Setting
|
||||
// message_id to NULL first turns them into ordinary orphaned
|
||||
// attachments the sweep already knows how to reclaim.
|
||||
if _, err := tx.ExecContext(ctx,
|
||||
`UPDATE attachments SET message_id = NULL
|
||||
WHERE message_id IN (SELECT id FROM messages WHERE channel_id = ?)
|
||||
AND EXISTS (
|
||||
SELECT 1 FROM channels
|
||||
WHERE channels.id = ? AND channels.type = 'dm'
|
||||
AND NOT EXISTS (SELECT 1 FROM dm_participants WHERE channel_id = channels.id)
|
||||
)`,
|
||||
chID, chID,
|
||||
); err != nil {
|
||||
return fmt.Errorf("DeleteAccount unlink dm attachments: %w", err)
|
||||
}
|
||||
if _, err := tx.ExecContext(ctx,
|
||||
`DELETE FROM channels WHERE id = ? AND type = 'dm'
|
||||
AND NOT EXISTS (SELECT 1 FROM dm_participants WHERE channel_id = channels.id)`,
|
||||
chID,
|
||||
); err != nil {
|
||||
return fmt.Errorf("DeleteAccount empty dm channel: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
+30
-19
@@ -383,25 +383,8 @@ func (d *DB) BackupToSafe(ctx context.Context, path, safeRoot string) error {
|
||||
return fmt.Errorf("BackupToSafe: path %q is not under safe root %q", absClean, absRoot)
|
||||
}
|
||||
|
||||
// Defence-in-depth: only allow safe characters (alphanumeric, path separators,
|
||||
// hyphen, underscore, dot, space, colon, tilde). This is a strict allowlist —
|
||||
// anything else is rejected to prevent SQL injection via the interpolated path.
|
||||
for _, ch := range absClean {
|
||||
switch {
|
||||
case ch >= 'a' && ch <= 'z',
|
||||
ch >= 'A' && ch <= 'Z',
|
||||
ch >= '0' && ch <= '9',
|
||||
ch == '/' || ch == '\\' || ch == '-' || ch == '_' || ch == '.' || ch == ' ' || ch == ':' || ch == '~':
|
||||
// allowed (colon for Windows drive letters, tilde for temp paths)
|
||||
default:
|
||||
return fmt.Errorf("BackupToSafe: path contains forbidden character %q", string(ch))
|
||||
}
|
||||
}
|
||||
|
||||
// Reject SQL comment sequences that could break the VACUUM INTO statement,
|
||||
// even though individual hyphens are allowed for filenames.
|
||||
if strings.Contains(absClean, "--") {
|
||||
return fmt.Errorf("BackupToSafe: path contains forbidden sequence %q", "--")
|
||||
if err := validateBackupPathChars(absClean); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// VACUUM INTO refuses to write over an existing destination on its own,
|
||||
@@ -431,6 +414,34 @@ func (d *DB) BackupToSafe(ctx context.Context, path, safeRoot string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateBackupPathChars is the strict character gate BackupToSafe applies to
|
||||
// the destination before it is interpolated into VACUUM INTO. It is a separate
|
||||
// function only so the allowlist loop's branch count does not dominate its
|
||||
// caller; the rules and the messages are unchanged.
|
||||
func validateBackupPathChars(absClean string) error {
|
||||
// Defence-in-depth: only allow safe characters (alphanumeric, path separators,
|
||||
// hyphen, underscore, dot, space, colon, tilde). This is a strict allowlist —
|
||||
// anything else is rejected to prevent SQL injection via the interpolated path.
|
||||
for _, ch := range absClean {
|
||||
switch {
|
||||
case ch >= 'a' && ch <= 'z',
|
||||
ch >= 'A' && ch <= 'Z',
|
||||
ch >= '0' && ch <= '9',
|
||||
ch == '/' || ch == '\\' || ch == '-' || ch == '_' || ch == '.' || ch == ' ' || ch == ':' || ch == '~':
|
||||
// allowed (colon for Windows drive letters, tilde for temp paths)
|
||||
default:
|
||||
return fmt.Errorf("BackupToSafe: path contains forbidden character %q", string(ch))
|
||||
}
|
||||
}
|
||||
|
||||
// Reject SQL comment sequences that could break the VACUUM INTO statement,
|
||||
// even though individual hyphens are allowed for filenames.
|
||||
if strings.Contains(absClean, "--") {
|
||||
return fmt.Errorf("BackupToSafe: path contains forbidden sequence %q", "--")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CheckBackupIntegrity opens the SQLite file at path read-only and runs
|
||||
// PRAGMA integrity_check against it. It returns nil only when SQLite reports
|
||||
// "ok". Use it to verify a backup right after it is written and again before
|
||||
|
||||
@@ -317,28 +317,43 @@ func (d *DB) GetUserIDsByUsernames(ctx context.Context, usernames []string) (map
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// ListMentionTargetsByRoles returns non-banned users holding any of the given
|
||||
// roles, with the presence status @here filters on.
|
||||
func (d *DB) ListMentionTargetsByRoles(ctx context.Context, roleIDs []int64) ([]MentionTarget, error) {
|
||||
if len(roleIDs) == 0 {
|
||||
// mentionTargetColumn is the users column a mention-target lookup matches its
|
||||
// id list against. It is a named type rather than a bare string so that every
|
||||
// call site has to name one of the two constants below instead of passing an
|
||||
// arbitrary string into the SELECT. Go named types are not closed, so this is
|
||||
// a convention the type makes visible, not one it enforces: do not introduce a
|
||||
// mentionTargetColumn(x) conversion from a runtime value.
|
||||
type mentionTargetColumn string
|
||||
|
||||
const (
|
||||
mentionTargetsByRole mentionTargetColumn = "role_id"
|
||||
mentionTargetsByUser mentionTargetColumn = "id"
|
||||
)
|
||||
|
||||
// listMentionTargets returns non-banned users whose column is in ids, with the
|
||||
// presence status @here filters on. It is the shared body of
|
||||
// ListMentionTargetsByRoles and ListMentionTargetsByUserIDs, which differ only
|
||||
// in the column they match and the name they report in errors.
|
||||
func (d *DB) listMentionTargets(ctx context.Context, column mentionTargetColumn, caller string, ids []int64) ([]MentionTarget, error) {
|
||||
if len(ids) == 0 {
|
||||
return []MentionTarget{}, nil
|
||||
}
|
||||
|
||||
placeholders := make([]string, len(roleIDs))
|
||||
args := make([]any, len(roleIDs))
|
||||
for i, id := range roleIDs {
|
||||
placeholders := make([]string, len(ids))
|
||||
args := make([]any, len(ids))
|
||||
for i, id := range ids {
|
||||
placeholders[i] = "?"
|
||||
args[i] = id
|
||||
}
|
||||
|
||||
rows, err := d.reader.QueryContext(ctx,
|
||||
fmt.Sprintf( //nolint:gosec // G201: placeholder interpolation, not user input
|
||||
`SELECT id, status, role_id FROM users WHERE %s AND role_id IN (%s)`,
|
||||
notBannedClause, strings.Join(placeholders, ",")),
|
||||
fmt.Sprintf( //nolint:gosec // G201: placeholders plus a named-type column constant, not user input
|
||||
`SELECT id, status, role_id FROM users WHERE %s AND %s IN (%s)`,
|
||||
notBannedClause, string(column), strings.Join(placeholders, ",")),
|
||||
args...,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ListMentionTargetsByRoles: %w", err)
|
||||
return nil, fmt.Errorf("%s: %w", caller, err)
|
||||
}
|
||||
defer rows.Close() //nolint:errcheck
|
||||
|
||||
@@ -346,56 +361,29 @@ func (d *DB) ListMentionTargetsByRoles(ctx context.Context, roleIDs []int64) ([]
|
||||
for rows.Next() {
|
||||
var t MentionTarget
|
||||
if scanErr := rows.Scan(&t.UserID, &t.Status, &t.RoleID); scanErr != nil {
|
||||
return nil, fmt.Errorf("ListMentionTargetsByRoles scan: %w", scanErr)
|
||||
return nil, fmt.Errorf("%s scan: %w", caller, scanErr)
|
||||
}
|
||||
targets = append(targets, t)
|
||||
}
|
||||
if rows.Err() != nil {
|
||||
return nil, fmt.Errorf("ListMentionTargetsByRoles rows: %w", rows.Err())
|
||||
return nil, fmt.Errorf("%s rows: %w", caller, rows.Err())
|
||||
}
|
||||
return targets, nil
|
||||
}
|
||||
|
||||
// ListMentionTargetsByRoles returns non-banned users holding any of the given
|
||||
// roles, with the presence status @here filters on.
|
||||
func (d *DB) ListMentionTargetsByRoles(ctx context.Context, roleIDs []int64) ([]MentionTarget, error) {
|
||||
return d.listMentionTargets(ctx, mentionTargetsByRole, "ListMentionTargetsByRoles", roleIDs)
|
||||
}
|
||||
|
||||
// ListMentionTargetsByUserIDs returns non-banned users by explicit id, with the
|
||||
// same fields ListMentionTargetsByRoles returns. It backs the additive half of
|
||||
// the per-user channel override layer: a member whose user override ALLOWs
|
||||
// READ_MESSAGES can read a channel their role cannot, so the role walk alone
|
||||
// would leave them out of an @everyone fan-out they are entitled to.
|
||||
func (d *DB) ListMentionTargetsByUserIDs(ctx context.Context, userIDs []int64) ([]MentionTarget, error) {
|
||||
if len(userIDs) == 0 {
|
||||
return []MentionTarget{}, nil
|
||||
}
|
||||
|
||||
placeholders := make([]string, len(userIDs))
|
||||
args := make([]any, len(userIDs))
|
||||
for i, id := range userIDs {
|
||||
placeholders[i] = "?"
|
||||
args[i] = id
|
||||
}
|
||||
|
||||
rows, err := d.reader.QueryContext(ctx,
|
||||
fmt.Sprintf( //nolint:gosec // G201: placeholder interpolation, not user input
|
||||
`SELECT id, status, role_id FROM users WHERE %s AND id IN (%s)`,
|
||||
notBannedClause, strings.Join(placeholders, ",")),
|
||||
args...,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ListMentionTargetsByUserIDs: %w", err)
|
||||
}
|
||||
defer rows.Close() //nolint:errcheck
|
||||
|
||||
targets := []MentionTarget{}
|
||||
for rows.Next() {
|
||||
var t MentionTarget
|
||||
if scanErr := rows.Scan(&t.UserID, &t.Status, &t.RoleID); scanErr != nil {
|
||||
return nil, fmt.Errorf("ListMentionTargetsByUserIDs scan: %w", scanErr)
|
||||
}
|
||||
targets = append(targets, t)
|
||||
}
|
||||
if rows.Err() != nil {
|
||||
return nil, fmt.Errorf("ListMentionTargetsByUserIDs rows: %w", rows.Err())
|
||||
}
|
||||
return targets, nil
|
||||
return d.listMentionTargets(ctx, mentionTargetsByUser, "ListMentionTargetsByUserIDs", userIDs)
|
||||
}
|
||||
|
||||
// ListBlockersOf returns the ids of users who have blocked the given user.
|
||||
|
||||
+345
-191
@@ -117,6 +117,122 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
bgCtx, bgCancel := context.WithCancel(context.Background())
|
||||
defer bgCancel()
|
||||
|
||||
runRemoveOldBinary(log)
|
||||
|
||||
// ── 1. Load configuration ──────────────────────────────────────────────
|
||||
cfg, err := runLoadConfig(log, levelVar, rc)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// ── 2. Ensure data directory exists ────────────────────────────────────
|
||||
if err := runPrepareDataDir(log, cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// ── 3. TLS ────────────────────────────────────────────────────────────
|
||||
tlsResult, err := auth.LoadOrGenerate(cfg.TLS)
|
||||
if err != nil {
|
||||
return fmt.Errorf("configuring TLS: %w", err)
|
||||
}
|
||||
tlsCfg := tlsResult.TLSConfig
|
||||
|
||||
// Print startup banner first so it appears above all init logs.
|
||||
printBanner(cfg, version, tlsCfg != nil)
|
||||
|
||||
// ── 4. Open database + run migrations ─────────────────────────────────
|
||||
database, err := runOpenDatabase(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer database.Close() //nolint:errcheck
|
||||
|
||||
if err := runInitDatabase(log, cfg, database, rc); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// ── 4b. Telemetry (Phase B Step 8) ─────────────────────────────────────
|
||||
telemetryStop := runInitTelemetry(log, cfg)
|
||||
defer telemetryStop()
|
||||
|
||||
// ── 5a. Construct plugin runtime BEFORE the router so the router can
|
||||
// wire the live registry into the plugin admin handler. ────────────────
|
||||
pluginRegistry := runInitPlugins(bgCtx, log, cfg, database)
|
||||
defer runClosePlugins(pluginRegistry)
|
||||
|
||||
// ── 5b. Build HTTP router ──────────────────────────────────────────────
|
||||
router, hub, routerCleanup := api.NewRouter(cfg, database, version, logBuf, pluginRegistry)
|
||||
defer routerCleanup()
|
||||
// Backstop for every early return below (serve error, ACME shutdown
|
||||
// failure, etc.): hub.GracefulStop is the only caller of
|
||||
// LiveKitProcess.Stop(), so skipping it orphans the companion
|
||||
// livekit-server process and leaves the hub's dispatch goroutine
|
||||
// running. gracefulOnce makes it idempotent alongside the explicit call
|
||||
// on the normal shutdown path below.
|
||||
defer hub.GracefulStop()
|
||||
|
||||
// ── 5c. Wire event persistence (Phase B Step 7) ────────────────────────
|
||||
persister, prunerDone := runStartEventPersistence(bgCtx, log, cfg, hub, database)
|
||||
defer runStopEventPersistence(log, bgCancel, persister, prunerDone)
|
||||
|
||||
// ── 5d. Async audit writer ─────────────────────────────────────────────
|
||||
// Moves audit-log INSERTs off the request path: once the writer is
|
||||
// installed, WriteAudit enqueues here and a background goroutine batches
|
||||
// the writes (same shape as the event persister above). Paths that never
|
||||
// install a writer — the token CLI, tests — keep the synchronous
|
||||
// behavior. This defer is registered after `defer database.Close()` so
|
||||
// LIFO ordering drains the queue before the database is torn down.
|
||||
auditWriter := runStartAuditWriter(bgCtx, database)
|
||||
defer runStopAuditWriter(auditWriter)
|
||||
|
||||
// ── 6. Start server ────────────────────────────────────────────────────
|
||||
addr := fmt.Sprintf(":%d", cfg.Server.Port)
|
||||
srv := &http.Server{
|
||||
Addr: addr,
|
||||
Handler: router,
|
||||
TLSConfig: tlsCfg,
|
||||
ReadTimeout: 30 * time.Second,
|
||||
WriteTimeout: 30 * time.Second,
|
||||
IdleTimeout: 120 * time.Second,
|
||||
ErrorLog: stdlog.New(io.Discard, "", 0), // suppress TLS handshake noise
|
||||
}
|
||||
|
||||
// ── 6b. ACME HTTP challenge server on :80 ─────────────────────────────
|
||||
// When using Let's Encrypt (tls.mode: acme), an HTTP server on port 80
|
||||
// is needed for HTTP-01 challenge validation and HTTP→HTTPS redirect.
|
||||
acmeSrv := runStartACME(log, tlsResult.HTTPHandler)
|
||||
|
||||
// ── 7. Background maintenance ────────────────────────────────────────
|
||||
maintenanceStop := runStartMaintenance(bgCtx, log, cfg, database)
|
||||
defer maintenanceStop()
|
||||
|
||||
// Listen for OS signals for graceful shutdown. The coordinator's context
|
||||
// is the parent, so a programmatic restart request (rc.Request) drains
|
||||
// exactly like a SIGTERM — including on Windows, where a process cannot
|
||||
// signal itself. Signals arriving mid-drain are swallowed until stop()
|
||||
// runs, same as on the real-signal path.
|
||||
ctx, stop := signal.NotifyContext(rc.Context(), os.Interrupt, syscall.SIGTERM)
|
||||
defer stop()
|
||||
|
||||
if err := runServeAndWait(ctx, log, rc, srv, tlsCfg, addr); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Graceful shutdown.
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := runShutdownServers(shutdownCtx, log, srv, acmeSrv, hub); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
log.Info("server stopped cleanly")
|
||||
return nil
|
||||
}
|
||||
|
||||
// runRemoveOldBinary deletes the binary a previous self-update left behind.
|
||||
// Extracted from run.
|
||||
func runRemoveOldBinary(log *slog.Logger) {
|
||||
// Clean up old binary from a previous update. Bounded retry: in spawn
|
||||
// mode the predecessor spawns this process as its very last act, so for
|
||||
// the first few hundred milliseconds it may not have fully exited — and
|
||||
@@ -125,30 +241,36 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
exePath, exeErr := os.Executable()
|
||||
if exeErr != nil {
|
||||
log.Warn("failed to determine executable path", "error", exeErr)
|
||||
} else {
|
||||
oldPath := exePath + ".old"
|
||||
if _, statErr := os.Stat(oldPath); statErr == nil {
|
||||
var rmErr error
|
||||
for attempt := range 5 {
|
||||
if attempt > 0 {
|
||||
time.Sleep(250 * time.Millisecond)
|
||||
}
|
||||
if rmErr = os.Remove(oldPath); rmErr == nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
if rmErr != nil {
|
||||
log.Warn("failed to remove old binary", "path", oldPath, "error", rmErr)
|
||||
} else {
|
||||
log.Info("removed old binary from previous update", "path", oldPath)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// ── 1. Load configuration ──────────────────────────────────────────────
|
||||
oldPath := exePath + ".old"
|
||||
if _, statErr := os.Stat(oldPath); statErr != nil {
|
||||
return
|
||||
}
|
||||
|
||||
var rmErr error
|
||||
for attempt := range 5 {
|
||||
if attempt > 0 {
|
||||
time.Sleep(250 * time.Millisecond)
|
||||
}
|
||||
if rmErr = os.Remove(oldPath); rmErr == nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
if rmErr != nil {
|
||||
log.Warn("failed to remove old binary", "path", oldPath, "error", rmErr)
|
||||
} else {
|
||||
log.Info("removed old binary from previous update", "path", oldPath)
|
||||
}
|
||||
}
|
||||
|
||||
// runLoadConfig loads the on-disk configuration, applies its logging level
|
||||
// and resolves the restart handoff mode. Extracted from run.
|
||||
func runLoadConfig(log *slog.Logger, levelVar *slog.LevelVar, rc *restartCoordinator) (*config.Config, error) {
|
||||
cfg, err := config.Load(config.DefaultPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("loading config: %w", err)
|
||||
return nil, fmt.Errorf("loading config: %w", err)
|
||||
}
|
||||
|
||||
// Apply the configured log level. The admin panel's live log view (ring
|
||||
@@ -165,7 +287,12 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
// after run() returns.
|
||||
rc.SetMode(resolveRestartMode(cfg.Server.RestartMode, log))
|
||||
|
||||
// ── 2. Ensure data directory exists ────────────────────────────────────
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// runPrepareDataDir creates the configured data directory and warns when the
|
||||
// volumes the server writes to are low on free space. Extracted from run.
|
||||
func runPrepareDataDir(log *slog.Logger, cfg *config.Config) error {
|
||||
if mkdirErr := os.MkdirAll(cfg.Server.DataDir, 0o750); mkdirErr != nil {
|
||||
return fmt.Errorf("creating data dir %s: %w", cfg.Server.DataDir, mkdirErr)
|
||||
}
|
||||
@@ -179,30 +306,31 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
warnLowDisk(log, "backup dir", cfg.Backup.Dir)
|
||||
}
|
||||
|
||||
// ── 3. TLS ────────────────────────────────────────────────────────────
|
||||
tlsResult, err := auth.LoadOrGenerate(cfg.TLS)
|
||||
if err != nil {
|
||||
return fmt.Errorf("configuring TLS: %w", err)
|
||||
}
|
||||
tlsCfg := tlsResult.TLSConfig
|
||||
return nil
|
||||
}
|
||||
|
||||
// Print startup banner first so it appears above all init logs.
|
||||
printBanner(cfg, version, tlsCfg != nil)
|
||||
|
||||
// ── 4. Open database + run migrations ─────────────────────────────────
|
||||
// runOpenDatabase validates the configured backend and opens the database.
|
||||
// Extracted from run.
|
||||
func runOpenDatabase(cfg *config.Config) (*db.DB, error) {
|
||||
// SQLite is the only supported backend; the unfinished Postgres
|
||||
// scaffolding (stubbed query layer, never wired into the runtime) was
|
||||
// removed rather than completed.
|
||||
if t := cfg.Database.Type; t != "" && t != "sqlite" {
|
||||
return fmt.Errorf("database.type=%q is not supported; set \"sqlite\" or omit it", t)
|
||||
return nil, fmt.Errorf("database.type=%q is not supported; set \"sqlite\" or omit it", t)
|
||||
}
|
||||
|
||||
database, err := db.OpenWithMaxReaders(cfg.Database.Path, cfg.Database.MaxReaders)
|
||||
if err != nil {
|
||||
return fmt.Errorf("opening database: %w", err)
|
||||
return nil, fmt.Errorf("opening database: %w", err)
|
||||
}
|
||||
defer database.Close() //nolint:errcheck
|
||||
|
||||
return database, nil
|
||||
}
|
||||
|
||||
// runInitDatabase points the admin panel at the live database, runs the
|
||||
// migrations and clears state left over from a previous run. Extracted from
|
||||
// run.
|
||||
func runInitDatabase(log *slog.Logger, cfg *config.Config, database *db.DB, rc *restartCoordinator) error {
|
||||
// The admin "Restore backup" handler needs the real database file path:
|
||||
// without this, it falls back to a hardcoded "data/chatserver.db" and
|
||||
// silently no-ops on any server with a configured database.path.
|
||||
@@ -233,7 +361,12 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
log.Info("cleared stale voice states")
|
||||
}
|
||||
|
||||
// ── 4b. Telemetry (Phase B Step 8) ─────────────────────────────────────
|
||||
return nil
|
||||
}
|
||||
|
||||
// runInitTelemetry initialises OpenTelemetry and returns the shutdown step
|
||||
// run defers. Extracted from run.
|
||||
func runInitTelemetry(log *slog.Logger, cfg *config.Config) func() {
|
||||
// Init can return (nil, err) when the otel build-tag skeleton hasn't been
|
||||
// finished wiring to the upstream SDK. Normalise to a no-op shutdown so
|
||||
// the deferred closure never calls a nil function.
|
||||
@@ -244,16 +377,19 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
if telemetryShutdown == nil {
|
||||
telemetryShutdown = func(context.Context) error { return nil }
|
||||
}
|
||||
defer func() {
|
||||
|
||||
return func() {
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
if err := telemetryShutdown(shutdownCtx); err != nil {
|
||||
log.Warn("telemetry shutdown returned error", "error", err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
// ── 5a. Construct plugin runtime BEFORE the router so the router can
|
||||
// wire the live registry into the plugin admin handler. ────────────────
|
||||
// runInitPlugins constructs the plugin runtime, returning nil when plugins
|
||||
// are disabled or failed to start. Extracted from run.
|
||||
func runInitPlugins(bgCtx context.Context, log *slog.Logger, cfg *config.Config, database *db.DB) *plugin.Registry {
|
||||
var pluginRegistry *plugin.Registry
|
||||
if cfg.Plugins.Enabled {
|
||||
registry, plugErr := plugin.NewRegistry(plugin.Config{
|
||||
@@ -270,95 +406,99 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
if err := registry.LoadAll(bgCtx); err != nil {
|
||||
log.Warn("plugin loader: failed to scan directory", "error", err)
|
||||
}
|
||||
defer func() {
|
||||
closeCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
_ = registry.Close(closeCtx)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
// ── 5b. Build HTTP router ──────────────────────────────────────────────
|
||||
router, hub, routerCleanup := api.NewRouter(cfg, database, version, logBuf, pluginRegistry)
|
||||
defer routerCleanup()
|
||||
// Backstop for every early return below (serve error, ACME shutdown
|
||||
// failure, etc.): hub.GracefulStop is the only caller of
|
||||
// LiveKitProcess.Stop(), so skipping it orphans the companion
|
||||
// livekit-server process and leaves the hub's dispatch goroutine
|
||||
// running. gracefulOnce makes it idempotent alongside the explicit call
|
||||
// on the normal shutdown path below.
|
||||
defer hub.GracefulStop()
|
||||
return pluginRegistry
|
||||
}
|
||||
|
||||
// ── 5c. Wire event persistence (Phase B Step 7) ────────────────────────
|
||||
if cfg.EventPersistence.Enabled && hub != nil {
|
||||
seedHubReplayState(bgCtx, hub, database, log)
|
||||
|
||||
persister := ws.NewEventPersister(
|
||||
database,
|
||||
4096,
|
||||
cfg.EventPersistence.BatchSize,
|
||||
time.Duration(cfg.EventPersistence.BatchFlushMs)*time.Millisecond,
|
||||
)
|
||||
persister.Start(bgCtx)
|
||||
hub.SetEventPersister(persister)
|
||||
hub.SetEventStore(database)
|
||||
|
||||
retention := time.Duration(cfg.EventPersistence.RetentionHours) * time.Hour
|
||||
prunerInterval := time.Duration(cfg.EventPersistence.PrunerIntervalMinutes) * time.Minute
|
||||
prunerDone := ws.StartEventPruner(bgCtx, database, retention, prunerInterval)
|
||||
defer func() {
|
||||
stopCtx, stopCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer stopCancel()
|
||||
persister.Stop(stopCtx)
|
||||
// Cancel the shared background context and JOIN the pruner before
|
||||
// the (LIFO-later) database.Close defer runs, so no prune is still
|
||||
// mid-query against a closing pool. Bounded: a stuck prune delays
|
||||
// shutdown by at most the timeout, then Close proceeds anyway.
|
||||
bgCancel()
|
||||
select {
|
||||
case <-prunerDone:
|
||||
case <-stopCtx.Done():
|
||||
log.Warn("event pruner did not exit before shutdown timeout")
|
||||
}
|
||||
}()
|
||||
// runClosePlugins shuts the plugin runtime down. Registered by run as a defer
|
||||
// only once the registry exists, so a nil registry is the disabled case and
|
||||
// has nothing to close. Extracted from run.
|
||||
func runClosePlugins(registry *plugin.Registry) {
|
||||
if registry == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// ── 5d. Async audit writer ─────────────────────────────────────────────
|
||||
// Moves audit-log INSERTs off the request path: once the writer is
|
||||
// installed, WriteAudit enqueues here and a background goroutine batches
|
||||
// the writes (same shape as the event persister above). Paths that never
|
||||
// install a writer — the token CLI, tests — keep the synchronous
|
||||
// behavior. This defer is registered after `defer database.Close()` so
|
||||
// LIFO ordering drains the queue before the database is torn down.
|
||||
closeCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
_ = registry.Close(closeCtx)
|
||||
}
|
||||
|
||||
// runStartEventPersistence starts the event persister and pruner, returning
|
||||
// both as (nil, nil) when event persistence is disabled. Extracted from run.
|
||||
func runStartEventPersistence(bgCtx context.Context, log *slog.Logger, cfg *config.Config, hub *ws.Hub, database *db.DB) (*ws.EventPersister, <-chan struct{}) {
|
||||
if !cfg.EventPersistence.Enabled || hub == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
seedHubReplayState(bgCtx, hub, database, log)
|
||||
|
||||
persister := ws.NewEventPersister(
|
||||
database,
|
||||
4096,
|
||||
cfg.EventPersistence.BatchSize,
|
||||
time.Duration(cfg.EventPersistence.BatchFlushMs)*time.Millisecond,
|
||||
)
|
||||
persister.Start(bgCtx)
|
||||
hub.SetEventPersister(persister)
|
||||
hub.SetEventStore(database)
|
||||
|
||||
retention := time.Duration(cfg.EventPersistence.RetentionHours) * time.Hour
|
||||
prunerInterval := time.Duration(cfg.EventPersistence.PrunerIntervalMinutes) * time.Minute
|
||||
prunerDone := ws.StartEventPruner(bgCtx, database, retention, prunerInterval)
|
||||
|
||||
return persister, prunerDone
|
||||
}
|
||||
|
||||
// runStopEventPersistence drains the event persister and pruner. Registered by
|
||||
// run as a defer unconditionally, so a nil persister is the disabled case and
|
||||
// must leave bgCtx alone — the LIFO backstop in run cancels it instead.
|
||||
// Extracted from run.
|
||||
func runStopEventPersistence(log *slog.Logger, bgCancel context.CancelFunc, persister *ws.EventPersister, prunerDone <-chan struct{}) {
|
||||
if persister == nil {
|
||||
return
|
||||
}
|
||||
|
||||
stopCtx, stopCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer stopCancel()
|
||||
persister.Stop(stopCtx)
|
||||
// Cancel the shared background context and JOIN the pruner before
|
||||
// the (LIFO-later) database.Close defer runs, so no prune is still
|
||||
// mid-query against a closing pool. Bounded: a stuck prune delays
|
||||
// shutdown by at most the timeout, then Close proceeds anyway.
|
||||
bgCancel()
|
||||
select {
|
||||
case <-prunerDone:
|
||||
case <-stopCtx.Done():
|
||||
log.Warn("event pruner did not exit before shutdown timeout")
|
||||
}
|
||||
}
|
||||
|
||||
// runStartAuditWriter installs the async audit writer. Extracted from run.
|
||||
func runStartAuditWriter(bgCtx context.Context, database *db.DB) *db.AuditWriter {
|
||||
auditWriter := db.NewAuditWriter(database, 1024, 50, 100*time.Millisecond)
|
||||
auditWriter.Start(bgCtx)
|
||||
database.SetAuditWriter(auditWriter)
|
||||
defer func() {
|
||||
stopCtx, stopCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer stopCancel()
|
||||
auditWriter.Stop(stopCtx)
|
||||
}()
|
||||
|
||||
// ── 6. Start server ────────────────────────────────────────────────────
|
||||
addr := fmt.Sprintf(":%d", cfg.Server.Port)
|
||||
srv := &http.Server{
|
||||
Addr: addr,
|
||||
Handler: router,
|
||||
TLSConfig: tlsCfg,
|
||||
ReadTimeout: 30 * time.Second,
|
||||
WriteTimeout: 30 * time.Second,
|
||||
IdleTimeout: 120 * time.Second,
|
||||
ErrorLog: stdlog.New(io.Discard, "", 0), // suppress TLS handshake noise
|
||||
}
|
||||
return auditWriter
|
||||
}
|
||||
|
||||
// ── 6b. ACME HTTP challenge server on :80 ─────────────────────────────
|
||||
// When using Let's Encrypt (tls.mode: acme), an HTTP server on port 80
|
||||
// is needed for HTTP-01 challenge validation and HTTP→HTTPS redirect.
|
||||
// runStopAuditWriter drains the async audit writer. Extracted from run.
|
||||
func runStopAuditWriter(auditWriter *db.AuditWriter) {
|
||||
stopCtx, stopCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer stopCancel()
|
||||
auditWriter.Stop(stopCtx)
|
||||
}
|
||||
|
||||
// runStartACME starts the ACME HTTP-01 challenge server when Let's Encrypt
|
||||
// is configured, and returns nil otherwise. Extracted from run.
|
||||
func runStartACME(log *slog.Logger, httpHandler http.Handler) *http.Server {
|
||||
var acmeSrv *http.Server
|
||||
if tlsResult.HTTPHandler != nil {
|
||||
if httpHandler != nil {
|
||||
acmeSrv = &http.Server{
|
||||
Addr: ":80",
|
||||
Handler: tlsResult.HTTPHandler,
|
||||
Handler: httpHandler,
|
||||
ReadTimeout: 10 * time.Second,
|
||||
WriteTimeout: 10 * time.Second,
|
||||
}
|
||||
@@ -370,7 +510,12 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
}()
|
||||
}
|
||||
|
||||
// ── 7. Background maintenance ────────────────────────────────────────
|
||||
return acmeSrv
|
||||
}
|
||||
|
||||
// runStartMaintenance starts the periodic maintenance loop and returns the
|
||||
// stop step run defers. Extracted from run.
|
||||
func runStartMaintenance(bgCtx context.Context, log *slog.Logger, cfg *config.Config, database *db.DB) func() {
|
||||
// Periodically purge expired sessions and orphaned attachments.
|
||||
fileStorage, fileStorageErr := storage.New(cfg.Upload.StorageDir, cfg.Upload.MaxSizeMB)
|
||||
if fileStorageErr != nil {
|
||||
@@ -379,7 +524,9 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
|
||||
stopMaintenance := make(chan struct{})
|
||||
maintenanceDone := make(chan struct{})
|
||||
defer func() {
|
||||
go runMaintenanceLoop(bgCtx, log, database, fileStorage, stopMaintenance, maintenanceDone)
|
||||
|
||||
return func() {
|
||||
// Backstop for early returns below (see hub.GracefulStop defer above),
|
||||
// and a bounded join so an in-flight tick (which can hold the writer —
|
||||
// scheduled backups run VACUUM INTO) isn't still using the database
|
||||
@@ -390,82 +537,87 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
case <-time.After(5 * time.Second):
|
||||
log.Warn("maintenance loop did not exit before shutdown timeout")
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
defer close(maintenanceDone)
|
||||
ticker := time.NewTicker(15 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
consecutiveFailures := 0
|
||||
const maxConsecutiveFailures = 5
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
if consecutiveFailures >= maxConsecutiveFailures {
|
||||
log.Error("maintenance loop: circuit breaker open, skipping tick",
|
||||
"consecutive_failures", consecutiveFailures)
|
||||
// Reset after one skip to allow retry next tick.
|
||||
consecutiveFailures = maxConsecutiveFailures - 1
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
tickFailed := false
|
||||
if err := database.DeleteExpiredSessions(bgCtx); err != nil {
|
||||
log.Warn("failed to delete expired sessions", "error", err)
|
||||
tickFailed = true
|
||||
}
|
||||
|
||||
// Scheduled backups + retention pruning, driven by the
|
||||
// backup_schedule / backup_retention admin settings.
|
||||
if err := admin.MaintainBackups(bgCtx, database); err != nil {
|
||||
log.Warn("backup maintenance failed", "error", err)
|
||||
tickFailed = true
|
||||
}
|
||||
|
||||
// Clean up orphaned attachments (uploaded but never linked to a message).
|
||||
//
|
||||
// Skipped entirely with no file storage configured: the delete is
|
||||
// atomic (row goes the instant it's selected, by design — see
|
||||
// db/attachment_queries.go), so with fileStorage nil the returned
|
||||
// stored_as names — the only remaining handle on those blobs —
|
||||
// would just be discarded and the files stranded on disk with no
|
||||
// query left able to name them. Leaving the rows in place keeps
|
||||
// them reclaimable once storage is available again.
|
||||
if fileStorage != nil {
|
||||
cutoff := time.Now().Add(-1 * time.Hour)
|
||||
orphanFiles, orphanErr := database.DeleteOrphanedAttachments(bgCtx, cutoff)
|
||||
if orphanErr != nil {
|
||||
log.Warn("failed to delete orphaned attachments", "error", orphanErr)
|
||||
tickFailed = true
|
||||
} else if len(orphanFiles) > 0 {
|
||||
// Best-effort file cleanup.
|
||||
for _, filename := range orphanFiles {
|
||||
if delErr := fileStorage.Delete(filename); delErr != nil {
|
||||
log.Warn("failed to delete orphan file", "file", filename, "error", delErr)
|
||||
}
|
||||
}
|
||||
log.Info("cleaned up orphaned attachments", "count", len(orphanFiles))
|
||||
}
|
||||
}
|
||||
|
||||
if tickFailed {
|
||||
consecutiveFailures++
|
||||
} else {
|
||||
consecutiveFailures = 0
|
||||
}
|
||||
case <-stopMaintenance:
|
||||
return
|
||||
// runMaintenanceLoop is the periodic maintenance goroutine started by
|
||||
// runStartMaintenance. Extracted from run.
|
||||
func runMaintenanceLoop(bgCtx context.Context, log *slog.Logger, database *db.DB, fileStorage *storage.Storage, stopMaintenance, maintenanceDone chan struct{}) {
|
||||
defer close(maintenanceDone)
|
||||
ticker := time.NewTicker(15 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
consecutiveFailures := 0
|
||||
const maxConsecutiveFailures = 5
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
if consecutiveFailures >= maxConsecutiveFailures {
|
||||
log.Error("maintenance loop: circuit breaker open, skipping tick",
|
||||
"consecutive_failures", consecutiveFailures)
|
||||
// Reset after one skip to allow retry next tick.
|
||||
consecutiveFailures = maxConsecutiveFailures - 1
|
||||
continue
|
||||
}
|
||||
|
||||
if runMaintenanceTick(bgCtx, log, database, fileStorage) {
|
||||
consecutiveFailures++
|
||||
} else {
|
||||
consecutiveFailures = 0
|
||||
}
|
||||
case <-stopMaintenance:
|
||||
return
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
// Listen for OS signals for graceful shutdown. The coordinator's context
|
||||
// is the parent, so a programmatic restart request (rc.Request) drains
|
||||
// exactly like a SIGTERM — including on Windows, where a process cannot
|
||||
// signal itself. Signals arriving mid-drain are swallowed until stop()
|
||||
// runs, same as on the real-signal path.
|
||||
ctx, stop := signal.NotifyContext(rc.Context(), os.Interrupt, syscall.SIGTERM)
|
||||
defer stop()
|
||||
// runMaintenanceTick runs one maintenance pass and reports whether any step
|
||||
// of it failed. Extracted from run.
|
||||
func runMaintenanceTick(bgCtx context.Context, log *slog.Logger, database *db.DB, fileStorage *storage.Storage) bool {
|
||||
tickFailed := false
|
||||
if err := database.DeleteExpiredSessions(bgCtx); err != nil {
|
||||
log.Warn("failed to delete expired sessions", "error", err)
|
||||
tickFailed = true
|
||||
}
|
||||
|
||||
// Scheduled backups + retention pruning, driven by the
|
||||
// backup_schedule / backup_retention admin settings.
|
||||
if err := admin.MaintainBackups(bgCtx, database); err != nil {
|
||||
log.Warn("backup maintenance failed", "error", err)
|
||||
tickFailed = true
|
||||
}
|
||||
|
||||
// Clean up orphaned attachments (uploaded but never linked to a message).
|
||||
//
|
||||
// Skipped entirely with no file storage configured: the delete is
|
||||
// atomic (row goes the instant it's selected, by design — see
|
||||
// db/attachment_queries.go), so with fileStorage nil the returned
|
||||
// stored_as names — the only remaining handle on those blobs —
|
||||
// would just be discarded and the files stranded on disk with no
|
||||
// query left able to name them. Leaving the rows in place keeps
|
||||
// them reclaimable once storage is available again.
|
||||
if fileStorage != nil {
|
||||
cutoff := time.Now().Add(-1 * time.Hour)
|
||||
orphanFiles, orphanErr := database.DeleteOrphanedAttachments(bgCtx, cutoff)
|
||||
if orphanErr != nil {
|
||||
log.Warn("failed to delete orphaned attachments", "error", orphanErr)
|
||||
tickFailed = true
|
||||
} else if len(orphanFiles) > 0 {
|
||||
// Best-effort file cleanup.
|
||||
for _, filename := range orphanFiles {
|
||||
if delErr := fileStorage.Delete(filename); delErr != nil {
|
||||
log.Warn("failed to delete orphan file", "file", filename, "error", delErr)
|
||||
}
|
||||
}
|
||||
log.Info("cleaned up orphaned attachments", "count", len(orphanFiles))
|
||||
}
|
||||
}
|
||||
|
||||
return tickFailed
|
||||
}
|
||||
|
||||
// runServeAndWait starts the listener and blocks until it fails or a
|
||||
// shutdown or restart signal arrives. Extracted from run.
|
||||
func runServeAndWait(ctx context.Context, log *slog.Logger, rc *restartCoordinator, srv *http.Server, tlsCfg *tls.Config, addr string) error {
|
||||
// Start serving in a goroutine.
|
||||
serveErr := make(chan error, 1)
|
||||
go func() {
|
||||
@@ -497,10 +649,13 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
}
|
||||
}
|
||||
|
||||
// Graceful shutdown.
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
return nil
|
||||
}
|
||||
|
||||
// runShutdownServers performs the ordered graceful shutdown: the ACME
|
||||
// server, then in-flight HTTP handlers, then the WebSocket hub. Extracted
|
||||
// from run.
|
||||
func runShutdownServers(shutdownCtx context.Context, log *slog.Logger, srv, acmeSrv *http.Server, hub *ws.Hub) error {
|
||||
if acmeSrv != nil {
|
||||
if err := acmeSrv.Shutdown(shutdownCtx); err != nil {
|
||||
log.Warn("ACME HTTP server shutdown error", "error", err)
|
||||
@@ -525,7 +680,6 @@ func run(log *slog.Logger, logBuf *admin.RingBuffer, levelVar *slog.LevelVar, rc
|
||||
return fmt.Errorf("graceful shutdown: %w", shutdownErr)
|
||||
}
|
||||
|
||||
log.Info("server stopped cleanly")
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
+184
-136
@@ -262,120 +262,23 @@ func (r *Registry) InstallFromZip(ctx context.Context, zipBytes []byte) (string,
|
||||
return "", fmt.Errorf("abs staging dir: %w", absErr)
|
||||
}
|
||||
|
||||
var totalUncompressed int64
|
||||
for _, f := range zr.File {
|
||||
// Reject symlinks, devices, and any non-regular file mode.
|
||||
if !f.Mode().IsRegular() && !f.Mode().IsDir() {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("plugin zip: refusing non-regular entry %q (mode=%v)", f.Name, f.Mode())
|
||||
}
|
||||
if f.Mode()&os.ModeSymlink != 0 {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("plugin zip: refusing symlink %q", f.Name)
|
||||
}
|
||||
// Reject zip-slip: cleaned absolute path must stay rooted at the
|
||||
// staging directory.
|
||||
clean := filepath.Clean(f.Name)
|
||||
if strings.HasPrefix(clean, "..") || filepath.IsAbs(clean) || strings.Contains(clean, "..\\") {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("plugin zip: refusing path-traversal entry %q", f.Name)
|
||||
}
|
||||
dest := filepath.Join(stageAbs, clean)
|
||||
destAbs, dErr := filepath.Abs(dest)
|
||||
if dErr != nil {
|
||||
cleanup()
|
||||
return "", dErr
|
||||
}
|
||||
rel, relErr := filepath.Rel(stageAbs, destAbs)
|
||||
if relErr != nil || strings.HasPrefix(rel, "..") || filepath.IsAbs(rel) {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("plugin zip: refusing escape %q", f.Name)
|
||||
}
|
||||
|
||||
if f.Mode().IsDir() {
|
||||
if err := os.MkdirAll(destAbs, 0o750); err != nil {
|
||||
cleanup()
|
||||
return "", err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(destAbs), 0o750); err != nil {
|
||||
cleanup()
|
||||
return "", err
|
||||
}
|
||||
rc, oErr := f.Open()
|
||||
if oErr != nil {
|
||||
cleanup()
|
||||
return "", oErr
|
||||
}
|
||||
out, cErr := os.OpenFile(destAbs, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o600)
|
||||
if cErr != nil {
|
||||
_ = rc.Close()
|
||||
cleanup()
|
||||
return "", cErr
|
||||
}
|
||||
// Cap each file at the remaining uncompressed budget so a zip bomb
|
||||
// can't OOM the host.
|
||||
remaining := maxUncompressedSum - totalUncompressed
|
||||
if remaining <= 0 {
|
||||
_ = rc.Close()
|
||||
_ = out.Close()
|
||||
cleanup()
|
||||
return "", fmt.Errorf("plugin zip: uncompressed total exceeds %d bytes", maxUncompressedSum)
|
||||
}
|
||||
n, copyErr := io.CopyN(out, rc, remaining+1)
|
||||
_ = rc.Close()
|
||||
_ = out.Close()
|
||||
if copyErr != nil && copyErr != io.EOF {
|
||||
cleanup()
|
||||
return "", copyErr
|
||||
}
|
||||
if n > remaining {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("plugin zip: uncompressed total exceeds %d bytes", maxUncompressedSum)
|
||||
}
|
||||
totalUncompressed += n
|
||||
if err := installZipExtract(zr, stageAbs); err != nil {
|
||||
cleanup()
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Stage 2: parse the manifest now that the staging dir is fully populated.
|
||||
manifestPath := filepath.Join(stageAbs, "plugin.json")
|
||||
raw, err := os.ReadFile(manifestPath)
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("plugin zip: missing plugin.json at root: %w", err)
|
||||
}
|
||||
manifest, err := ParseManifest(raw)
|
||||
manifest, err := installZipStagedManifest(stageAbs)
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return "", err
|
||||
}
|
||||
// Validate the staged contents the same way scanPluginDirectory does.
|
||||
if err := rejectSymlinksUnder(stageAbs); err != nil {
|
||||
cleanup()
|
||||
return "", err
|
||||
}
|
||||
wasmPath := filepath.Join(stageAbs, manifest.Entrypoint)
|
||||
if info, statErr := os.Lstat(wasmPath); statErr != nil {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("entrypoint %s missing: %w", manifest.Entrypoint, statErr)
|
||||
} else if info.Mode()&os.ModeSymlink != 0 {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("entrypoint %s is a symlink", manifest.Entrypoint)
|
||||
}
|
||||
|
||||
// Stage 3: atomically rename into the canonical plugin name directory.
|
||||
finalDir := filepath.Join(r.cfg.Directory, manifest.Name)
|
||||
// If a previous version exists, remove it. The store row is replaced by
|
||||
// installFromDisk via the existing UPSERT path.
|
||||
if _, err := os.Stat(finalDir); err == nil {
|
||||
if err := os.RemoveAll(finalDir); err != nil {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("remove existing plugin dir: %w", err)
|
||||
}
|
||||
}
|
||||
if err := os.Rename(stageAbs, finalDir); err != nil {
|
||||
if err := installZipPromote(stageAbs, finalDir); err != nil {
|
||||
cleanup()
|
||||
return "", fmt.Errorf("install rename: %w", err)
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Stage 4: register via the existing on-disk install path.
|
||||
@@ -387,42 +290,187 @@ func (r *Registry) InstallFromZip(ctx context.Context, zipBytes []byte) (string,
|
||||
return manifest.Name, fmt.Errorf("installFromDisk: %w", err)
|
||||
}
|
||||
|
||||
// installFromDisk always registers the fresh instance as disabled and
|
||||
// InstallPlugin's upsert never touches the `enabled` column, so a plugin
|
||||
// that was enabled before this upgrade would otherwise come out the
|
||||
// other side with the store row still saying enabled while the runtime
|
||||
// instance sits inactive. LoadAll's startup path avoids this because it
|
||||
// always runs activateAll afterward; this is the one caller of
|
||||
// installFromDisk that doesn't, so it has to reactivate for itself.
|
||||
// EnablePlugin already rolls the DB flag back if activation fails, so
|
||||
// the two can no longer disagree.
|
||||
if row, err := r.cfg.Store.GetPluginByName(ctx, manifest.Name); err == nil && row != nil && row.Enabled {
|
||||
if err := r.EnablePlugin(ctx, row.ID); err != nil {
|
||||
if errors.Is(err, ErrRuntimeUnavailable) {
|
||||
// Default (non-wazero) build: nothing can activate here, and
|
||||
// leaving EnablePlugin's rollback in place would persistently
|
||||
// disable a plugin the admin left enabled — after a rebuild
|
||||
// with -tags wazero it would silently stay off. Preserve the
|
||||
// enabled intent instead; the next wazero-tagged start's
|
||||
// activateAll does the real activation.
|
||||
if reErr := r.cfg.Store.EnablePlugin(ctx, row.ID); reErr != nil {
|
||||
slog.Warn("plugin: could not preserve enabled flag across runtime-less upgrade",
|
||||
"name", manifest.Name, "err", reErr)
|
||||
} else {
|
||||
r.mu.Lock()
|
||||
if inst, ok := r.byName[manifest.Name]; ok {
|
||||
inst.Enabled = true
|
||||
}
|
||||
r.mu.Unlock()
|
||||
slog.Info("plugin: runtime unavailable, enabled flag preserved across upgrade",
|
||||
"name", manifest.Name)
|
||||
}
|
||||
} else {
|
||||
slog.Warn("plugin: reactivate after upgrade failed", "name", manifest.Name, "err", err)
|
||||
r.installZipReactivate(ctx, manifest.Name)
|
||||
return manifest.Name, nil
|
||||
}
|
||||
|
||||
// installZipExtract writes every entry of zr into the already-created staging
|
||||
// directory stageAbs, enforcing the zip-slip, symlink and uncompressed-size
|
||||
// caps entry by entry before each write. The caller owns stageAbs and removes
|
||||
// it on any error returned here.
|
||||
func installZipExtract(zr *zip.Reader, stageAbs string) error {
|
||||
var totalUncompressed int64
|
||||
for _, f := range zr.File {
|
||||
destAbs, entryErr := installZipEntryDest(f, stageAbs)
|
||||
if entryErr != nil {
|
||||
return entryErr
|
||||
}
|
||||
|
||||
if f.Mode().IsDir() {
|
||||
if err := os.MkdirAll(destAbs, 0o750); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(destAbs), 0o750); err != nil {
|
||||
return err
|
||||
}
|
||||
// Cap each file at the remaining uncompressed budget so a zip bomb
|
||||
// can't OOM the host.
|
||||
remaining := maxUncompressedSum - totalUncompressed
|
||||
n, writeErr := installZipWriteEntry(f, destAbs, remaining)
|
||||
if writeErr != nil {
|
||||
return writeErr
|
||||
}
|
||||
totalUncompressed += n
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// installZipEntryDest validates one zip entry's mode and name and returns the
|
||||
// absolute path it may be written to under stageAbs. Every rejection here is a
|
||||
// hard stop: non-regular modes, symlinks, and any name that escapes stageAbs.
|
||||
func installZipEntryDest(f *zip.File, stageAbs string) (string, error) {
|
||||
// Reject symlinks, devices, and any non-regular file mode.
|
||||
if !f.Mode().IsRegular() && !f.Mode().IsDir() {
|
||||
return "", fmt.Errorf("plugin zip: refusing non-regular entry %q (mode=%v)", f.Name, f.Mode())
|
||||
}
|
||||
if f.Mode()&os.ModeSymlink != 0 {
|
||||
return "", fmt.Errorf("plugin zip: refusing symlink %q", f.Name)
|
||||
}
|
||||
// Reject zip-slip: cleaned absolute path must stay rooted at the
|
||||
// staging directory.
|
||||
clean := filepath.Clean(f.Name)
|
||||
if strings.HasPrefix(clean, "..") || filepath.IsAbs(clean) || strings.Contains(clean, "..\\") {
|
||||
return "", fmt.Errorf("plugin zip: refusing path-traversal entry %q", f.Name)
|
||||
}
|
||||
dest := filepath.Join(stageAbs, clean)
|
||||
destAbs, dErr := filepath.Abs(dest)
|
||||
if dErr != nil {
|
||||
return "", dErr
|
||||
}
|
||||
rel, relErr := filepath.Rel(stageAbs, destAbs)
|
||||
if relErr != nil || strings.HasPrefix(rel, "..") || filepath.IsAbs(rel) {
|
||||
return "", fmt.Errorf("plugin zip: refusing escape %q", f.Name)
|
||||
}
|
||||
return destAbs, nil
|
||||
}
|
||||
|
||||
// installZipWriteEntry copies one regular entry to destAbs, refusing to write
|
||||
// more than remaining bytes — this entry's share of the maxUncompressedSum
|
||||
// budget — and returns how many bytes it wrote.
|
||||
func installZipWriteEntry(f *zip.File, destAbs string, remaining int64) (int64, error) {
|
||||
rc, oErr := f.Open()
|
||||
if oErr != nil {
|
||||
return 0, oErr
|
||||
}
|
||||
out, cErr := os.OpenFile(destAbs, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o600)
|
||||
if cErr != nil {
|
||||
_ = rc.Close()
|
||||
return 0, cErr
|
||||
}
|
||||
if remaining <= 0 {
|
||||
_ = rc.Close()
|
||||
_ = out.Close()
|
||||
return 0, fmt.Errorf("plugin zip: uncompressed total exceeds %d bytes", maxUncompressedSum)
|
||||
}
|
||||
n, copyErr := io.CopyN(out, rc, remaining+1)
|
||||
_ = rc.Close()
|
||||
_ = out.Close()
|
||||
if copyErr != nil && copyErr != io.EOF {
|
||||
return 0, copyErr
|
||||
}
|
||||
if n > remaining {
|
||||
return 0, fmt.Errorf("plugin zip: uncompressed total exceeds %d bytes", maxUncompressedSum)
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// installZipStagedManifest parses the staged plugin.json and holds the staged
|
||||
// tree to the same rules scanPluginDirectory applies to an on-disk plugin (no
|
||||
// symlinks anywhere, entrypoint present and not a symlink).
|
||||
func installZipStagedManifest(stageAbs string) (*Manifest, error) {
|
||||
manifestPath := filepath.Join(stageAbs, "plugin.json")
|
||||
raw, err := os.ReadFile(manifestPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("plugin zip: missing plugin.json at root: %w", err)
|
||||
}
|
||||
manifest, err := ParseManifest(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Validate the staged contents the same way scanPluginDirectory does.
|
||||
if err := rejectSymlinksUnder(stageAbs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
wasmPath := filepath.Join(stageAbs, manifest.Entrypoint)
|
||||
if info, statErr := os.Lstat(wasmPath); statErr != nil {
|
||||
return nil, fmt.Errorf("entrypoint %s missing: %w", manifest.Entrypoint, statErr)
|
||||
} else if info.Mode()&os.ModeSymlink != 0 {
|
||||
return nil, fmt.Errorf("entrypoint %s is a symlink", manifest.Entrypoint)
|
||||
}
|
||||
return manifest, nil
|
||||
}
|
||||
|
||||
// installZipPromote moves the fully validated staging directory into its
|
||||
// canonical plugin-name directory.
|
||||
func installZipPromote(stageAbs, finalDir string) error {
|
||||
// If a previous version exists, remove it. The store row is replaced by
|
||||
// installFromDisk via the existing UPSERT path.
|
||||
if _, err := os.Stat(finalDir); err == nil {
|
||||
if err := os.RemoveAll(finalDir); err != nil {
|
||||
return fmt.Errorf("remove existing plugin dir: %w", err)
|
||||
}
|
||||
}
|
||||
return manifest.Name, nil
|
||||
if err := os.Rename(stageAbs, finalDir); err != nil {
|
||||
return fmt.Errorf("install rename: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// installZipReactivate restores the enabled state of a plugin that was already
|
||||
// enabled before this upgrade.
|
||||
//
|
||||
// installFromDisk always registers the fresh instance as disabled and
|
||||
// InstallPlugin's upsert never touches the `enabled` column, so a plugin
|
||||
// that was enabled before this upgrade would otherwise come out the
|
||||
// other side with the store row still saying enabled while the runtime
|
||||
// instance sits inactive. LoadAll's startup path avoids this because it
|
||||
// always runs activateAll afterward; this is the one caller of
|
||||
// installFromDisk that doesn't, so it has to reactivate for itself.
|
||||
// EnablePlugin already rolls the DB flag back if activation fails, so
|
||||
// the two can no longer disagree.
|
||||
func (r *Registry) installZipReactivate(ctx context.Context, name string) {
|
||||
row, rowErr := r.cfg.Store.GetPluginByName(ctx, name)
|
||||
if rowErr != nil || row == nil || !row.Enabled {
|
||||
return
|
||||
}
|
||||
err := r.EnablePlugin(ctx, row.ID)
|
||||
if err == nil {
|
||||
return
|
||||
}
|
||||
if !errors.Is(err, ErrRuntimeUnavailable) {
|
||||
slog.Warn("plugin: reactivate after upgrade failed", "name", name, "err", err)
|
||||
return
|
||||
}
|
||||
// Default (non-wazero) build: nothing can activate here, and
|
||||
// leaving EnablePlugin's rollback in place would persistently
|
||||
// disable a plugin the admin left enabled — after a rebuild
|
||||
// with -tags wazero it would silently stay off. Preserve the
|
||||
// enabled intent instead; the next wazero-tagged start's
|
||||
// activateAll does the real activation.
|
||||
if reErr := r.cfg.Store.EnablePlugin(ctx, row.ID); reErr != nil {
|
||||
slog.Warn("plugin: could not preserve enabled flag across runtime-less upgrade",
|
||||
"name", name, "err", reErr)
|
||||
return
|
||||
}
|
||||
r.mu.Lock()
|
||||
if inst, ok := r.byName[name]; ok {
|
||||
inst.Enabled = true
|
||||
}
|
||||
r.mu.Unlock()
|
||||
slog.Info("plugin: runtime unavailable, enabled flag preserved across upgrade",
|
||||
"name", name)
|
||||
}
|
||||
|
||||
// bytesReaderAt is a tiny wrapper that satisfies io.ReaderAt for a byte
|
||||
|
||||
+213
-165
@@ -28,55 +28,11 @@ func (s *MessageService) SendMessage(ctx context.Context, p SendMessageParams) (
|
||||
span.End()
|
||||
}()
|
||||
|
||||
// Rate limit.
|
||||
ratKey := auth.Key("chat", p.UserID)
|
||||
if s.limiter != nil && !s.limiter.Allow(ratKey, 10, time.Second) {
|
||||
return nil, ErrRateLimited
|
||||
}
|
||||
|
||||
if p.ChannelID <= 0 {
|
||||
return nil, fmt.Errorf("%w: channel_id must be a positive integer", ErrBadRequest)
|
||||
}
|
||||
|
||||
ch, err := s.st.GetChannel(ctx, p.ChannelID)
|
||||
if err != nil || ch == nil {
|
||||
return nil, fmt.Errorf("%w: channel not found", ErrNotFound)
|
||||
}
|
||||
|
||||
isDM := ch.Type == "dm"
|
||||
|
||||
// Permission check. Also refuses a write against an archived channel — see
|
||||
// requireChannelWritable in message_perms.go, the shared gate every
|
||||
// message write sink routes through.
|
||||
if err := s.checkSendPermission(ctx, p.UserID, ch); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Validate and sanitize content.
|
||||
content, err := sanitizeContent(p.Content, len(p.AttachmentIDs) > 0)
|
||||
ch, content, err := s.sendMessagePrecheck(ctx, p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Attachment permission (non-DM).
|
||||
if !isDM && len(p.AttachmentIDs) > 0 {
|
||||
if !s.perms.HasChannelPerm(ctx, p.UserID, p.ChannelID, permissions.AttachFiles) {
|
||||
return nil, fmt.Errorf("%w: missing ATTACH_FILES permission", ErrForbidden)
|
||||
}
|
||||
}
|
||||
|
||||
// Slow mode (non-DM only). Deliberately checked last, after content and
|
||||
// attachment validation: Allow() below records the cooldown timestamp the
|
||||
// instant it returns true, so a send that fails validation after this
|
||||
// point must not have already spent the once-per-window token — that
|
||||
// would lock the composer for up to ch.SlowMode seconds for a send that
|
||||
// never actually posted anything.
|
||||
if !isDM && ch.SlowMode > 0 && !s.perms.HasChannelPerm(ctx, p.UserID, p.ChannelID, permissions.ManageMessages) {
|
||||
slowKey := auth.Key(auth.Key("slow", p.UserID), p.ChannelID)
|
||||
if s.limiter != nil && !s.limiter.Allow(slowKey, 1, time.Duration(ch.SlowMode)*time.Second) {
|
||||
return nil, fmt.Errorf("%w: channel has %ds slow mode", ErrSlowMode, ch.SlowMode)
|
||||
}
|
||||
}
|
||||
isDM := ch.Type == "dm"
|
||||
|
||||
// Resolve mentions against the sanitized content, before the insert, so the
|
||||
// row and its mention set are written together. Unknown @words and an
|
||||
@@ -93,51 +49,9 @@ func (s *MessageService) SendMessage(ctx context.Context, p SendMessageParams) (
|
||||
}
|
||||
msgID := msg.ID
|
||||
|
||||
// Link attachments. Ownership is enforced atomically inside the link
|
||||
// UPDATE itself (uploader match + still unlinked), so another user's
|
||||
// upload, an already-linked attachment, or a nonexistent id is skipped by
|
||||
// the statement — no check-then-link race and no N+1 pre-verification.
|
||||
var attachments []db.AttachmentInfo
|
||||
if len(p.AttachmentIDs) > 0 {
|
||||
linked, linkErr := s.st.LinkAttachmentsToMessage(ctx, msgID, p.UserID, p.AttachmentIDs)
|
||||
if linkErr != nil {
|
||||
slog.Error("MessageService.SendMessage LinkAttachments", "err", linkErr, "msg_id", msgID)
|
||||
// Cleanup: soft-delete the message. The compensating delete must run
|
||||
// even when the link failed because the request ctx was canceled.
|
||||
if delErr := s.st.DeleteMessage(context.WithoutCancel(ctx), msgID, p.UserID, true); delErr != nil {
|
||||
slog.Error("MessageService.SendMessage DeleteMessage (cleanup)", "err", delErr, "msg_id", msgID)
|
||||
}
|
||||
return nil, fmt.Errorf("%w: failed to send message with attachments", ErrInternal)
|
||||
}
|
||||
if linked < int64(len(p.AttachmentIDs)) {
|
||||
slog.Warn("MessageService.SendMessage: skipped attachments (not owned, already linked, or missing)",
|
||||
"msg_id", msgID, "user_id", p.UserID, "requested", len(p.AttachmentIDs), "linked", linked)
|
||||
}
|
||||
if linked == 0 && content == "" {
|
||||
// sanitizeContent waived the empty-content check purely on the
|
||||
// requested attachment count, before any link attempt. None of
|
||||
// them actually linked (all missing, foreign, or already
|
||||
// linked — e.g. a retry of a partially-completed send), so the
|
||||
// row that just committed has no content and no attachments.
|
||||
// Compensate the same way the linkErr path above does, rather
|
||||
// than broadcasting a blank message.
|
||||
if delErr := s.st.DeleteMessage(context.WithoutCancel(ctx), msgID, p.UserID, true); delErr != nil {
|
||||
slog.Error("MessageService.SendMessage DeleteMessage (empty-after-link cleanup)", "err", delErr, "msg_id", msgID)
|
||||
}
|
||||
return nil, fmt.Errorf("%w: message content cannot be empty", ErrBadRequest)
|
||||
}
|
||||
if linked > 0 {
|
||||
// Detached from ctx for the same reason as the compensating deletes
|
||||
// above: the link already committed, so a request ctx canceled the
|
||||
// instant it returns (sender disconnects right after) must not turn
|
||||
// a successful attachment-only send into a blank broadcast bubble.
|
||||
attMap, attErr := s.st.GetAttachmentsByMessageIDs(context.WithoutCancel(ctx), []int64{msgID})
|
||||
if attErr != nil {
|
||||
slog.Error("MessageService.SendMessage GetAttachments", "err", attErr)
|
||||
} else {
|
||||
attachments = attMap[msgID]
|
||||
}
|
||||
}
|
||||
attachments, err := s.sendMessageLinkAttachments(ctx, p, msgID, content)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Advance the author's own read state past the message they just sent.
|
||||
@@ -171,60 +85,8 @@ func (s *MessageService) SendMessage(ctx context.Context, p SendMessageParams) (
|
||||
}
|
||||
|
||||
// DM path: open DM for recipients.
|
||||
if isDM {
|
||||
// The message is already committed, so everything below must survive
|
||||
// the sender's connection dropping the instant the write commits — the
|
||||
// same reason the compensating deletes and applyMentionCounts below
|
||||
// detach from ctx. Without WithoutCancel, a canceled request ctx here
|
||||
// silently drops every recipient from the fan-out (ParticipantIDs
|
||||
// stays nil), skips re-opening the recipient's dm_open_state, and
|
||||
// degrades the payload shape — with no error surfaced to anyone: the
|
||||
// sender sees chat_send_ok and the other participant never gets the
|
||||
// message live.
|
||||
bgCtx := context.WithoutCancel(ctx)
|
||||
participantIDs, pErr := s.st.GetDMParticipantIDs(bgCtx, p.ChannelID)
|
||||
if pErr != nil {
|
||||
slog.Error("MessageService.SendMessage GetDMParticipantIDs", "err", pErr, "channel_id", p.ChannelID)
|
||||
return result, nil // Message saved, skip DM side effects.
|
||||
}
|
||||
result.ParticipantIDs = participantIDs
|
||||
|
||||
sender, _ := s.st.GetUserByID(bgCtx, p.UserID)
|
||||
result.SenderUser = sender
|
||||
|
||||
// Viewer-neutral (viewerID 0 matches nobody, so every status is
|
||||
// broadcast-collapsed); the ws layer re-derives "who is the recipient"
|
||||
// per addressee. A read failure is non-fatal — the message is already
|
||||
// committed, and the caller falls back to the 1:1 shape.
|
||||
if participants, partErr := s.st.GetDMParticipants(bgCtx, p.ChannelID, 0); partErr == nil {
|
||||
result.DMParticipants = participants
|
||||
} else {
|
||||
slog.Warn("MessageService.SendMessage GetDMParticipants", "err", partErr, "channel_id", p.ChannelID)
|
||||
}
|
||||
if isGroup, gErr := s.st.IsGroupDM(bgCtx, p.ChannelID); gErr == nil {
|
||||
result.DMIsGroup = isGroup
|
||||
}
|
||||
|
||||
for _, pid := range participantIDs {
|
||||
if pid == p.UserID {
|
||||
continue
|
||||
}
|
||||
// OpenDM is INSERT OR IGNORE and idempotent: opened reports whether
|
||||
// this call actually inserted the row. Only a genuine (re)open goes
|
||||
// into OpenedDMFor — the ws layer emits a dm_channel_open per id in
|
||||
// that slice, and each one bumps the hub's global visibility
|
||||
// watermark, forcing every other connected client's next reconnect
|
||||
// onto a full resync. An already-open DM must not pay that cost on
|
||||
// every single message.
|
||||
opened, openErr := s.st.OpenDM(bgCtx, pid, p.ChannelID)
|
||||
if openErr != nil {
|
||||
slog.Error("MessageService.SendMessage OpenDM", "err", openErr, "recipient_id", pid, "channel_id", p.ChannelID)
|
||||
continue
|
||||
}
|
||||
if opened {
|
||||
result.OpenedDMFor = append(result.OpenedDMFor, pid)
|
||||
}
|
||||
}
|
||||
if isDM && !s.sendMessageDMSideEffects(ctx, p, result) {
|
||||
return result, nil // Message saved, skip DM side effects.
|
||||
}
|
||||
|
||||
// Mention badges run off the send path: the message is already committed, so
|
||||
@@ -245,6 +107,182 @@ func (s *MessageService) SendMessage(ctx context.Context, p SendMessageParams) (
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// sendMessagePrecheck runs every gate a send must clear before anything is
|
||||
// written: rate limit, channel lookup, send permission, content sanitization,
|
||||
// attachment permission and slow mode. It returns the resolved channel and the
|
||||
// sanitized content for the caller to persist.
|
||||
func (s *MessageService) sendMessagePrecheck(ctx context.Context, p SendMessageParams) (*db.Channel, string, error) {
|
||||
// Rate limit.
|
||||
ratKey := auth.Key("chat", p.UserID)
|
||||
if s.limiter != nil && !s.limiter.Allow(ratKey, 10, time.Second) {
|
||||
return nil, "", ErrRateLimited
|
||||
}
|
||||
|
||||
if p.ChannelID <= 0 {
|
||||
return nil, "", fmt.Errorf("%w: channel_id must be a positive integer", ErrBadRequest)
|
||||
}
|
||||
|
||||
ch, err := s.st.GetChannel(ctx, p.ChannelID)
|
||||
if err != nil || ch == nil {
|
||||
return nil, "", fmt.Errorf("%w: channel not found", ErrNotFound)
|
||||
}
|
||||
|
||||
isDM := ch.Type == "dm"
|
||||
|
||||
// Permission check. Also refuses a write against an archived channel — see
|
||||
// requireChannelWritable in message_perms.go, the shared gate every
|
||||
// message write sink routes through.
|
||||
if err := s.checkSendPermission(ctx, p.UserID, ch); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
|
||||
// Validate and sanitize content.
|
||||
content, err := sanitizeContent(p.Content, len(p.AttachmentIDs) > 0)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
|
||||
// Attachment permission (non-DM).
|
||||
if !isDM && len(p.AttachmentIDs) > 0 {
|
||||
if !s.perms.HasChannelPerm(ctx, p.UserID, p.ChannelID, permissions.AttachFiles) {
|
||||
return nil, "", fmt.Errorf("%w: missing ATTACH_FILES permission", ErrForbidden)
|
||||
}
|
||||
}
|
||||
|
||||
// Slow mode (non-DM only). Deliberately checked last, after content and
|
||||
// attachment validation: Allow() below records the cooldown timestamp the
|
||||
// instant it returns true, so a send that fails validation after this
|
||||
// point must not have already spent the once-per-window token — that
|
||||
// would lock the composer for up to ch.SlowMode seconds for a send that
|
||||
// never actually posted anything.
|
||||
if !isDM && ch.SlowMode > 0 && !s.perms.HasChannelPerm(ctx, p.UserID, p.ChannelID, permissions.ManageMessages) {
|
||||
slowKey := auth.Key(auth.Key("slow", p.UserID), p.ChannelID)
|
||||
if s.limiter != nil && !s.limiter.Allow(slowKey, 1, time.Duration(ch.SlowMode)*time.Second) {
|
||||
return nil, "", fmt.Errorf("%w: channel has %ds slow mode", ErrSlowMode, ch.SlowMode)
|
||||
}
|
||||
}
|
||||
|
||||
return ch, content, nil
|
||||
}
|
||||
|
||||
// sendMessageLinkAttachments links the requested uploads to the message row
|
||||
// that was just committed and returns the attachment data the broadcast needs.
|
||||
// Nothing requested is a no-op. A failed link, or a link that attached nothing
|
||||
// to a message with no content of its own, compensates by soft-deleting the
|
||||
// row and returns the error the caller surfaces to the sender.
|
||||
func (s *MessageService) sendMessageLinkAttachments(ctx context.Context, p SendMessageParams, msgID int64, content string) ([]db.AttachmentInfo, error) {
|
||||
if len(p.AttachmentIDs) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// Link attachments. Ownership is enforced atomically inside the link
|
||||
// UPDATE itself (uploader match + still unlinked), so another user's
|
||||
// upload, an already-linked attachment, or a nonexistent id is skipped by
|
||||
// the statement — no check-then-link race and no N+1 pre-verification.
|
||||
var attachments []db.AttachmentInfo
|
||||
linked, linkErr := s.st.LinkAttachmentsToMessage(ctx, msgID, p.UserID, p.AttachmentIDs)
|
||||
if linkErr != nil {
|
||||
slog.Error("MessageService.SendMessage LinkAttachments", "err", linkErr, "msg_id", msgID)
|
||||
// Cleanup: soft-delete the message. The compensating delete must run
|
||||
// even when the link failed because the request ctx was canceled.
|
||||
if delErr := s.st.DeleteMessage(context.WithoutCancel(ctx), msgID, p.UserID, true); delErr != nil {
|
||||
slog.Error("MessageService.SendMessage DeleteMessage (cleanup)", "err", delErr, "msg_id", msgID)
|
||||
}
|
||||
return nil, fmt.Errorf("%w: failed to send message with attachments", ErrInternal)
|
||||
}
|
||||
if linked < int64(len(p.AttachmentIDs)) {
|
||||
slog.Warn("MessageService.SendMessage: skipped attachments (not owned, already linked, or missing)",
|
||||
"msg_id", msgID, "user_id", p.UserID, "requested", len(p.AttachmentIDs), "linked", linked)
|
||||
}
|
||||
if linked == 0 && content == "" {
|
||||
// sanitizeContent waived the empty-content check purely on the
|
||||
// requested attachment count, before any link attempt. None of
|
||||
// them actually linked (all missing, foreign, or already
|
||||
// linked — e.g. a retry of a partially-completed send), so the
|
||||
// row that just committed has no content and no attachments.
|
||||
// Compensate the same way the linkErr path above does, rather
|
||||
// than broadcasting a blank message.
|
||||
if delErr := s.st.DeleteMessage(context.WithoutCancel(ctx), msgID, p.UserID, true); delErr != nil {
|
||||
slog.Error("MessageService.SendMessage DeleteMessage (empty-after-link cleanup)", "err", delErr, "msg_id", msgID)
|
||||
}
|
||||
return nil, fmt.Errorf("%w: message content cannot be empty", ErrBadRequest)
|
||||
}
|
||||
if linked > 0 {
|
||||
// Detached from ctx for the same reason as the compensating deletes
|
||||
// above: the link already committed, so a request ctx canceled the
|
||||
// instant it returns (sender disconnects right after) must not turn
|
||||
// a successful attachment-only send into a blank broadcast bubble.
|
||||
attMap, attErr := s.st.GetAttachmentsByMessageIDs(context.WithoutCancel(ctx), []int64{msgID})
|
||||
if attErr != nil {
|
||||
slog.Error("MessageService.SendMessage GetAttachments", "err", attErr)
|
||||
} else {
|
||||
attachments = attMap[msgID]
|
||||
}
|
||||
}
|
||||
return attachments, nil
|
||||
}
|
||||
|
||||
// sendMessageDMSideEffects fills in the DM-specific fields of result and
|
||||
// (re)opens the DM for every other participant. It reports false when the
|
||||
// participant lookup failed, which is the one case where the caller returns
|
||||
// the already-saved message without the remaining side effects.
|
||||
func (s *MessageService) sendMessageDMSideEffects(ctx context.Context, p SendMessageParams, result *SendMessageResult) bool {
|
||||
// The message is already committed, so everything below must survive
|
||||
// the sender's connection dropping the instant the write commits — the
|
||||
// same reason the compensating deletes and applyMentionCounts below
|
||||
// detach from ctx. Without WithoutCancel, a canceled request ctx here
|
||||
// silently drops every recipient from the fan-out (ParticipantIDs
|
||||
// stays nil), skips re-opening the recipient's dm_open_state, and
|
||||
// degrades the payload shape — with no error surfaced to anyone: the
|
||||
// sender sees chat_send_ok and the other participant never gets the
|
||||
// message live.
|
||||
bgCtx := context.WithoutCancel(ctx)
|
||||
participantIDs, pErr := s.st.GetDMParticipantIDs(bgCtx, p.ChannelID)
|
||||
if pErr != nil {
|
||||
slog.Error("MessageService.SendMessage GetDMParticipantIDs", "err", pErr, "channel_id", p.ChannelID)
|
||||
return false
|
||||
}
|
||||
result.ParticipantIDs = participantIDs
|
||||
|
||||
sender, _ := s.st.GetUserByID(bgCtx, p.UserID)
|
||||
result.SenderUser = sender
|
||||
|
||||
// Viewer-neutral (viewerID 0 matches nobody, so every status is
|
||||
// broadcast-collapsed); the ws layer re-derives "who is the recipient"
|
||||
// per addressee. A read failure is non-fatal — the message is already
|
||||
// committed, and the caller falls back to the 1:1 shape.
|
||||
if participants, partErr := s.st.GetDMParticipants(bgCtx, p.ChannelID, 0); partErr == nil {
|
||||
result.DMParticipants = participants
|
||||
} else {
|
||||
slog.Warn("MessageService.SendMessage GetDMParticipants", "err", partErr, "channel_id", p.ChannelID)
|
||||
}
|
||||
if isGroup, gErr := s.st.IsGroupDM(bgCtx, p.ChannelID); gErr == nil {
|
||||
result.DMIsGroup = isGroup
|
||||
}
|
||||
|
||||
for _, pid := range participantIDs {
|
||||
if pid == p.UserID {
|
||||
continue
|
||||
}
|
||||
// OpenDM is INSERT OR IGNORE and idempotent: opened reports whether
|
||||
// this call actually inserted the row. Only a genuine (re)open goes
|
||||
// into OpenedDMFor — the ws layer emits a dm_channel_open per id in
|
||||
// that slice, and each one bumps the hub's global visibility
|
||||
// watermark, forcing every other connected client's next reconnect
|
||||
// onto a full resync. An already-open DM must not pay that cost on
|
||||
// every single message.
|
||||
opened, openErr := s.st.OpenDM(bgCtx, pid, p.ChannelID)
|
||||
if openErr != nil {
|
||||
slog.Error("MessageService.SendMessage OpenDM", "err", openErr, "recipient_id", pid, "channel_id", p.ChannelID)
|
||||
continue
|
||||
}
|
||||
if opened {
|
||||
result.OpenedDMFor = append(result.OpenedDMFor, pid)
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// EditMessage validates and persists a message edit.
|
||||
func (s *MessageService) EditMessage(ctx context.Context, userID, msgID int64, rawContent string) (*EditMessageResult, error) {
|
||||
// Rate limit.
|
||||
@@ -283,26 +321,8 @@ func (s *MessageService) EditMessage(ctx context.Context, userID, msgID int64, r
|
||||
chanType := ch.Type
|
||||
isDM := chanType == "dm"
|
||||
|
||||
if isDM {
|
||||
ok, dmErr := s.st.IsDMParticipant(ctx, userID, msg.ChannelID)
|
||||
if dmErr != nil || !ok {
|
||||
return nil, fmt.Errorf("%w: cannot edit this message", ErrForbidden)
|
||||
}
|
||||
if blkErr := requireDMNotBlocked(ctx, s.st, userID, msg.ChannelID); blkErr != nil {
|
||||
return nil, blkErr
|
||||
}
|
||||
} else if permErr := s.checkSendPermission(ctx, userID, ch); permErr != nil {
|
||||
// An edit injects new text into the channel and is fanned out to every
|
||||
// reader, so it must clear the same gate as a send rather than
|
||||
// SEND_MESSAGES alone: READ_MESSAGES so a role locked out of a private
|
||||
// channel (the panel's "Can access" toggle denies
|
||||
// READ_MESSAGES|CONNECT_VOICE and leaves SEND_MESSAGES intact) cannot
|
||||
// rewrite its old posts, and the announcement rule so a demoted
|
||||
// moderator cannot rewrite a trusted broadcast. Mirrors DeleteMessage,
|
||||
// SetMessagePinned and handleReaction, which already require
|
||||
// READ_MESSAGES. The reason is collapsed into this sink's single opaque
|
||||
// error so the reply stays an ownership/permission non-oracle.
|
||||
return nil, fmt.Errorf("%w: cannot edit this message", ErrForbidden)
|
||||
if accessErr := s.editMessageCheckAccess(ctx, userID, msg.ChannelID, ch, isDM); accessErr != nil {
|
||||
return nil, accessErr
|
||||
}
|
||||
|
||||
// EditMessage checks ownership internally and returns the updated row via
|
||||
@@ -354,6 +374,34 @@ func (s *MessageService) EditMessage(ctx context.Context, userID, msgID int64, r
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// editMessageCheckAccess gates an edit on the channel the message lives in:
|
||||
// participation plus the block check for a DM, the shared send gate for
|
||||
// everything else. isDM is the caller's already-computed ch.Type == "dm".
|
||||
func (s *MessageService) editMessageCheckAccess(ctx context.Context, userID, channelID int64, ch *db.Channel, isDM bool) error {
|
||||
if isDM {
|
||||
ok, dmErr := s.st.IsDMParticipant(ctx, userID, channelID)
|
||||
if dmErr != nil || !ok {
|
||||
return fmt.Errorf("%w: cannot edit this message", ErrForbidden)
|
||||
}
|
||||
if blkErr := requireDMNotBlocked(ctx, s.st, userID, channelID); blkErr != nil {
|
||||
return blkErr
|
||||
}
|
||||
} else if permErr := s.checkSendPermission(ctx, userID, ch); permErr != nil {
|
||||
// An edit injects new text into the channel and is fanned out to every
|
||||
// reader, so it must clear the same gate as a send rather than
|
||||
// SEND_MESSAGES alone: READ_MESSAGES so a role locked out of a private
|
||||
// channel (the panel's "Can access" toggle denies
|
||||
// READ_MESSAGES|CONNECT_VOICE and leaves SEND_MESSAGES intact) cannot
|
||||
// rewrite its old posts, and the announcement rule so a demoted
|
||||
// moderator cannot rewrite a trusted broadcast. Mirrors DeleteMessage,
|
||||
// SetMessagePinned and handleReaction, which already require
|
||||
// READ_MESSAGES. The reason is collapsed into this sink's single opaque
|
||||
// error so the reply stays an ownership/permission non-oracle.
|
||||
return fmt.Errorf("%w: cannot edit this message", ErrForbidden)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteMessage validates and soft-deletes a message.
|
||||
func (s *MessageService) DeleteMessage(ctx context.Context, userID, msgID int64) (*DeleteMessageResult, error) {
|
||||
// Rate limit.
|
||||
|
||||
@@ -102,53 +102,11 @@ func (s *MessageService) handleReaction(ctx context.Context, userID, msgID int64
|
||||
return nil, fmt.Errorf("%w: cannot react to deleted message", ErrBadRequest)
|
||||
}
|
||||
|
||||
// Fail closed, mirroring EditMessage/DeleteMessage (message_crud.go): a
|
||||
// lookup failure must not fall through to the non-DM permission branch
|
||||
// below. That branch passes on the base role mask alone
|
||||
// (READ_MESSAGES|ADD_REACTIONS, no per-channel override exists for a DM),
|
||||
// skipping both IsDMParticipant and requireDMNotBlocked entirely.
|
||||
ch, chErr := s.st.GetChannel(ctx, msg.ChannelID)
|
||||
if chErr != nil || ch == nil {
|
||||
return nil, fmt.Errorf("%w: cannot react to this message", ErrForbidden)
|
||||
}
|
||||
isDM := ch.Type == "dm"
|
||||
|
||||
// Archived channels are read-only. handleReaction bypasses
|
||||
// checkSendPermission (it runs its own DM/permission branch below), so it
|
||||
// needs the shared gate directly — see requireChannelWritable in
|
||||
// message_perms.go.
|
||||
if err := requireChannelWritable(ch); err != nil {
|
||||
participantIDs, isDM, err := s.reactionAudience(ctx, userID, msg.ChannelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var participantIDs []int64
|
||||
if isDM {
|
||||
ok, dmErr := s.st.IsDMParticipant(ctx, userID, msg.ChannelID)
|
||||
if dmErr != nil || !ok {
|
||||
return nil, fmt.Errorf("%w: not a DM participant", ErrBadRequest)
|
||||
}
|
||||
if blkErr := requireDMNotBlocked(ctx, s.st, userID, msg.ChannelID); blkErr != nil {
|
||||
return nil, blkErr
|
||||
}
|
||||
// Resolve the fan-out audience before mutating anything. Participants
|
||||
// are unaffected by the reaction itself, so failing here is cheap;
|
||||
// fetching this after AddReaction/RemoveReaction commits (as this
|
||||
// used to) risked a reaction persisted with no participant list to
|
||||
// broadcast it to, which reactionV2Handler would then fan out to
|
||||
// nobody while reporting success to the caller.
|
||||
ids, pErr := s.st.GetDMParticipantIDs(ctx, msg.ChannelID)
|
||||
if pErr != nil {
|
||||
slog.Error("MessageService.handleReaction GetDMParticipantIDs", "err", pErr, "channel_id", msg.ChannelID)
|
||||
return nil, fmt.Errorf("%w: failed to resolve DM participants", ErrInternal)
|
||||
}
|
||||
participantIDs = ids
|
||||
} else if !s.perms.HasChannelPerm(ctx, userID, msg.ChannelID, permissions.ReadMessages|permissions.AddReactions) {
|
||||
// Require READ_MESSAGES in addition to ADD_REACTIONS so a user cannot
|
||||
// react in a channel they cannot read. Mirrors checkSendPermission,
|
||||
// which requires ReadMessages|SendMessages for non-DM sends.
|
||||
return nil, fmt.Errorf("%w: missing ADD_REACTIONS permission", ErrForbidden)
|
||||
}
|
||||
|
||||
action := "add"
|
||||
if add {
|
||||
if err := s.st.AddReaction(ctx, msgID, userID, emoji); err != nil {
|
||||
@@ -178,3 +136,61 @@ func (s *MessageService) handleReaction(ctx context.Context, userID, msgID int64
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// reactionAudience resolves the channel a message lives in and enforces the
|
||||
// channel-scoped gates on reacting in it — archived, DM participation, DM
|
||||
// block, and the non-DM READ_MESSAGES|ADD_REACTIONS check. The gates its
|
||||
// caller keeps (rate limit, message id, emoji validity, deleted message) stay
|
||||
// in handleReaction and still run first. It also returns the DM participant
|
||||
// ids, resolved here so they exist before anything is mutated. The order of
|
||||
// the checks is load-bearing and unchanged.
|
||||
func (s *MessageService) reactionAudience(ctx context.Context, userID, channelID int64) ([]int64, bool, error) {
|
||||
// Fail closed, mirroring EditMessage/DeleteMessage (message_crud.go): a
|
||||
// lookup failure must not fall through to the non-DM permission branch
|
||||
// below. That branch passes on the base role mask alone
|
||||
// (READ_MESSAGES|ADD_REACTIONS, no per-channel override exists for a DM),
|
||||
// skipping both IsDMParticipant and requireDMNotBlocked entirely.
|
||||
ch, chErr := s.st.GetChannel(ctx, channelID)
|
||||
if chErr != nil || ch == nil {
|
||||
return nil, false, fmt.Errorf("%w: cannot react to this message", ErrForbidden)
|
||||
}
|
||||
isDM := ch.Type == "dm"
|
||||
|
||||
// Archived channels are read-only. handleReaction bypasses
|
||||
// checkSendPermission (it runs its own DM/permission branch below), so it
|
||||
// needs the shared gate directly — see requireChannelWritable in
|
||||
// message_perms.go.
|
||||
if err := requireChannelWritable(ch); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
var participantIDs []int64
|
||||
if isDM {
|
||||
ok, dmErr := s.st.IsDMParticipant(ctx, userID, channelID)
|
||||
if dmErr != nil || !ok {
|
||||
return nil, false, fmt.Errorf("%w: not a DM participant", ErrBadRequest)
|
||||
}
|
||||
if blkErr := requireDMNotBlocked(ctx, s.st, userID, channelID); blkErr != nil {
|
||||
return nil, false, blkErr
|
||||
}
|
||||
// Resolve the fan-out audience before mutating anything. Participants
|
||||
// are unaffected by the reaction itself, so failing here is cheap;
|
||||
// fetching this after AddReaction/RemoveReaction commits (as this
|
||||
// used to) risked a reaction persisted with no participant list to
|
||||
// broadcast it to, which reactionV2Handler would then fan out to
|
||||
// nobody while reporting success to the caller.
|
||||
ids, pErr := s.st.GetDMParticipantIDs(ctx, channelID)
|
||||
if pErr != nil {
|
||||
slog.Error("MessageService.handleReaction GetDMParticipantIDs", "err", pErr, "channel_id", channelID)
|
||||
return nil, false, fmt.Errorf("%w: failed to resolve DM participants", ErrInternal)
|
||||
}
|
||||
participantIDs = ids
|
||||
} else if !s.perms.HasChannelPerm(ctx, userID, channelID, permissions.ReadMessages|permissions.AddReactions) {
|
||||
// Require READ_MESSAGES in addition to ADD_REACTIONS so a user cannot
|
||||
// react in a channel they cannot read. Mirrors checkSendPermission,
|
||||
// which requires ReadMessages|SendMessages for non-DM sends.
|
||||
return nil, false, fmt.Errorf("%w: missing ADD_REACTIONS permission", ErrForbidden)
|
||||
}
|
||||
|
||||
return participantIDs, isDM, nil
|
||||
}
|
||||
|
||||
@@ -86,7 +86,7 @@ func TestHandleVoiceTokenRefresh_NilUser(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// ─── rollbackVoiceJoin (voice_join.go:239) ──────────────────────────────────
|
||||
// ─── rollbackVoiceJoin (voice_join.go) ──────────────────────────────────
|
||||
|
||||
func TestRollbackVoiceJoin_ClearsVoiceStateAndBroadcasts(t *testing.T) {
|
||||
hub, database := newCoverageHub(t)
|
||||
@@ -152,7 +152,7 @@ func TestRollbackVoiceJoin_NoDBState_DoesNotPanic(t *testing.T) {
|
||||
// "DELETE ... WHERE user_id = ?" with no channel/token condition wipes
|
||||
// whatever the user is currently in, not just the failed join.
|
||||
//
|
||||
// This covers the token-generation-failure call site (voice_join.go:300),
|
||||
// This covers the token-generation-failure call site (voiceJoinGrantToken),
|
||||
// which already holds the failed join's own JoinedAt by the time it rolls
|
||||
// back — that value must scope the delete instead of being discarded.
|
||||
func TestRollbackVoiceJoin_StaleTokenDoesNotDeleteNewerJoin(t *testing.T) {
|
||||
@@ -193,7 +193,7 @@ func TestRollbackVoiceJoin_StaleTokenDoesNotDeleteNewerJoin(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// OC-0044: mirrors the GetVoiceState-failure call site (voice_join.go:208),
|
||||
// OC-0044: mirrors the GetVoiceState-failure call site (voiceJoinPersist),
|
||||
// which never learns the failed join's own JoinedAt and so rolls back with
|
||||
// an empty token. That must not degrade to the old unconditional
|
||||
// "DELETE ... WHERE user_id = ?" — it must re-read the row and refuse to
|
||||
|
||||
+89
-60
@@ -26,70 +26,14 @@ func (h *Hub) HandleVoiceLeaveForTest(c *Client) {
|
||||
|
||||
// handleMessage parses the envelope and dispatches to the appropriate handler.
|
||||
func (h *Hub) handleMessage(c *Client, raw []byte) {
|
||||
// Periodic session expiry check: every SessionCheckInterval messages,
|
||||
// re-validate the session token. This catches sessions that are revoked or
|
||||
// expire while the WebSocket connection is still open.
|
||||
c.mu.Lock()
|
||||
c.msgCount++
|
||||
shouldCheck := c.msgCount >= SessionCheckInterval
|
||||
if shouldCheck {
|
||||
c.msgCount = 0
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
if shouldCheck && c.tokenHash != "" {
|
||||
result, dbErr := h.db.GetSessionWithBanStatus(c.ctx, c.tokenHash)
|
||||
if dbErr != nil || result == nil || auth.IsSessionExpired(result.ExpiresAt) {
|
||||
slog.Info("ws session expired, closing connection", "user_id", c.userID)
|
||||
h.kickClient(c)
|
||||
return
|
||||
}
|
||||
tempUser := &db.User{Banned: result.Banned, BanExpires: result.BanExpires}
|
||||
if auth.IsEffectivelyBanned(tempUser) {
|
||||
slog.Info("ws user banned, closing connection", "user_id", c.userID)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeBanned, "you are banned"))
|
||||
h.kickClient(c)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
var env envelope
|
||||
if err := json.Unmarshal(raw, &env); err != nil {
|
||||
c.mu.Lock()
|
||||
c.invalidCount++
|
||||
count := c.invalidCount
|
||||
c.mu.Unlock()
|
||||
|
||||
slog.Warn("ws handleMessage invalid JSON", "user_id", c.userID, "err", err, "invalid_count", count)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInvalidJSON, "message must be valid JSON"))
|
||||
|
||||
if count >= 10 {
|
||||
slog.Warn("ws too many invalid messages, closing connection", "user_id", c.userID, "invalid_count", count)
|
||||
h.kickClient(c)
|
||||
}
|
||||
if h.handleMessageSessionRecheck(c) {
|
||||
return
|
||||
}
|
||||
|
||||
// Valid parse — reset consecutive invalid counter.
|
||||
c.mu.Lock()
|
||||
c.invalidCount = 0
|
||||
c.mu.Unlock()
|
||||
|
||||
// Cap client-controlled fields before logging to prevent log injection
|
||||
// and unbounded log entries.
|
||||
msgType := env.Type
|
||||
if len(msgType) > 64 {
|
||||
msgType = msgType[:64]
|
||||
env, msgType, reqID, ok := h.handleMessageDecode(c, raw)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
reqID := env.ID
|
||||
if len(reqID) > 64 {
|
||||
reqID = reqID[:64]
|
||||
}
|
||||
|
||||
// Correlation attrs (user_id/msg_type/req_id) are inlined at each log site
|
||||
// below rather than bound via slog.With — the With clone allocated a new
|
||||
// handler chain per message even when nothing ended up being logged.
|
||||
slog.Debug("ws ← client message", "user_id", c.userID, "msg_type", msgType, "req_id", reqID)
|
||||
|
||||
// ── Typed command dispatch ───────────────────────────────────────────
|
||||
// Every message type parses through its constructor into a typed Command,
|
||||
@@ -155,6 +99,91 @@ func (h *Hub) handleMessage(c *Client, raw []byte) {
|
||||
return
|
||||
}
|
||||
|
||||
h.handleMessageApply(c, env, result)
|
||||
}
|
||||
|
||||
// handleMessageSessionRecheck performs handleMessage's periodic session
|
||||
// revalidation. It reports whether the connection was closed, in which case
|
||||
// the caller must stop processing the frame.
|
||||
func (h *Hub) handleMessageSessionRecheck(c *Client) bool {
|
||||
// Periodic session expiry check: every SessionCheckInterval messages,
|
||||
// re-validate the session token. This catches sessions that are revoked or
|
||||
// expire while the WebSocket connection is still open.
|
||||
c.mu.Lock()
|
||||
c.msgCount++
|
||||
shouldCheck := c.msgCount >= SessionCheckInterval
|
||||
if shouldCheck {
|
||||
c.msgCount = 0
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
if shouldCheck && c.tokenHash != "" {
|
||||
result, dbErr := h.db.GetSessionWithBanStatus(c.ctx, c.tokenHash)
|
||||
if dbErr != nil || result == nil || auth.IsSessionExpired(result.ExpiresAt) {
|
||||
slog.Info("ws session expired, closing connection", "user_id", c.userID)
|
||||
h.kickClient(c)
|
||||
return true
|
||||
}
|
||||
tempUser := &db.User{Banned: result.Banned, BanExpires: result.BanExpires}
|
||||
if auth.IsEffectivelyBanned(tempUser) {
|
||||
slog.Info("ws user banned, closing connection", "user_id", c.userID)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeBanned, "you are banned"))
|
||||
h.kickClient(c)
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// handleMessageDecode parses handleMessage's inbound frame into an envelope,
|
||||
// maintaining the consecutive-invalid-JSON counter, and returns the capped
|
||||
// msg_type / req_id used for logging. It reports false when the frame was
|
||||
// rejected and the caller must stop.
|
||||
func (h *Hub) handleMessageDecode(c *Client, raw []byte) (envelope, string, string, bool) {
|
||||
var env envelope
|
||||
if err := json.Unmarshal(raw, &env); err != nil {
|
||||
c.mu.Lock()
|
||||
c.invalidCount++
|
||||
count := c.invalidCount
|
||||
c.mu.Unlock()
|
||||
|
||||
slog.Warn("ws handleMessage invalid JSON", "user_id", c.userID, "err", err, "invalid_count", count)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInvalidJSON, "message must be valid JSON"))
|
||||
|
||||
if count >= 10 {
|
||||
slog.Warn("ws too many invalid messages, closing connection", "user_id", c.userID, "invalid_count", count)
|
||||
h.kickClient(c)
|
||||
}
|
||||
return env, "", "", false
|
||||
}
|
||||
|
||||
// Valid parse — reset consecutive invalid counter.
|
||||
c.mu.Lock()
|
||||
c.invalidCount = 0
|
||||
c.mu.Unlock()
|
||||
|
||||
// Cap client-controlled fields before logging to prevent log injection
|
||||
// and unbounded log entries.
|
||||
msgType := env.Type
|
||||
if len(msgType) > 64 {
|
||||
msgType = msgType[:64]
|
||||
}
|
||||
reqID := env.ID
|
||||
if len(reqID) > 64 {
|
||||
reqID = reqID[:64]
|
||||
}
|
||||
|
||||
// Correlation attrs (user_id/msg_type/req_id) are inlined at each log site
|
||||
// below rather than bound via slog.With — the With clone allocated a new
|
||||
// handler chain per message even when nothing ended up being logged.
|
||||
slog.Debug("ws ← client message", "user_id", c.userID, "msg_type", msgType, "req_id", reqID)
|
||||
|
||||
return env, msgType, reqID, true
|
||||
}
|
||||
|
||||
// handleMessageApply applies the client state mutations and side effects that a
|
||||
// successful V2 Result asks for.
|
||||
func (h *Hub) handleMessageApply(c *Client, env envelope, result Result) {
|
||||
// Apply client state mutations and side effects.
|
||||
if result.SetChannelID != nil {
|
||||
h.applySetChannelID(c, *result.SetChannelID)
|
||||
|
||||
+55
-47
@@ -245,23 +245,7 @@ func (h *Hub) channelReadAudienceImpl(ctx context.Context, channelID int64, igno
|
||||
return []int64{}
|
||||
}
|
||||
if ch != nil && ch.Type == "dm" {
|
||||
participantIDs, err := h.db.GetDMParticipantIDs(ctx, channelID)
|
||||
if err != nil {
|
||||
slog.Error("ws: channelReadAudience GetDMParticipantIDs failed, denying",
|
||||
"channel_id", channelID, "err", err)
|
||||
return []int64{}
|
||||
}
|
||||
connected := make(map[int64]struct{}, len(userIDs))
|
||||
for _, uid := range userIDs {
|
||||
connected[uid] = struct{}{}
|
||||
}
|
||||
audience := make([]int64, 0, len(participantIDs))
|
||||
for _, uid := range participantIDs {
|
||||
if _, ok := connected[uid]; ok {
|
||||
audience = append(audience, uid)
|
||||
}
|
||||
}
|
||||
return audience
|
||||
return h.channelReadAudienceDM(ctx, channelID, userIDs)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -293,6 +277,30 @@ func (h *Hub) channelReadAudienceImpl(ctx context.Context, channelID int64, igno
|
||||
return audience
|
||||
}
|
||||
|
||||
// channelReadAudienceDM resolves the audience of a DM channel: the DM's
|
||||
// participants, intersected with the connected userIDs. Split verbatim out of
|
||||
// channelReadAudienceImpl; the reason a DM must not fall through to the role
|
||||
// scan is on the call site.
|
||||
func (h *Hub) channelReadAudienceDM(ctx context.Context, channelID int64, userIDs []int64) []int64 {
|
||||
participantIDs, err := h.db.GetDMParticipantIDs(ctx, channelID)
|
||||
if err != nil {
|
||||
slog.Error("ws: channelReadAudience GetDMParticipantIDs failed, denying",
|
||||
"channel_id", channelID, "err", err)
|
||||
return []int64{}
|
||||
}
|
||||
connected := make(map[int64]struct{}, len(userIDs))
|
||||
for _, uid := range userIDs {
|
||||
connected[uid] = struct{}{}
|
||||
}
|
||||
audience := make([]int64, 0, len(participantIDs))
|
||||
for _, uid := range participantIDs {
|
||||
if _, ok := connected[uid]; ok {
|
||||
audience = append(audience, uid)
|
||||
}
|
||||
}
|
||||
return audience
|
||||
}
|
||||
|
||||
// BroadcastServerRestart sends a server_restart message to all connected clients.
|
||||
// reason describes why the server is restarting (e.g., "update").
|
||||
// delaySeconds tells clients how long until the server actually shuts down.
|
||||
@@ -404,35 +412,6 @@ func (h *Hub) RefreshChannelVisibility(ch *db.Channel) {
|
||||
return h.permChecker.HasChannelPerm(ctx, role.Permissions, roleID, userID, ch.ID, permissions.ReadMessages)
|
||||
}
|
||||
|
||||
// userCanSend mirrors channelCanSend (serve_ready.go) — the value the ready
|
||||
// payload ships per channel — but expressed as per-user permission checks
|
||||
// so it works in both the service and bare-hub branches without needing a
|
||||
// resolved *db.Role. HasChannelPerm already bypasses for admins and fails
|
||||
// closed on a lookup error, matching channelCanSend's own admin shortcut.
|
||||
//
|
||||
// Without this, can_send is only ever computed at connect time, so a role
|
||||
// edit or override edit leaves every connected client's composer stuck on
|
||||
// its stale connect-time verdict until the socket is rebuilt.
|
||||
userCanSend := func(userID, roleID int64) bool {
|
||||
has := func(perm int64) bool {
|
||||
if h.perms != nil {
|
||||
return h.perms.HasChannelPerm(ctx, userID, ch.ID, perm)
|
||||
}
|
||||
role, err := h.db.GetRoleByID(ctx, roleID)
|
||||
if err != nil || role == nil {
|
||||
return false
|
||||
}
|
||||
return h.permChecker.HasChannelPerm(ctx, role.Permissions, roleID, userID, ch.ID, perm)
|
||||
}
|
||||
if !has(permissions.ReadMessages) || !has(permissions.SendMessages) {
|
||||
return false
|
||||
}
|
||||
if ch.Type == "announcement" {
|
||||
return has(permissions.ManageMessages)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
for _, c := range clients {
|
||||
if c.user == nil {
|
||||
continue
|
||||
@@ -486,7 +465,7 @@ func (h *Hub) RefreshChannelVisibility(ch *db.Channel) {
|
||||
// Addressed per client so it can carry this recipient's own
|
||||
// can_send verdict — the whole point of this fan-out is that a
|
||||
// permission change just made those verdicts diverge.
|
||||
live.sendMsg(buildChannelCreateFor(ch, userCanSend(c.user.ID, c.user.RoleID)))
|
||||
live.sendMsg(buildChannelCreateFor(ch, h.refreshChannelVisibilityCanSend(ctx, ch, c.user.ID, c.user.RoleID)))
|
||||
continue
|
||||
}
|
||||
live.sendMsg(buildChannelDelete(ch.ID))
|
||||
@@ -507,6 +486,35 @@ func (h *Hub) RefreshChannelVisibility(ch *db.Channel) {
|
||||
h.bumpVisibilityWatermark()
|
||||
}
|
||||
|
||||
// refreshChannelVisibilityCanSend mirrors channelCanSend (serve_ready.go) — the value the ready
|
||||
// payload ships per channel — but expressed as per-user permission checks
|
||||
// so it works in both the service and bare-hub branches without needing a
|
||||
// resolved *db.Role. HasChannelPerm already bypasses for admins and fails
|
||||
// closed on a lookup error, matching channelCanSend's own admin shortcut.
|
||||
//
|
||||
// Without this, can_send is only ever computed at connect time, so a role
|
||||
// edit or override edit leaves every connected client's composer stuck on
|
||||
// its stale connect-time verdict until the socket is rebuilt.
|
||||
func (h *Hub) refreshChannelVisibilityCanSend(ctx context.Context, ch *db.Channel, userID, roleID int64) bool {
|
||||
has := func(perm int64) bool {
|
||||
if h.perms != nil {
|
||||
return h.perms.HasChannelPerm(ctx, userID, ch.ID, perm)
|
||||
}
|
||||
role, err := h.db.GetRoleByID(ctx, roleID)
|
||||
if err != nil || role == nil {
|
||||
return false
|
||||
}
|
||||
return h.permChecker.HasChannelPerm(ctx, role.Permissions, roleID, userID, ch.ID, perm)
|
||||
}
|
||||
if !has(permissions.ReadMessages) || !has(permissions.SendMessages) {
|
||||
return false
|
||||
}
|
||||
if ch.Type == "announcement" {
|
||||
return has(permissions.ManageMessages)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// RefreshAllChannelVisibility re-runs RefreshChannelVisibility for every
|
||||
// non-DM channel. A role's permission mask is the base every channel's
|
||||
// effective permission is computed from, so editing or deleting a role can
|
||||
|
||||
+18
-11
@@ -160,17 +160,10 @@ func (h *Hub) sweepRevokedSessions() {
|
||||
}
|
||||
}
|
||||
|
||||
// sweepStaleVoiceStates queries all voice_states rows and removes any that
|
||||
// don't match a connected client's voiceChID. This catches ghost users that
|
||||
// slip through the primary cleanup paths (registerNow, readPump defer,
|
||||
// LiveKit webhook).
|
||||
func (h *Hub) sweepStaleVoiceStates() {
|
||||
if h.db == nil {
|
||||
return
|
||||
}
|
||||
// Hub run-loop sweeper — no request tie.
|
||||
ctx := context.Background()
|
||||
|
||||
// sweepStaleVoiceEvictRevoked is sweepStaleVoiceStates' permission stage: it
|
||||
// re-checks CONNECT_VOICE for every client currently in voice and evicts the
|
||||
// ones who no longer hold it.
|
||||
func (h *Hub) sweepStaleVoiceEvictRevoked(ctx context.Context) {
|
||||
// Revocation must evict a live session, not merely block the next join.
|
||||
// Nothing else in ws re-validates voice permissions for a connection that
|
||||
// stays open, so a user stripped of CONNECT_VOICE kept their SFU session
|
||||
@@ -216,6 +209,20 @@ func (h *Hub) sweepStaleVoiceStates() {
|
||||
"user_id", c.userID, "channel_id", chID)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "missing CONNECT_VOICE permission"))
|
||||
}
|
||||
}
|
||||
|
||||
// sweepStaleVoiceStates queries all voice_states rows and removes any that
|
||||
// don't match a connected client's voiceChID. This catches ghost users that
|
||||
// slip through the primary cleanup paths (registerNow, readPump defer,
|
||||
// LiveKit webhook).
|
||||
func (h *Hub) sweepStaleVoiceStates() {
|
||||
if h.db == nil {
|
||||
return
|
||||
}
|
||||
// Hub run-loop sweeper — no request tie.
|
||||
ctx := context.Background()
|
||||
|
||||
h.sweepStaleVoiceEvictRevoked(ctx)
|
||||
|
||||
allStates, err := h.db.GetAllVoiceStates(ctx)
|
||||
if err != nil {
|
||||
|
||||
@@ -54,7 +54,7 @@ func TestSweepStaleVoiceStates_JoinCatchesUpDuringDeleteWindow(t *testing.T) {
|
||||
|
||||
sweepStaleVoiceJoinRaceHook = func(userID, channelID int64, joinedAt string) {
|
||||
if userID == uid {
|
||||
// Simulates voice_join.go:253's c.setVoiceState landing inside
|
||||
// Simulates voiceJoinPersist's c.setVoiceState landing inside
|
||||
// the sweep's snapshot-to-delete window.
|
||||
c.setVoiceState(channelID, joinedAt)
|
||||
}
|
||||
|
||||
@@ -140,27 +140,8 @@ func EnsureLiveKitBinary(ctx context.Context, dataDir, version string) (string,
|
||||
return "", fmt.Errorf("livekit archive checksum mismatch for %s: expected %s, got %s", asset, expectedHash, actual)
|
||||
}
|
||||
|
||||
tmpBin := dest + ".tmp"
|
||||
_ = os.Remove(tmpBin)
|
||||
if strings.HasSuffix(asset, ".zip") {
|
||||
err = extractLiveKitFromZip(f, size, tmpBin)
|
||||
} else {
|
||||
if _, seekErr := f.Seek(0, io.SeekStart); seekErr != nil {
|
||||
return "", fmt.Errorf("rewinding archive: %w", seekErr)
|
||||
}
|
||||
err = extractLiveKitFromTarGz(f, tmpBin)
|
||||
}
|
||||
if err != nil {
|
||||
_ = os.Remove(tmpBin)
|
||||
return "", fmt.Errorf("extracting %s: %w", asset, err)
|
||||
}
|
||||
if err := os.Chmod(tmpBin, 0o755); err != nil { //nolint:gosec // G302: must be executable
|
||||
_ = os.Remove(tmpBin)
|
||||
return "", fmt.Errorf("chmod binary: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpBin, dest); err != nil {
|
||||
_ = os.Remove(tmpBin)
|
||||
return "", fmt.Errorf("staging binary: %w", err)
|
||||
if err := ensureLiveKitStageBinary(f, size, asset, dest); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
cleanupOldLiveKitBinaries(dir, filepath.Base(dest))
|
||||
@@ -168,6 +149,36 @@ func EnsureLiveKitBinary(ctx context.Context, dataDir, version string) (string,
|
||||
return dest, nil
|
||||
}
|
||||
|
||||
// ensureLiveKitStageBinary extracts the already-verified archive f (asset's
|
||||
// suffix picks zip vs tar.gz) into a temp file beside dest, makes it
|
||||
// executable and renames it into place. Every failure removes the temp file.
|
||||
func ensureLiveKitStageBinary(f *os.File, size int64, asset, dest string) error {
|
||||
tmpBin := dest + ".tmp"
|
||||
_ = os.Remove(tmpBin)
|
||||
var err error
|
||||
if strings.HasSuffix(asset, ".zip") {
|
||||
err = extractLiveKitFromZip(f, size, tmpBin)
|
||||
} else {
|
||||
if _, seekErr := f.Seek(0, io.SeekStart); seekErr != nil {
|
||||
return fmt.Errorf("rewinding archive: %w", seekErr)
|
||||
}
|
||||
err = extractLiveKitFromTarGz(f, tmpBin)
|
||||
}
|
||||
if err != nil {
|
||||
_ = os.Remove(tmpBin)
|
||||
return fmt.Errorf("extracting %s: %w", asset, err)
|
||||
}
|
||||
if err := os.Chmod(tmpBin, 0o755); err != nil { //nolint:gosec // G302: must be executable
|
||||
_ = os.Remove(tmpBin)
|
||||
return fmt.Errorf("chmod binary: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpBin, dest); err != nil {
|
||||
_ = os.Remove(tmpBin)
|
||||
return fmt.Errorf("staging binary: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// livekitBinaryEntry reports whether an archive entry name is the
|
||||
// livekit-server binary (archives contain it at the top level plus LICENSE).
|
||||
func livekitBinaryEntry(name string) bool {
|
||||
|
||||
+119
-94
@@ -139,42 +139,51 @@ func (h *Hub) handleWebhookParticipantJoined(ctx context.Context, event *livekit
|
||||
// A replayed token from a previous session will not have a matching row,
|
||||
// so we remove the rogue participant from LiveKit.
|
||||
if h.db != nil {
|
||||
state, stateErr := h.db.GetVoiceState(ctx, userID)
|
||||
if stateErr != nil {
|
||||
// A transient read failure (I/O error, lock contention, a
|
||||
// maintenance window) is not proof of a rogue participant —
|
||||
// treating it as one would eject a legitimate participant from
|
||||
// the SFU on a single bad read. Mirrors sweepStaleVoiceStates'
|
||||
// hasChannelPermChecked guard: skip and let the participant be;
|
||||
// a later webhook retry or sweep tick resolves it.
|
||||
slog.Error("livekit webhook: GetVoiceState failed, skipping rogue-participant check",
|
||||
"error", stateErr, "user_id", userID, "channel_id", channelID)
|
||||
return
|
||||
}
|
||||
if state == nil || state.ChannelID != channelID {
|
||||
slog.Warn("livekit webhook: rogue participant_joined — no matching voice state, removing",
|
||||
"user_id", userID, "channel_id", channelID)
|
||||
if h.livekit != nil {
|
||||
if rmErr := h.livekit.RemoveParticipant(ctx, channelID, userID, joinToken); rmErr != nil {
|
||||
slog.Error("livekit webhook: failed to remove rogue participant",
|
||||
"error", rmErr, "user_id", userID, "channel_id", channelID)
|
||||
}
|
||||
h.webhookJoinedEnforceVoiceState(ctx, userID, channelID, joinToken)
|
||||
}
|
||||
}
|
||||
|
||||
// webhookJoinedEnforceVoiceState is the voice_states reconciliation stage of
|
||||
// handleWebhookParticipantJoined: it matches the joining participant against
|
||||
// their DB row and removes them from the SFU when the row is missing, points
|
||||
// at another channel, or carries a different join token. Callers guarantee
|
||||
// h.db != nil.
|
||||
func (h *Hub) webhookJoinedEnforceVoiceState(ctx context.Context, userID, channelID int64, joinToken string) {
|
||||
state, stateErr := h.db.GetVoiceState(ctx, userID)
|
||||
if stateErr != nil {
|
||||
// A transient read failure (I/O error, lock contention, a
|
||||
// maintenance window) is not proof of a rogue participant —
|
||||
// treating it as one would eject a legitimate participant from
|
||||
// the SFU on a single bad read. Mirrors sweepStaleVoiceStates'
|
||||
// hasChannelPermChecked guard: skip and let the participant be;
|
||||
// a later webhook retry or sweep tick resolves it.
|
||||
slog.Error("livekit webhook: GetVoiceState failed, skipping rogue-participant check",
|
||||
"error", stateErr, "user_id", userID, "channel_id", channelID)
|
||||
return
|
||||
}
|
||||
if state == nil || state.ChannelID != channelID {
|
||||
slog.Warn("livekit webhook: rogue participant_joined — no matching voice state, removing",
|
||||
"user_id", userID, "channel_id", channelID)
|
||||
if h.livekit != nil {
|
||||
if rmErr := h.livekit.RemoveParticipant(ctx, channelID, userID, joinToken); rmErr != nil {
|
||||
slog.Error("livekit webhook: failed to remove rogue participant",
|
||||
"error", rmErr, "user_id", userID, "channel_id", channelID)
|
||||
}
|
||||
return
|
||||
}
|
||||
// Verify join token matches to prevent token replay from old sessions.
|
||||
if joinToken != "" && state.JoinedAt != joinToken {
|
||||
slog.Warn("livekit webhook: stale join token on participant_joined, removing",
|
||||
"user_id", userID, "channel_id", channelID,
|
||||
"expected_token", state.JoinedAt, "got_token", joinToken)
|
||||
if h.livekit != nil {
|
||||
if rmErr := h.livekit.RemoveParticipant(ctx, channelID, userID, joinToken); rmErr != nil {
|
||||
slog.Error("livekit webhook: failed to remove stale participant",
|
||||
"error", rmErr, "user_id", userID, "channel_id", channelID)
|
||||
}
|
||||
return
|
||||
}
|
||||
// Verify join token matches to prevent token replay from old sessions.
|
||||
if joinToken != "" && state.JoinedAt != joinToken {
|
||||
slog.Warn("livekit webhook: stale join token on participant_joined, removing",
|
||||
"user_id", userID, "channel_id", channelID,
|
||||
"expected_token", state.JoinedAt, "got_token", joinToken)
|
||||
if h.livekit != nil {
|
||||
if rmErr := h.livekit.RemoveParticipant(ctx, channelID, userID, joinToken); rmErr != nil {
|
||||
slog.Error("livekit webhook: failed to remove stale participant",
|
||||
"error", rmErr, "user_id", userID, "channel_id", channelID)
|
||||
}
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
@@ -225,68 +234,7 @@ func (h *Hub) handleWebhookParticipantLeft(ctx context.Context, event *livekit.W
|
||||
h.mu.RUnlock()
|
||||
|
||||
if exists {
|
||||
// Atomic compare-and-clear under c.voiceMu, replacing the previous
|
||||
// read-then-read-then-clear: two independent unlocked getVoiceState
|
||||
// snapshots followed by an unconditional clearVoiceState is not a
|
||||
// guard at all — no lock spans the second read and the clear, so a
|
||||
// voice_join committed on the readPump goroutine in between (a
|
||||
// channel switch, or a same-channel rejoin with a fresh token) is
|
||||
// wiped out from under the new session, dropping its VoiceTopic
|
||||
// subscription along with it. client.go's clearVoiceStateIfMatch
|
||||
// only compares the channel, not the token, so it would still be
|
||||
// fooled by a same-channel rejoin — this compares both, inlined here
|
||||
// via direct field access (same package as client.go) under the
|
||||
// client's own voiceMu.
|
||||
c.voiceMu.Lock()
|
||||
matched := c.voiceChID == channelID && c.voiceJoinToken != "" && c.voiceJoinToken == joinToken
|
||||
if matched {
|
||||
c.voiceChID = 0
|
||||
c.voiceJoinToken = ""
|
||||
c.e2eePubKey = ""
|
||||
c.e2eeSignature = ""
|
||||
}
|
||||
c.voiceMu.Unlock()
|
||||
|
||||
if matched {
|
||||
h.pubsub.Unsubscribe(c, VoiceTopic(channelID))
|
||||
|
||||
if h.db != nil {
|
||||
if err := leaveVoiceChannelWithRetry(ctx, h, userID, channelID, joinToken); err != nil {
|
||||
slog.Error("livekit webhook: LeaveVoiceChannel exhausted retries",
|
||||
"error", err, "user_id", userID, "channel_id", channelID)
|
||||
}
|
||||
}
|
||||
|
||||
// This participant is out of voice, so the E2EE key holder may
|
||||
// need to move. Without this the map keeps naming the departed
|
||||
// user and the real lowest-uid participant's rekey offers are
|
||||
// rejected with NOT_KEY_HOLDER. Safe here: no locks are held.
|
||||
h.updateKeyHolder(channelID)
|
||||
|
||||
// The leaver's own client state was just cleared above, so
|
||||
// broadcastVoiceEvent's still-in-the-room union can no longer see
|
||||
// them — without broadcastVoiceEventWithLeaver's extra term, a
|
||||
// participant without READ_MESSAGES on this channel (voice
|
||||
// membership needs only CONNECT_VOICE) never learns the server
|
||||
// already tore down their call. Mirrors finishVoiceLeave and
|
||||
// CleanupVoiceForChannel, which add the leaver for the same reason.
|
||||
h.broadcastVoiceEventWithLeaver(ctx, channelID, buildVoiceLeave(channelID, userID), userID)
|
||||
slog.Info("livekit webhook: cleaned up stale voice state",
|
||||
"user_id", userID,
|
||||
"channel_id", channelID)
|
||||
} else if h.db != nil {
|
||||
// Client has voiceChID=0 or moved to a different channel (e.g.
|
||||
// after F5 reload), or this webhook is for an older join instance.
|
||||
deleted, dbErr := h.db.LeaveVoiceChannelIfMatch(ctx, userID, channelID, joinToken)
|
||||
if dbErr != nil {
|
||||
slog.Error("livekit webhook: LeaveVoiceChannelIfMatch failed (stale DB row)",
|
||||
"error", dbErr, "user_id", userID, "channel_id", channelID)
|
||||
} else if deleted {
|
||||
h.broadcastVoiceEvent(ctx, channelID, buildVoiceLeave(channelID, userID))
|
||||
slog.Info("livekit webhook: cleaned stale DB voice row after reconnect",
|
||||
"user_id", userID, "channel_id", channelID)
|
||||
}
|
||||
}
|
||||
h.webhookLeftCleanupClient(ctx, c, userID, channelID, joinToken)
|
||||
} else if h.db != nil {
|
||||
// Client already disconnected from WS — use channel-conditional delete
|
||||
// to avoid wiping a newer row if the user reconnected and rejoined.
|
||||
@@ -300,6 +248,83 @@ func (h *Hub) handleWebhookParticipantLeft(ctx context.Context, event *livekit.W
|
||||
}
|
||||
}
|
||||
|
||||
// webhookLeftCleanupClient is the still-connected-client stage of
|
||||
// handleWebhookParticipantLeft: it compare-and-clears the client's voice
|
||||
// fields for this exact join instance, then either finishes the leave or, when
|
||||
// the client has already moved on, clears the stale DB row.
|
||||
func (h *Hub) webhookLeftCleanupClient(ctx context.Context, c *Client, userID, channelID int64, joinToken string) {
|
||||
// Atomic compare-and-clear under c.voiceMu, replacing the previous
|
||||
// read-then-read-then-clear: two independent unlocked getVoiceState
|
||||
// snapshots followed by an unconditional clearVoiceState is not a
|
||||
// guard at all — no lock spans the second read and the clear, so a
|
||||
// voice_join committed on the readPump goroutine in between (a
|
||||
// channel switch, or a same-channel rejoin with a fresh token) is
|
||||
// wiped out from under the new session, dropping its VoiceTopic
|
||||
// subscription along with it. client.go's clearVoiceStateIfMatch
|
||||
// only compares the channel, not the token, so it would still be
|
||||
// fooled by a same-channel rejoin — this compares both, inlined here
|
||||
// via direct field access (same package as client.go) under the
|
||||
// client's own voiceMu.
|
||||
c.voiceMu.Lock()
|
||||
matched := c.voiceChID == channelID && c.voiceJoinToken != "" && c.voiceJoinToken == joinToken
|
||||
if matched {
|
||||
c.voiceChID = 0
|
||||
c.voiceJoinToken = ""
|
||||
c.e2eePubKey = ""
|
||||
c.e2eeSignature = ""
|
||||
}
|
||||
c.voiceMu.Unlock()
|
||||
|
||||
if matched {
|
||||
h.webhookLeftFinishLeave(ctx, c, userID, channelID, joinToken)
|
||||
} else if h.db != nil {
|
||||
// Client has voiceChID=0 or moved to a different channel (e.g.
|
||||
// after F5 reload), or this webhook is for an older join instance.
|
||||
deleted, dbErr := h.db.LeaveVoiceChannelIfMatch(ctx, userID, channelID, joinToken)
|
||||
if dbErr != nil {
|
||||
slog.Error("livekit webhook: LeaveVoiceChannelIfMatch failed (stale DB row)",
|
||||
"error", dbErr, "user_id", userID, "channel_id", channelID)
|
||||
} else if deleted {
|
||||
h.broadcastVoiceEvent(ctx, channelID, buildVoiceLeave(channelID, userID))
|
||||
slog.Info("livekit webhook: cleaned stale DB voice row after reconnect",
|
||||
"user_id", userID, "channel_id", channelID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// webhookLeftFinishLeave is the tear-down stage of webhookLeftCleanupClient,
|
||||
// reached once the client's voice fields matched this join instance and were
|
||||
// cleared: drop the voice subscription, clear the DB row, move the E2EE key
|
||||
// holder on, and broadcast the leave.
|
||||
func (h *Hub) webhookLeftFinishLeave(ctx context.Context, c *Client, userID, channelID int64, joinToken string) {
|
||||
h.pubsub.Unsubscribe(c, VoiceTopic(channelID))
|
||||
|
||||
if h.db != nil {
|
||||
if err := leaveVoiceChannelWithRetry(ctx, h, userID, channelID, joinToken); err != nil {
|
||||
slog.Error("livekit webhook: LeaveVoiceChannel exhausted retries",
|
||||
"error", err, "user_id", userID, "channel_id", channelID)
|
||||
}
|
||||
}
|
||||
|
||||
// This participant is out of voice, so the E2EE key holder may
|
||||
// need to move. Without this the map keeps naming the departed
|
||||
// user and the real lowest-uid participant's rekey offers are
|
||||
// rejected with NOT_KEY_HOLDER. Safe here: no locks are held.
|
||||
h.updateKeyHolder(channelID)
|
||||
|
||||
// The leaver's own client state was just cleared above, so
|
||||
// broadcastVoiceEvent's still-in-the-room union can no longer see
|
||||
// them — without broadcastVoiceEventWithLeaver's extra term, a
|
||||
// participant without READ_MESSAGES on this channel (voice
|
||||
// membership needs only CONNECT_VOICE) never learns the server
|
||||
// already tore down their call. Mirrors finishVoiceLeave and
|
||||
// CleanupVoiceForChannel, which add the leaver for the same reason.
|
||||
h.broadcastVoiceEventWithLeaver(ctx, channelID, buildVoiceLeave(channelID, userID), userID)
|
||||
slog.Info("livekit webhook: cleaned up stale voice state",
|
||||
"user_id", userID,
|
||||
"channel_id", channelID)
|
||||
}
|
||||
|
||||
// MountWebhookRoute is a helper for the router to mount the webhook endpoint.
|
||||
func MountWebhookRoute(h *Hub, apiKey, apiSecret string) http.HandlerFunc {
|
||||
return h.NewLiveKitWebhookHandler(apiKey, apiSecret)
|
||||
|
||||
+286
-207
@@ -156,22 +156,8 @@ var handleReconnectPreRegisterRaceHook func()
|
||||
func (h *Hub) handleReconnect(
|
||||
ctx context.Context, conn *websocket.Conn, c *Client, database *db.DB, lastSeq uint64,
|
||||
) (handled, startPumps bool) {
|
||||
// Channel-visibility changes are delivered as targeted, unsequenced
|
||||
// messages, so replay cannot bring a client that missed one back into a
|
||||
// coherent state — force the full-ready path instead.
|
||||
if h.mustFullResync(lastSeq) {
|
||||
slog.Info("ws replay skipped (visibility changed since last_seq), sending full ready",
|
||||
"user_id", c.userID, "last_seq", lastSeq)
|
||||
h.reconnectTierFull.Add(1)
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", "full"))
|
||||
return false, false
|
||||
}
|
||||
// Compute the set of channel IDs the reconnecting user can access so that
|
||||
// channel-scoped replay events are filtered by current permissions (M3).
|
||||
allowedChannelIDs, err := h.computeAllowedChannels(ctx, database, c.user)
|
||||
if err != nil {
|
||||
slog.Warn("ws handleReconnect: computeAllowedChannels failed, falling back to full ready",
|
||||
"user_id", c.userID, "err", err)
|
||||
allowedChannelIDs, ok := h.reconnectPrecheck(ctx, database, c, lastSeq)
|
||||
if !ok {
|
||||
return false, false
|
||||
}
|
||||
|
||||
@@ -190,123 +176,11 @@ func (h *Hub) handleReconnect(
|
||||
liveVoiceChID = old.getVoiceChID()
|
||||
}
|
||||
|
||||
var (
|
||||
events [][]byte
|
||||
replaySource = "buffer"
|
||||
persistedTail [][]byte // cold-tier rows only; re-merged with a fresh buffer tail below
|
||||
maxPersistedSeq uint64
|
||||
)
|
||||
if buf := h.ReplayBuffer().EventsSinceFiltered(lastSeq, allowedChannelIDs); buf != nil {
|
||||
events = buf
|
||||
} else {
|
||||
// Phase B Step 7 — try cold-tier replay from the EventStore before
|
||||
// giving up and forcing a full ready re-sync.
|
||||
if esp := h.eventStore.Load(); esp != nil {
|
||||
es := *esp
|
||||
channelIDs := make([]int64, 0, len(allowedChannelIDs))
|
||||
for cid := range allowedChannelIDs {
|
||||
channelIDs = append(channelIDs, cid)
|
||||
}
|
||||
coldCap := h.maxColdReplayLimit()
|
||||
persisted, dbErr := es.GetEventsSinceForChannels(ctx, int64(lastSeq), channelIDs, coldCap) //nolint:gosec // lastSeq is a sequence counter bounded well below MaxInt64
|
||||
switch {
|
||||
case dbErr != nil:
|
||||
slog.Warn("ws handleReconnect: cold-tier replay query failed",
|
||||
"user_id", c.userID, "err", dbErr)
|
||||
case len(persisted) >= coldCap:
|
||||
// The query is "ORDER BY seq ASC LIMIT maxColdReplay", so a full
|
||||
// result means the gap exceeds the cap and the NEWEST events were
|
||||
// dropped. Replaying it would look like a complete resume to the
|
||||
// client — it tracks only max(seq) and cannot detect the hole —
|
||||
// silently losing state events that REST history never repairs.
|
||||
// Leave events nil so the fall-through forces a full ready.
|
||||
slog.Warn("ws handleReconnect: cold-tier replay hit the row cap, forcing full ready",
|
||||
"user_id", c.userID, "last_seq", lastSeq, "cap", coldCap)
|
||||
case len(persisted) > 0:
|
||||
// Retention pruning (PruneEventsOlderThan) deletes purely by
|
||||
// created_at with no seq-floor coordination, so this
|
||||
// channel-filtered result can be a surviving suffix left behind
|
||||
// after the events between lastSeq and persisted[0] were
|
||||
// pruned. Accepting it as-is would present a hole as a complete
|
||||
// resume, since the client tracks only max(seq). Probe the
|
||||
// store's oldest surviving seq UNFILTERED before trusting it —
|
||||
// a channel-filtered contiguity check on persisted itself can't
|
||||
// work, since a sparse per-channel result is legitimately
|
||||
// non-contiguous.
|
||||
oldest, oldestErr := es.GetEventsSince(ctx, 0, 1)
|
||||
switch {
|
||||
case oldestErr != nil:
|
||||
slog.Warn("ws handleReconnect: cold-tier oldest-seq probe failed, forcing full ready",
|
||||
"user_id", c.userID, "err", oldestErr)
|
||||
case len(oldest) == 0 || uint64(oldest[0].Seq) > lastSeq+1: //nolint:gosec // seq is a counter bounded well below MaxInt64
|
||||
var oldestSeq int64
|
||||
if len(oldest) > 0 {
|
||||
oldestSeq = oldest[0].Seq
|
||||
}
|
||||
slog.Warn("ws handleReconnect: retention pruning left a gap before last_seq, forcing full ready",
|
||||
"user_id", c.userID, "last_seq", lastSeq, "oldest_seq", oldestSeq)
|
||||
default:
|
||||
persistedTail = make([][]byte, 0, len(persisted))
|
||||
for _, p := range persisted {
|
||||
persistedTail = append(persistedTail, p.Payload)
|
||||
}
|
||||
maxPersistedSeq = uint64(persisted[len(persisted)-1].Seq) //nolint:gosec // seq is a counter bounded well below MaxInt64
|
||||
|
||||
// persisted is channel-filtered, so a hole in a channel
|
||||
// outside allowedChannelIDs would slip past a contiguity
|
||||
// check on persisted itself — and EventPersister can lose a
|
||||
// row outright (a full queue drops silently in Enqueue, a
|
||||
// per-row insert failure inside a batch flush is logged but
|
||||
// never surfaced here; see event_persister.go). Count the
|
||||
// UNFILTERED range (lastSeq, maxPersistedSeq] and require
|
||||
// every seq in it to be present. seq is the events table's
|
||||
// primary key, so the count can only come up short, never
|
||||
// over.
|
||||
expectedCount := maxPersistedSeq - lastSeq
|
||||
switch gapCount, gapErr := es.CountEventsInRange(ctx, int64(lastSeq), int64(maxPersistedSeq)); { //nolint:gosec // bounded well below MaxInt64
|
||||
case gapErr != nil:
|
||||
slog.Warn("ws handleReconnect: cold-tier contiguity probe failed, forcing full ready",
|
||||
"user_id", c.userID, "err", gapErr)
|
||||
persistedTail = nil
|
||||
case uint64(gapCount) != expectedCount: //nolint:gosec // bounded well below MaxInt64
|
||||
slog.Warn("ws handleReconnect: cold-tier replay has an interior gap, forcing full ready",
|
||||
"user_id", c.userID, "last_seq", lastSeq, "max_persisted_seq", maxPersistedSeq,
|
||||
"expected", expectedCount, "found", gapCount)
|
||||
persistedTail = nil
|
||||
}
|
||||
|
||||
if persistedTail != nil {
|
||||
// The EventPersister flushes asynchronously, so cold rows can
|
||||
// lag the live seq: events broadcast after the last flush sit
|
||||
// only in the ring buffer. Confirm the buffer can cover
|
||||
// everything above the newest persisted row — the
|
||||
// authoritative re-read happens atomically with registerNow
|
||||
// below, but a hole here must still force a full ready
|
||||
// rather than a replay with a silent gap at its end.
|
||||
switch tail := h.ReplayBuffer().EventsSinceFiltered(maxPersistedSeq, allowedChannelIDs); {
|
||||
case tail != nil:
|
||||
case atomic.LoadUint64(&h.seq) == maxPersistedSeq:
|
||||
// Post-restart empty buffer with the hub seq seeded from
|
||||
// the store max: nothing was broadcast after the last
|
||||
// persisted row, so the cold rows alone are complete.
|
||||
default:
|
||||
slog.Warn("ws handleReconnect: ring buffer cannot cover the post-flush tail, forcing full ready",
|
||||
"user_id", c.userID, "max_persisted_seq", maxPersistedSeq)
|
||||
persistedTail = nil
|
||||
}
|
||||
}
|
||||
if persistedTail != nil {
|
||||
events = persistedTail
|
||||
replaySource = "db"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if events == nil {
|
||||
h.reconnectTierFull.Add(1)
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", "full"))
|
||||
return false, false
|
||||
}
|
||||
events, replaySource, persistedTail, maxPersistedSeq := h.reconnectSelectReplay(ctx, c, lastSeq, allowedChannelIDs)
|
||||
if events == nil {
|
||||
h.reconnectTierFull.Add(1)
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", "full"))
|
||||
return false, false
|
||||
}
|
||||
|
||||
// Register BEFORE writing replay data so broadcasts that arrive during
|
||||
@@ -324,7 +198,8 @@ func (h *Hub) handleReconnect(
|
||||
// means it can never be requested again once a later frame arrives.
|
||||
// Close the window by re-reading the ring-buffer-derived portion of
|
||||
// `events` and calling registerNow inside the SAME h.seqMu critical
|
||||
// section deliverBroadcast uses, so no seq can be allocated in between.
|
||||
// section deliverBroadcast uses, so no seq can be allocated in between
|
||||
// (reconnectRegister below).
|
||||
// Restore the client's channel subscription BEFORE registration.
|
||||
//
|
||||
// registerNow copies the channel subscription from the OLD client entry,
|
||||
@@ -354,6 +229,218 @@ func (h *Hub) handleReconnect(
|
||||
}
|
||||
}
|
||||
|
||||
events, ok = h.reconnectRegister(ctx, c, lastSeq, allowedChannelIDs, replaySource, persistedTail, maxPersistedSeq)
|
||||
if !ok {
|
||||
return false, false
|
||||
}
|
||||
|
||||
switch replaySource {
|
||||
case "buffer":
|
||||
h.reconnectTierBuf.Add(1)
|
||||
case "db":
|
||||
h.reconnectTierDB.Add(1)
|
||||
}
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", replaySource))
|
||||
|
||||
// Best-effort supplement: the user's own live voice room may sit outside
|
||||
// allowedChannelIDs (see the capture of liveVoiceChID above), so its
|
||||
// voice_state/voice_leave would otherwise never reach this replay at all.
|
||||
// Tries the ring buffer first, then the cold-tier store; a miss on both
|
||||
// just leaves this one supplement as a no-op, not a regression versus the
|
||||
// pre-fix behaviour.
|
||||
if liveVoiceChID != 0 && !allowedChannelIDs[liveVoiceChID] {
|
||||
events = append(events, h.liveVoiceEventsSince(ctx, lastSeq, liveVoiceChID)...)
|
||||
}
|
||||
|
||||
if !h.reconnectWriteReplay(ctx, conn, c, lastSeq, events, replaySource) {
|
||||
// startPumps=false: the teardown inside reconnectWriteReplay already ran
|
||||
// in full. Starting readPump on this closed conn would hit an immediate
|
||||
// Read error and its defer would run the identical teardown a second
|
||||
// time (OC-0051).
|
||||
return true, false
|
||||
}
|
||||
|
||||
// Update presence but skip member_join — user was already known.
|
||||
applyConnectStatus(ctx, database, c)
|
||||
h.announceConnectPresence(c)
|
||||
|
||||
return true, true
|
||||
}
|
||||
|
||||
// reconnectPrecheck runs handleReconnect's two entry guards and, when replay is
|
||||
// still on the table, returns the read-permission set replay is filtered by.
|
||||
// ok=false means the caller must fall through to a full ready.
|
||||
func (h *Hub) reconnectPrecheck(
|
||||
ctx context.Context, database *db.DB, c *Client, lastSeq uint64,
|
||||
) (map[int64]bool, bool) {
|
||||
// Channel-visibility changes are delivered as targeted, unsequenced
|
||||
// messages, so replay cannot bring a client that missed one back into a
|
||||
// coherent state — force the full-ready path instead.
|
||||
if h.mustFullResync(lastSeq) {
|
||||
slog.Info("ws replay skipped (visibility changed since last_seq), sending full ready",
|
||||
"user_id", c.userID, "last_seq", lastSeq)
|
||||
h.reconnectTierFull.Add(1)
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", "full"))
|
||||
return nil, false
|
||||
}
|
||||
// Compute the set of channel IDs the reconnecting user can access so that
|
||||
// channel-scoped replay events are filtered by current permissions (M3).
|
||||
allowedChannelIDs, err := h.computeAllowedChannels(ctx, database, c.user)
|
||||
if err != nil {
|
||||
slog.Warn("ws handleReconnect: computeAllowedChannels failed, falling back to full ready",
|
||||
"user_id", c.userID, "err", err)
|
||||
return nil, false
|
||||
}
|
||||
return allowedChannelIDs, true
|
||||
}
|
||||
|
||||
// reconnectSelectReplay picks the tier that serves this resume — the ring
|
||||
// buffer when it still covers lastSeq, otherwise the cold-tier EventStore — and
|
||||
// returns the events found, the tier name, and (cold tier only) the persisted
|
||||
// rows plus their highest seq, which reconnectRegister needs for its re-read.
|
||||
// A nil events return means neither tier can replay and the caller must fall
|
||||
// through to a full ready.
|
||||
func (h *Hub) reconnectSelectReplay(
|
||||
ctx context.Context, c *Client, lastSeq uint64, allowedChannelIDs map[int64]bool,
|
||||
) ([][]byte, string, [][]byte, uint64) {
|
||||
var (
|
||||
events [][]byte
|
||||
replaySource = "buffer"
|
||||
persistedTail [][]byte // cold-tier rows only; re-merged with a fresh buffer tail below
|
||||
maxPersistedSeq uint64
|
||||
)
|
||||
if buf := h.ReplayBuffer().EventsSinceFiltered(lastSeq, allowedChannelIDs); buf != nil {
|
||||
events = buf
|
||||
return events, replaySource, persistedTail, maxPersistedSeq
|
||||
}
|
||||
// Phase B Step 7 — try cold-tier replay from the EventStore before
|
||||
// giving up and forcing a full ready re-sync.
|
||||
if esp := h.eventStore.Load(); esp != nil {
|
||||
es := *esp
|
||||
channelIDs := make([]int64, 0, len(allowedChannelIDs))
|
||||
for cid := range allowedChannelIDs {
|
||||
channelIDs = append(channelIDs, cid)
|
||||
}
|
||||
coldCap := h.maxColdReplayLimit()
|
||||
persisted, dbErr := es.GetEventsSinceForChannels(ctx, int64(lastSeq), channelIDs, coldCap) //nolint:gosec // lastSeq is a sequence counter bounded well below MaxInt64
|
||||
switch {
|
||||
case dbErr != nil:
|
||||
slog.Warn("ws handleReconnect: cold-tier replay query failed",
|
||||
"user_id", c.userID, "err", dbErr)
|
||||
case len(persisted) >= coldCap:
|
||||
// The query is "ORDER BY seq ASC LIMIT maxColdReplay", so a full
|
||||
// result means the gap exceeds the cap and the NEWEST events were
|
||||
// dropped. Replaying it would look like a complete resume to the
|
||||
// client — it tracks only max(seq) and cannot detect the hole —
|
||||
// silently losing state events that REST history never repairs.
|
||||
// Leave events nil so the fall-through forces a full ready.
|
||||
slog.Warn("ws handleReconnect: cold-tier replay hit the row cap, forcing full ready",
|
||||
"user_id", c.userID, "last_seq", lastSeq, "cap", coldCap)
|
||||
case len(persisted) > 0:
|
||||
// Retention pruning (PruneEventsOlderThan) deletes purely by
|
||||
// created_at with no seq-floor coordination, so this
|
||||
// channel-filtered result can be a surviving suffix left behind
|
||||
// after the events between lastSeq and persisted[0] were
|
||||
// pruned. Accepting it as-is would present a hole as a complete
|
||||
// resume, since the client tracks only max(seq). Probe the
|
||||
// store's oldest surviving seq UNFILTERED before trusting it —
|
||||
// a channel-filtered contiguity check on persisted itself can't
|
||||
// work, since a sparse per-channel result is legitimately
|
||||
// non-contiguous.
|
||||
oldest, oldestErr := es.GetEventsSince(ctx, 0, 1)
|
||||
switch {
|
||||
case oldestErr != nil:
|
||||
slog.Warn("ws handleReconnect: cold-tier oldest-seq probe failed, forcing full ready",
|
||||
"user_id", c.userID, "err", oldestErr)
|
||||
case len(oldest) == 0 || uint64(oldest[0].Seq) > lastSeq+1: //nolint:gosec // seq is a counter bounded well below MaxInt64
|
||||
var oldestSeq int64
|
||||
if len(oldest) > 0 {
|
||||
oldestSeq = oldest[0].Seq
|
||||
}
|
||||
slog.Warn("ws handleReconnect: retention pruning left a gap before last_seq, forcing full ready",
|
||||
"user_id", c.userID, "last_seq", lastSeq, "oldest_seq", oldestSeq)
|
||||
default:
|
||||
persistedTail, maxPersistedSeq = h.reconnectVetColdTail(ctx, c, es, lastSeq, persisted, allowedChannelIDs)
|
||||
if persistedTail != nil {
|
||||
events = persistedTail
|
||||
replaySource = "db"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return events, replaySource, persistedTail, maxPersistedSeq
|
||||
}
|
||||
|
||||
// reconnectVetColdTail turns a cold-tier result into a replayable tail, or
|
||||
// returns nil when it cannot be trusted: the range it covers must have no
|
||||
// interior gap, and the ring buffer must cover everything newer than its last
|
||||
// row. The returned seq is the highest one in persisted.
|
||||
func (h *Hub) reconnectVetColdTail(
|
||||
ctx context.Context, c *Client, es EventStore, lastSeq uint64,
|
||||
persisted []db.PersistedEvent, allowedChannelIDs map[int64]bool,
|
||||
) ([][]byte, uint64) {
|
||||
persistedTail := make([][]byte, 0, len(persisted))
|
||||
for _, p := range persisted {
|
||||
persistedTail = append(persistedTail, p.Payload)
|
||||
}
|
||||
maxPersistedSeq := uint64(persisted[len(persisted)-1].Seq) //nolint:gosec // seq is a counter bounded well below MaxInt64
|
||||
|
||||
// persisted is channel-filtered, so a hole in a channel
|
||||
// outside allowedChannelIDs would slip past a contiguity
|
||||
// check on persisted itself — and EventPersister can lose a
|
||||
// row outright (a full queue drops silently in Enqueue, a
|
||||
// per-row insert failure inside a batch flush is logged but
|
||||
// never surfaced here; see event_persister.go). Count the
|
||||
// UNFILTERED range (lastSeq, maxPersistedSeq] and require
|
||||
// every seq in it to be present. seq is the events table's
|
||||
// primary key, so the count can only come up short, never
|
||||
// over.
|
||||
expectedCount := maxPersistedSeq - lastSeq
|
||||
switch gapCount, gapErr := es.CountEventsInRange(ctx, int64(lastSeq), int64(maxPersistedSeq)); { //nolint:gosec // bounded well below MaxInt64
|
||||
case gapErr != nil:
|
||||
slog.Warn("ws handleReconnect: cold-tier contiguity probe failed, forcing full ready",
|
||||
"user_id", c.userID, "err", gapErr)
|
||||
persistedTail = nil
|
||||
case uint64(gapCount) != expectedCount: //nolint:gosec // bounded well below MaxInt64
|
||||
slog.Warn("ws handleReconnect: cold-tier replay has an interior gap, forcing full ready",
|
||||
"user_id", c.userID, "last_seq", lastSeq, "max_persisted_seq", maxPersistedSeq,
|
||||
"expected", expectedCount, "found", gapCount)
|
||||
persistedTail = nil
|
||||
}
|
||||
|
||||
if persistedTail != nil {
|
||||
// The EventPersister flushes asynchronously, so cold rows can
|
||||
// lag the live seq: events broadcast after the last flush sit
|
||||
// only in the ring buffer. Confirm the buffer can cover
|
||||
// everything above the newest persisted row — the
|
||||
// authoritative re-read happens atomically with registerNow
|
||||
// below, but a hole here must still force a full ready
|
||||
// rather than a replay with a silent gap at its end.
|
||||
switch tail := h.ReplayBuffer().EventsSinceFiltered(maxPersistedSeq, allowedChannelIDs); {
|
||||
case tail != nil:
|
||||
case atomic.LoadUint64(&h.seq) == maxPersistedSeq:
|
||||
// Post-restart empty buffer with the hub seq seeded from
|
||||
// the store max: nothing was broadcast after the last
|
||||
// persisted row, so the cold rows alone are complete.
|
||||
default:
|
||||
slog.Warn("ws handleReconnect: ring buffer cannot cover the post-flush tail, forcing full ready",
|
||||
"user_id", c.userID, "max_persisted_seq", maxPersistedSeq)
|
||||
persistedTail = nil
|
||||
}
|
||||
}
|
||||
return persistedTail, maxPersistedSeq
|
||||
}
|
||||
|
||||
// reconnectRegister re-reads the ring-buffer-derived portion of the replay and
|
||||
// registers c inside the SAME h.seqMu critical section deliverBroadcast uses,
|
||||
// so no seq can be allocated in between (see the comment in handleReconnect).
|
||||
// It returns the events to actually send; ok=false means one of the re-checks
|
||||
// tripped and the caller must fall through to a full ready.
|
||||
func (h *Hub) reconnectRegister(
|
||||
ctx context.Context, c *Client, lastSeq uint64, allowedChannelIDs map[int64]bool,
|
||||
replaySource string, persistedTail [][]byte, maxPersistedSeq uint64,
|
||||
) ([][]byte, bool) {
|
||||
var events [][]byte
|
||||
h.seqMu.Lock()
|
||||
switch replaySource {
|
||||
case "buffer":
|
||||
@@ -367,7 +454,7 @@ func (h *Hub) handleReconnect(
|
||||
"user_id", c.userID, "last_seq", lastSeq)
|
||||
h.reconnectTierFull.Add(1)
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", "full"))
|
||||
return false, false
|
||||
return nil, false
|
||||
}
|
||||
events = fresh
|
||||
case "db":
|
||||
@@ -382,7 +469,7 @@ func (h *Hub) handleReconnect(
|
||||
"user_id", c.userID, "max_persisted_seq", maxPersistedSeq)
|
||||
h.reconnectTierFull.Add(1)
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", "full"))
|
||||
return false, false
|
||||
return nil, false
|
||||
}
|
||||
}
|
||||
if handleReconnectPreRegisterRaceHook != nil {
|
||||
@@ -404,29 +491,21 @@ func (h *Hub) handleReconnect(
|
||||
"user_id", c.userID, "last_seq", lastSeq)
|
||||
h.reconnectTierFull.Add(1)
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", "full"))
|
||||
return false, false
|
||||
return nil, false
|
||||
}
|
||||
h.registerNow(c, allowedChannelIDs)
|
||||
h.seqMu.Unlock()
|
||||
return events, true
|
||||
}
|
||||
|
||||
switch replaySource {
|
||||
case "buffer":
|
||||
h.reconnectTierBuf.Add(1)
|
||||
case "db":
|
||||
h.reconnectTierDB.Add(1)
|
||||
}
|
||||
telemetry.NewAppMetrics().WSReconnectTierTotal.Add(ctx, 1, telemetry.String("tier", replaySource))
|
||||
|
||||
// Best-effort supplement: the user's own live voice room may sit outside
|
||||
// allowedChannelIDs (see the capture of liveVoiceChID above), so its
|
||||
// voice_state/voice_leave would otherwise never reach this replay at all.
|
||||
// Tries the ring buffer first, then the cold-tier store; a miss on both
|
||||
// just leaves this one supplement as a no-op, not a regression versus the
|
||||
// pre-fix behaviour.
|
||||
if liveVoiceChID != 0 && !allowedChannelIDs[liveVoiceChID] {
|
||||
events = append(events, h.liveVoiceEventsSince(ctx, lastSeq, liveVoiceChID)...)
|
||||
}
|
||||
|
||||
// reconnectWriteReplay writes the resume handshake: auth_ok followed by the
|
||||
// replayed events. A false return means a write failed, in which case the full
|
||||
// unregisterFailedHandshake teardown has already run and conn is closed, so the
|
||||
// caller must not start any pump (OC-0051).
|
||||
func (h *Hub) reconnectWriteReplay(
|
||||
ctx context.Context, conn *websocket.Conn, c *Client, lastSeq uint64,
|
||||
events [][]byte, replaySource string,
|
||||
) bool {
|
||||
// Replay succeeded — send auth_ok then missed events. The replay tier
|
||||
// is included in the payload so the client can attribute reconnect
|
||||
// behaviour without separate metric scraping.
|
||||
@@ -435,26 +514,18 @@ func (h *Hub) handleReconnect(
|
||||
slog.Warn("ws: failed to send auth_ok (reconnect)", "user_id", c.userID, "err", err)
|
||||
h.unregisterFailedHandshake(ctx, c)
|
||||
_ = conn.Close(websocket.StatusInternalError, "handshake failed")
|
||||
// startPumps=false: the teardown above already ran in full. Starting
|
||||
// readPump on this closed conn would hit an immediate Read error and
|
||||
// its defer would run the identical teardown a second time (OC-0051).
|
||||
return true, false
|
||||
return false
|
||||
}
|
||||
for _, evt := range events {
|
||||
if err := conn.Write(ctx, websocket.MessageText, evt); err != nil {
|
||||
slog.Warn("ws: failed to send replay event", "user_id", c.userID, "err", err)
|
||||
h.unregisterFailedHandshake(ctx, c)
|
||||
_ = conn.Close(websocket.StatusInternalError, "handshake failed")
|
||||
return true, false
|
||||
return false
|
||||
}
|
||||
}
|
||||
slog.Info("ws replay completed", "user_id", c.userID, "events_replayed", len(events), "from_seq", lastSeq, "source", replaySource)
|
||||
|
||||
// Update presence but skip member_join — user was already known.
|
||||
applyConnectStatus(ctx, database, c)
|
||||
h.announceConnectPresence(c)
|
||||
|
||||
return true, true
|
||||
return true
|
||||
}
|
||||
|
||||
// liveVoiceEventsSince returns voice_state/voice_leave events for chID at or
|
||||
@@ -620,47 +691,7 @@ func (h *Hub) handleFreshConnect(
|
||||
// session must be removed so the ready payload doesn't include it and
|
||||
// other clients see a voice_leave broadcast.
|
||||
if vs, err := database.GetVoiceState(ctx, c.userID); err == nil && vs != nil {
|
||||
// Replay-failure fallback (lastSeq > 0): registerNow below transfers
|
||||
// the still-registered old connection's live voice state into this
|
||||
// client. Deleting the DB row here — and the LiveKit participant,
|
||||
// whose removal token is the very JoinedAt being transferred — would
|
||||
// leave the user "in voice" on the hub only: voice_join bounces off
|
||||
// ALREADY_JOINED and sweepStaleVoiceStates never heals
|
||||
// memory-without-row. Keep the row so ready stays consistent. If the
|
||||
// old client unregisters before registerNow runs, the transfer is
|
||||
// skipped and the next sweep reaps the then-truly-stale row.
|
||||
if old := h.GetClient(c.userID); c.lastSeq > 0 && old != nil && old.getVoiceChID() == vs.ChannelID {
|
||||
slog.Info("ws fresh connect: keeping voice state for replay-failure fallback",
|
||||
"user_id", c.userID, "channel_id", vs.ChannelID)
|
||||
} else {
|
||||
slog.Info("ws fresh connect: cleaning stale voice state",
|
||||
"user_id", c.userID, "channel_id", vs.ChannelID)
|
||||
if _, delErr := database.LeaveVoiceChannelIfMatch(ctx, c.userID, vs.ChannelID, vs.JoinedAt); delErr != nil {
|
||||
slog.Warn("ws fresh connect: LeaveVoiceChannelIfMatch failed", "err", delErr)
|
||||
}
|
||||
h.broadcastVoiceEvent(ctx, vs.ChannelID, buildVoiceLeave(vs.ChannelID, c.userID))
|
||||
if h.livekit != nil {
|
||||
// BUG-089: Capture stale join token so the goroutine only removes
|
||||
// the exact stale participant. The identity includes joinedAt, so
|
||||
// even if the user rejoins voice quickly, the new session has a
|
||||
// different identity and won't be removed. The removal must
|
||||
// complete even if this connection drops mid-handshake, so detach
|
||||
// from cancellation (values kept); shutdown is handled via h.stop.
|
||||
staleChID, staleUserID, staleJoinToken := vs.ChannelID, c.userID, vs.JoinedAt
|
||||
lkCtx := context.WithoutCancel(ctx)
|
||||
go func() {
|
||||
select {
|
||||
case <-h.stop:
|
||||
return
|
||||
default:
|
||||
}
|
||||
if err := h.livekit.RemoveParticipant(lkCtx, staleChID, staleUserID, staleJoinToken); err != nil {
|
||||
slog.Warn("ws fresh connect: RemoveParticipant failed (may already be gone)",
|
||||
"err", err, "user_id", staleUserID, "channel_id", staleChID)
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
h.freshConnectCleanStaleVoice(ctx, database, c, vs)
|
||||
}
|
||||
|
||||
// Look up role for permission-filtered ready payload.
|
||||
@@ -743,3 +774,51 @@ func (h *Hub) handleFreshConnect(
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// freshConnectCleanStaleVoice removes the voice state left behind by this
|
||||
// user's previous session, unless that session is the still-registered
|
||||
// connection this one is about to inherit from.
|
||||
func (h *Hub) freshConnectCleanStaleVoice(ctx context.Context, database *db.DB, c *Client, vs *db.VoiceState) {
|
||||
// Replay-failure fallback (lastSeq > 0): registerNow below transfers
|
||||
// the still-registered old connection's live voice state into this
|
||||
// client. Deleting the DB row here — and the LiveKit participant,
|
||||
// whose removal token is the very JoinedAt being transferred — would
|
||||
// leave the user "in voice" on the hub only: voice_join bounces off
|
||||
// ALREADY_JOINED and sweepStaleVoiceStates never heals
|
||||
// memory-without-row. Keep the row so ready stays consistent. If the
|
||||
// old client unregisters before registerNow runs, the transfer is
|
||||
// skipped and the next sweep reaps the then-truly-stale row.
|
||||
if old := h.GetClient(c.userID); c.lastSeq > 0 && old != nil && old.getVoiceChID() == vs.ChannelID {
|
||||
slog.Info("ws fresh connect: keeping voice state for replay-failure fallback",
|
||||
"user_id", c.userID, "channel_id", vs.ChannelID)
|
||||
return
|
||||
}
|
||||
slog.Info("ws fresh connect: cleaning stale voice state",
|
||||
"user_id", c.userID, "channel_id", vs.ChannelID)
|
||||
if _, delErr := database.LeaveVoiceChannelIfMatch(ctx, c.userID, vs.ChannelID, vs.JoinedAt); delErr != nil {
|
||||
slog.Warn("ws fresh connect: LeaveVoiceChannelIfMatch failed", "err", delErr)
|
||||
}
|
||||
h.broadcastVoiceEvent(ctx, vs.ChannelID, buildVoiceLeave(vs.ChannelID, c.userID))
|
||||
if h.livekit == nil {
|
||||
return
|
||||
}
|
||||
// BUG-089: Capture stale join token so the goroutine only removes
|
||||
// the exact stale participant. The identity includes joinedAt, so
|
||||
// even if the user rejoins voice quickly, the new session has a
|
||||
// different identity and won't be removed. The removal must
|
||||
// complete even if this connection drops mid-handshake, so detach
|
||||
// from cancellation (values kept); shutdown is handled via h.stop.
|
||||
staleChID, staleUserID, staleJoinToken := vs.ChannelID, c.userID, vs.JoinedAt
|
||||
lkCtx := context.WithoutCancel(ctx)
|
||||
go func() {
|
||||
select {
|
||||
case <-h.stop:
|
||||
return
|
||||
default:
|
||||
}
|
||||
if err := h.livekit.RemoveParticipant(lkCtx, staleChID, staleUserID, staleJoinToken); err != nil {
|
||||
slog.Warn("ws fresh connect: RemoveParticipant failed (may already be gone)",
|
||||
"err", err, "user_id", staleUserID, "channel_id", staleChID)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
+60
-71
@@ -10,61 +10,70 @@ import (
|
||||
"github.com/owncord/server/db"
|
||||
)
|
||||
|
||||
// writePumpWrite writes one frame to the WebSocket under writeTimeout.
|
||||
// Returns false only when the write failed.
|
||||
func writePumpWrite(ctx context.Context, conn *websocket.Conn, c *Client, msg []byte) bool {
|
||||
wCtx, cancel := context.WithTimeout(ctx, writeTimeout)
|
||||
err := conn.Write(wCtx, websocket.MessageText, msg)
|
||||
cancel()
|
||||
if err != nil {
|
||||
slog.Warn("ws writePump error", "user_id", c.userID, "err", err)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// writePumpDrainChannel writes every message still buffered on ch without blocking.
|
||||
// Returns false only when a write failed; empty or closed is true.
|
||||
func writePumpDrainChannel(ctx context.Context, conn *websocket.Conn, c *Client, ch chan []byte) bool {
|
||||
for {
|
||||
select {
|
||||
case msg, ok := <-ch:
|
||||
if !ok {
|
||||
return true
|
||||
}
|
||||
if !writePumpWrite(ctx, conn, c, msg) {
|
||||
return false
|
||||
}
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// writePumpDrainAndClose flushes whatever the kick paths queued before closing the
|
||||
// send channels (e.g. the BANNED error frame that makes the client clear
|
||||
// its credentials) — serve.go and hub_broadcast.go both document that
|
||||
// writePump drains remaining messages after closeSend. Returning on the
|
||||
// first closed channel would drop those frames.
|
||||
func writePumpDrainAndClose(ctx context.Context, conn *websocket.Conn, c *Client) {
|
||||
if writePumpDrainChannel(ctx, conn, c, c.sendHigh) && writePumpDrainChannel(ctx, conn, c, c.send) {
|
||||
writePumpDrainChannel(ctx, conn, c, c.sendLow)
|
||||
}
|
||||
_ = conn.Close(websocket.StatusNormalClosure, "")
|
||||
}
|
||||
|
||||
// writePumpDeliver handles one frame received from a send channel: a closed
|
||||
// channel drains and closes the connection, a failed write ends the pump
|
||||
// without draining. Returns false when writePump must return.
|
||||
func writePumpDeliver(ctx context.Context, conn *websocket.Conn, c *Client, msg []byte, ok bool) bool {
|
||||
if !ok {
|
||||
writePumpDrainAndClose(ctx, conn, c)
|
||||
return false
|
||||
}
|
||||
return writePumpWrite(ctx, conn, c, msg)
|
||||
}
|
||||
|
||||
// writePump drains the client's send channels and writes to the WebSocket.
|
||||
// Priority ordering: high > normal > low. High-priority messages (DMs, mentions)
|
||||
// are drained first. Normal messages (chat, reactions) come next. Low-priority
|
||||
// messages (typing, presence) are only sent when no higher-priority work is pending.
|
||||
func writePump(ctx context.Context, conn *websocket.Conn, c *Client) {
|
||||
writeMsg := func(msg []byte) bool {
|
||||
wCtx, cancel := context.WithTimeout(ctx, writeTimeout)
|
||||
err := conn.Write(wCtx, websocket.MessageText, msg)
|
||||
cancel()
|
||||
if err != nil {
|
||||
slog.Warn("ws writePump error", "user_id", c.userID, "err", err)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// drainChannel writes every message still buffered on ch without blocking.
|
||||
// Returns false only when a write failed; empty or closed is true.
|
||||
drainChannel := func(ch chan []byte) bool {
|
||||
for {
|
||||
select {
|
||||
case msg, ok := <-ch:
|
||||
if !ok {
|
||||
return true
|
||||
}
|
||||
if !writeMsg(msg) {
|
||||
return false
|
||||
}
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// drainAndClose flushes whatever the kick paths queued before closing the
|
||||
// send channels (e.g. the BANNED error frame that makes the client clear
|
||||
// its credentials) — serve.go and hub_broadcast.go both document that
|
||||
// writePump drains remaining messages after closeSend. Returning on the
|
||||
// first closed channel would drop those frames.
|
||||
drainAndClose := func() {
|
||||
if drainChannel(c.sendHigh) && drainChannel(c.send) {
|
||||
drainChannel(c.sendLow)
|
||||
}
|
||||
_ = conn.Close(websocket.StatusNormalClosure, "")
|
||||
}
|
||||
|
||||
for {
|
||||
// Priority 1: drain all pending high-priority messages first.
|
||||
select {
|
||||
case msg, ok := <-c.sendHigh:
|
||||
if !ok {
|
||||
drainAndClose()
|
||||
return
|
||||
}
|
||||
if !writeMsg(msg) {
|
||||
if !writePumpDeliver(ctx, conn, c, msg, ok) {
|
||||
return
|
||||
}
|
||||
continue
|
||||
@@ -81,20 +90,12 @@ func writePump(ctx context.Context, conn *websocket.Conn, c *Client) {
|
||||
// neither high nor normal has anything ready right now.
|
||||
select {
|
||||
case msg, ok := <-c.sendHigh:
|
||||
if !ok {
|
||||
drainAndClose()
|
||||
return
|
||||
}
|
||||
if !writeMsg(msg) {
|
||||
if !writePumpDeliver(ctx, conn, c, msg, ok) {
|
||||
return
|
||||
}
|
||||
continue
|
||||
case msg, ok := <-c.send:
|
||||
if !ok {
|
||||
drainAndClose()
|
||||
return
|
||||
}
|
||||
if !writeMsg(msg) {
|
||||
if !writePumpDeliver(ctx, conn, c, msg, ok) {
|
||||
return
|
||||
}
|
||||
continue
|
||||
@@ -106,27 +107,15 @@ func writePump(ctx context.Context, conn *websocket.Conn, c *Client) {
|
||||
// frames instead of busy-looping.
|
||||
select {
|
||||
case msg, ok := <-c.sendHigh:
|
||||
if !ok {
|
||||
drainAndClose()
|
||||
return
|
||||
}
|
||||
if !writeMsg(msg) {
|
||||
if !writePumpDeliver(ctx, conn, c, msg, ok) {
|
||||
return
|
||||
}
|
||||
case msg, ok := <-c.send:
|
||||
if !ok {
|
||||
drainAndClose()
|
||||
return
|
||||
}
|
||||
if !writeMsg(msg) {
|
||||
if !writePumpDeliver(ctx, conn, c, msg, ok) {
|
||||
return
|
||||
}
|
||||
case msg, ok := <-c.sendLow:
|
||||
if !ok {
|
||||
drainAndClose()
|
||||
return
|
||||
}
|
||||
if !writeMsg(msg) {
|
||||
if !writePumpDeliver(ctx, conn, c, msg, ok) {
|
||||
return
|
||||
}
|
||||
case <-ctx.Done():
|
||||
|
||||
+69
-33
@@ -159,26 +159,10 @@ func channelCanSend(role *db.Role, o db.ChannelOverride, chanType string) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// buildReady constructs the ready server→client message.
|
||||
// Per docs/protocol.md, channels include unread_count and last_message_id per
|
||||
// user, and only protocol-specified fields (no slow_mode, archived, voice_*
|
||||
// extras).
|
||||
func (h *Hub) buildReady(ctx context.Context, database *db.DB, userID int64, role *db.Role) ([]byte, error) {
|
||||
channels, err := database.ListChannels(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady ListChannels: %w", err)
|
||||
}
|
||||
roles, err := database.ListRoles(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady ListRoles: %w", err)
|
||||
}
|
||||
|
||||
members, err := database.ListMembers(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady ListMembers: %w", err)
|
||||
}
|
||||
members = h.presentableMembers(members, userID)
|
||||
|
||||
// readyVisibleChannels resolves the channels the user may see for the ready
|
||||
// payload, returning the per-channel override map it fetched alongside them so
|
||||
// buildReady can reuse it for the can_send affordance without a second query.
|
||||
func (h *Hub) readyVisibleChannels(ctx context.Context, database *db.DB, userID int64, role *db.Role, channels []db.Channel) ([]db.Channel, map[int64]db.ChannelOverride, error) {
|
||||
// Filter channels by READ_MESSAGES through the single permissions.Checker
|
||||
// predicate shared with REST ListVisibleChannels and reconnect replay
|
||||
// filtering (computeAllowedChannels). The overrides map is fetched once and
|
||||
@@ -189,7 +173,7 @@ func (h *Hub) buildReady(ctx context.Context, database *db.DB, userID int64, rol
|
||||
var oErr error
|
||||
overrides, oErr = database.GetChannelOverridesFor(ctx, role.ID, userID)
|
||||
if oErr != nil {
|
||||
return nil, fmt.Errorf("buildReady GetChannelOverridesFor: %w", oErr)
|
||||
return nil, nil, fmt.Errorf("buildReady GetChannelOverridesFor: %w", oErr)
|
||||
}
|
||||
}
|
||||
var visibleChannels []db.Channel
|
||||
@@ -205,14 +189,12 @@ func (h *Hub) buildReady(ctx context.Context, database *db.DB, userID int64, rol
|
||||
if visibleChannels == nil {
|
||||
visibleChannels = []db.Channel{}
|
||||
}
|
||||
return visibleChannels, overrides, nil
|
||||
}
|
||||
|
||||
// Per-user unread counts.
|
||||
unreadMap, err := database.GetChannelUnreadCounts(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady GetChannelUnreadCounts: %w", err)
|
||||
}
|
||||
|
||||
// Build protocol-compliant channel objects (strip extra fields).
|
||||
// readyChannelPayloads builds the ready payload's channel objects — one entry
|
||||
// per visible channel, with the per-user unread fields folded in.
|
||||
func readyChannelPayloads(visibleChannels []db.Channel, overrides map[int64]db.ChannelOverride, unreadMap map[int64]db.ChannelUnread, role *db.Role) []map[string]any {
|
||||
channelPayloads := make([]map[string]any, 0, len(visibleChannels))
|
||||
for i := range visibleChannels {
|
||||
entry := map[string]any{
|
||||
@@ -254,12 +236,13 @@ func (h *Hub) buildReady(ctx context.Context, database *db.DB, userID int64, rol
|
||||
}
|
||||
channelPayloads = append(channelPayloads, entry)
|
||||
}
|
||||
return channelPayloads
|
||||
}
|
||||
|
||||
// Load open DM channels for this user. Hoisted above the voice-state
|
||||
// filter below so DM channel IDs can seed visibleSet — permissions.Checker
|
||||
// (and therefore visibleChannels) deliberately skips DM channels, since
|
||||
// their visibility is membership-based rather than role-based, so without
|
||||
// this a DM voice call's voice_state rows would never make it into ready.
|
||||
// readyDMChannels loads the user's open DM channels and reconciles them with
|
||||
// the rest of the ready payload: mention counts from unreadMap, and the same
|
||||
// presence rule presentableMembers applies to the members array.
|
||||
func (h *Hub) readyDMChannels(ctx context.Context, database *db.DB, userID int64, unreadMap map[int64]db.ChannelUnread) ([]db.DMChannelInfo, error) {
|
||||
dmChannels, err := database.GetUserDMChannels(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady GetUserDMChannels: %w", err)
|
||||
@@ -281,7 +264,12 @@ func (h *Hub) buildReady(ctx context.Context, database *db.DB, userID int64, rol
|
||||
// presentableMembers above; apply the same half here so dm_channels
|
||||
// cannot disagree with members about the same user within one payload.
|
||||
dmChannels = h.presentableDMChannels(dmChannels)
|
||||
return dmChannels, nil
|
||||
}
|
||||
|
||||
// readyVoiceStates gathers the voice states the ready payload may expose to
|
||||
// this user. A collect failure is non-fatal, so this returns no error.
|
||||
func (h *Hub) readyVoiceStates(ctx context.Context, database *db.DB, channels []db.Channel, visibleChannels []db.Channel, dmChannels []db.DMChannelInfo, userID int64) []db.VoiceState {
|
||||
// Collect voice states, filtered to visible channels (BUG-095) plus the
|
||||
// user's own open DM channels — mirroring computeAllowedChannels, which
|
||||
// layers DM IDs onto the same checker result for reconnect replay
|
||||
@@ -318,6 +306,54 @@ func (h *Hub) buildReady(ctx context.Context, database *db.DB, userID int64, rol
|
||||
voiceStates = append(voiceStates, allVoiceStates[i])
|
||||
}
|
||||
}
|
||||
return voiceStates
|
||||
}
|
||||
|
||||
// buildReady constructs the ready server→client message.
|
||||
// Per docs/protocol.md, channels include unread_count and last_message_id per
|
||||
// user, and only protocol-specified fields (no slow_mode, archived, voice_*
|
||||
// extras).
|
||||
func (h *Hub) buildReady(ctx context.Context, database *db.DB, userID int64, role *db.Role) ([]byte, error) {
|
||||
channels, err := database.ListChannels(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady ListChannels: %w", err)
|
||||
}
|
||||
roles, err := database.ListRoles(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady ListRoles: %w", err)
|
||||
}
|
||||
|
||||
members, err := database.ListMembers(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady ListMembers: %w", err)
|
||||
}
|
||||
members = h.presentableMembers(members, userID)
|
||||
|
||||
visibleChannels, overrides, err := h.readyVisibleChannels(ctx, database, userID, role, channels)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Per-user unread counts.
|
||||
unreadMap, err := database.GetChannelUnreadCounts(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("buildReady GetChannelUnreadCounts: %w", err)
|
||||
}
|
||||
|
||||
// Build protocol-compliant channel objects (strip extra fields).
|
||||
channelPayloads := readyChannelPayloads(visibleChannels, overrides, unreadMap, role)
|
||||
|
||||
// Load open DM channels for this user. Hoisted above the voice-state
|
||||
// filter below so DM channel IDs can seed visibleSet — permissions.Checker
|
||||
// (and therefore visibleChannels) deliberately skips DM channels, since
|
||||
// their visibility is membership-based rather than role-based, so without
|
||||
// this a DM voice call's voice_state rows would never make it into ready.
|
||||
dmChannels, err := h.readyDMChannels(ctx, database, userID, unreadMap)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
voiceStates := h.readyVoiceStates(ctx, database, channels, visibleChannels, dmChannels, userID)
|
||||
|
||||
serverName, motd := h.getCachedSettings(ctx)
|
||||
|
||||
|
||||
+151
-115
@@ -4,70 +4,164 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/owncord/server/auth"
|
||||
"github.com/owncord/server/permissions"
|
||||
)
|
||||
|
||||
// handleVoiceMuteV2 processes a voice_mute command.
|
||||
func handleVoiceMuteV2(ctx context.Context, cmd Command, info ClientInfo, deps any) Result {
|
||||
d := deps.(VoiceDeps)
|
||||
muteCmd := cmd.(VoiceMuteCmd)
|
||||
// voiceSelfToggle parameterises the two self-toggle handlers, voice_mute and
|
||||
// voice_deafen. They were verbatim duplicates of each other, differing only in
|
||||
// the fields below; a fix landing on one and not the other — the asymmetric
|
||||
// moderator gate is exactly such a fix — is the failure mode this collapse
|
||||
// removes.
|
||||
type voiceSelfToggle struct {
|
||||
rateKey string // auth.Key namespace, "voice_mute" / "voice_deafen"
|
||||
rateLimit int
|
||||
rateWindow time.Duration
|
||||
rateMsg string
|
||||
// serverDeafen picks which moderator flag refuseIfServerSilenced consults:
|
||||
// false = ServerMuted (blocks a self-unmute), true = ServerDeafened (blocks a
|
||||
// self-undeafen). A server deafen is the moderator's to lift, same as a mute.
|
||||
serverDeafen bool
|
||||
update func(ctx context.Context, userID int64, on bool) error
|
||||
updateLog string // slog.Error message when update fails
|
||||
failMsg string
|
||||
changedLog string // slog.Debug message once applied
|
||||
stateKey string // slog key carrying the new value
|
||||
}
|
||||
|
||||
// voiceSelfToggleV2 is the shared body of handleVoiceMuteV2 and
|
||||
// handleVoiceDeafenV2.
|
||||
func voiceSelfToggleV2(ctx context.Context, d VoiceDeps, info ClientInfo, on bool, t voiceSelfToggle) Result {
|
||||
userID := info.UserID
|
||||
|
||||
ratKey := auth.Key("voice_mute", userID)
|
||||
if d.Limiter != nil && !d.Limiter.Allow(ratKey, voiceMuteRateLimit, voiceMuteWindow) {
|
||||
return Result{Error: ClientError{Code: ErrCodeRateLimited, Message: "too many mute toggles"}}
|
||||
ratKey := auth.Key(t.rateKey, userID)
|
||||
if d.Limiter != nil && !d.Limiter.Allow(ratKey, t.rateLimit, t.rateWindow) {
|
||||
return Result{Error: ClientError{Code: ErrCodeRateLimited, Message: t.rateMsg}}
|
||||
}
|
||||
|
||||
if info.VoiceChannelID == 0 {
|
||||
return Result{Error: ClientError{Code: ErrCodeVoiceError, Message: "not in a voice channel"}}
|
||||
}
|
||||
|
||||
// A moderator-imposed mute is not the user's to lift. Only the unmute
|
||||
// direction reads the row: muting oneself is always allowed.
|
||||
if !muteCmd.Muted() {
|
||||
if r := refuseIfServerSilenced(ctx, d, userID, false); r != nil {
|
||||
// A moderator-imposed mute/deafen is not the user's to lift. Only the
|
||||
// clearing direction reads the row: silencing oneself is always allowed.
|
||||
if !on {
|
||||
if r := refuseIfServerSilenced(ctx, d, userID, t.serverDeafen); r != nil {
|
||||
return *r
|
||||
}
|
||||
}
|
||||
|
||||
if err := d.DB.UpdateVoiceMute(ctx, userID, muteCmd.Muted()); err != nil {
|
||||
slog.Error("ws handleVoiceMuteV2 UpdateVoiceMute", "err", err, "user_id", userID)
|
||||
return Result{Error: ClientError{Code: ErrCodeInternal, Message: "failed to update mute state"}}
|
||||
if err := t.update(ctx, userID, on); err != nil {
|
||||
slog.Error(t.updateLog, "err", err, "user_id", userID)
|
||||
return Result{Error: ClientError{Code: ErrCodeInternal, Message: t.failMsg}}
|
||||
}
|
||||
slog.Debug("voice mute changed", "user_id", userID, "muted", muteCmd.Muted(), "channel_id", info.VoiceChannelID)
|
||||
slog.Debug(t.changedLog, "user_id", userID, t.stateKey, on, "channel_id", info.VoiceChannelID)
|
||||
|
||||
return voiceStateBroadcast(ctx, d, userID)
|
||||
}
|
||||
|
||||
// handleVoiceMuteV2 processes a voice_mute command.
|
||||
func handleVoiceMuteV2(ctx context.Context, cmd Command, info ClientInfo, deps any) Result {
|
||||
d := deps.(VoiceDeps)
|
||||
muteCmd := cmd.(VoiceMuteCmd)
|
||||
return voiceSelfToggleV2(ctx, d, info, muteCmd.Muted(), voiceSelfToggle{
|
||||
rateKey: "voice_mute",
|
||||
rateLimit: voiceMuteRateLimit,
|
||||
rateWindow: voiceMuteWindow,
|
||||
rateMsg: "too many mute toggles",
|
||||
serverDeafen: false,
|
||||
update: d.DB.UpdateVoiceMute,
|
||||
updateLog: "ws handleVoiceMuteV2 UpdateVoiceMute",
|
||||
failMsg: "failed to update mute state",
|
||||
changedLog: "voice mute changed",
|
||||
stateKey: "muted",
|
||||
})
|
||||
}
|
||||
|
||||
// handleVoiceDeafenV2 processes a voice_deafen command.
|
||||
func handleVoiceDeafenV2(ctx context.Context, cmd Command, info ClientInfo, deps any) Result {
|
||||
d := deps.(VoiceDeps)
|
||||
deafenCmd := cmd.(VoiceDeafenCmd)
|
||||
userID := info.UserID
|
||||
return voiceSelfToggleV2(ctx, d, info, deafenCmd.Deafened(), voiceSelfToggle{
|
||||
rateKey: "voice_deafen",
|
||||
rateLimit: voiceDeafenRateLimit,
|
||||
rateWindow: voiceDeafenWindow,
|
||||
rateMsg: "too many deafen toggles",
|
||||
serverDeafen: true,
|
||||
update: d.DB.UpdateVoiceDeafen,
|
||||
updateLog: "ws handleVoiceDeafenV2 UpdateVoiceDeafen",
|
||||
failMsg: "failed to update deafen state",
|
||||
changedLog: "voice deafen changed",
|
||||
stateKey: "deafened",
|
||||
})
|
||||
}
|
||||
|
||||
ratKey := auth.Key("voice_deafen", userID)
|
||||
if d.Limiter != nil && !d.Limiter.Allow(ratKey, voiceDeafenRateLimit, voiceDeafenWindow) {
|
||||
return Result{Error: ClientError{Code: ErrCodeRateLimited, Message: "too many deafen toggles"}}
|
||||
// voiceStreamToggle parameterises the two video-stream handlers, voice_camera
|
||||
// and voice_screenshare. Like the self-toggles above they were verbatim
|
||||
// duplicates. Keeping them one function is not only tidier: camera and
|
||||
// screenshare draw from a single per-channel voice_max_video budget (OC-0023),
|
||||
// and that bug existed precisely because the two paths had drifted apart.
|
||||
type voiceStreamToggle struct {
|
||||
rateKey string // auth.Key namespace, "voice_camera" / "voice_screenshare"
|
||||
rateLimit int
|
||||
rateWindow time.Duration
|
||||
rateMsg string
|
||||
perm int64 // permission required to ENABLE the stream
|
||||
permLabel string // its name, for the refusal
|
||||
// tryReserve is the atomic under-cap check-and-set for this stream's
|
||||
// column; update is its plain unconditional write. Both are handed to
|
||||
// enableVideoSlot, which owns the shared-budget rule.
|
||||
tryReserve func(ctx context.Context, userID, channelID int64, maxVideo int) (bool, error)
|
||||
update func(ctx context.Context, userID int64, enabled bool) error
|
||||
logPrefix string // handler name, for slog messages
|
||||
kind string // "camera" / "screenshare", used in operator-facing text
|
||||
disableLog string // slog.Error message when the disable update fails
|
||||
changedLog string // slog.Debug message once applied
|
||||
}
|
||||
|
||||
// voiceStreamToggleV2 is the shared body of handleVoiceCameraV2 and
|
||||
// handleVoiceScreenshareV2.
|
||||
//
|
||||
// Only the enable direction is gated on the permission, mirroring
|
||||
// handleVoiceMuteV2/handleVoiceDeafenV2's asymmetric gate — once a moderator
|
||||
// revokes it mid-call the user must still be able to turn the stream off, or
|
||||
// the column (voice_states.camera / voice_states.screenshare) stays stuck at 1:
|
||||
// for camera that permanently burns a voice_max_video slot, and for screenshare
|
||||
// every subsequent voice_state keeps advertising a stream nobody can watch.
|
||||
// Nothing else ever clears either one short of leaving voice.
|
||||
func voiceStreamToggleV2(ctx context.Context, d VoiceDeps, info ClientInfo, enabled bool, t voiceStreamToggle) Result {
|
||||
userID := info.UserID
|
||||
voiceChID := info.VoiceChannelID
|
||||
|
||||
ratKey := auth.Key(t.rateKey, userID)
|
||||
if d.Limiter != nil && !d.Limiter.Allow(ratKey, t.rateLimit, t.rateWindow) {
|
||||
return Result{Error: ClientError{Code: ErrCodeRateLimited, Message: t.rateMsg}}
|
||||
}
|
||||
|
||||
if info.VoiceChannelID == 0 {
|
||||
if voiceChID == 0 {
|
||||
return Result{Error: ClientError{Code: ErrCodeVoiceError, Message: "not in a voice channel"}}
|
||||
}
|
||||
|
||||
// See handleVoiceMuteV2: server deafen is the moderator's to lift.
|
||||
if !deafenCmd.Deafened() {
|
||||
if r := refuseIfServerSilenced(ctx, d, userID, true); r != nil {
|
||||
if enabled {
|
||||
if r := requirePerm(ctx, d.DB, d.Permissions, d.PermSvc, userID, voiceChID, t.perm, t.permLabel); r != nil {
|
||||
return *r
|
||||
}
|
||||
// Enforce the channel's shared voice_max_video budget atomically (OC-0023:
|
||||
// enableVideoSlot's query counts camera = 1 OR screenshare = 1 rows, so
|
||||
// neither kind can occupy a slot the cap meant to deny it nor hide from
|
||||
// the other's count).
|
||||
if r := enableVideoSlot(ctx, d, userID, voiceChID, t.tryReserve, t.update, t.logPrefix, t.kind); r != nil {
|
||||
return *r
|
||||
}
|
||||
} else {
|
||||
if err := t.update(ctx, userID, false); err != nil {
|
||||
slog.Error(t.disableLog, "err", err, "user_id", userID)
|
||||
return Result{Error: ClientError{Code: ErrCodeInternal, Message: "failed to update " + t.kind + " state"}}
|
||||
}
|
||||
}
|
||||
|
||||
if err := d.DB.UpdateVoiceDeafen(ctx, userID, deafenCmd.Deafened()); err != nil {
|
||||
slog.Error("ws handleVoiceDeafenV2 UpdateVoiceDeafen", "err", err, "user_id", userID)
|
||||
return Result{Error: ClientError{Code: ErrCodeInternal, Message: "failed to update deafen state"}}
|
||||
}
|
||||
slog.Debug("voice deafen changed", "user_id", userID, "deafened", deafenCmd.Deafened(), "channel_id", info.VoiceChannelID)
|
||||
slog.Debug(t.changedLog, "user_id", userID, "enabled", enabled, "channel_id", voiceChID)
|
||||
|
||||
return voiceStateBroadcast(ctx, d, userID)
|
||||
}
|
||||
@@ -76,98 +170,40 @@ func handleVoiceDeafenV2(ctx context.Context, cmd Command, info ClientInfo, deps
|
||||
func handleVoiceCameraV2(ctx context.Context, cmd Command, info ClientInfo, deps any) Result {
|
||||
d := deps.(VoiceDeps)
|
||||
cameraCmd := cmd.(VoiceCameraCmd)
|
||||
userID := info.UserID
|
||||
voiceChID := info.VoiceChannelID
|
||||
|
||||
ratKey := auth.Key("voice_camera", userID)
|
||||
if d.Limiter != nil && !d.Limiter.Allow(ratKey, voiceCameraRateLimit, voiceCameraWindow) {
|
||||
return Result{Error: ClientError{Code: ErrCodeRateLimited, Message: "too many camera toggles"}}
|
||||
}
|
||||
|
||||
if voiceChID == 0 {
|
||||
return Result{Error: ClientError{Code: ErrCodeVoiceError, Message: "not in a voice channel"}}
|
||||
}
|
||||
|
||||
enabled := cameraCmd.Enabled()
|
||||
|
||||
// Only the enable direction is gated on USE_VIDEO — mirrors
|
||||
// handleVoiceMuteV2/handleVoiceDeafenV2's asymmetric gate: once a
|
||||
// moderator revokes the permission mid-call, the user must still be able
|
||||
// to turn their camera off, or voice_states.camera stays stuck at 1 —
|
||||
// permanently burning a voice_max_video slot — until they leave voice,
|
||||
// since nothing else ever clears it.
|
||||
if enabled {
|
||||
if r := requirePerm(ctx, d.DB, d.Permissions, d.PermSvc, userID, voiceChID, permissions.UseVideo, "USE_VIDEO"); r != nil {
|
||||
return *r
|
||||
}
|
||||
}
|
||||
|
||||
// Enforce MaxVideo limit when enabling camera using an atomic check-and-update.
|
||||
// Camera and screenshare draw from the same voice_max_video budget
|
||||
// (OC-0023), so this gate is shared with handleVoiceScreenshareV2 below.
|
||||
if enabled {
|
||||
if r := enableVideoSlot(ctx, d, userID, voiceChID, d.DB.EnableCameraIfUnderLimit, d.DB.UpdateVoiceCamera, "handleVoiceCameraV2", "camera"); r != nil {
|
||||
return *r
|
||||
}
|
||||
} else {
|
||||
if err := d.DB.UpdateVoiceCamera(ctx, userID, false); err != nil {
|
||||
slog.Error("ws handleVoiceCameraV2 UpdateVoiceCamera", "err", err, "user_id", userID)
|
||||
return Result{Error: ClientError{Code: ErrCodeInternal, Message: "failed to update camera state"}}
|
||||
}
|
||||
}
|
||||
slog.Debug("voice camera changed", "user_id", userID, "enabled", enabled, "channel_id", voiceChID)
|
||||
|
||||
return voiceStateBroadcast(ctx, d, userID)
|
||||
return voiceStreamToggleV2(ctx, d, info, cameraCmd.Enabled(), voiceStreamToggle{
|
||||
rateKey: "voice_camera",
|
||||
rateLimit: voiceCameraRateLimit,
|
||||
rateWindow: voiceCameraWindow,
|
||||
rateMsg: "too many camera toggles",
|
||||
perm: permissions.UseVideo,
|
||||
permLabel: "USE_VIDEO",
|
||||
tryReserve: d.DB.EnableCameraIfUnderLimit,
|
||||
update: d.DB.UpdateVoiceCamera,
|
||||
logPrefix: "handleVoiceCameraV2",
|
||||
kind: "camera",
|
||||
disableLog: "ws handleVoiceCameraV2 UpdateVoiceCamera",
|
||||
changedLog: "voice camera changed",
|
||||
})
|
||||
}
|
||||
|
||||
// handleVoiceScreenshareV2 processes a voice_screenshare command.
|
||||
func handleVoiceScreenshareV2(ctx context.Context, cmd Command, info ClientInfo, deps any) Result {
|
||||
d := deps.(VoiceDeps)
|
||||
ssCmd := cmd.(VoiceScreenshareCmd)
|
||||
userID := info.UserID
|
||||
voiceChID := info.VoiceChannelID
|
||||
|
||||
ratKey := auth.Key("voice_screenshare", userID)
|
||||
if d.Limiter != nil && !d.Limiter.Allow(ratKey, voiceScreenshareRateLimit, voiceScreenshareWindow) {
|
||||
return Result{Error: ClientError{Code: ErrCodeRateLimited, Message: "too many screenshare toggles"}}
|
||||
}
|
||||
|
||||
if voiceChID == 0 {
|
||||
return Result{Error: ClientError{Code: ErrCodeVoiceError, Message: "not in a voice channel"}}
|
||||
}
|
||||
|
||||
enabled := ssCmd.Enabled()
|
||||
|
||||
// Only the enable direction is gated on SHARE_SCREEN — mirrors
|
||||
// handleVoiceMuteV2/handleVoiceDeafenV2's asymmetric gate: once a
|
||||
// moderator revokes the permission mid-share, the user must still be able
|
||||
// to stop sharing, or voice_states.screenshare stays stuck at 1 — every
|
||||
// subsequent voice_state keeps advertising a stream nobody can watch —
|
||||
// until they leave voice.
|
||||
if enabled {
|
||||
if r := requirePerm(ctx, d.DB, d.Permissions, d.PermSvc, userID, voiceChID, permissions.ShareScreen, "SHARE_SCREEN"); r != nil {
|
||||
return *r
|
||||
}
|
||||
}
|
||||
|
||||
// Enforce the same voice_max_video budget handleVoiceCameraV2 enforces —
|
||||
// camera and screenshare are both "video streams" against one cap
|
||||
// (OC-0023): a screenshare must not be able to occupy a slot the cap
|
||||
// intended to deny it, and must not be invisible to the camera gate's
|
||||
// count either (enableVideoSlot's atomic query counts both fields).
|
||||
if enabled {
|
||||
if r := enableVideoSlot(ctx, d, userID, voiceChID, d.DB.EnableScreenshareIfUnderLimit, d.DB.UpdateVoiceScreenshare, "handleVoiceScreenshareV2", "screenshare"); r != nil {
|
||||
return *r
|
||||
}
|
||||
} else {
|
||||
if err := d.DB.UpdateVoiceScreenshare(ctx, userID, false); err != nil {
|
||||
slog.Error("ws handleVoiceScreenshareV2 UpdateVoiceScreenshare", "err", err, "user_id", userID)
|
||||
return Result{Error: ClientError{Code: ErrCodeInternal, Message: "failed to update screenshare state"}}
|
||||
}
|
||||
}
|
||||
slog.Debug("voice screenshare changed", "user_id", userID, "enabled", enabled, "channel_id", voiceChID)
|
||||
|
||||
return voiceStateBroadcast(ctx, d, userID)
|
||||
return voiceStreamToggleV2(ctx, d, info, ssCmd.Enabled(), voiceStreamToggle{
|
||||
rateKey: "voice_screenshare",
|
||||
rateLimit: voiceScreenshareRateLimit,
|
||||
rateWindow: voiceScreenshareWindow,
|
||||
rateMsg: "too many screenshare toggles",
|
||||
perm: permissions.ShareScreen,
|
||||
permLabel: "SHARE_SCREEN",
|
||||
tryReserve: d.DB.EnableScreenshareIfUnderLimit,
|
||||
update: d.DB.UpdateVoiceScreenshare,
|
||||
logPrefix: "handleVoiceScreenshareV2",
|
||||
kind: "screenshare",
|
||||
disableLog: "ws handleVoiceScreenshareV2 UpdateVoiceScreenshare",
|
||||
changedLog: "voice screenshare changed",
|
||||
})
|
||||
}
|
||||
|
||||
// enableVideoSlot enforces the channel's shared voice_max_video budget
|
||||
|
||||
+125
-52
@@ -54,26 +54,56 @@ var voiceJoinPostTokenRaceHook func(*Client)
|
||||
// 8. Broadcasts voice_state to all clients.
|
||||
// 9. Sends voice_config to the joiner.
|
||||
func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMessage) {
|
||||
channelID, ch, ok := h.voiceJoinPrecheck(ctx, c, payload)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
wasServerMuted, wasServerDeafened, ok := h.voiceJoinLeaveCurrent(ctx, c, channelID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
state, ok := h.voiceJoinPersist(ctx, c, ch, channelID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
state = h.voiceJoinRestoreModFlags(ctx, c, channelID, state, wasServerMuted, wasServerDeafened)
|
||||
|
||||
if !h.voiceJoinGrantToken(ctx, c, channelID, state) {
|
||||
return
|
||||
}
|
||||
|
||||
h.voiceJoinComplete(ctx, c, ch, channelID, state)
|
||||
}
|
||||
|
||||
// voiceJoinPrecheck runs every gate that must pass before handleVoiceJoin
|
||||
// mutates any state: rate limit, payload parse, CONNECT_VOICE, channel
|
||||
// existence, channel type, DM block, archive, authenticated user and LiveKit
|
||||
// availability. It reports the target channel id and row when the join may
|
||||
// proceed; on refusal it has already sent the error frame and returns false.
|
||||
func (h *Hub) voiceJoinPrecheck(ctx context.Context, c *Client, payload json.RawMessage) (int64, *db.Channel, bool) {
|
||||
// Rate limit: voice_join broadcasts a voice_state update to every connected
|
||||
// client, so cap how often a single user can trigger the fan-out. Mirrors the
|
||||
// Limiter.Allow(...) idiom used by the voice control handlers.
|
||||
ratKey := auth.Key("voice_join", c.userID)
|
||||
if h.limiter != nil && !h.limiter.Allow(ratKey, voiceJoinRateLimit, voiceJoinWindow) {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeRateLimited, "too many voice join attempts"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
channelID, err := parseChannelID(payload)
|
||||
if err != nil || channelID <= 0 {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "channel_id must be a positive integer"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
// channel_id is attacker-controlled, so the gate must be channel-TYPE aware:
|
||||
// a role-only check passes for any DM channel id (DMs have no overrides), and
|
||||
// the token minted below carries RoomJoin+CanSubscribe for that DM's room.
|
||||
if !h.requireChannelAccess(ctx, c, channelID, permissions.ConnectVoice, "CONNECT_VOICE") {
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
// Validate the target channel exists before any state changes (leaving
|
||||
@@ -81,7 +111,7 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
ch, err := h.db.GetChannel(ctx, channelID)
|
||||
if err != nil || ch == nil {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeNotFound, "channel not found"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
// channel_id is attacker-controlled and requireChannelAccess above only
|
||||
@@ -92,7 +122,7 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
// group voice calls join through this same handler.
|
||||
if ch.Type != "voice" && ch.Type != "dm" {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "not a voice channel"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
// A blocked user is still a DM participant — blocking never touches
|
||||
@@ -106,7 +136,7 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
if ch.Type == "dm" {
|
||||
if err := service.RequireDMNotBlocked(ctx, h.db, c.userID, channelID); err != nil {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot join voice: blocked"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
}
|
||||
|
||||
@@ -117,7 +147,7 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
// archive transition also evicts whoever is already inside.
|
||||
if ch.Archived {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "channel is archived"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
// Ensure authenticated user is present before any state changes.
|
||||
@@ -126,14 +156,14 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
if c.user == nil {
|
||||
slog.Error("handleVoiceJoin: nil user on client", "user_id", c.userID)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInternal, "not authenticated"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
// Hard-fail when LiveKit is not configured — without an SFU the client
|
||||
// cannot connect to voice, so persisting state would create a ghost.
|
||||
if h.livekit == nil {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeVoiceError, "voice is not configured on this server"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
// Guard: reject voice join if the companion LiveKit process is not running
|
||||
@@ -141,15 +171,25 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
if h.lkProcess != nil && !h.lkProcess.IsRunning() {
|
||||
slog.Warn("handleVoiceJoin: LiveKit process not running", "user_id", c.userID)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeVoiceError, "voice is temporarily unavailable — LiveKit is not running"))
|
||||
return
|
||||
return 0, nil, false
|
||||
}
|
||||
|
||||
return channelID, ch, true
|
||||
}
|
||||
|
||||
// voiceJoinLeaveCurrent handles the case where the client is already in a
|
||||
// voice channel: it no-ops a re-join of the same channel, and for a switch it
|
||||
// snapshots the moderator-imposed mute/deafen flags, leaves the old channel
|
||||
// and verifies the old row is really gone. The two booleans are the
|
||||
// snapshotted flags for voiceJoinRestoreModFlags; false in the third position
|
||||
// means the join must not proceed (the error frame has already been sent).
|
||||
func (h *Hub) voiceJoinLeaveCurrent(ctx context.Context, c *Client, channelID int64) (bool, bool, bool) {
|
||||
currentChID := c.getVoiceChID()
|
||||
|
||||
// If user is already in the same voice channel, no-op.
|
||||
if currentChID == channelID {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeAlreadyJoined, "already in this voice channel"))
|
||||
return
|
||||
return false, false, false
|
||||
}
|
||||
|
||||
// A moderator-imposed mute/deafen must survive a channel switch.
|
||||
@@ -186,7 +226,7 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
slog.Warn("handleVoiceJoin: could not verify voice state cleared",
|
||||
"user_id", c.userID, "err", err)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInternal, "voice channel switch failed — please try again"))
|
||||
return
|
||||
return false, false, false
|
||||
}
|
||||
if vs != nil {
|
||||
slog.Warn("handleVoiceJoin: stale voice state persists after leave, aborting switch",
|
||||
@@ -206,28 +246,35 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
// (re-broadcasting voice_leave, harmlessly) within one tick, and
|
||||
// the user_id-PK upsert lets the user rejoin immediately.
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInternal, "voice channel switch failed — please try again"))
|
||||
return
|
||||
return false, false, false
|
||||
}
|
||||
}
|
||||
|
||||
return wasServerMuted, wasServerDeafened, true
|
||||
}
|
||||
|
||||
// voiceJoinPersist commits the join to the DB under the channel's capacity
|
||||
// limit, loads back the persisted row and publishes the client's in-memory
|
||||
// voice state. Returns false once the error frame has been sent.
|
||||
func (h *Hub) voiceJoinPersist(ctx context.Context, c *Client, ch *db.Channel, channelID int64) (*db.VoiceState, bool) {
|
||||
// Check channel capacity and persist to DB atomically.
|
||||
maxUsers := ch.VoiceMaxUsers
|
||||
if maxUsers > 0 {
|
||||
if err := h.db.JoinVoiceChannelIfCapacity(ctx, c.userID, channelID, maxUsers); err != nil {
|
||||
if errors.Is(err, db.ErrChannelFull) {
|
||||
c.sendMsg(buildErrorMsg(ErrCodeChannelFull, "voice channel is full"))
|
||||
return
|
||||
return nil, false
|
||||
}
|
||||
slog.Error("ws handleVoiceJoin JoinVoiceChannelIfCapacity", "err", err, "user_id", c.userID)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to join voice channel"))
|
||||
return
|
||||
return nil, false
|
||||
}
|
||||
} else {
|
||||
// No capacity limit — use standard join.
|
||||
if err := h.db.JoinVoiceChannel(ctx, c.userID, channelID); err != nil {
|
||||
slog.Error("ws handleVoiceJoin JoinVoiceChannel", "err", err, "user_id", c.userID)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to join voice channel"))
|
||||
return
|
||||
return nil, false
|
||||
}
|
||||
}
|
||||
|
||||
@@ -238,7 +285,7 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
slog.Error("ws handleVoiceJoin GetVoiceState", "err", err, "user_id", c.userID)
|
||||
h.rollbackVoiceJoin(ctx, c, channelID, "", false)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to join voice channel"))
|
||||
return
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// BUG-088: set the client's voice channel as soon as the DB row is
|
||||
@@ -252,6 +299,13 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
// c.clearVoiceChID(), same as before.
|
||||
c.setVoiceState(channelID, state.JoinedAt)
|
||||
|
||||
return state, true
|
||||
}
|
||||
|
||||
// voiceJoinRestoreModFlags re-applies a moderator-imposed mute/deafen that
|
||||
// predates a channel switch and returns the voice state the caller should
|
||||
// broadcast — the re-read row when the restore ran, the original otherwise.
|
||||
func (h *Hub) voiceJoinRestoreModFlags(ctx context.Context, c *Client, channelID int64, state *db.VoiceState, wasServerMuted, wasServerDeafened bool) *db.VoiceState {
|
||||
// Restore a moderator-imposed mute/deafen that predates this switch (see
|
||||
// the snapshot above). Best-effort: a failure here is logged but does not
|
||||
// fail the join, matching every other SetVoiceServerMute/Deafen call site.
|
||||
@@ -283,6 +337,49 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
}
|
||||
}
|
||||
|
||||
return state
|
||||
}
|
||||
|
||||
// voiceJoinPublishPerms derives the SFU publish permissions from role —
|
||||
// prevents SFU-level bypass when the client connects directly via direct_url
|
||||
// (BUG-128). With a PermissionService the three bits come from the per-user
|
||||
// cache; the bare-hub fallback answers them from one role fetch + one
|
||||
// overrides fetch via HasChannelPermBatch instead of three hasChannelPerm
|
||||
// round trips. Both branches fail closed: an unresolved role or override map
|
||||
// yields no publish grants (admins bypass overrides, so an override fetch
|
||||
// error cannot demote them).
|
||||
func (h *Hub) voiceJoinPublishPerms(ctx context.Context, userID, channelID int64) (canPublish, canVideo, canScreenShare bool) {
|
||||
if h.perms != nil {
|
||||
// PermissionService answers all three bits from one cached
|
||||
// role+overrides snapshot (populated by the CONNECT_VOICE gate
|
||||
// above, so these are cache hits). Same fail-closed posture: an
|
||||
// unresolved role or override map yields no publish grants.
|
||||
canPublish = h.perms.HasChannelPerm(ctx, userID, channelID, permissions.SpeakVoice)
|
||||
canVideo = h.perms.HasChannelPerm(ctx, userID, channelID, permissions.UseVideo)
|
||||
canScreenShare = h.perms.HasChannelPerm(ctx, userID, channelID, permissions.ShareScreen)
|
||||
} else if role, roleErr := h.db.GetRoleForUser(ctx, userID); roleErr == nil && role != nil {
|
||||
// Admins bypass overrides, so skip the fetch for them (mirrors
|
||||
// computeAllowedChannels); HasChannelPermBatch answers true from
|
||||
// the role bits alone.
|
||||
var overrides map[int64]db.ChannelOverride
|
||||
var oErr error
|
||||
if !permissions.HasAdmin(role.Permissions) {
|
||||
overrides, oErr = h.db.GetChannelOverridesFor(ctx, role.ID, userID)
|
||||
}
|
||||
if oErr == nil {
|
||||
po := permOverrides(overrides)
|
||||
canPublish = h.permChecker.HasChannelPermBatch(role.Permissions, po, channelID, permissions.SpeakVoice)
|
||||
canVideo = h.permChecker.HasChannelPermBatch(role.Permissions, po, channelID, permissions.UseVideo)
|
||||
canScreenShare = h.permChecker.HasChannelPermBatch(role.Permissions, po, channelID, permissions.ShareScreen)
|
||||
}
|
||||
}
|
||||
return canPublish, canVideo, canScreenShare
|
||||
}
|
||||
|
||||
// voiceJoinGrantToken mints the LiveKit credential and delivers it, withholding
|
||||
// it if the join was superseded in the meantime. Returns false once the join
|
||||
// has been abandoned (rolled back, or superseded) and must not complete.
|
||||
func (h *Hub) voiceJoinGrantToken(ctx context.Context, c *Client, channelID int64, state *db.VoiceState) bool {
|
||||
// Generate LiveKit token if LiveKit client is available.
|
||||
// Token generation failure is fatal — without a token the client cannot
|
||||
// connect to the SFU, so we must roll back the DB join.
|
||||
@@ -291,46 +388,14 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
// is still called with broadcast=false, so a failure here does not
|
||||
// broadcast a spurious voice_leave for a join no other client ever saw.
|
||||
if h.livekit != nil {
|
||||
// Derive publish permissions from role — prevents SFU-level bypass
|
||||
// when client connects directly via direct_url (BUG-128). With a
|
||||
// PermissionService the three bits come from the per-user cache; the
|
||||
// bare-hub fallback answers them from one role fetch + one overrides
|
||||
// fetch via HasChannelPermBatch instead of three hasChannelPerm round
|
||||
// trips. Both branches fail closed: an unresolved role or override map
|
||||
// yields no publish grants (admins bypass overrides, so an override
|
||||
// fetch error cannot demote them).
|
||||
var canPublish, canVideo, canScreenShare bool
|
||||
canPublish, canVideo, canScreenShare := h.voiceJoinPublishPerms(ctx, c.userID, channelID)
|
||||
canSubscribe := true
|
||||
if h.perms != nil {
|
||||
// PermissionService answers all three bits from one cached
|
||||
// role+overrides snapshot (populated by the CONNECT_VOICE gate
|
||||
// above, so these are cache hits). Same fail-closed posture: an
|
||||
// unresolved role or override map yields no publish grants.
|
||||
canPublish = h.perms.HasChannelPerm(ctx, c.userID, channelID, permissions.SpeakVoice)
|
||||
canVideo = h.perms.HasChannelPerm(ctx, c.userID, channelID, permissions.UseVideo)
|
||||
canScreenShare = h.perms.HasChannelPerm(ctx, c.userID, channelID, permissions.ShareScreen)
|
||||
} else if role, roleErr := h.db.GetRoleForUser(ctx, c.userID); roleErr == nil && role != nil {
|
||||
// Admins bypass overrides, so skip the fetch for them (mirrors
|
||||
// computeAllowedChannels); HasChannelPermBatch answers true from
|
||||
// the role bits alone.
|
||||
var overrides map[int64]db.ChannelOverride
|
||||
var oErr error
|
||||
if !permissions.HasAdmin(role.Permissions) {
|
||||
overrides, oErr = h.db.GetChannelOverridesFor(ctx, role.ID, c.userID)
|
||||
}
|
||||
if oErr == nil {
|
||||
po := permOverrides(overrides)
|
||||
canPublish = h.permChecker.HasChannelPermBatch(role.Permissions, po, channelID, permissions.SpeakVoice)
|
||||
canVideo = h.permChecker.HasChannelPermBatch(role.Permissions, po, channelID, permissions.UseVideo)
|
||||
canScreenShare = h.permChecker.HasChannelPermBatch(role.Permissions, po, channelID, permissions.ShareScreen)
|
||||
}
|
||||
}
|
||||
token, tokenErr := h.livekit.GenerateToken(c.userID, c.user.Username, channelID, state.JoinedAt, canPublish, canSubscribe, canVideo, canScreenShare)
|
||||
if tokenErr != nil {
|
||||
slog.Error("ws handleVoiceJoin GenerateToken", "err", tokenErr, "user_id", c.userID)
|
||||
h.rollbackVoiceJoin(ctx, c, channelID, state.JoinedAt, false)
|
||||
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to generate voice token"))
|
||||
return
|
||||
return false
|
||||
}
|
||||
if voiceJoinPostTokenRaceHook != nil {
|
||||
voiceJoinPostTokenRaceHook(c)
|
||||
@@ -361,7 +426,7 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
slog.Warn("ws handleVoiceJoin: RemoveParticipant after supersession failed (may already be gone)",
|
||||
"err", err, "user_id", c.userID, "channel_id", channelID)
|
||||
}
|
||||
return
|
||||
return false
|
||||
}
|
||||
// Send both proxy path and direct URL. The client uses direct_url
|
||||
// when on localhost (avoids self-signed TLS issues with WebView
|
||||
@@ -374,6 +439,13 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
c.sendMsg(buildVoiceToken(channelID, token, "/livekit", h.livekit.URL(), isKeyHolder))
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// voiceJoinComplete finishes a join that survived every guard: voice topic
|
||||
// subscription, key-holder election, the joiner's own voice_state fan-out, the
|
||||
// existing participants' states and E2EE keys, and voice_config.
|
||||
func (h *Hub) voiceJoinComplete(ctx context.Context, c *Client, ch *db.Channel, channelID int64, state *db.VoiceState) {
|
||||
// Voice channel state itself was already set above (BUG-088), immediately
|
||||
// after the DB row committed — which also means a concurrent eviction (the
|
||||
// revocation sweep, a participant_left webhook, a moderator kick/move) can
|
||||
@@ -434,6 +506,7 @@ func (h *Hub) handleVoiceJoin(ctx context.Context, c *Client, payload json.RawMe
|
||||
"quality", q, "channel_id", channelID)
|
||||
}
|
||||
}
|
||||
maxUsers := ch.VoiceMaxUsers
|
||||
bitrate := qualityBitrate(quality)
|
||||
c.sendMsg(buildVoiceConfig(channelID, quality, bitrate, maxUsers))
|
||||
|
||||
|
||||
@@ -259,56 +259,7 @@ func handleVoiceModDeafenV2(ctx context.Context, cmd Command, info ClientInfo, d
|
||||
if err != nil {
|
||||
slog.Error("ws handleVoiceModDeafenV2 SetVoiceServerMute", "err", err, "target_id", c.TargetID())
|
||||
}
|
||||
// The deafen write above already committed as its own statement (no
|
||||
// transaction spans the two — a single UPDATE covering both columns
|
||||
// needs a db-change; see cross_batch). Best-effort undo it rather
|
||||
// than leave server_deafened=1 with server_muted=0: that combination
|
||||
// is not SFU-muted yet still refuses the target's own undeafen
|
||||
// (refuseIfServerSilenced), for a deafen nobody was ever told about.
|
||||
// Detached from ctx — the cancellation that most likely caused the
|
||||
// failure above (the moderator's socket dropping mid-request, or the
|
||||
// target moving off state.ChannelID between the two writes) must not
|
||||
// also abort the rollback.
|
||||
//
|
||||
// Re-read the row's CURRENT channel rather than reusing the stale
|
||||
// state.ChannelID snapshot: when the mismatch above was caused by
|
||||
// the target switching channels (not leaving voice), the row is no
|
||||
// longer on state.ChannelID, so a rollback scoped to that stale
|
||||
// channel matches zero rows and silently no-ops -- exactly the case
|
||||
// this rollback exists to handle (OC-0034). Clearing a restriction
|
||||
// is safe on whatever channel the row is actually on now; if the
|
||||
// row is gone entirely (target left voice), there is nothing left
|
||||
// to roll back.
|
||||
//
|
||||
// The rollback value is the OPPOSITE of the request (!c.Deafened()),
|
||||
// so which channel it is safe to scope to depends on which
|
||||
// direction it runs:
|
||||
// - request was a DEAFEN (c.Deafened()==true): rollback CLEARS.
|
||||
// Clearing a restriction can never authorize anything the
|
||||
// target wasn't already free of, so following the row to
|
||||
// cur.ChannelID is safe -- this is the OC-0034 case above.
|
||||
// - request was an UNDEAFEN (c.Deafened()==false): rollback
|
||||
// APPLIES a restriction. Scoping an apply to cur.ChannelID
|
||||
// would stamp it onto whatever channel the row now points at,
|
||||
// including one voiceModTarget never authorized the actor
|
||||
// against (OC-0036) -- the exact hazard channel-scoping exists
|
||||
// to prevent for the ordinary write path. Scope to
|
||||
// state.ChannelID (the channel that WAS authorized) instead,
|
||||
// so a moved/rejoined target simply matches zero rows.
|
||||
compCtx := context.WithoutCancel(ctx)
|
||||
if cur, gErr := d.DB.GetVoiceState(compCtx, c.TargetID()); gErr != nil {
|
||||
slog.Error("ws handleVoiceModDeafenV2 GetVoiceState for rollback",
|
||||
"err", gErr, "target_id", c.TargetID())
|
||||
} else if cur != nil {
|
||||
rollbackChannelID := cur.ChannelID
|
||||
if !c.Deafened() {
|
||||
rollbackChannelID = state.ChannelID
|
||||
}
|
||||
if _, compErr := d.DB.SetVoiceServerDeafen(compCtx, c.TargetID(), rollbackChannelID, !c.Deafened()); compErr != nil {
|
||||
slog.Error("ws handleVoiceModDeafenV2 SetVoiceServerDeafen rollback failed",
|
||||
"err", compErr, "target_id", c.TargetID())
|
||||
}
|
||||
}
|
||||
voiceModDeafenRollback(ctx, d, c, state)
|
||||
if err != nil {
|
||||
return Result{Error: ClientError{Code: ErrCodeInternal, Message: "failed to update server deafen"}}
|
||||
}
|
||||
@@ -329,6 +280,63 @@ func handleVoiceModDeafenV2(ctx context.Context, cmd Command, info ClientInfo, d
|
||||
return voiceStateBroadcast(ctx, d, c.TargetID())
|
||||
}
|
||||
|
||||
// voiceModDeafenRollback best-effort undoes the server_deafened write
|
||||
// handleVoiceModDeafenV2 committed just before the implied server_muted write
|
||||
// failed to land.
|
||||
//
|
||||
// The deafen write above already committed as its own statement (no
|
||||
// transaction spans the two — a single UPDATE covering both columns
|
||||
// needs a db-change; see cross_batch). Best-effort undo it rather
|
||||
// than leave server_deafened=1 with server_muted=0: that combination
|
||||
// is not SFU-muted yet still refuses the target's own undeafen
|
||||
// (refuseIfServerSilenced), for a deafen nobody was ever told about.
|
||||
// Detached from ctx — the cancellation that most likely caused the
|
||||
// failure above (the moderator's socket dropping mid-request, or the
|
||||
// target moving off state.ChannelID between the two writes) must not
|
||||
// also abort the rollback.
|
||||
//
|
||||
// Re-read the row's CURRENT channel rather than reusing the stale
|
||||
// state.ChannelID snapshot: when the mismatch above was caused by
|
||||
// the target switching channels (not leaving voice), the row is no
|
||||
// longer on state.ChannelID, so a rollback scoped to that stale
|
||||
// channel matches zero rows and silently no-ops -- exactly the case
|
||||
// this rollback exists to handle (OC-0034). Clearing a restriction
|
||||
// is safe on whatever channel the row is actually on now; if the
|
||||
// row is gone entirely (target left voice), there is nothing left
|
||||
// to roll back.
|
||||
//
|
||||
// The rollback value is the OPPOSITE of the request (!c.Deafened()),
|
||||
// so which channel it is safe to scope to depends on which
|
||||
// direction it runs:
|
||||
// - request was a DEAFEN (c.Deafened()==true): rollback CLEARS.
|
||||
// Clearing a restriction can never authorize anything the
|
||||
// target wasn't already free of, so following the row to
|
||||
// cur.ChannelID is safe -- this is the OC-0034 case above.
|
||||
// - request was an UNDEAFEN (c.Deafened()==false): rollback
|
||||
// APPLIES a restriction. Scoping an apply to cur.ChannelID
|
||||
// would stamp it onto whatever channel the row now points at,
|
||||
// including one voiceModTarget never authorized the actor
|
||||
// against (OC-0036) -- the exact hazard channel-scoping exists
|
||||
// to prevent for the ordinary write path. Scope to
|
||||
// state.ChannelID (the channel that WAS authorized) instead,
|
||||
// so a moved/rejoined target simply matches zero rows.
|
||||
func voiceModDeafenRollback(ctx context.Context, d VoiceDeps, c VoiceModDeafenCmd, state *db.VoiceState) {
|
||||
compCtx := context.WithoutCancel(ctx)
|
||||
if cur, gErr := d.DB.GetVoiceState(compCtx, c.TargetID()); gErr != nil {
|
||||
slog.Error("ws handleVoiceModDeafenV2 GetVoiceState for rollback",
|
||||
"err", gErr, "target_id", c.TargetID())
|
||||
} else if cur != nil {
|
||||
rollbackChannelID := cur.ChannelID
|
||||
if !c.Deafened() {
|
||||
rollbackChannelID = state.ChannelID
|
||||
}
|
||||
if _, compErr := d.DB.SetVoiceServerDeafen(compCtx, c.TargetID(), rollbackChannelID, !c.Deafened()); compErr != nil {
|
||||
slog.Error("ws handleVoiceModDeafenV2 SetVoiceServerDeafen rollback failed",
|
||||
"err", compErr, "target_id", c.TargetID())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// handleVoiceModMoveV2 processes a voice_mod_move command.
|
||||
//
|
||||
// The move is a server-driven leave followed by a client-driven re-join: the
|
||||
|
||||
Reference in New Issue
Block a user