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:
J3vb
2026-08-18 20:39:45 +02:00
committed by GitHub
co-authored by Claude Opus 5
parent 7f87be6306
commit 39551de4a6
31 changed files with 3429 additions and 2498 deletions
+168 -118
View File
@@ -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
View File
@@ -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
View File
@@ -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.
+27
View File
@@ -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
View File
@@ -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
// (' -> &#39;, & -> &amp;, " -> &#34;), 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
// (' -> &#39;, & -> &amp;, " -> &#34;), 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 110 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 110 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
View File
@@ -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)
+57 -43
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
+35 -47
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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.
+60 -44
View File
@@ -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
}
+3 -3
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 {
+1 -1
View File
@@ -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)
}
+32 -21
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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))
+58 -50
View File
@@ -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