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" ||
|
if got := loadConfig().Discord; got.ClientID != "client-1" || got.ClientSecret != "client-secret-1" ||
|
||||||
got.GuildID != "guild-1" || got.RequiredRole != "role-9" ||
|
got.GuildID != "guild-1" || got.RequiredRole != "role-9" ||
|
||||||
got.RedirectURI != "https://bm.example.com/auth/discord/callback" ||
|
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)
|
t.Fatalf("Discord config = %+v, want every field set", got)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -568,8 +568,8 @@ func TestLoadConfigDiscord(t *testing.T) {
|
|||||||
if got.RequiredRole != "" {
|
if got.RequiredRole != "" {
|
||||||
t.Fatalf("RequiredRole = %q, want empty by default", got.RequiredRole)
|
t.Fatalf("RequiredRole = %q, want empty by default", got.RequiredRole)
|
||||||
}
|
}
|
||||||
if got.APIBBase != "https://discord.com/api/v10" {
|
if got.APIBase != "https://discord.com/api/v10" {
|
||||||
t.Fatalf("APIBBase = %q, want the Discord default", got.APIBBase)
|
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
|
// cannot shorten anyone's cooldown; it only makes the poller wake up and find
|
||||||
// nothing due more often.
|
// nothing due more often.
|
||||||
type Poller struct {
|
type Poller struct {
|
||||||
Store *store.Store
|
Store *store.Store
|
||||||
Fetch Fetcher
|
Fetch Fetcher
|
||||||
// BrowserFetch handles sites behind a JavaScript challenge that Fetch
|
// BrowserFetch handles sites behind a JavaScript challenge that Fetch
|
||||||
// cannot clear. Nil disables those sites entirely rather than falling back
|
// cannot clear. Nil disables those sites entirely rather than falling back
|
||||||
// to Fetch, which would only ever retrieve a challenge page.
|
// 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
|
// 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
|
// 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
|
// 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
|
// 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
|
// 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 {
|
for _, tc := range cases {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
r := httptest.NewRequest(http.MethodPost, "/login", nil)
|
r := httptest.NewRequest(http.MethodPost, "/", nil)
|
||||||
if tc.tls {
|
if tc.tls {
|
||||||
r.TLS = &tls.ConnectionState{}
|
r.TLS = &tls.ConnectionState{}
|
||||||
}
|
}
|
||||||
@@ -119,7 +119,7 @@ func TestClientIP(t *testing.T) {
|
|||||||
}
|
}
|
||||||
for _, tc := range cases {
|
for _, tc := range cases {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
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
|
r.RemoteAddr = tc.remoteAddr
|
||||||
for _, v := range tc.xff {
|
for _, v := range tc.xff {
|
||||||
r.Header.Add("X-Forwarded-For", v)
|
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
|
-- 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
|
-- 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
|
-- 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 (
|
CREATE TABLE sessions (
|
||||||
id text PRIMARY KEY,
|
id text PRIMARY KEY,
|
||||||
reader_id bigint NOT NULL REFERENCES readers (id) ON DELETE CASCADE,
|
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
|
// 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
|
// 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
|
// were never looked up are swept in the same transaction: this is the one
|
||||||
// makes, so the table stays bounded without a background job.
|
// 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) {
|
func (s *Store) CreateSession(id string, readerID int64, ttl time.Duration) (Session, error) {
|
||||||
expires := time.Now().Add(ttl)
|
tx, err := s.db.Begin()
|
||||||
_, err := s.db.Exec(`INSERT INTO sessions (id, reader_id, expires_at) VALUES ($1, $2, $3)`,
|
|
||||||
id, readerID, expires)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return Session{}, err
|
return Session{}, err
|
||||||
}
|
}
|
||||||
_, err = s.db.Exec(`DELETE FROM sessions WHERE expires_at < now()`)
|
defer tx.Rollback()
|
||||||
if err != nil {
|
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{}, err
|
||||||
}
|
}
|
||||||
return Session{ID: id, ReaderID: readerID, ExpiresAt: expires}, nil
|
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
|
return Session{}, false, err
|
||||||
}
|
}
|
||||||
if !sess.ExpiresAt.After(now) {
|
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)
|
_, _ = s.db.Exec(`DELETE FROM sessions WHERE id = $1`, id)
|
||||||
return Session{}, false, nil
|
return Session{}, false, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"log"
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -43,9 +44,9 @@ type DiscordConfig struct {
|
|||||||
// RequiredRole, when non-empty, is a role ID a member must hold on top of
|
// RequiredRole, when non-empty, is a role ID a member must hold on top of
|
||||||
// guild membership. Empty by default: membership alone suffices.
|
// guild membership. Empty by default: membership alone suffices.
|
||||||
RequiredRole string
|
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.
|
// flow against a local stub.
|
||||||
APIBBase string
|
APIBase string
|
||||||
// RedirectURI is the full public URL of the callback — Discord requires
|
// RedirectURI is the full public URL of the callback — Discord requires
|
||||||
// the exact string, so it is configured, never derived from headers.
|
// the exact string, so it is configured, never derived from headers.
|
||||||
RedirectURI string
|
RedirectURI string
|
||||||
@@ -59,7 +60,6 @@ type DiscordConfig struct {
|
|||||||
type oauthStates struct {
|
type oauthStates struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
expiry map[string]time.Time
|
expiry map[string]time.Time
|
||||||
order []string // FIFO for eviction when the map is full
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func newOAuthStates() *oauthStates {
|
func newOAuthStates() *oauthStates {
|
||||||
@@ -75,18 +75,19 @@ func (s *oauthStates) put(state string, expires time.Time) {
|
|||||||
delete(s.expiry, k)
|
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 {
|
if len(s.expiry) >= maxStates {
|
||||||
// Evict from the front until under the cap. The front may already
|
var oldest string
|
||||||
// have been consumed by take (which removes from the map, not the
|
var oldestAt time.Time
|
||||||
// order slice), so pop until the count actually drops.
|
for k, at := range s.expiry {
|
||||||
for len(s.expiry) >= maxStates && len(s.order) > 0 {
|
if oldest == "" || at.Before(oldestAt) {
|
||||||
oldest := s.order[0]
|
oldest, oldestAt = k, at
|
||||||
s.order = s.order[1:]
|
}
|
||||||
delete(s.expiry, oldest)
|
|
||||||
}
|
}
|
||||||
|
delete(s.expiry, oldest)
|
||||||
}
|
}
|
||||||
s.expiry[state] = expires
|
s.expiry[state] = expires
|
||||||
s.order = append(s.order, state)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// take validates and consumes a state in one step: a state works exactly
|
// 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) {
|
func (h *Handler) discordStart(w http.ResponseWriter, r *http.Request) {
|
||||||
state := session.NewID()
|
state := session.NewID()
|
||||||
h.states.put(state, time.Now().Add(oauthStateTTL))
|
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},
|
"client_id": {h.discord.ClientID},
|
||||||
"redirect_uri": {h.discord.RedirectURI},
|
"redirect_uri": {h.discord.RedirectURI},
|
||||||
"response_type": {"code"},
|
"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
|
// 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
|
// required role, and it names neither the guild nor its id: an outsider
|
||||||
// cannot tell whether the guild exists, let alone which one gates.
|
// 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.limiter.Fail(ip, time.Now())
|
||||||
h.renderLogin(w, http.StatusForbidden,
|
h.renderLogin(w, http.StatusForbidden,
|
||||||
"This Discord account is not a member of this community.")
|
"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},
|
"redirect_uri": {h.discord.RedirectURI},
|
||||||
}
|
}
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
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 {
|
if err != nil {
|
||||||
return discordToken{}, err
|
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.
|
// discordUserID fetches the signed-in user's id via the identify scope.
|
||||||
func (h *Handler) discordUserID(ctx context.Context, accessToken string) (string, error) {
|
func (h *Handler) discordUserID(ctx context.Context, accessToken string) (string, error) {
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
|
||||||
h.discord.APIBBase+"/users/@me", nil)
|
h.discord.APIBase+"/users/@me", nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
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
|
// 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.
|
// 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) {
|
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)
|
"/members/" + url.PathEscape(userID)
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -305,12 +306,3 @@ func (h *Handler) discordMember(ctx context.Context, accessToken, userID string)
|
|||||||
type discordToken struct {
|
type discordToken struct {
|
||||||
AccessToken string `json:"access_token"`
|
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
|
// The map is capped: a flood of starts evicts older states instead of growing,
|
||||||
// states (which leave the FIFO behind) must not defeat the cap.
|
// and consumed states must not change that.
|
||||||
func TestOAuthStateEviction(t *testing.T) {
|
func TestOAuthStateEviction(t *testing.T) {
|
||||||
s := newOAuthStates()
|
s := newOAuthStates()
|
||||||
key := func(i, salt int) string {
|
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)
|
t.Fatalf("states after a flood = %d, want %d", got, maxStates)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Consume everything, then flood again: the map stays bounded and the
|
// Consume everything, then flood again: the map stays bounded.
|
||||||
// eviction loop pops the stale FIFO entries instead of stalling.
|
|
||||||
for state := range s.expiry {
|
for state := range s.expiry {
|
||||||
s.take(state)
|
s.take(state)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -790,7 +790,7 @@ button { cursor: pointer; }
|
|||||||
.login-card button:hover {
|
.login-card button:hover {
|
||||||
background: var(--ember);
|
background: var(--ember);
|
||||||
border-color: var(--ember);
|
border-color: var(--ember);
|
||||||
color: #fff;
|
color: var(--ember-ink);
|
||||||
}
|
}
|
||||||
.login-card .login-note {
|
.login-card .login-note {
|
||||||
margin: 14px 0 0;
|
margin: 14px 0 0;
|
||||||
|
|||||||
@@ -41,7 +41,7 @@ type Handler struct {
|
|||||||
states *oauthStates
|
states *oauthStates
|
||||||
limiter *session.LoginLimiter
|
limiter *session.LoginLimiter
|
||||||
// httpClient is the plain stdlib client that talks to Discord. It is not
|
// 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
|
httpClient *http.Client
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -162,7 +162,7 @@ func loadConfig() Config {
|
|||||||
ClientSecret: os.Getenv("DISCORD_CLIENT_SECRET"),
|
ClientSecret: os.Getenv("DISCORD_CLIENT_SECRET"),
|
||||||
GuildID: os.Getenv("DISCORD_GUILD_ID"),
|
GuildID: os.Getenv("DISCORD_GUILD_ID"),
|
||||||
RequiredRole: os.Getenv("DISCORD_REQUIRED_ROLE"),
|
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"),
|
RedirectURI: os.Getenv("DISCORD_REDIRECT_URI"),
|
||||||
OwnerDiscordID: c.OwnerDiscordID,
|
OwnerDiscordID: c.OwnerDiscordID,
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-1
@@ -134,7 +134,7 @@ func discordConfig(stubURL string) web.DiscordConfig {
|
|||||||
ClientID: "client-1",
|
ClientID: "client-1",
|
||||||
ClientSecret: "client-secret-1",
|
ClientSecret: "client-secret-1",
|
||||||
GuildID: "guild-1",
|
GuildID: "guild-1",
|
||||||
APIBBase: stubURL,
|
APIBase: stubURL,
|
||||||
RedirectURI: "https://bm.example.com/auth/discord/callback",
|
RedirectURI: "https://bm.example.com/auth/discord/callback",
|
||||||
OwnerDiscordID: testOwnerID,
|
OwnerDiscordID: testOwnerID,
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user