feat(backend): Discord OAuth login with DB-backed sessions (#23) #31

Merged
sulthan merged 2 commits from feat/discord-login into main 2026-08-08 08:51:23 +07:00
12 changed files with 51 additions and 49 deletions
Showing only changes of commit 1a7e130f9f - Show all commits
+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" || 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)
} }
} }
+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 // 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.
+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 // 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
+2 -2
View File
@@ -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,
+17 -7
View File
@@ -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
} }
+17 -25
View File
@@ -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
}
+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 // 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)
} }
+1 -1
View File
@@ -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;
+1 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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,
} }