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
+3 -3
View File
@@ -557,7 +557,7 @@ func TestLoadConfigDiscord(t *testing.T) {
if got := loadConfig().Discord; got.ClientID != "client-1" || got.ClientSecret != "client-secret-1" ||
got.GuildID != "guild-1" || got.RequiredRole != "role-9" ||
got.RedirectURI != "https://bm.example.com/auth/discord/callback" ||
got.APIBBase != "https://stub.example/api" {
got.APIBase != "https://stub.example/api" {
t.Fatalf("Discord config = %+v, want every field set", got)
}
@@ -568,8 +568,8 @@ func TestLoadConfigDiscord(t *testing.T) {
if got.RequiredRole != "" {
t.Fatalf("RequiredRole = %q, want empty by default", got.RequiredRole)
}
if got.APIBBase != "https://discord.com/api/v10" {
t.Fatalf("APIBBase = %q, want the Discord default", got.APIBBase)
if got.APIBase != "https://discord.com/api/v10" {
t.Fatalf("APIBase = %q, want the Discord default", got.APIBase)
}
}
+2 -2
View File
@@ -30,8 +30,8 @@ type Fetcher interface {
// cannot shorten anyone's cooldown; it only makes the poller wake up and find
// nothing due more often.
type Poller struct {
Store *store.Store
Fetch Fetcher
Store *store.Store
Fetch Fetcher
// BrowserFetch handles sites behind a JavaScript challenge that Fetch
// cannot clear. Nil disables those sites entirely rather than falling back
// to Fetch, which would only ever retrieve a challenge page.
+1 -1
View File
@@ -95,7 +95,7 @@ func ClientIP(r *http.Request) string {
// Behind carrier-grade NAT this budget is shared with every other subscriber on
// the same public address, so a stranger can lock the owner out for up to one
// window. That is accepted: the block self-heals, and ten attempts is generous
// for a mistyped password.
// for the occasional fumbled sign-in.
//
// State is in memory and per-process, so a restart clears it. Entries are
// pruned lazily on access; for a single-user deployment the map cannot grow
+2 -2
View File
@@ -39,7 +39,7 @@ func TestSetSessionCookieAttributes(t *testing.T) {
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
r := httptest.NewRequest(http.MethodPost, "/login", nil)
r := httptest.NewRequest(http.MethodPost, "/", nil)
if tc.tls {
r.TLS = &tls.ConnectionState{}
}
@@ -119,7 +119,7 @@ func TestClientIP(t *testing.T) {
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
r := httptest.NewRequest(http.MethodPost, "/login", nil)
r := httptest.NewRequest(http.MethodPost, "/", nil)
r.RemoteAddr = tc.remoteAddr
for _, v := range tc.xff {
r.Header.Add("X-Forwarded-For", v)
@@ -1,7 +1,8 @@
-- One row per browser session. The id is an opaque random value the cookie
-- carries verbatim; a request is authenticated by looking the row up, and
-- deleting the row is how a session is revoked. Expired rows are removed
-- lazily on lookup, so nothing sweeps them.
-- lazily on lookup and swept by the next login, so nothing runs a background
-- cleanup.
CREATE TABLE sessions (
id text PRIMARY KEY,
reader_id bigint NOT NULL REFERENCES readers (id) ON DELETE CASCADE,
+17 -7
View File
@@ -15,17 +15,24 @@ type Session struct {
// CreateSession stores a new session row for reader. The id is generated by
// the caller (session.NewID) — the store only persists it. Expired rows that
// were never looked up are swept here: this is the one write every login
// makes, so the table stays bounded without a background job.
// were never looked up are swept in the same transaction: this is the one
// write every login makes, so the table stays bounded without a background
// job.
func (s *Store) CreateSession(id string, readerID int64, ttl time.Duration) (Session, error) {
expires := time.Now().Add(ttl)
_, err := s.db.Exec(`INSERT INTO sessions (id, reader_id, expires_at) VALUES ($1, $2, $3)`,
id, readerID, expires)
tx, err := s.db.Begin()
if err != nil {
return Session{}, err
}
_, err = s.db.Exec(`DELETE FROM sessions WHERE expires_at < now()`)
if err != nil {
defer tx.Rollback()
expires := time.Now().Add(ttl)
if _, err := tx.Exec(`INSERT INTO sessions (id, reader_id, expires_at) VALUES ($1, $2, $3)`,
id, readerID, expires); err != nil {
return Session{}, err
}
if _, err := tx.Exec(`DELETE FROM sessions WHERE expires_at < now()`); err != nil {
return Session{}, err
}
if err := tx.Commit(); err != nil {
return Session{}, err
}
return Session{ID: id, ReaderID: readerID, ExpiresAt: expires}, nil
@@ -46,6 +53,9 @@ func (s *Store) GetSession(id string, now time.Time) (Session, bool, error) {
return Session{}, false, err
}
if !sess.ExpiresAt.After(now) {
// Best-effort: the row is dead either way; failing the request over a
// cleanup delete would only hide the real error. CreateSession's
// sweep catches anything this misses.
_, _ = s.db.Exec(`DELETE FROM sessions WHERE id = $1`, id)
return Session{}, false, nil
}
+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
}
+3 -4
View File
@@ -27,8 +27,8 @@ func TestOAuthStateUnknownOrExpired(t *testing.T) {
}
}
// The map is capped: a flood of starts evicts the oldest states, and consumed
// states (which leave the FIFO behind) must not defeat the cap.
// The map is capped: a flood of starts evicts older states instead of growing,
// and consumed states must not change that.
func TestOAuthStateEviction(t *testing.T) {
s := newOAuthStates()
key := func(i, salt int) string {
@@ -42,8 +42,7 @@ func TestOAuthStateEviction(t *testing.T) {
t.Fatalf("states after a flood = %d, want %d", got, maxStates)
}
// Consume everything, then flood again: the map stays bounded and the
// eviction loop pops the stale FIFO entries instead of stalling.
// Consume everything, then flood again: the map stays bounded.
for state := range s.expiry {
s.take(state)
}
+1 -1
View File
@@ -790,7 +790,7 @@ button { cursor: pointer; }
.login-card button:hover {
background: var(--ember);
border-color: var(--ember);
color: #fff;
color: var(--ember-ink);
}
.login-card .login-note {
margin: 14px 0 0;
+1 -1
View File
@@ -41,7 +41,7 @@ type Handler struct {
states *oauthStates
limiter *session.LoginLimiter
// httpClient is the plain stdlib client that talks to Discord. It is not
// an injected interface: tests point APIBBase at a stub server instead.
// an injected interface: tests point APIBase at a stub server instead.
httpClient *http.Client
}
+1 -1
View File
@@ -162,7 +162,7 @@ func loadConfig() Config {
ClientSecret: os.Getenv("DISCORD_CLIENT_SECRET"),
GuildID: os.Getenv("DISCORD_GUILD_ID"),
RequiredRole: os.Getenv("DISCORD_REQUIRED_ROLE"),
APIBBase: envOr("DISCORD_API_BASE", "https://discord.com/api/v10"),
APIBase: envOr("DISCORD_API_BASE", "https://discord.com/api/v10"),
RedirectURI: os.Getenv("DISCORD_REDIRECT_URI"),
OwnerDiscordID: c.OwnerDiscordID,
}
+1 -1
View File
@@ -134,7 +134,7 @@ func discordConfig(stubURL string) web.DiscordConfig {
ClientID: "client-1",
ClientSecret: "client-secret-1",
GuildID: "guild-1",
APIBBase: stubURL,
APIBase: stubURL,
RedirectURI: "https://bm.example.com/auth/discord/callback",
OwnerDiscordID: testOwnerID,
}