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:
+3
-3
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user