Implements #23 per ADR-0002. - Discord authorization code grant (identify + guilds.members.read), form-encoded token exchange - Guild membership gate via the single-guild endpoint; optional DISCORD_REQUIRED_ROLE (empty default) - Owner Discord ID is the only identity allowed to sign in - Sessions are DB rows with opaque random ids; cookie carries only the id; expiry enforced; delete = revoke - HMAC session signing, derived key, and WEB_PASSWORD removed; no replacement signing secret - Login rate limiting preserved on the callback - Full flow tested through the real router against a local Discord stub (DISCORD_API_BASE) - Env: DISCORD_CLIENT_ID/_CLIENT_SECRET/_GUILD_ID/_REQUIRED_ROLE/_API_BASE/_REDIRECT_URI; docs updated go test ./... passes. Reviewed-on: #31 Co-authored-by: Sulthan Zaki <sultankiki05@gmail.com> Co-committed-by: Sulthan Zaki <sultankiki05@gmail.com>
This commit was merged in pull request #31.
This commit is contained in:
+12
-7
@@ -38,14 +38,17 @@ Guidance for OpenCode (and Claude Code) working under `backend/`. See root `AGEN
|
||||
and enforces the ownership rule: client `title`/`series_url`/`cover` are
|
||||
written only when the series row is new (ADR-0003).
|
||||
- **Endpoints:** `GET /bookmarks`, `PUT /bookmarks/{key}` (upsert; see `updated_at` rule below), `DELETE /bookmarks/{key}`, `GET /healthz` (no auth).
|
||||
- **Web UI:** same binary serve password-gated browser UI on second
|
||||
- **Web UI:** same binary serve the browser UI on a second
|
||||
hostname — `GET /` (list, or login page when no session),
|
||||
`POST /login`, `POST /logout`, `GET /static/*`, htmx fragment endpoints
|
||||
`GET /auth/discord` + `GET /auth/discord/callback` (Discord OAuth,
|
||||
ADR-0002), `POST /logout`, `GET /static/*`, htmx fragment endpoints
|
||||
under `/ui/*`. Templates + assets `go:embed`-ed under
|
||||
`backend/internal/web/`, so `backend/Dockerfile` must copy the whole
|
||||
`internal/` tree, not just `*.go`. Sessions stateless
|
||||
HMAC cookies keyed off `API_TOKEN`; `WEB_PASSWORD` gates them, and when empty,
|
||||
web routes not registered at all. UI mutations read-modify-write
|
||||
`internal/` tree, not just `*.go`. Sessions are rows in the `sessions`
|
||||
table: the cookie carries only an opaque id, looked up (and expiry-
|
||||
checked) on every request, and deleting the row revokes the session.
|
||||
The owner's Discord ID is the only identity that can sign in while
|
||||
registration is closed. UI mutations read-modify-write
|
||||
through `Store.Get` + `Store.Upsert` so `updated_at` rule stays one
|
||||
place. See `docs/superpowers/specs/2026-07-25-web-ui-design.md`.
|
||||
**Design-tool caveat:** templates link `/static/style.css` root-absolutely
|
||||
@@ -101,8 +104,10 @@ Guidance for OpenCode (and Claude Code) working under `backend/`. See root `AGEN
|
||||
- **Config via env:** `API_TOKEN`, `OWNER_DISCORD_ID` (seeds the owner Reader;
|
||||
required), `ALLOWED_ORIGINS` (comma list),
|
||||
`DATABASE_URL` (Postgres connection URL, required — no default),
|
||||
`PORT` (default `8080`), `WEB_PASSWORD`
|
||||
(gates browser UI; unset disable it),
|
||||
`PORT` (default `8080`), `DISCORD_CLIENT_ID`/`_CLIENT_SECRET`/`_GUILD_ID`/
|
||||
`_REDIRECT_URI` (required; Discord OAuth for the browser UI),
|
||||
`DISCORD_REQUIRED_ROLE` (optional role gate, empty by default),
|
||||
`DISCORD_API_BASE` (default `https://discord.com/api/v10`),
|
||||
`LATEST_CHAPTER_POLL_ENABLED`/`_COOLDOWN`/`_INTERVAL`/`_BATCH`/`_STAGGER`
|
||||
(background latest-chapter poller; defaults on, `1h`/`10m`/`14`/`20s`).
|
||||
`USERSCRIPT_PATH` and `NOVEL_USERSCRIPT_PATH` (files served at
|
||||
|
||||
+37
-15
@@ -36,14 +36,23 @@ func newTestServer(t *testing.T) http.Handler {
|
||||
|
||||
func newTestStore(t *testing.T) *store.Store {
|
||||
t.Helper()
|
||||
s, err := store.Open(pgtest.URL(t), store.Owner{
|
||||
s, _ := newTestStoreURL(t)
|
||||
return s
|
||||
}
|
||||
|
||||
// newTestStoreURL is newTestStore plus the database URL, for tests that need
|
||||
// to reach the same database directly.
|
||||
func newTestStoreURL(t *testing.T) (*store.Store, string) {
|
||||
t.Helper()
|
||||
url := pgtest.URL(t)
|
||||
s, err := store.Open(url, store.Owner{
|
||||
DiscordID: "test-owner", TokenHash: sha256.Sum256([]byte("owner-token-hash")),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("store.Open: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { s.Close() })
|
||||
return s
|
||||
return s, url
|
||||
}
|
||||
|
||||
func auth(req *http.Request) *http.Request {
|
||||
@@ -538,16 +547,29 @@ func TestLatestChapterNullable(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadConfigWebPassword(t *testing.T) {
|
||||
t.Setenv("API_TOKEN", "token-abc")
|
||||
t.Setenv("WEB_PASSWORD", "hunter2")
|
||||
if got := loadConfig().WebPassword; got != "hunter2" {
|
||||
t.Fatalf("WebPassword = %q, want hunter2", got)
|
||||
func TestLoadConfigDiscord(t *testing.T) {
|
||||
t.Setenv("DISCORD_CLIENT_ID", "client-1")
|
||||
t.Setenv("DISCORD_CLIENT_SECRET", "client-secret-1")
|
||||
t.Setenv("DISCORD_GUILD_ID", "guild-1")
|
||||
t.Setenv("DISCORD_REQUIRED_ROLE", "role-9")
|
||||
t.Setenv("DISCORD_REDIRECT_URI", "https://bm.example.com/auth/discord/callback")
|
||||
t.Setenv("DISCORD_API_BASE", "https://stub.example/api")
|
||||
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.APIBase != "https://stub.example/api" {
|
||||
t.Fatalf("Discord config = %+v, want every field set", got)
|
||||
}
|
||||
|
||||
t.Setenv("WEB_PASSWORD", "")
|
||||
if got := loadConfig().WebPassword; got != "" {
|
||||
t.Fatalf("WebPassword = %q with the variable unset, want empty", got)
|
||||
// API base falls back to the Discord default; the role is optional.
|
||||
t.Setenv("DISCORD_REQUIRED_ROLE", "")
|
||||
t.Setenv("DISCORD_API_BASE", "")
|
||||
got := loadConfig().Discord
|
||||
if got.RequiredRole != "" {
|
||||
t.Fatalf("RequiredRole = %q, want empty by default", got.RequiredRole)
|
||||
}
|
||||
if got.APIBase != "https://discord.com/api/v10" {
|
||||
t.Fatalf("APIBase = %q, want the Discord default", got.APIBase)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -578,9 +600,9 @@ func TestPutDoesNotClobberLatestCheckedAt(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// The userscript route is registered outside the `if cfg.WebPassword != ""`
|
||||
// block in newRouter, so it must keep working on a deployment that never set
|
||||
// WEB_PASSWORD — see internal/userscript for the handler's own behaviour.
|
||||
// The userscript route is registered outside the web UI's Discord auth, so it
|
||||
// must keep working whatever the web config — see internal/userscript for the
|
||||
// handler's own behaviour.
|
||||
func TestUserscriptServedWithWebUIDisabled(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "manga-bookmark.user.js")
|
||||
if err := os.WriteFile(path, []byte("console.log(1);\n"), 0o644); err != nil {
|
||||
@@ -588,7 +610,7 @@ func TestUserscriptServedWithWebUIDisabled(t *testing.T) {
|
||||
}
|
||||
|
||||
s := newTestStore(t)
|
||||
cfg := testConfig() // WebPassword empty
|
||||
cfg := testConfig() // no Discord config needed for the userscript route
|
||||
cfg.UserscriptPath = path
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
@@ -600,7 +622,7 @@ func TestUserscriptServedWithWebUIDisabled(t *testing.T) {
|
||||
}
|
||||
|
||||
// Both scripts are served from the same handler on the same token, outside the
|
||||
// WEB_PASSWORD gate — a wrong token is a 404, never a 401.
|
||||
// web UI's auth — a wrong token is a 404, never a 401.
|
||||
func TestNovelUserscriptServed(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
novelPath := filepath.Join(dir, "novel-bookmark.user.js")
|
||||
|
||||
@@ -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,13 +1,10 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"net"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -16,49 +13,18 @@ import (
|
||||
const (
|
||||
CookieName = "bmgr_session"
|
||||
// 60 days: long enough that a phone stays logged in between reading spells.
|
||||
sessionTTL = 60 * 24 * time.Hour
|
||||
// Domain separation, so the session key can never collide with any other
|
||||
// use of the secrets it is derived from. Changing this string logs
|
||||
// everyone out.
|
||||
sessionKeyPurpose = "bmgr-web-session-v1"
|
||||
SessionTTL = 60 * 24 * time.Hour
|
||||
)
|
||||
|
||||
// Key derives the cookie-signing key from both secrets. Sessions are
|
||||
// stateless — there is no session table — so rotating either API_TOKEN or
|
||||
// WEB_PASSWORD invalidates every outstanding cookie at once. The \x00
|
||||
// separator prevents the concatenation ambiguity a bare apiToken+webPassword
|
||||
// would have (e.g. "ab"+"c" colliding with "a"+"bc").
|
||||
func Key(apiToken, webPassword string) []byte {
|
||||
sum := sha256.Sum256([]byte(apiToken + "\x00" + webPassword + sessionKeyPurpose))
|
||||
return sum[:]
|
||||
}
|
||||
|
||||
// Sign encodes "<expiryMs>.<base64url HMAC(expiryMs)>".
|
||||
func Sign(key []byte, expiryMs int64) string {
|
||||
payload := strconv.FormatInt(expiryMs, 10)
|
||||
return payload + "." + sessionMAC(key, payload)
|
||||
}
|
||||
|
||||
func sessionMAC(key []byte, payload string) string {
|
||||
mac := hmac.New(sha256.New, key)
|
||||
mac.Write([]byte(payload))
|
||||
return base64.RawURLEncoding.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
// Verify checks shape, then expiry, then the signature — in that order.
|
||||
// The signature comparison is constant-time; the checks before it only look at
|
||||
// data the holder already supplied, so their timing leaks nothing.
|
||||
func Verify(key []byte, value string, nowMs int64) bool {
|
||||
payload, sig, ok := strings.Cut(value, ".")
|
||||
if !ok {
|
||||
return false
|
||||
// NewID returns an opaque session id: 32 random bytes, hex-encoded. The id is
|
||||
// all the cookie carries and all the sessions table keys on, so its entropy is
|
||||
// what stops a guessed id from being someone else's session.
|
||||
func NewID() string {
|
||||
var b [32]byte
|
||||
if _, err := rand.Read(b[:]); err != nil {
|
||||
panic("session id: " + err.Error())
|
||||
}
|
||||
expiry, err := strconv.ParseInt(payload, 10, 64)
|
||||
if err != nil || expiry <= nowMs {
|
||||
return false
|
||||
}
|
||||
want := sessionMAC(key, payload)
|
||||
return subtle.ConstantTimeCompare([]byte(sig), []byte(want)) == 1
|
||||
return hex.EncodeToString(b[:])
|
||||
}
|
||||
|
||||
// isHTTPS reports whether the browser's connection is encrypted. Behind Traefik
|
||||
@@ -69,12 +35,14 @@ func isHTTPS(r *http.Request) bool {
|
||||
return r.TLS != nil || r.Header.Get("X-Forwarded-Proto") == "https"
|
||||
}
|
||||
|
||||
func SetCookie(w http.ResponseWriter, r *http.Request, key []byte) {
|
||||
// SetCookie writes the session cookie. The value is the session id and nothing
|
||||
// else; the row behind it is looked up on every request.
|
||||
func SetCookie(w http.ResponseWriter, r *http.Request, id string) {
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: CookieName,
|
||||
Value: Sign(key, time.Now().Add(sessionTTL).UnixMilli()),
|
||||
Value: id,
|
||||
Path: "/",
|
||||
MaxAge: int(sessionTTL / time.Second),
|
||||
MaxAge: int(SessionTTL / time.Second),
|
||||
HttpOnly: true,
|
||||
Secure: isHTTPS(r),
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
@@ -120,14 +88,14 @@ func ClientIP(r *http.Request) string {
|
||||
return host
|
||||
}
|
||||
|
||||
// LoginLimiter throttles password guessing: MaxFailures failures inside a
|
||||
// rolling Window blocks further attempts from that IP until the oldest one
|
||||
// LoginLimiter throttles failed sign-in attempts: MaxFailures failures inside
|
||||
// a rolling Window blocks further attempts from that IP until the oldest one
|
||||
// ages out. There is no permanent ban and no unlock step.
|
||||
//
|
||||
// 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
|
||||
|
||||
@@ -9,66 +9,19 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestSessionRoundTrip(t *testing.T) {
|
||||
key := Key("token-abc", "pw-abc")
|
||||
now := time.Now().UnixMilli()
|
||||
value := Sign(key, now+60_000)
|
||||
if !Verify(key, value, now) {
|
||||
t.Fatal("Verify = false for a freshly signed cookie, want true")
|
||||
func TestNewID(t *testing.T) {
|
||||
a := NewID()
|
||||
b := NewID()
|
||||
if a == b {
|
||||
t.Fatal("NewID returned the same value twice")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionRejects(t *testing.T) {
|
||||
key := Key("token-abc", "pw-abc")
|
||||
now := time.Now().UnixMilli()
|
||||
valid := Sign(key, now+60_000)
|
||||
payload, sig, _ := strings.Cut(valid, ".")
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
value string
|
||||
}{
|
||||
{"empty", ""},
|
||||
{"no separator", payload + sig},
|
||||
{"unparseable expiry", "notanumber." + sig},
|
||||
{"expired", Sign(key, now-1)},
|
||||
{"tampered signature", payload + "." + flipLastChar(sig)},
|
||||
{"tampered expiry", "99999999999999." + sig},
|
||||
{"signed with another key", Sign(Key("other-token", "pw-abc"), now+60_000)},
|
||||
if len(a) != 64 { // 32 random bytes, hex
|
||||
t.Fatalf("NewID() length = %d, want 64", len(a))
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if Verify(key, tc.value, now) {
|
||||
t.Fatalf("Verify(%q) = true, want false", tc.value)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func flipLastChar(s string) string {
|
||||
if s == "" {
|
||||
return "x"
|
||||
}
|
||||
last := s[len(s)-1]
|
||||
if last == 'A' {
|
||||
return s[:len(s)-1] + "B"
|
||||
}
|
||||
return s[:len(s)-1] + "A"
|
||||
}
|
||||
|
||||
func TestSessionKeyDependsOnToken(t *testing.T) {
|
||||
a := Key("token-a", "pw-abc")
|
||||
b := Key("token-b", "pw-abc")
|
||||
if string(a) == string(b) {
|
||||
t.Fatal("Key collided for different API tokens")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionKeyDependsOnWebPassword(t *testing.T) {
|
||||
a := Key("token-abc", "pw-a")
|
||||
b := Key("token-abc", "pw-b")
|
||||
if string(a) == string(b) {
|
||||
t.Fatal("Key collided for different web passwords with the same API token")
|
||||
for _, r := range a {
|
||||
if !strings.ContainsRune("0123456789abcdef", r) {
|
||||
t.Fatalf("NewID() = %q, want hex", a)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -86,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{}
|
||||
}
|
||||
@@ -94,7 +47,7 @@ func TestSetSessionCookieAttributes(t *testing.T) {
|
||||
r.Header.Set("X-Forwarded-Proto", tc.forwarded)
|
||||
}
|
||||
rr := httptest.NewRecorder()
|
||||
SetCookie(rr, r, Key("token-abc", "pw-abc"))
|
||||
SetCookie(rr, r, "abc123")
|
||||
|
||||
cookies := rr.Result().Cookies()
|
||||
if len(cookies) != 1 {
|
||||
@@ -104,6 +57,9 @@ func TestSetSessionCookieAttributes(t *testing.T) {
|
||||
if c.Name != CookieName {
|
||||
t.Fatalf("cookie name = %q, want %q", c.Name, CookieName)
|
||||
}
|
||||
if c.Value != "abc123" {
|
||||
t.Fatalf("cookie value = %q, want the session id verbatim", c.Value)
|
||||
}
|
||||
if !c.HttpOnly {
|
||||
t.Fatal("cookie HttpOnly = false, want true")
|
||||
}
|
||||
@@ -116,8 +72,8 @@ func TestSetSessionCookieAttributes(t *testing.T) {
|
||||
if c.Secure != tc.wantSecure {
|
||||
t.Fatalf("cookie Secure = %v, want %v", c.Secure, tc.wantSecure)
|
||||
}
|
||||
if c.MaxAge != int(sessionTTL/time.Second) {
|
||||
t.Fatalf("cookie MaxAge = %d, want %d", c.MaxAge, int(sessionTTL/time.Second))
|
||||
if c.MaxAge != int(SessionTTL/time.Second) {
|
||||
t.Fatalf("cookie MaxAge = %d, want %d", c.MaxAge, int(SessionTTL/time.Second))
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -163,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)
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
-- 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 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,
|
||||
created_at timestamptz NOT NULL DEFAULT now(),
|
||||
expires_at timestamptz NOT NULL
|
||||
);
|
||||
@@ -0,0 +1,69 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Session is one browser login: an opaque id the cookie carries verbatim,
|
||||
// the Reader it belongs to, and when it stops being valid.
|
||||
type Session struct {
|
||||
ID string
|
||||
ReaderID int64
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
// 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 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) {
|
||||
tx, err := s.db.Begin()
|
||||
if err != nil {
|
||||
return Session{}, err
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
// GetSession returns the live session row for id, or ok=false when the id is
|
||||
// unknown or expired. An expired row is deleted on the way out, so the table
|
||||
// never grows past sessions that are still valid.
|
||||
func (s *Store) GetSession(id string, now time.Time) (Session, bool, error) {
|
||||
var sess Session
|
||||
err := s.db.QueryRow(
|
||||
`SELECT id, reader_id, expires_at FROM sessions WHERE id = $1`, id,
|
||||
).Scan(&sess.ID, &sess.ReaderID, &sess.ExpiresAt)
|
||||
if err == sql.ErrNoRows {
|
||||
return Session{}, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
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
|
||||
}
|
||||
return sess, true, nil
|
||||
}
|
||||
|
||||
// DeleteSession revokes one session. Deleting an unknown id is not an error.
|
||||
func (s *Store) DeleteSession(id string) error {
|
||||
_, err := s.db.Exec(`DELETE FROM sessions WHERE id = $1`, id)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCreateAndGetSession(t *testing.T) {
|
||||
s := newTestStore(t)
|
||||
owner := s.OwnerID()
|
||||
|
||||
sess, err := s.CreateSession("sess-1", owner, time.Hour)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateSession: %v", err)
|
||||
}
|
||||
if sess.ID != "sess-1" || sess.ReaderID != owner {
|
||||
t.Fatalf("CreateSession returned %+v, want id sess-1 reader %d", sess, owner)
|
||||
}
|
||||
|
||||
got, ok, err := s.GetSession("sess-1", time.Now())
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("GetSession: ok=%v err=%v, want ok", ok, err)
|
||||
}
|
||||
if got.ReaderID != owner {
|
||||
t.Fatalf("session reader = %d, want %d", got.ReaderID, owner)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetSessionUnknownID(t *testing.T) {
|
||||
s := newTestStore(t)
|
||||
if _, ok, err := s.GetSession("nope", time.Now()); err != nil || ok {
|
||||
t.Fatalf("GetSession(unknown) = ok=%v err=%v, want ok=false", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExpiredSessionIsGone(t *testing.T) {
|
||||
s := newTestStore(t)
|
||||
owner := s.OwnerID()
|
||||
if _, err := s.CreateSession("sess-exp", owner, -time.Minute); err != nil {
|
||||
t.Fatalf("CreateSession: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
if _, ok, err := s.GetSession("sess-exp", now); err != nil || ok {
|
||||
t.Fatalf("GetSession(expired) = ok=%v err=%v, want ok=false", ok, err)
|
||||
}
|
||||
// The expired row is deleted on lookup, so the next call cannot revive it.
|
||||
if _, ok, err := s.GetSession("sess-exp", now.Add(-time.Hour)); err != nil || ok {
|
||||
t.Fatalf("GetSession(expired again) = ok=%v err=%v, want ok=false", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteSessionRevokes(t *testing.T) {
|
||||
s := newTestStore(t)
|
||||
owner := s.OwnerID()
|
||||
if _, err := s.CreateSession("sess-del", owner, time.Hour); err != nil {
|
||||
t.Fatalf("CreateSession: %v", err)
|
||||
}
|
||||
if err := s.DeleteSession("sess-del"); err != nil {
|
||||
t.Fatalf("DeleteSession: %v", err)
|
||||
}
|
||||
if _, ok, err := s.GetSession("sess-del", time.Now()); err != nil || ok {
|
||||
t.Fatalf("GetSession after delete = ok=%v err=%v, want ok=false", ok, err)
|
||||
}
|
||||
// Deleting twice is not an error.
|
||||
if err := s.DeleteSession("sess-del"); err != nil {
|
||||
t.Fatalf("DeleteSession twice: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteSessionIsPerReader(t *testing.T) {
|
||||
s := newTestStore(t)
|
||||
other := secondReader(t, s)
|
||||
if _, err := s.CreateSession("sess-other", other, time.Hour); err != nil {
|
||||
t.Fatalf("CreateSession: %v", err)
|
||||
}
|
||||
got, ok, err := s.GetSession("sess-other", time.Now())
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("GetSession: ok=%v err=%v, want ok", ok, err)
|
||||
}
|
||||
if got.ReaderID != other {
|
||||
t.Fatalf("session reader = %d, want %d", got.ReaderID, other)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,308 @@
|
||||
package web
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"bookmarkmanager/backend/internal/session"
|
||||
)
|
||||
|
||||
const (
|
||||
// oauthStateTTL bounds how long a started sign-in stays valid. Ten
|
||||
// minutes is generous for Discord's round trip and short enough that a
|
||||
// captured state is stale before it is worth replaying.
|
||||
oauthStateTTL = 10 * time.Minute
|
||||
// maxStates caps the state map so a flood of /auth/discord hits cannot
|
||||
// grow memory; past the cap the oldest state is evicted, which at worst
|
||||
// invalidates an in-flight sign-in.
|
||||
maxStates = 256
|
||||
// maxResponseBytes caps Discord API bodies; they are small, and an
|
||||
// unbounded read is an OOM handed to Discord's CDN.
|
||||
maxResponseBytes = 1 << 20
|
||||
// discordTimeout keeps a hung Discord request from hanging the login
|
||||
// callback forever.
|
||||
discordTimeout = 15 * time.Second
|
||||
)
|
||||
|
||||
// DiscordConfig is the OAuth application this service registers as, plus the
|
||||
// guild that gates access.
|
||||
type DiscordConfig struct {
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
GuildID string
|
||||
// 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
|
||||
// APIBase is the Discord API root; configurable so tests run the whole
|
||||
// flow against a local stub.
|
||||
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
|
||||
// OwnerDiscordID is the only Discord identity allowed to sign in until
|
||||
// registration exists (issue #23).
|
||||
OwnerDiscordID string
|
||||
}
|
||||
|
||||
// oauthStates stores one-time sign-in states. A state is generated at
|
||||
// /auth/discord, echoed back by Discord at the callback, and consumed there.
|
||||
type oauthStates struct {
|
||||
mu sync.Mutex
|
||||
expiry map[string]time.Time
|
||||
}
|
||||
|
||||
func newOAuthStates() *oauthStates {
|
||||
return &oauthStates{expiry: make(map[string]time.Time)}
|
||||
}
|
||||
|
||||
func (s *oauthStates) put(state string, expires time.Time) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
now := time.Now()
|
||||
for k, at := range s.expiry {
|
||||
if !at.After(now) {
|
||||
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 {
|
||||
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
|
||||
}
|
||||
|
||||
// take validates and consumes a state in one step: a state works exactly
|
||||
// once, which is what makes a replayed callback useless.
|
||||
func (s *oauthStates) take(state string) bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
expires, ok := s.expiry[state]
|
||||
if !ok || !expires.After(time.Now()) {
|
||||
return false
|
||||
}
|
||||
delete(s.expiry, state)
|
||||
return true
|
||||
}
|
||||
|
||||
// discordStart begins the authorization code grant: a fresh state, then a
|
||||
// redirect to Discord's authorize page.
|
||||
func (h *Handler) discordStart(w http.ResponseWriter, r *http.Request) {
|
||||
state := session.NewID()
|
||||
h.states.put(state, time.Now().Add(oauthStateTTL))
|
||||
u := h.discord.APIBase + "/oauth2/authorize?" + url.Values{
|
||||
"client_id": {h.discord.ClientID},
|
||||
"redirect_uri": {h.discord.RedirectURI},
|
||||
"response_type": {"code"},
|
||||
"scope": {"identify guilds.members.read"},
|
||||
"state": {state},
|
||||
}.Encode()
|
||||
http.Redirect(w, r, u, http.StatusSeeOther)
|
||||
}
|
||||
|
||||
// discordCallback completes the grant: exchange the code, verify identity,
|
||||
// membership and role, then mint a session. Every failure path renders the
|
||||
// login page with an author-written message — nothing Discord supplied is
|
||||
// ever interpolated into a page, and no secret reaches a log line.
|
||||
func (h *Handler) discordCallback(w http.ResponseWriter, r *http.Request) {
|
||||
ip := session.ClientIP(r)
|
||||
if wait := h.limiter.RetryAfter(ip, time.Now()); wait > 0 {
|
||||
secs := int(wait.Seconds()) + 1
|
||||
w.Header().Set("Retry-After", strconv.Itoa(secs))
|
||||
h.renderLogin(w, http.StatusTooManyRequests,
|
||||
"Too many attempts. Try again in "+strconv.Itoa((secs+59)/60)+" min.")
|
||||
return
|
||||
}
|
||||
|
||||
// Discord refuses the grant (the reader hit cancel, or the application
|
||||
// was misconfigured). The state is consumed so the flow is cleanly over;
|
||||
// this makes no Discord calls, so it is not a failure the limiter counts.
|
||||
if oerr := r.URL.Query().Get("error"); oerr != "" {
|
||||
h.states.take(r.URL.Query().Get("state"))
|
||||
h.renderLogin(w, http.StatusBadRequest, "Sign-in was cancelled.")
|
||||
return
|
||||
}
|
||||
|
||||
code := r.URL.Query().Get("code")
|
||||
if code == "" || !h.states.take(r.URL.Query().Get("state")) {
|
||||
h.limiter.Fail(ip, time.Now())
|
||||
h.renderLogin(w, http.StatusBadRequest,
|
||||
"This sign-in link was invalid or already used. Start again.")
|
||||
return
|
||||
}
|
||||
|
||||
tok, err := h.exchangeToken(r.Context(), code)
|
||||
if err != nil {
|
||||
h.limiter.Fail(ip, time.Now())
|
||||
log.Printf("discord token exchange: %v", err)
|
||||
h.renderLogin(w, http.StatusBadGateway,
|
||||
"Discord sign-in is unavailable right now. Try again in a moment.")
|
||||
return
|
||||
}
|
||||
|
||||
userID, err := h.discordUserID(r.Context(), tok.AccessToken)
|
||||
if err != nil {
|
||||
h.limiter.Fail(ip, time.Now())
|
||||
log.Printf("discord users/@me: %v", err)
|
||||
h.renderLogin(w, http.StatusBadGateway,
|
||||
"Discord sign-in is unavailable right now. Try again in a moment.")
|
||||
return
|
||||
}
|
||||
|
||||
if userID != h.discord.OwnerDiscordID {
|
||||
h.limiter.Fail(ip, time.Now())
|
||||
h.renderLogin(w, http.StatusForbidden,
|
||||
"This Discord account is not the library owner.")
|
||||
return
|
||||
}
|
||||
|
||||
member, isMember, err := h.discordMember(r.Context(), tok.AccessToken, userID)
|
||||
if err != nil {
|
||||
h.limiter.Fail(ip, time.Now())
|
||||
log.Printf("discord member check: %v", err)
|
||||
h.renderLogin(w, http.StatusBadGateway,
|
||||
"Discord sign-in is unavailable right now. Try again in a moment.")
|
||||
return
|
||||
}
|
||||
// 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 != "" && !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.")
|
||||
return
|
||||
}
|
||||
|
||||
h.limiter.Reset(ip)
|
||||
sess, err := h.store.CreateSession(session.NewID(), h.readerID, session.SessionTTL)
|
||||
if err != nil {
|
||||
log.Printf("create session: %v", err)
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
session.SetCookie(w, r, sess.ID)
|
||||
http.Redirect(w, r, "/", http.StatusSeeOther)
|
||||
}
|
||||
|
||||
// exchangeToken trades an authorization code for an access token. The body is
|
||||
// form-encoded because that is what Discord accepts — it rejects a JSON
|
||||
// payload — so the wire format is fixed here, not in a client library.
|
||||
func (h *Handler) exchangeToken(ctx context.Context, code string) (discordToken, error) {
|
||||
form := url.Values{
|
||||
"client_id": {h.discord.ClientID},
|
||||
"client_secret": {h.discord.ClientSecret},
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {code},
|
||||
"redirect_uri": {h.discord.RedirectURI},
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
||||
h.discord.APIBase+"/oauth2/token", strings.NewReader(form.Encode()))
|
||||
if err != nil {
|
||||
return discordToken{}, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
resp, err := h.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return discordToken{}, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return discordToken{}, fmt.Errorf("status %d", resp.StatusCode)
|
||||
}
|
||||
var tok discordToken
|
||||
if err := json.NewDecoder(io.LimitReader(resp.Body, maxResponseBytes)).Decode(&tok); err != nil {
|
||||
return discordToken{}, err
|
||||
}
|
||||
if tok.AccessToken == "" {
|
||||
return discordToken{}, errors.New("empty access token")
|
||||
}
|
||||
return tok, nil
|
||||
}
|
||||
|
||||
// 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.APIBase+"/users/@me", nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||
req.Header.Set("Accept", "application/json")
|
||||
resp, err := h.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("status %d", resp.StatusCode)
|
||||
}
|
||||
var u struct {
|
||||
ID string `json:"id"`
|
||||
}
|
||||
if err := json.NewDecoder(io.LimitReader(resp.Body, maxResponseBytes)).Decode(&u); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if u.ID == "" {
|
||||
return "", errors.New("empty user id")
|
||||
}
|
||||
return u.ID, nil
|
||||
}
|
||||
|
||||
type discordMember struct {
|
||||
Roles []string `json:"roles"`
|
||||
}
|
||||
|
||||
// discordMember fetches the user's membership in the configured guild — the
|
||||
// single-guild endpoint, not the list of every guild the user is in, so the
|
||||
// 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.APIBase + "/guilds/" + url.PathEscape(h.discord.GuildID) +
|
||||
"/members/" + url.PathEscape(userID)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
|
||||
if err != nil {
|
||||
return discordMember{}, false, err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||
req.Header.Set("Accept", "application/json")
|
||||
resp, err := h.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return discordMember{}, false, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusForbidden {
|
||||
return discordMember{}, false, nil
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return discordMember{}, false, fmt.Errorf("status %d", resp.StatusCode)
|
||||
}
|
||||
var m discordMember
|
||||
if err := json.NewDecoder(io.LimitReader(resp.Body, maxResponseBytes)).Decode(&m); err != nil {
|
||||
return discordMember{}, false, err
|
||||
}
|
||||
return m, true, nil
|
||||
}
|
||||
|
||||
type discordToken struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package web
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestOAuthStateSingleUse(t *testing.T) {
|
||||
s := newOAuthStates()
|
||||
s.put("st", time.Now().Add(time.Minute))
|
||||
if !s.take("st") {
|
||||
t.Fatal("take of a fresh state = false, want true")
|
||||
}
|
||||
if s.take("st") {
|
||||
t.Fatal("take of a consumed state = true, want false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthStateUnknownOrExpired(t *testing.T) {
|
||||
s := newOAuthStates()
|
||||
if s.take("never-seen") {
|
||||
t.Fatal("take of an unknown state = true, want false")
|
||||
}
|
||||
s.put("stale", time.Now().Add(-time.Minute))
|
||||
if s.take("stale") {
|
||||
t.Fatal("take of an expired state = true, want false")
|
||||
}
|
||||
}
|
||||
|
||||
// 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 {
|
||||
return string(rune('a'+i%26)) + string(rune('0'+i/26+salt*16))
|
||||
}
|
||||
now := time.Now().Add(time.Hour)
|
||||
for i := 0; i < maxStates*2; i++ {
|
||||
s.put(key(i, 0), now)
|
||||
}
|
||||
if got := len(s.expiry); got != maxStates {
|
||||
t.Fatalf("states after a flood = %d, want %d", got, maxStates)
|
||||
}
|
||||
|
||||
// Consume everything, then flood again: the map stays bounded.
|
||||
for state := range s.expiry {
|
||||
s.take(state)
|
||||
}
|
||||
for i := 0; i < maxStates; i++ {
|
||||
s.put(key(i, 1), now)
|
||||
}
|
||||
if got := len(s.expiry); got != maxStates {
|
||||
t.Fatalf("states after consume+flood = %d, want %d", got, maxStates)
|
||||
}
|
||||
}
|
||||
@@ -774,26 +774,6 @@ button { cursor: pointer; }
|
||||
filter: drop-shadow(0 0 34px var(--ember-wash)) drop-shadow(0 18px 24px rgba(0,0,0,.5));
|
||||
}
|
||||
.login-card form { display: flex; flex-direction: column; gap: 18px; }
|
||||
.login-card label {
|
||||
font: 500 10px/1 var(--font-mono);
|
||||
letter-spacing: .16em;
|
||||
text-transform: uppercase;
|
||||
color: var(--mute);
|
||||
}
|
||||
.login-card input {
|
||||
width: 100%;
|
||||
height: 54px;
|
||||
margin-top: 9px;
|
||||
padding: 0 2px;
|
||||
border: none;
|
||||
border-bottom: 1px solid var(--field-line);
|
||||
background: transparent;
|
||||
color: var(--paper);
|
||||
font: 500 20px var(--font-mono);
|
||||
letter-spacing: .16em;
|
||||
outline: none;
|
||||
}
|
||||
.login-card input:focus { border-bottom-color: var(--paper); }
|
||||
.login-card .error {
|
||||
margin: 0;
|
||||
min-height: 20px;
|
||||
@@ -810,7 +790,13 @@ 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;
|
||||
text-align: center;
|
||||
font: 400 12px/1.4 var(--font-body);
|
||||
color: var(--mute);
|
||||
}
|
||||
|
||||
/* ---- laptop and up: the whole sheet is drawn 20% larger, which is what
|
||||
|
||||
@@ -19,17 +19,13 @@
|
||||
<figure class="login-art" aria-hidden="true">
|
||||
<img src="/static/login-art.png" alt="">
|
||||
</figure>
|
||||
<form method="post" action="/login">
|
||||
<div>
|
||||
<label for="password">Password</label>
|
||||
<input id="password" name="password" type="password"
|
||||
autocomplete="current-password" autofocus required>
|
||||
</div>
|
||||
<form method="get" action="/auth/discord">
|
||||
{{/* The page reloads on a failed sign-in, so the message is present from
|
||||
the start; role=alert is what gets it announced anyway. */}}
|
||||
<p class="error" role="alert">{{.Error}}</p>
|
||||
<button type="submit">Sign in</button>
|
||||
<button type="submit">Continue with Discord</button>
|
||||
</form>
|
||||
<p class="login-note">Guild membership is required to sign in.</p>
|
||||
</main>
|
||||
</body>
|
||||
</html>
|
||||
|
||||
+70
-57
@@ -1,7 +1,7 @@
|
||||
package web
|
||||
|
||||
import (
|
||||
"crypto/subtle"
|
||||
"context"
|
||||
"embed"
|
||||
"html/template"
|
||||
"io/fs"
|
||||
@@ -31,14 +31,18 @@ const RecentCount = 5
|
||||
// It is a separate handler from api.Handler because the two speak different
|
||||
// representations (HTML versus JSON) to different clients under different auth.
|
||||
type Handler struct {
|
||||
store *store.Store
|
||||
// readerID is the Reader this UI acts as — the seeded owner, while the web
|
||||
// password is still the only credential (issue #22).
|
||||
store *store.Store
|
||||
// readerID is the owner Reader's id, the only Reader that can exist
|
||||
// while registration is closed (issue #23). Every session row points at
|
||||
// it, so it is also the Reader the UI acts as.
|
||||
readerID int64
|
||||
tmpl *template.Template
|
||||
key []byte
|
||||
password string
|
||||
discord DiscordConfig
|
||||
states *oauthStates
|
||||
limiter *session.LoginLimiter
|
||||
// httpClient is the plain stdlib client that talks to Discord. It is not
|
||||
// an injected interface: tests point APIBase at a stub server instead.
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
// listView is what every list-rendering template receives.
|
||||
@@ -46,7 +50,7 @@ type listView struct {
|
||||
// Lib is the library this view renders: store.KindManga or store.KindNovel.
|
||||
// Manga is the default and carries no query parameter, so every pre-novel
|
||||
// URL keeps meaning exactly what it did.
|
||||
Lib string
|
||||
Lib string
|
||||
Tab string // "all", "fav", or "new"
|
||||
Recent []store.Bookmark
|
||||
Items []store.Bookmark
|
||||
@@ -84,24 +88,26 @@ type loginView struct {
|
||||
|
||||
// New parses every template up front so a broken one kills the process at
|
||||
// startup rather than the first request that touches it.
|
||||
func New(s *store.Store, readerID int64, apiToken, webPassword string) (*Handler, error) {
|
||||
func New(s *store.Store, readerID int64, discord DiscordConfig) (*Handler, error) {
|
||||
tmpl, err := template.ParseFS(templateFS, "templates/*.html")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Handler{
|
||||
store: s,
|
||||
readerID: readerID,
|
||||
tmpl: tmpl,
|
||||
key: session.Key(apiToken, webPassword),
|
||||
password: webPassword,
|
||||
limiter: session.NewLoginLimiter(),
|
||||
store: s,
|
||||
readerID: readerID,
|
||||
tmpl: tmpl,
|
||||
discord: discord,
|
||||
states: newOAuthStates(),
|
||||
limiter: session.NewLoginLimiter(),
|
||||
httpClient: &http.Client{Timeout: discordTimeout},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("GET /{$}", h.index)
|
||||
mux.HandleFunc("POST /login", h.login)
|
||||
mux.HandleFunc("GET /auth/discord", h.discordStart)
|
||||
mux.HandleFunc("GET /auth/discord/callback", h.discordCallback)
|
||||
mux.HandleFunc("POST /logout", h.logout)
|
||||
mux.Handle("GET /static/", staticHandler())
|
||||
|
||||
@@ -135,10 +141,26 @@ func staticHandler() http.Handler {
|
||||
}))
|
||||
}
|
||||
|
||||
// authed reports whether the request carries a valid session cookie.
|
||||
func (h *Handler) authed(r *http.Request) bool {
|
||||
type ctxKey int
|
||||
|
||||
// readerCtxKey is where requireSession stashes the authenticated Reader id.
|
||||
const readerCtxKey ctxKey = iota
|
||||
|
||||
// sessionReader reports whether the request carries a live session, and for
|
||||
// whom. The cookie holds only the session id; the row behind it is looked up
|
||||
// on every request, so deleting a session takes effect immediately. Expiry is
|
||||
// enforced here, in the store, which also removes rows that have lapsed.
|
||||
func (h *Handler) sessionReader(r *http.Request) (int64, bool) {
|
||||
c, err := r.Cookie(session.CookieName)
|
||||
return err == nil && session.Verify(h.key, c.Value, time.Now().UnixMilli())
|
||||
if err != nil {
|
||||
return 0, false
|
||||
}
|
||||
sess, ok, err := h.store.GetSession(c.Value, time.Now())
|
||||
if err != nil {
|
||||
log.Printf("session lookup: %v", err)
|
||||
return 0, false
|
||||
}
|
||||
return sess.ReaderID, ok
|
||||
}
|
||||
|
||||
// requireSession guards the fragment endpoints. It answers 401 rather than
|
||||
@@ -146,14 +168,18 @@ func (h *Handler) authed(r *http.Request) bool {
|
||||
// redirected login page would be spliced into the card list.
|
||||
func (h *Handler) requireSession(next http.HandlerFunc) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
if !h.authed(r) {
|
||||
readerID, ok := h.sessionReader(r)
|
||||
if !ok {
|
||||
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
next(w, r)
|
||||
next(w, r.WithContext(context.WithValue(r.Context(), readerCtxKey, readerID)))
|
||||
}
|
||||
}
|
||||
|
||||
// readerOf returns the authenticated Reader id requireSession stashed.
|
||||
func readerOf(r *http.Request) int64 { return r.Context().Value(readerCtxKey).(int64) }
|
||||
|
||||
func (h *Handler) render(w http.ResponseWriter, status int, name string, data any) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
w.WriteHeader(status)
|
||||
@@ -167,11 +193,12 @@ func (h *Handler) render(w http.ResponseWriter, status int, name string, data an
|
||||
// page is served at / with status 200 rather than as a redirect to a separate
|
||||
// URL: one page, no redirect loop to reason about.
|
||||
func (h *Handler) index(w http.ResponseWriter, r *http.Request) {
|
||||
if !h.authed(r) {
|
||||
readerID, ok := h.sessionReader(r)
|
||||
if !ok {
|
||||
h.render(w, http.StatusOK, "login", loginView{})
|
||||
return
|
||||
}
|
||||
view, err := h.buildListView(libOf(r.URL.Query().Get("lib")), r.URL.Query().Get("tab"))
|
||||
view, err := h.buildListView(readerID, libOf(r.URL.Query().Get("lib")), r.URL.Query().Get("tab"))
|
||||
if err != nil {
|
||||
log.Printf("index: %v", err)
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
@@ -212,15 +239,15 @@ func libOf(q string) string {
|
||||
return store.KindManga
|
||||
}
|
||||
|
||||
// buildListView loads the list once and derives both the tab-filtered items and
|
||||
// the recent strip from it.
|
||||
// buildListView loads one reader's list once and derives both the tab-filtered
|
||||
// items and the recent strip from it.
|
||||
//
|
||||
// Archived and finished series appear in their own tab and nowhere else — not
|
||||
// in All, not in Updated, not in Favourites, and not in the recent strip. An
|
||||
// archived favourite therefore shows only under Archived: Favourites means
|
||||
// "favourites I am currently reading".
|
||||
func (h *Handler) buildListView(lib, tab string) (listView, error) {
|
||||
all, err := h.store.List(h.readerID) // already ordered updated_at DESC
|
||||
func (h *Handler) buildListView(readerID int64, lib, tab string) (listView, error) {
|
||||
all, err := h.store.List(readerID) // already ordered updated_at DESC
|
||||
if err != nil {
|
||||
return listView{}, err
|
||||
}
|
||||
@@ -271,7 +298,7 @@ func (h *Handler) buildListView(lib, tab string) (listView, error) {
|
||||
}
|
||||
|
||||
func (h *Handler) uiList(w http.ResponseWriter, r *http.Request) {
|
||||
view, err := h.buildListView(libOf(r.URL.Query().Get("lib")), r.URL.Query().Get("tab"))
|
||||
view, err := h.buildListView(readerOf(r), libOf(r.URL.Query().Get("lib")), r.URL.Query().Get("tab"))
|
||||
if err != nil {
|
||||
log.Printf("ui list: %v", err)
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
@@ -329,7 +356,7 @@ func (h *Handler) writeChromeOOB(w http.ResponseWriter, view listView) {
|
||||
// refreshChrome rebuilds the chrome for the reader's current tab after a
|
||||
// mutation and appends it to the response.
|
||||
func (h *Handler) refreshChrome(w http.ResponseWriter, r *http.Request) {
|
||||
view, err := h.buildListView(currentLib(r), currentTab(r))
|
||||
view, err := h.buildListView(readerOf(r), currentLib(r), currentTab(r))
|
||||
if err != nil {
|
||||
log.Printf("ui chrome: %v", err)
|
||||
return
|
||||
@@ -337,35 +364,21 @@ func (h *Handler) refreshChrome(w http.ResponseWriter, r *http.Request) {
|
||||
h.writeChromeOOB(w, view)
|
||||
}
|
||||
|
||||
func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
||||
ip := session.ClientIP(r)
|
||||
if wait := h.limiter.RetryAfter(ip, time.Now()); wait > 0 {
|
||||
secs := int(wait.Seconds()) + 1
|
||||
w.Header().Set("Retry-After", strconv.Itoa(secs))
|
||||
h.render(w, http.StatusTooManyRequests, "login", loginView{
|
||||
Error: "Too many attempts. Try again in " +
|
||||
strconv.Itoa((secs+59)/60) + " min.",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if err := r.ParseForm(); err != nil {
|
||||
http.Error(w, "invalid form", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
got := r.PostFormValue("password")
|
||||
if subtle.ConstantTimeCompare([]byte(got), []byte(h.password)) != 1 {
|
||||
h.limiter.Fail(ip, time.Now())
|
||||
h.render(w, http.StatusUnauthorized, "login", loginView{Error: "Wrong password."})
|
||||
return
|
||||
}
|
||||
|
||||
h.limiter.Reset(ip)
|
||||
session.SetCookie(w, r, h.key)
|
||||
http.Redirect(w, r, "/", http.StatusSeeOther)
|
||||
// renderLogin renders the login page with an error message, for refused or
|
||||
// failed sign-ins. Every message is author-written text — nothing Discord
|
||||
// supplied is ever interpolated into a page.
|
||||
func (h *Handler) renderLogin(w http.ResponseWriter, status int, msg string) {
|
||||
h.render(w, status, "login", loginView{Error: msg})
|
||||
}
|
||||
|
||||
// logout revokes the session row and clears the cookie in one step: the next
|
||||
// request finds no row and is rejected.
|
||||
func (h *Handler) logout(w http.ResponseWriter, r *http.Request) {
|
||||
if c, err := r.Cookie(session.CookieName); err == nil {
|
||||
if err := h.store.DeleteSession(c.Value); err != nil {
|
||||
log.Printf("delete session: %v", err)
|
||||
}
|
||||
}
|
||||
session.ClearCookie(w, r)
|
||||
http.Redirect(w, r, "/", http.StatusSeeOther)
|
||||
}
|
||||
@@ -378,7 +391,7 @@ func (h *Handler) loadForMutation(w http.ResponseWriter, r *http.Request) (store
|
||||
http.Error(w, "missing key", http.StatusBadRequest)
|
||||
return store.Bookmark{}, false
|
||||
}
|
||||
b, ok, err := h.store.Get(h.readerID, key)
|
||||
b, ok, err := h.store.Get(readerOf(r), key)
|
||||
if err != nil {
|
||||
log.Printf("ui get %q: %v", key, err)
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
@@ -401,7 +414,7 @@ func (h *Handler) loadForMutation(w http.ResponseWriter, r *http.Request) (store
|
||||
// describe the whole library, so they are rebuilt out of band on every
|
||||
// mutation, at the cost of one extra list read per toggle.
|
||||
func (h *Handler) saveAndRenderCard(w http.ResponseWriter, r *http.Request, b store.Bookmark) {
|
||||
stored, err := h.store.Upsert(h.readerID, b)
|
||||
stored, err := h.store.Upsert(readerOf(r), b)
|
||||
if err != nil {
|
||||
log.Printf("ui upsert %q: %v", b.Key, err)
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
@@ -493,7 +506,7 @@ func (h *Handler) uiDelete(w http.ResponseWriter, r *http.Request) {
|
||||
http.Error(w, "missing key", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if err := h.store.Delete(h.readerID, key); err != nil {
|
||||
if err := h.store.Delete(readerOf(r), key); err != nil {
|
||||
log.Printf("ui delete %q: %v", key, err)
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
return
|
||||
|
||||
+32
-14
@@ -29,11 +29,12 @@ type Config struct {
|
||||
// because a wrong guess would silently start on an empty database.
|
||||
DatabaseURL string
|
||||
Port string
|
||||
// WebPassword gates the browser UI. Empty disables the web routes entirely.
|
||||
WebPassword string
|
||||
// OwnerDiscordID identifies the seeded owner Reader (issue #22). Required:
|
||||
// bookmarks are scoped to a Reader, and without an owner there is none.
|
||||
// It is also the only Discord identity allowed to sign in (issue #23).
|
||||
OwnerDiscordID string
|
||||
// Discord is the OAuth application the browser UI signs in with.
|
||||
Discord web.DiscordConfig
|
||||
// UserscriptPath is the file served at /u/{token}/manga-bookmark.user.js.
|
||||
// Supplied by a bindmount so the script can be edited without a rebuild.
|
||||
UserscriptPath string
|
||||
@@ -151,12 +152,20 @@ func loadConfig() Config {
|
||||
Token: os.Getenv("API_TOKEN"),
|
||||
DatabaseURL: os.Getenv("DATABASE_URL"),
|
||||
Port: envOr("PORT", "8080"),
|
||||
WebPassword: os.Getenv("WEB_PASSWORD"),
|
||||
OwnerDiscordID: os.Getenv("OWNER_DISCORD_ID"),
|
||||
UserscriptPath: envOr("USERSCRIPT_PATH", "/userscript/manga-bookmark.user.js"),
|
||||
NovelUserscriptPath: envOr("NOVEL_USERSCRIPT_PATH", "/userscript/novel-bookmark.user.js"),
|
||||
LatestPoll: loadLatestPoll(),
|
||||
}
|
||||
c.Discord = web.DiscordConfig{
|
||||
ClientID: os.Getenv("DISCORD_CLIENT_ID"),
|
||||
ClientSecret: os.Getenv("DISCORD_CLIENT_SECRET"),
|
||||
GuildID: os.Getenv("DISCORD_GUILD_ID"),
|
||||
RequiredRole: os.Getenv("DISCORD_REQUIRED_ROLE"),
|
||||
APIBase: envOr("DISCORD_API_BASE", "https://discord.com/api/v10"),
|
||||
RedirectURI: os.Getenv("DISCORD_REDIRECT_URI"),
|
||||
OwnerDiscordID: c.OwnerDiscordID,
|
||||
}
|
||||
for _, o := range strings.Split(os.Getenv("ALLOWED_ORIGINS"), ",") {
|
||||
if o = strings.TrimSpace(o); o != "" {
|
||||
c.AllowedOrigins = append(c.AllowedOrigins, o)
|
||||
@@ -173,8 +182,8 @@ func newRouter(s *store.Store, cfg Config) http.Handler {
|
||||
mux.HandleFunc("GET /healthz", api.Healthz)
|
||||
|
||||
// Outside httpmw.Auth (the updater sends no Authorization header) and
|
||||
// outside the WEB_PASSWORD gate (the script must be installable either
|
||||
// way). The path segment carries the token instead.
|
||||
// outside the web UI's Discord auth (the script must be installable
|
||||
// without a browser session). The path segment carries the token instead.
|
||||
mux.HandleFunc("GET /u/{token}/manga-bookmark.user.js", userscript.Handler(cfg.Token, cfg.UserscriptPath))
|
||||
mux.HandleFunc("GET /u/{token}/novel-bookmark.user.js", userscript.Handler(cfg.Token, cfg.NovelUserscriptPath))
|
||||
|
||||
@@ -188,16 +197,13 @@ func newRouter(s *store.Store, cfg Config) http.Handler {
|
||||
mux.Handle("/bookmarks", auth)
|
||||
mux.Handle("/bookmarks/", auth)
|
||||
|
||||
// The browser UI is registered only when a password is configured, so a
|
||||
// deployment that forgets WEB_PASSWORD exposes nothing rather than
|
||||
// exposing an unprotected list.
|
||||
if cfg.WebPassword != "" {
|
||||
wh, err := web.New(s, s.OwnerID(), cfg.Token, cfg.WebPassword)
|
||||
if err != nil {
|
||||
log.Fatalf("web handler: %v", err)
|
||||
}
|
||||
wh.Register(mux)
|
||||
// The browser UI is always registered; signing in is Discord OAuth, so
|
||||
// there is no password to forget and no gate to leave unset.
|
||||
wh, err := web.New(s, s.OwnerID(), cfg.Discord)
|
||||
if err != nil {
|
||||
log.Fatalf("web handler: %v", err)
|
||||
}
|
||||
wh.Register(mux)
|
||||
|
||||
return httpmw.CORS(cfg.AllowedOrigins, httpmw.Gzip(guardEmptyUserscriptToken(mux)))
|
||||
}
|
||||
@@ -228,6 +234,18 @@ func main() {
|
||||
if cfg.DatabaseURL == "" {
|
||||
log.Fatal("DATABASE_URL is required")
|
||||
}
|
||||
// The web UI signs in through Discord, so a deployment without the OAuth
|
||||
// application is misconfigured rather than passwordless.
|
||||
for key, v := range map[string]string{
|
||||
"DISCORD_CLIENT_ID": cfg.Discord.ClientID,
|
||||
"DISCORD_CLIENT_SECRET": cfg.Discord.ClientSecret,
|
||||
"DISCORD_GUILD_ID": cfg.Discord.GuildID,
|
||||
"DISCORD_REDIRECT_URI": cfg.Discord.RedirectURI,
|
||||
} {
|
||||
if v == "" {
|
||||
log.Fatalf("%s is required", key)
|
||||
}
|
||||
}
|
||||
|
||||
// The owner's userscript token is the global API token today (issue #22);
|
||||
// the readers row carries its SHA-256, not the token itself.
|
||||
|
||||
+482
-85
@@ -1,10 +1,14 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -13,13 +17,17 @@ import (
|
||||
"bookmarkmanager/backend/internal/session"
|
||||
"bookmarkmanager/backend/internal/store"
|
||||
"bookmarkmanager/backend/internal/web"
|
||||
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
)
|
||||
|
||||
const testPassword = "hunter2"
|
||||
const testOwnerID = "owner-snowflake"
|
||||
|
||||
// webConfig returns a config whose web UI is usable: Discord identity is set,
|
||||
// though the OAuth endpoints still need the stub URL from discordConfig.
|
||||
func webConfig() Config {
|
||||
cfg := testConfig()
|
||||
cfg.WebPassword = testPassword
|
||||
cfg.Discord.OwnerDiscordID = testOwnerID
|
||||
return cfg
|
||||
}
|
||||
|
||||
@@ -31,13 +39,165 @@ func newWebTestServer(t *testing.T, cfg Config) (http.Handler, *store.Store) {
|
||||
return newRouter(st, cfg), st
|
||||
}
|
||||
|
||||
// sessionCookie returns a cookie a handler will accept for cfg's API token.
|
||||
func sessionCookie(t *testing.T, cfg Config) *http.Cookie {
|
||||
// sessionCookie mints a live session row for the owner and returns the cookie
|
||||
// carrying its id — the only credential the UI accepts.
|
||||
func sessionCookie(t *testing.T, st *store.Store) *http.Cookie {
|
||||
t.Helper()
|
||||
return &http.Cookie{
|
||||
Name: session.CookieName,
|
||||
Value: session.Sign(session.Key(cfg.Token, cfg.WebPassword), time.Now().Add(time.Hour).UnixMilli()),
|
||||
sess, err := st.CreateSession(session.NewID(), st.OwnerID(), session.SessionTTL)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateSession: %v", err)
|
||||
}
|
||||
return &http.Cookie{Name: session.CookieName, Value: sess.ID}
|
||||
}
|
||||
|
||||
// discordStub is a minimal Discord API. The router is pointed at it through
|
||||
// the configured API base URL, so the real request construction — including
|
||||
// the form-encoded token exchange — is what the tests exercise, not an
|
||||
// injected client interface.
|
||||
type discordStub struct {
|
||||
ownerID string // id /users/@me answers
|
||||
member bool // whether the member endpoint reports membership
|
||||
roles []string // roles the member holds
|
||||
tokenStatus int // status the token endpoint answers; 0 = 200
|
||||
userStatus int // status users/@me answers; 0 = 200
|
||||
memberStatus int // status the member endpoint answers; 0 = member ? 200 : 404
|
||||
|
||||
tokenRequests []tokenRequest // recorded token exchanges
|
||||
userAuth []string // Authorization headers seen on users/@me
|
||||
memberAuth []string // Authorization headers seen on the member endpoint
|
||||
memberPaths []string
|
||||
}
|
||||
|
||||
type tokenRequest struct {
|
||||
contentType string
|
||||
form url.Values
|
||||
}
|
||||
|
||||
func newDiscordStub(t *testing.T) (*discordStub, *httptest.Server) {
|
||||
t.Helper()
|
||||
st := &discordStub{ownerID: testOwnerID, member: true}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
case r.URL.Path == "/oauth2/token":
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
form, _ := url.ParseQuery(string(body))
|
||||
st.tokenRequests = append(st.tokenRequests, tokenRequest{
|
||||
contentType: r.Header.Get("Content-Type"),
|
||||
form: form,
|
||||
})
|
||||
status := st.tokenStatus
|
||||
if status == 0 {
|
||||
status = http.StatusOK
|
||||
}
|
||||
w.WriteHeader(status)
|
||||
if status == http.StatusOK {
|
||||
fmt.Fprintf(w, `{"access_token":"tok-%d","token_type":"Bearer"}`, len(st.tokenRequests))
|
||||
}
|
||||
case r.URL.Path == "/users/@me":
|
||||
st.userAuth = append(st.userAuth, r.Header.Get("Authorization"))
|
||||
status := st.userStatus
|
||||
if status == 0 {
|
||||
status = http.StatusOK
|
||||
}
|
||||
w.WriteHeader(status)
|
||||
if status == http.StatusOK {
|
||||
fmt.Fprintf(w, `{"id":%q,"username":"owner"}`, st.ownerID)
|
||||
}
|
||||
case strings.HasPrefix(r.URL.Path, "/guilds/"):
|
||||
st.memberPaths = append(st.memberPaths, r.URL.Path)
|
||||
st.memberAuth = append(st.memberAuth, r.Header.Get("Authorization"))
|
||||
status := st.memberStatus
|
||||
if status == 0 {
|
||||
if st.member {
|
||||
status = http.StatusOK
|
||||
} else {
|
||||
status = http.StatusNotFound
|
||||
}
|
||||
}
|
||||
w.WriteHeader(status)
|
||||
if status == http.StatusOK {
|
||||
roles, _ := json.Marshal(st.roles)
|
||||
fmt.Fprintf(w, `{"roles":%s}`, roles)
|
||||
}
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
return st, srv
|
||||
}
|
||||
|
||||
// discordConfig is the OAuth application config every sign-in test uses, with
|
||||
// the API base pointed at a stub.
|
||||
func discordConfig(stubURL string) web.DiscordConfig {
|
||||
return web.DiscordConfig{
|
||||
ClientID: "client-1",
|
||||
ClientSecret: "client-secret-1",
|
||||
GuildID: "guild-1",
|
||||
APIBase: stubURL,
|
||||
RedirectURI: "https://bm.example.com/auth/discord/callback",
|
||||
OwnerDiscordID: testOwnerID,
|
||||
}
|
||||
}
|
||||
|
||||
// oauthWebTestServer returns the full router, its store, and a Discord stub
|
||||
// wired as the configured API — the starting point for sign-in tests.
|
||||
func oauthWebTestServer(t *testing.T) (http.Handler, *store.Store, *discordStub) {
|
||||
t.Helper()
|
||||
stub, srv := newDiscordStub(t)
|
||||
cfg := webConfig()
|
||||
cfg.Discord = discordConfig(srv.URL)
|
||||
router, st := newWebTestServer(t, cfg)
|
||||
return router, st, stub
|
||||
}
|
||||
|
||||
// startSignIn runs GET /auth/discord and returns the state Discord would echo
|
||||
// back. A failed start fails the test.
|
||||
func startSignIn(t *testing.T, srv http.Handler) string {
|
||||
t.Helper()
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, "/auth/discord", nil))
|
||||
if rr.Code != http.StatusSeeOther {
|
||||
t.Fatalf("GET /auth/discord status = %d, want 303", rr.Code)
|
||||
}
|
||||
loc, err := url.Parse(rr.Header().Get("Location"))
|
||||
if err != nil {
|
||||
t.Fatalf("Location %q: %v", rr.Header().Get("Location"), err)
|
||||
}
|
||||
if loc.Path != "/oauth2/authorize" {
|
||||
t.Fatalf("redirect path = %q, want /oauth2/authorize", loc.Path)
|
||||
}
|
||||
if state := loc.Query().Get("state"); state != "" {
|
||||
return state
|
||||
}
|
||||
t.Fatal("authorize URL carries no state")
|
||||
return ""
|
||||
}
|
||||
|
||||
// completeSignIn drives the callback with a fresh code for state.
|
||||
func completeSignIn(t *testing.T, srv http.Handler, state string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodGet,
|
||||
"/auth/discord/callback?code=discord-code-1&state="+url.QueryEscape(state), nil)
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
return rr
|
||||
}
|
||||
|
||||
// readerCount pokes the readers table directly — the refusal contract is that
|
||||
// nothing was created, which the store API would not show.
|
||||
func readerCount(t *testing.T, url string) int {
|
||||
t.Helper()
|
||||
db, err := sql.Open("pgx", url)
|
||||
if err != nil {
|
||||
t.Fatalf("open: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
var n int
|
||||
if err := db.QueryRow(`SELECT count(*) FROM readers`).Scan(&n); err != nil {
|
||||
t.Fatalf("count readers: %v", err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func TestIndexWithoutSessionShowsLogin(t *testing.T) {
|
||||
@@ -48,8 +208,8 @@ func TestIndexWithoutSessionShowsLogin(t *testing.T) {
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("GET / status = %d, want 200", rr.Code)
|
||||
}
|
||||
if !strings.Contains(rr.Body.String(), `type="password"`) {
|
||||
t.Fatal("GET / without a session did not render the password field")
|
||||
if !strings.Contains(rr.Body.String(), "Continue with Discord") {
|
||||
t.Fatal("GET / without a session did not render the Discord sign-in button")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -65,7 +225,7 @@ func TestIndexWithSessionShowsList(t *testing.T) {
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(sessionCookie(t, st))
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
|
||||
@@ -77,56 +237,265 @@ func TestIndexWithSessionShowsList(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginSuccessSetsCookie(t *testing.T) {
|
||||
srv, _ := newWebTestServer(t, webConfig())
|
||||
req := httptest.NewRequest(http.MethodPost, "/login",
|
||||
strings.NewReader(url.Values{"password": {testPassword}}.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
func TestDiscordLoginFullFlow(t *testing.T) {
|
||||
srv, st, stub := oauthWebTestServer(t)
|
||||
|
||||
// The authorize redirect carries the app, the scopes the gate needs, and
|
||||
// a fresh state.
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, "/auth/discord", nil))
|
||||
if rr.Code != http.StatusSeeOther {
|
||||
t.Fatalf("POST /login status = %d, want 303", rr.Code)
|
||||
t.Fatalf("GET /auth/discord status = %d, want 303", rr.Code)
|
||||
}
|
||||
loc, err := url.Parse(rr.Header().Get("Location"))
|
||||
if err != nil {
|
||||
t.Fatalf("Location: %v", err)
|
||||
}
|
||||
q := loc.Query()
|
||||
if q.Get("client_id") != "client-1" || q.Get("response_type") != "code" {
|
||||
t.Fatalf("authorize query = %v, want client_id client-1 and response_type code", q)
|
||||
}
|
||||
if q.Get("redirect_uri") != "https://bm.example.com/auth/discord/callback" {
|
||||
t.Fatalf("redirect_uri = %q, want the configured callback", q.Get("redirect_uri"))
|
||||
}
|
||||
for _, want := range []string{"identify", "guilds.members.read"} {
|
||||
if !strings.Contains(q.Get("scope"), want) {
|
||||
t.Fatalf("scope %q missing %s", q.Get("scope"), want)
|
||||
}
|
||||
}
|
||||
state := q.Get("state")
|
||||
if state == "" {
|
||||
t.Fatal("authorize URL carries no state")
|
||||
}
|
||||
|
||||
// The callback lands the reader logged in.
|
||||
rr = completeSignIn(t, srv, state)
|
||||
if rr.Code != http.StatusSeeOther {
|
||||
t.Fatalf("callback status = %d, want 303 (body %s)", rr.Code, rr.Body.String())
|
||||
}
|
||||
cookies := rr.Result().Cookies()
|
||||
if len(cookies) != 1 || cookies[0].Name != session.CookieName || cookies[0].Value == "" {
|
||||
t.Fatalf("POST /login cookies = %+v, want one non-empty %s", cookies, session.CookieName)
|
||||
t.Fatalf("callback cookies = %+v, want one non-empty %s", cookies, session.CookieName)
|
||||
}
|
||||
|
||||
// The token exchange went out form-encoded — the wire format Discord
|
||||
// rejects if JSON — with every field Discord requires.
|
||||
if len(stub.tokenRequests) != 1 {
|
||||
t.Fatalf("token exchanges = %d, want 1", len(stub.tokenRequests))
|
||||
}
|
||||
tr := stub.tokenRequests[0]
|
||||
if !strings.HasPrefix(tr.contentType, "application/x-www-form-urlencoded") {
|
||||
t.Fatalf("token exchange Content-Type = %q, want form-urlencoded", tr.contentType)
|
||||
}
|
||||
wantForm := url.Values{
|
||||
"client_id": {"client-1"},
|
||||
"client_secret": {"client-secret-1"},
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {"discord-code-1"},
|
||||
"redirect_uri": {"https://bm.example.com/auth/discord/callback"},
|
||||
}
|
||||
if !reflect.DeepEqual(tr.form, wantForm) {
|
||||
t.Fatalf("token form = %v, want %v", tr.form, wantForm)
|
||||
}
|
||||
|
||||
// Identity and membership were fetched with the exchanged token, and the
|
||||
// membership check used the single-guild endpoint.
|
||||
if len(stub.userAuth) != 1 || stub.userAuth[0] != "Bearer tok-1" {
|
||||
t.Fatalf("users/@me Authorization = %v, want [Bearer tok-1]", stub.userAuth)
|
||||
}
|
||||
if len(stub.memberPaths) != 1 || stub.memberPaths[0] != "/guilds/guild-1/members/owner-snowflake" {
|
||||
t.Fatalf("member requests = %v, want the single-guild endpoint", stub.memberPaths)
|
||||
}
|
||||
if len(stub.memberAuth) != 1 || stub.memberAuth[0] != "Bearer tok-1" {
|
||||
t.Fatalf("member Authorization = %v, want [Bearer tok-1]", stub.memberAuth)
|
||||
}
|
||||
|
||||
// The session row exists, and the cookie it minted opens the library.
|
||||
if _, ok, err := st.GetSession(cookies[0].Value, time.Now()); err != nil || !ok {
|
||||
t.Fatalf("session row: ok=%v err=%v, want ok", ok, err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
req.AddCookie(cookies[0])
|
||||
rr = httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
if rr.Code != http.StatusOK || strings.Contains(rr.Body.String(), "Continue with Discord") {
|
||||
t.Fatalf("GET / with the new cookie = %d, still showing the login page", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginWrongPassword(t *testing.T) {
|
||||
srv, _ := newWebTestServer(t, webConfig())
|
||||
req := httptest.NewRequest(http.MethodPost, "/login",
|
||||
strings.NewReader(url.Values{"password": {"wrong"}}.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
func TestDiscordCallbackRejectsMissingState(t *testing.T) {
|
||||
srv, _, stub := oauthWebTestServer(t)
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
|
||||
if rr.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("POST /login status = %d, want 401", rr.Code)
|
||||
srv.ServeHTTP(rr, httptest.NewRequest(http.MethodGet,
|
||||
"/auth/discord/callback?code=discord-code-1", nil))
|
||||
if rr.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want 400", rr.Code)
|
||||
}
|
||||
if len(rr.Result().Cookies()) != 0 {
|
||||
t.Fatal("a failed login set a cookie")
|
||||
t.Fatal("a refused callback set a cookie")
|
||||
}
|
||||
if len(stub.tokenRequests) != 0 || len(stub.userAuth) != 0 {
|
||||
t.Fatal("a state-less callback still called Discord")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginRateLimited(t *testing.T) {
|
||||
srv, _ := newWebTestServer(t, webConfig())
|
||||
post := func() *httptest.ResponseRecorder {
|
||||
req := httptest.NewRequest(http.MethodPost, "/login",
|
||||
strings.NewReader(url.Values{"password": {"wrong"}}.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
func TestDiscordCallbackRejectsMismatchedState(t *testing.T) {
|
||||
srv, _, stub := oauthWebTestServer(t)
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, httptest.NewRequest(http.MethodGet,
|
||||
"/auth/discord/callback?code=discord-code-1&state=not-the-state", nil))
|
||||
if rr.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want 400", rr.Code)
|
||||
}
|
||||
if len(rr.Result().Cookies()) != 0 {
|
||||
t.Fatal("a refused callback set a cookie")
|
||||
}
|
||||
if len(stub.tokenRequests) != 0 || len(stub.userAuth) != 0 {
|
||||
t.Fatal("a mismatched-state callback still called Discord")
|
||||
}
|
||||
}
|
||||
|
||||
// A state is single-use: replaying a consumed callback is refused.
|
||||
func TestDiscordCallbackStateIsSingleUse(t *testing.T) {
|
||||
srv, _, _ := oauthWebTestServer(t)
|
||||
state := startSignIn(t, srv)
|
||||
if rr := completeSignIn(t, srv, state); rr.Code != http.StatusSeeOther {
|
||||
t.Fatalf("first use status = %d, want 303", rr.Code)
|
||||
}
|
||||
rr := completeSignIn(t, srv, state)
|
||||
if rr.Code != http.StatusBadRequest {
|
||||
t.Fatalf("replayed state status = %d, want 400", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDiscordLoginRefusesNonMember covers the refusals that must read the
|
||||
// same: no membership, membership without the required role, and a member
|
||||
// endpoint that answers 403 (token lacking the scope). Neither may leak the
|
||||
// guild's existence or id, and neither may create anything.
|
||||
func TestDiscordLoginRefusesNonMember(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
member bool
|
||||
memberStatus int
|
||||
roles []string
|
||||
require string
|
||||
}{
|
||||
{"not a member", false, 0, nil, ""},
|
||||
{"missing the required role", true, 0, []string{"role-1"}, "role-2"},
|
||||
{"member endpoint 403", true, http.StatusForbidden, nil, ""},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
stub, srv := newDiscordStub(t)
|
||||
stub.member = tc.member
|
||||
stub.memberStatus = tc.memberStatus
|
||||
stub.roles = tc.roles
|
||||
cfg := webConfig()
|
||||
cfg.Discord = discordConfig(srv.URL)
|
||||
cfg.Discord.RequiredRole = tc.require
|
||||
st, dbURL := newTestStoreURL(t)
|
||||
router := newRouter(st, cfg)
|
||||
|
||||
rr := completeSignIn(t, router, startSignIn(t, router))
|
||||
if rr.Code != http.StatusForbidden {
|
||||
t.Fatalf("status = %d, want 403", rr.Code)
|
||||
}
|
||||
if !strings.Contains(rr.Body.String(), "not a member of this community") {
|
||||
t.Fatalf("refusal body = %q, want the clear non-member explanation", rr.Body.String())
|
||||
}
|
||||
if strings.Contains(rr.Body.String(), "guild-1") {
|
||||
t.Fatalf("refusal body = %q, leaks the guild id", rr.Body.String())
|
||||
}
|
||||
if len(rr.Result().Cookies()) != 0 {
|
||||
t.Fatal("a refused sign-in set a cookie")
|
||||
}
|
||||
// The seed owner is still the only Reader, and no session exists.
|
||||
if n := readerCount(t, dbURL); n != 1 {
|
||||
t.Fatalf("readers = %d after a refusal, want 1", n)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// The positive role-gated path: a member holding the required role signs in.
|
||||
func TestDiscordLoginRequiresRolePositive(t *testing.T) {
|
||||
stub, srv := newDiscordStub(t)
|
||||
stub.roles = []string{"role-1"}
|
||||
cfg := webConfig()
|
||||
cfg.Discord = discordConfig(srv.URL)
|
||||
cfg.Discord.RequiredRole = "role-1"
|
||||
router, st := newWebTestServer(t, cfg)
|
||||
|
||||
rr := completeSignIn(t, router, startSignIn(t, router))
|
||||
if rr.Code != http.StatusSeeOther {
|
||||
t.Fatalf("status = %d, want 303 (body %s)", rr.Code, rr.Body.String())
|
||||
}
|
||||
cookies := rr.Result().Cookies()
|
||||
if len(cookies) != 1 || cookies[0].Value == "" {
|
||||
t.Fatalf("cookies = %+v, want a session cookie", cookies)
|
||||
}
|
||||
if _, ok, err := st.GetSession(cookies[0].Value, time.Now()); err != nil || !ok {
|
||||
t.Fatalf("session row: ok=%v err=%v, want ok", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscordLoginRefusesNonOwner(t *testing.T) {
|
||||
stub, srv := newDiscordStub(t)
|
||||
stub.ownerID = "someone-elses-snowflake"
|
||||
cfg := webConfig()
|
||||
cfg.Discord = discordConfig(srv.URL)
|
||||
router, st := newWebTestServer(t, cfg)
|
||||
|
||||
rr := completeSignIn(t, router, startSignIn(t, router))
|
||||
if rr.Code != http.StatusForbidden {
|
||||
t.Fatalf("status = %d, want 403", rr.Code)
|
||||
}
|
||||
if !strings.Contains(rr.Body.String(), "not the library owner") {
|
||||
t.Fatalf("refusal body = %q, want the owner-only explanation", rr.Body.String())
|
||||
}
|
||||
if len(rr.Result().Cookies()) != 0 {
|
||||
t.Fatal("a refused sign-in set a cookie")
|
||||
}
|
||||
if sess, ok, _ := st.GetSession("anything", time.Now()); ok && sess.ID != "" {
|
||||
t.Fatal("a refused sign-in created a session")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscordLoginTokenEndpointDown(t *testing.T) {
|
||||
stub, srv := newDiscordStub(t)
|
||||
stub.tokenStatus = http.StatusInternalServerError
|
||||
cfg := webConfig()
|
||||
cfg.Discord = discordConfig(srv.URL)
|
||||
router, _ := newWebTestServer(t, cfg)
|
||||
|
||||
rr := completeSignIn(t, router, startSignIn(t, router))
|
||||
if rr.Code != http.StatusBadGateway {
|
||||
t.Fatalf("status = %d, want 502", rr.Code)
|
||||
}
|
||||
if !strings.Contains(rr.Body.String(), "unavailable") {
|
||||
t.Fatalf("body = %q, want the unavailable message", rr.Body.String())
|
||||
}
|
||||
if len(rr.Result().Cookies()) != 0 {
|
||||
t.Fatal("a failed sign-in set a cookie")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackRateLimited(t *testing.T) {
|
||||
srv, _, _ := oauthWebTestServer(t)
|
||||
call := func() *httptest.ResponseRecorder {
|
||||
req := httptest.NewRequest(http.MethodGet,
|
||||
"/auth/discord/callback?code=x&state=not-the-state", nil)
|
||||
req.Header.Set("X-Forwarded-For", "203.0.113.9")
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
return rr
|
||||
}
|
||||
for i := 0; i < session.MaxFailures; i++ {
|
||||
if code := post().Code; code != http.StatusUnauthorized {
|
||||
t.Fatalf("attempt %d status = %d, want 401", i+1, code)
|
||||
if code := call().Code; code != http.StatusBadRequest {
|
||||
t.Fatalf("attempt %d status = %d, want 400", i+1, code)
|
||||
}
|
||||
}
|
||||
rr := post()
|
||||
rr := call()
|
||||
if rr.Code != http.StatusTooManyRequests {
|
||||
t.Fatalf("attempt %d status = %d, want 429", session.MaxFailures+1, rr.Code)
|
||||
}
|
||||
@@ -137,11 +506,13 @@ func TestLoginRateLimited(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogoutClearsCookie(t *testing.T) {
|
||||
func TestLogoutDeletesSession(t *testing.T) {
|
||||
cfg := webConfig()
|
||||
srv, _ := newWebTestServer(t, cfg)
|
||||
srv, st := newWebTestServer(t, cfg)
|
||||
cookie := sessionCookie(t, st)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/logout", nil)
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(cookie)
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
|
||||
@@ -152,26 +523,52 @@ func TestLogoutClearsCookie(t *testing.T) {
|
||||
if len(cookies) != 1 || cookies[0].MaxAge >= 0 {
|
||||
t.Fatalf("POST /logout cookies = %+v, want one expiring cookie", cookies)
|
||||
}
|
||||
// The row is gone, so the same cookie is dead on the next request.
|
||||
if _, ok, _ := st.GetSession(cookie.Value, time.Now()); ok {
|
||||
t.Fatal("session row still present after logout")
|
||||
}
|
||||
req = httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
req.AddCookie(cookie)
|
||||
rr = httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
if !strings.Contains(rr.Body.String(), "Continue with Discord") {
|
||||
t.Fatal("GET / after logout still rendered the library")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebDisabledWhenNoPassword(t *testing.T) {
|
||||
cfg := testConfig() // WebPassword empty
|
||||
srv, _ := newWebTestServer(t, cfg)
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, "/", nil))
|
||||
func TestExpiredSessionRejected(t *testing.T) {
|
||||
cfg := webConfig()
|
||||
srv, st := newWebTestServer(t, cfg)
|
||||
sess, err := st.CreateSession(session.NewID(), st.OwnerID(), -time.Minute)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateSession: %v", err)
|
||||
}
|
||||
cookie := &http.Cookie{Name: session.CookieName, Value: sess.ID}
|
||||
|
||||
if rr.Code != http.StatusNotFound {
|
||||
t.Fatalf("GET / with WEB_PASSWORD unset = %d, want 404", rr.Code)
|
||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
req.AddCookie(cookie)
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
if rr.Code != http.StatusOK || !strings.Contains(rr.Body.String(), "Continue with Discord") {
|
||||
t.Fatalf("GET / with an expired session = %d, want the login page", rr.Code)
|
||||
}
|
||||
|
||||
req = httptest.NewRequest(http.MethodGet, "/ui/list", nil)
|
||||
req.AddCookie(cookie)
|
||||
rr = httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
if rr.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("GET /ui/list with an expired session = %d, want 401", rr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBookmarksAPIStillBearerOnly(t *testing.T) {
|
||||
cfg := webConfig()
|
||||
srv, _ := newWebTestServer(t, cfg)
|
||||
srv, st := newWebTestServer(t, cfg)
|
||||
|
||||
// A session cookie must not grant access to the userscript's JSON API.
|
||||
req := httptest.NewRequest(http.MethodGet, "/bookmarks", nil)
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(sessionCookie(t, st))
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
if rr.Code != http.StatusUnauthorized {
|
||||
@@ -210,7 +607,7 @@ func seed(t *testing.T, st *store.Store, b store.Bookmark) store.Bookmark {
|
||||
return stored
|
||||
}
|
||||
|
||||
func uiRequest(t *testing.T, cfg Config, method, path string, form url.Values) *http.Request {
|
||||
func uiRequest(t *testing.T, st *store.Store, method, path string, form url.Values) *http.Request {
|
||||
t.Helper()
|
||||
var req *http.Request
|
||||
if form == nil {
|
||||
@@ -219,7 +616,7 @@ func uiRequest(t *testing.T, cfg Config, method, path string, form url.Values) *
|
||||
req = httptest.NewRequest(method, path, strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
}
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(sessionCookie(t, st))
|
||||
return req
|
||||
}
|
||||
|
||||
@@ -252,7 +649,7 @@ func TestFavoriteTogglesWithoutReordering(t *testing.T) {
|
||||
})
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodPost, "/ui/bookmarks/asura:solo/favorite", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodPost, "/ui/bookmarks/asura:solo/favorite", nil))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("favorite status = %d, want 200", rr.Code)
|
||||
}
|
||||
@@ -274,7 +671,7 @@ func TestFavoriteTogglesWithoutReordering(t *testing.T) {
|
||||
|
||||
// Toggling again turns it back off.
|
||||
rr = httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodPost, "/ui/bookmarks/asura:solo/favorite", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodPost, "/ui/bookmarks/asura:solo/favorite", nil))
|
||||
back, _, _ := st.Get(st.OwnerID(), "asura:solo")
|
||||
if back.Favorite {
|
||||
t.Fatal("Favorite = true after a second toggle, want false")
|
||||
@@ -301,7 +698,7 @@ func TestCardHxTargetIsValidSelectorForColonKey(t *testing.T) {
|
||||
})
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodGet, "/ui/list", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodGet, "/ui/list", nil))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
@@ -328,7 +725,7 @@ func TestChapterOverrideMovesUpdatedAt(t *testing.T) {
|
||||
})
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodPost,
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodPost,
|
||||
"/ui/bookmarks/asura:solo/chapter", url.Values{"chapter": {"60"}}))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("chapter override status = %d, want 200", rr.Code)
|
||||
@@ -368,7 +765,7 @@ func TestChapterOverrideNoOpPreservesURLAndUpdatedAt(t *testing.T) {
|
||||
// or move updated_at. The seed stores "45.0" against 45 so the display
|
||||
// string differs from what the form submits back.
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodPost,
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodPost,
|
||||
"/ui/bookmarks/asura:solo/chapter", url.Values{"chapter": {"45"}}))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("chapter no-op status = %d, want 200", rr.Code)
|
||||
@@ -403,7 +800,7 @@ func TestChapterOverrideRejectsBadInput(t *testing.T) {
|
||||
for _, bad := range []string{"", "abc", "-3", "NaN", "Infinity", "-Inf"} {
|
||||
t.Run("input "+bad, func(t *testing.T) {
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodPost,
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodPost,
|
||||
"/ui/bookmarks/asura:solo/chapter", url.Values{"chapter": {bad}}))
|
||||
if rr.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want 400", rr.Code)
|
||||
@@ -418,13 +815,13 @@ func TestChapterOverrideRejectsBadInput(t *testing.T) {
|
||||
|
||||
func TestMutationsOnMissingKey(t *testing.T) {
|
||||
cfg := webConfig()
|
||||
srv, _ := newWebTestServer(t, cfg)
|
||||
srv, st := newWebTestServer(t, cfg)
|
||||
cases := []struct {
|
||||
name string
|
||||
req *http.Request
|
||||
}{
|
||||
{"favorite", uiRequest(t, cfg, http.MethodPost, "/ui/bookmarks/asura:nope/favorite", nil)},
|
||||
{"chapter", uiRequest(t, cfg, http.MethodPost, "/ui/bookmarks/asura:nope/chapter", url.Values{"chapter": {"1"}})},
|
||||
{"favorite", uiRequest(t, st, http.MethodPost, "/ui/bookmarks/asura:nope/favorite", nil)},
|
||||
{"chapter", uiRequest(t, st, http.MethodPost, "/ui/bookmarks/asura:nope/chapter", url.Values{"chapter": {"1"}})},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
@@ -446,7 +843,7 @@ func TestUIDeleteRemovesRow(t *testing.T) {
|
||||
})
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodDelete, "/ui/bookmarks/asura:solo", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodDelete, "/ui/bookmarks/asura:solo", nil))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("delete status = %d, want 200", rr.Code)
|
||||
}
|
||||
@@ -477,7 +874,7 @@ func TestUIListFavouritesTab(t *testing.T) {
|
||||
})
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodGet, "/ui/list?tab=fav", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodGet, "/ui/list?tab=fav", nil))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
@@ -507,7 +904,7 @@ func TestUIListNewTab(t *testing.T) {
|
||||
})
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodGet, "/ui/list?tab=new", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodGet, "/ui/list?tab=new", nil))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
@@ -562,7 +959,7 @@ func TestTabsShowOnlyTheirBucket(t *testing.T) {
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.tab, func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/ui/list?tab="+tc.tab, nil)
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(sessionCookie(t, st))
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
|
||||
@@ -586,10 +983,10 @@ func TestTabsShowOnlyTheirBucket(t *testing.T) {
|
||||
|
||||
// stripOf returns everything above the list, which is where the recent section
|
||||
// renders.
|
||||
func stripOf(t *testing.T, srv http.Handler, cfg Config, tab string) string {
|
||||
func stripOf(t *testing.T, srv http.Handler, st *store.Store, tab string) string {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodGet, "/?tab="+tab, nil)
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(sessionCookie(t, st))
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
body := rr.Body.String()
|
||||
@@ -616,7 +1013,7 @@ func TestRecentStripCarriesUnreadOnlyAndOnlyOnAll(t *testing.T) {
|
||||
t.Fatalf("seed %s: %v", caught.Key, err)
|
||||
}
|
||||
|
||||
strip := stripOf(t, srv, cfg, "all")
|
||||
strip := stripOf(t, srv, st, "all")
|
||||
if !strings.Contains(strip, "ReadingOne") {
|
||||
t.Fatal("strip dropped the series with an unread chapter")
|
||||
}
|
||||
@@ -626,7 +1023,7 @@ func TestRecentStripCarriesUnreadOnlyAndOnlyOnAll(t *testing.T) {
|
||||
}
|
||||
}
|
||||
for _, tab := range []string{"new", "fav", "archived", "finished"} {
|
||||
if strings.Contains(stripOf(t, srv, cfg, tab), "ReadingOne") {
|
||||
if strings.Contains(stripOf(t, srv, st, tab), "ReadingOne") {
|
||||
t.Fatalf("tab %s rendered the strip", tab)
|
||||
}
|
||||
}
|
||||
@@ -642,7 +1039,7 @@ func TestRecentStripCarriesUnreadOnlyAndOnlyOnAll(t *testing.T) {
|
||||
}
|
||||
// The section still ships (an out-of-band swap needs the id to exist) but
|
||||
// carries no cards and is hidden.
|
||||
empty := stripOf(t, srv, cfg, "all")
|
||||
empty := stripOf(t, srv, st, "all")
|
||||
if strings.Contains(empty, "recent-card") {
|
||||
t.Fatal("strip rendered cards with no unread chapters anywhere")
|
||||
}
|
||||
@@ -666,18 +1063,18 @@ func TestRecentStripCapped(t *testing.T) {
|
||||
t.Fatalf("seed %s: %v", b.Key, err)
|
||||
}
|
||||
}
|
||||
if got := strings.Count(stripOf(t, srv, cfg, "all"), "recent-card"); got != web.RecentCount {
|
||||
if got := strings.Count(stripOf(t, srv, st, "all"), "recent-card"); got != web.RecentCount {
|
||||
t.Fatalf("strip rendered %d cards, want %d", got, web.RecentCount)
|
||||
}
|
||||
}
|
||||
|
||||
func postStatus(t *testing.T, srv http.Handler, cfg Config, key, status string) *httptest.ResponseRecorder {
|
||||
func postStatus(t *testing.T, srv http.Handler, st *store.Store, key, status string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
form := url.Values{"status": {status}}
|
||||
req := httptest.NewRequest(http.MethodPost, "/ui/bookmarks/"+key+"/status",
|
||||
strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(sessionCookie(t, st))
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
return rr
|
||||
@@ -689,7 +1086,7 @@ func TestUIStatusSetsBucket(t *testing.T) {
|
||||
seedStatusRows(t, st)
|
||||
|
||||
for _, want := range []string{store.StatusArchived, store.StatusFinished, store.StatusReading} {
|
||||
if rr := postStatus(t, srv, cfg, "asura:reading", want); rr.Code != http.StatusOK {
|
||||
if rr := postStatus(t, srv, st, "asura:reading", want); rr.Code != http.StatusOK {
|
||||
t.Fatalf("set %s: status = %d, body %s", want, rr.Code, rr.Body.String())
|
||||
}
|
||||
b, ok, err := st.Get(st.OwnerID(), "asura:reading")
|
||||
@@ -707,7 +1104,7 @@ func TestUIStatusRejectsUnknownValue(t *testing.T) {
|
||||
srv, st := newWebTestServer(t, cfg)
|
||||
seedStatusRows(t, st)
|
||||
|
||||
if rr := postStatus(t, srv, cfg, "asura:reading", "dropped"); rr.Code != http.StatusBadRequest {
|
||||
if rr := postStatus(t, srv, st, "asura:reading", "dropped"); rr.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want 400", rr.Code)
|
||||
}
|
||||
b, _, _ := st.Get(st.OwnerID(), "asura:reading")
|
||||
@@ -738,7 +1135,7 @@ func TestUIStatusDoesNotReorderList(t *testing.T) {
|
||||
|
||||
before, _, _ := st.Get(st.OwnerID(), "asura:reading")
|
||||
time.Sleep(2 * time.Millisecond)
|
||||
if rr := postStatus(t, srv, cfg, "asura:reading", store.StatusArchived); rr.Code != http.StatusOK {
|
||||
if rr := postStatus(t, srv, st, "asura:reading", store.StatusArchived); rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", rr.Code)
|
||||
}
|
||||
after, _, _ := st.Get(st.OwnerID(), "asura:reading")
|
||||
@@ -766,7 +1163,7 @@ func TestCardShowsStatusControls(t *testing.T) {
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.tab, func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/ui/list?tab="+tc.tab, nil)
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(sessionCookie(t, st))
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
|
||||
@@ -791,7 +1188,7 @@ func TestAppRendersNewTabs(t *testing.T) {
|
||||
seedStatusRows(t, st)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
req.AddCookie(sessionCookie(t, cfg))
|
||||
req.AddCookie(sessionCookie(t, st))
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, req)
|
||||
|
||||
@@ -813,12 +1210,12 @@ func TestMutationRefreshesChromeOutOfBand(t *testing.T) {
|
||||
LatestChapterNum: floatPtr(11), UpdatedAt: time.Now().UnixMilli(),
|
||||
})
|
||||
|
||||
before := stripOf(t, srv, cfg, "all")
|
||||
before := stripOf(t, srv, st, "all")
|
||||
if !strings.Contains(before, "Solo Leveling") || !strings.Contains(before, `id="new-count"`) {
|
||||
t.Fatalf("expected the series in the strip to start with: %q", before)
|
||||
}
|
||||
|
||||
req := uiRequest(t, cfg, http.MethodPost, "/ui/bookmarks/asura:solo/status",
|
||||
req := uiRequest(t, st, http.MethodPost, "/ui/bookmarks/asura:solo/status",
|
||||
url.Values{"status": {store.StatusArchived}})
|
||||
req.Header.Set("HX-Current-URL", "http://localhost/?tab=all")
|
||||
rr := httptest.NewRecorder()
|
||||
@@ -865,7 +1262,7 @@ func TestLibrariesAreDisjoint(t *testing.T) {
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodGet, tc.path, nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodGet, tc.path, nil))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
@@ -890,7 +1287,7 @@ func TestKindlessRowShowsInMangaLibrary(t *testing.T) {
|
||||
})
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodGet, "/ui/list?tab=all", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodGet, "/ui/list?tab=all", nil))
|
||||
if !strings.Contains(rr.Body.String(), "Legacy Series") {
|
||||
t.Fatal("a row with no kind must appear in the manga library")
|
||||
}
|
||||
@@ -902,7 +1299,7 @@ func TestNovelPageOmitsUpdatedTab(t *testing.T) {
|
||||
seedLibraries(t, st)
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodGet, "/?lib=novel&tab=all", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodGet, "/?lib=novel&tab=all", nil))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
@@ -933,7 +1330,7 @@ func TestMangaPageKeepsUpdatedTab(t *testing.T) {
|
||||
seedLibraries(t, st)
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodGet, "/?tab=all", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodGet, "/?tab=all", nil))
|
||||
body := rr.Body.String()
|
||||
if !strings.Contains(body, "/?tab=new") {
|
||||
t.Fatal("manga page must keep the Updated tab")
|
||||
@@ -951,7 +1348,7 @@ func TestNovelNewTabFallsBackToAll(t *testing.T) {
|
||||
seedLibraries(t, st)
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
srv.ServeHTTP(rr, uiRequest(t, cfg, http.MethodGet, "/ui/list?lib=novel&tab=new", nil))
|
||||
srv.ServeHTTP(rr, uiRequest(t, st, http.MethodGet, "/ui/list?lib=novel&tab=new", nil))
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rr.Code)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user