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:
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user