refactor(web): harden oauth state store and session writes after review

- Drop the FIFO from oauthStates: consumed states left entries behind, so
  an unrate-limited start/cancel cycle grew the slice without bound.
  Evict by oldest expiry instead — the map alone now bounds memory.
- CreateSession runs INSERT + expiry sweep in one transaction.
- slices.Contains replaces a hand-rolled contains; APIBase typo fixed.
- Stale comments and test paths updated; login hover uses --ember-ink.
This commit is contained in:
2026-08-08 08:47:09 +07:00
parent 13e8e73da7
commit 1a7e130f9f
12 changed files with 51 additions and 49 deletions
+17 -25
View File
@@ -9,6 +9,7 @@ import (
"log"
"net/http"
"net/url"
"slices"
"strconv"
"strings"
"sync"
@@ -43,9 +44,9 @@ type DiscordConfig struct {
// RequiredRole, when non-empty, is a role ID a member must hold on top of
// guild membership. Empty by default: membership alone suffices.
RequiredRole string
// APIBBase is the Discord API root; configurable so tests run the whole
// APIBase is the Discord API root; configurable so tests run the whole
// flow against a local stub.
APIBBase string
APIBase string
// RedirectURI is the full public URL of the callback — Discord requires
// the exact string, so it is configured, never derived from headers.
RedirectURI string
@@ -59,7 +60,6 @@ type DiscordConfig struct {
type oauthStates struct {
mu sync.Mutex
expiry map[string]time.Time
order []string // FIFO for eviction when the map is full
}
func newOAuthStates() *oauthStates {
@@ -75,18 +75,19 @@ func (s *oauthStates) put(state string, expires time.Time) {
delete(s.expiry, k)
}
}
// Evict the state closest to expiring when full, so a flood of starts
// cannot grow memory; at worst it invalidates an in-flight sign-in.
if len(s.expiry) >= maxStates {
// Evict from the front until under the cap. The front may already
// have been consumed by take (which removes from the map, not the
// order slice), so pop until the count actually drops.
for len(s.expiry) >= maxStates && len(s.order) > 0 {
oldest := s.order[0]
s.order = s.order[1:]
delete(s.expiry, oldest)
var oldest string
var oldestAt time.Time
for k, at := range s.expiry {
if oldest == "" || at.Before(oldestAt) {
oldest, oldestAt = k, at
}
}
delete(s.expiry, oldest)
}
s.expiry[state] = expires
s.order = append(s.order, state)
}
// take validates and consumes a state in one step: a state works exactly
@@ -107,7 +108,7 @@ func (s *oauthStates) take(state string) bool {
func (h *Handler) discordStart(w http.ResponseWriter, r *http.Request) {
state := session.NewID()
h.states.put(state, time.Now().Add(oauthStateTTL))
u := h.discord.APIBBase + "/oauth2/authorize?" + url.Values{
u := h.discord.APIBase + "/oauth2/authorize?" + url.Values{
"client_id": {h.discord.ClientID},
"redirect_uri": {h.discord.RedirectURI},
"response_type": {"code"},
@@ -184,7 +185,7 @@ func (h *Handler) discordCallback(w http.ResponseWriter, r *http.Request) {
// The refusal is the same for a non-member and a member without the
// required role, and it names neither the guild nor its id: an outsider
// cannot tell whether the guild exists, let alone which one gates.
if !isMember || (h.discord.RequiredRole != "" && !contains(member.Roles, h.discord.RequiredRole)) {
if !isMember || (h.discord.RequiredRole != "" && !slices.Contains(member.Roles, h.discord.RequiredRole)) {
h.limiter.Fail(ip, time.Now())
h.renderLogin(w, http.StatusForbidden,
"This Discord account is not a member of this community.")
@@ -214,7 +215,7 @@ func (h *Handler) exchangeToken(ctx context.Context, code string) (discordToken,
"redirect_uri": {h.discord.RedirectURI},
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
h.discord.APIBBase+"/oauth2/token", strings.NewReader(form.Encode()))
h.discord.APIBase+"/oauth2/token", strings.NewReader(form.Encode()))
if err != nil {
return discordToken{}, err
}
@@ -241,7 +242,7 @@ func (h *Handler) exchangeToken(ctx context.Context, code string) (discordToken,
// discordUserID fetches the signed-in user's id via the identify scope.
func (h *Handler) discordUserID(ctx context.Context, accessToken string) (string, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
h.discord.APIBBase+"/users/@me", nil)
h.discord.APIBase+"/users/@me", nil)
if err != nil {
return "", err
}
@@ -276,7 +277,7 @@ type discordMember struct {
// gate asks exactly the question it names. A 404 or 403 (not in the guild, or
// the token lacks the scope) is a non-member, not an error.
func (h *Handler) discordMember(ctx context.Context, accessToken, userID string) (discordMember, bool, error) {
u := h.discord.APIBBase + "/guilds/" + url.PathEscape(h.discord.GuildID) +
u := h.discord.APIBase + "/guilds/" + url.PathEscape(h.discord.GuildID) +
"/members/" + url.PathEscape(userID)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
if err != nil {
@@ -305,12 +306,3 @@ func (h *Handler) discordMember(ctx context.Context, accessToken, userID string)
type discordToken struct {
AccessToken string `json:"access_token"`
}
func contains(ss []string, want string) bool {
for _, s := range ss {
if s == want {
return true
}
}
return false
}