diff --git a/backend/api_test.go b/backend/api_test.go index bf2d597..187e9f1 100644 --- a/backend/api_test.go +++ b/backend/api_test.go @@ -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) } } diff --git a/backend/internal/latest/poller.go b/backend/internal/latest/poller.go index 2498c49..ca77339 100644 --- a/backend/internal/latest/poller.go +++ b/backend/internal/latest/poller.go @@ -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. diff --git a/backend/internal/session/session.go b/backend/internal/session/session.go index ee61633..86a8400 100644 --- a/backend/internal/session/session.go +++ b/backend/internal/session/session.go @@ -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 diff --git a/backend/internal/session/session_test.go b/backend/internal/session/session_test.go index 51c0106..7934a23 100644 --- a/backend/internal/session/session_test.go +++ b/backend/internal/session/session_test.go @@ -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) diff --git a/backend/internal/store/migrations/0005_sessions.sql b/backend/internal/store/migrations/0005_sessions.sql index c29b5f7..03b25ed 100644 --- a/backend/internal/store/migrations/0005_sessions.sql +++ b/backend/internal/store/migrations/0005_sessions.sql @@ -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, diff --git a/backend/internal/store/sessions.go b/backend/internal/store/sessions.go index 43ab483..f4578b7 100644 --- a/backend/internal/store/sessions.go +++ b/backend/internal/store/sessions.go @@ -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 } diff --git a/backend/internal/web/discord.go b/backend/internal/web/discord.go index e04b1bf..930c72c 100644 --- a/backend/internal/web/discord.go +++ b/backend/internal/web/discord.go @@ -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 -} diff --git a/backend/internal/web/oauth_test.go b/backend/internal/web/oauth_test.go index 4b4d4f8..5cdbf74 100644 --- a/backend/internal/web/oauth_test.go +++ b/backend/internal/web/oauth_test.go @@ -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) } diff --git a/backend/internal/web/static/style.css b/backend/internal/web/static/style.css index 4b1d911..846e7df 100644 --- a/backend/internal/web/static/style.css +++ b/backend/internal/web/static/style.css @@ -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; diff --git a/backend/internal/web/web.go b/backend/internal/web/web.go index 13697d2..eb477a9 100644 --- a/backend/internal/web/web.go +++ b/backend/internal/web/web.go @@ -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 } diff --git a/backend/main.go b/backend/main.go index 9d36599..0d65dbf 100644 --- a/backend/main.go +++ b/backend/main.go @@ -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, } diff --git a/backend/web_test.go b/backend/web_test.go index caa29ea..8240c48 100644 --- a/backend/web_test.go +++ b/backend/web_test.go @@ -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, }