Commit a6f7827f2e
Verified · cmc
Layout: unified · split
internal/store/migrations/0062_web_session_idle.down.sql added +3
| @@ -0,0 +1,3 @@ | ||
| 1 | UPDATE web_sessions SET expires_at = absolute_expires_at; | |
| 2 | ALTER TABLE web_sessions DROP COLUMN last_used_at; | |
| 3 | ALTER TABLE web_sessions DROP COLUMN absolute_expires_at; | |
internal/store/migrations/0062_web_session_idle.up.sql added +8
| @@ -0,0 +1,8 @@ | ||
| 1 | -- expires_at slides forward on use, never past absolute_expires_at. | |
| 2 | -- Sessions open now keep their cap and get a full idle window from here. | |
| 3 | ALTER TABLE web_sessions ADD COLUMN absolute_expires_at TEXT; | |
| 4 | ALTER TABLE web_sessions ADD COLUMN last_used_at TEXT; | |
| 5 | UPDATE web_sessions SET | |
| 6 | absolute_expires_at = expires_at, | |
| 7 | last_used_at = strftime('%Y-%m-%dT%H:%M:%fZ','now'), | |
| 8 | expires_at = min(expires_at, strftime('%Y-%m-%dT%H:%M:%fZ','now','+12 hours')); | |
internal/store/sessions.go +25 −9
| @@ -64,25 +64,40 @@ func (s *Store) ConsumeLoginToken(hash string) (int64, error) { | ||
| 64 | 64 | return userID, err |
| 65 | 65 | } |
| 66 | 66 | |
| 67 | // WebSessionIdle is how long a browser session lasts without a request. | |
| 68 | // Each use moves its expiry this far ahead, never past the cap it was | |
| 69 | // created with. Migration 0062 repeats the value for sessions it | |
| 70 | // converts. | |
| 71 | const WebSessionIdle = 12 * time.Hour | |
| 72 | ||
| 73 | // CreateWebSession stores a session that lapses after WebSessionIdle | |
| 74 | // without use, and after ttl regardless. | |
| 67 | 75 | func (s *Store) CreateWebSession(hash string, userID int64, ttl time.Duration) error { |
| 76 | now := time.Now() | |
| 68 | 77 | _, err := s.DB.Exec( |
| 69 | "INSERT INTO web_sessions (token_hash, user_id, expires_at) VALUES (?, ?, ?)", | |
| 70 | hash, userID, fmtTime(time.Now().Add(ttl))) | |
| 78 | "INSERT INTO web_sessions (token_hash, user_id, expires_at, absolute_expires_at, last_used_at) VALUES (?, ?, ?, ?, ?)", | |
| 79 | hash, userID, fmtTime(now.Add(min(ttl, WebSessionIdle))), fmtTime(now.Add(ttl)), fmtTime(now)) | |
| 71 | 80 | return err |
| 72 | 81 | } |
| 73 | 82 | |
| 74 | // WebSessionUser resolves a session cookie hash to its user. | |
| 83 | // WebSessionUser resolves a session cookie hash to its user and renews | |
| 84 | // the session's idle expiry. A session is written at most once a | |
| 85 | // minute, so a burst of requests costs one UPDATE. | |
| 75 | 86 | func (s *Store) WebSessionUser(hash string) (User, error) { |
| 87 | now := time.Now() | |
| 76 | 88 | var userID int64 |
| 77 | 89 | err := s.DB.QueryRow( |
| 78 | 90 | "SELECT user_id FROM web_sessions WHERE token_hash = ? AND expires_at > ?", |
| 79 | hash, fmtTime(time.Now())).Scan(&userID) | |
| 91 | hash, fmtTime(now)).Scan(&userID) | |
| 80 | 92 | if errors.Is(err, sql.ErrNoRows) { |
| 81 | 93 | return User{}, ErrNotFound |
| 82 | 94 | } |
| 83 | 95 | if err != nil { |
| 84 | 96 | return User{}, err |
| 85 | 97 | } |
| 98 | s.DB.Exec(`UPDATE web_sessions SET last_used_at = ?, expires_at = min(absolute_expires_at, ?) | |
| 99 | WHERE token_hash = ? AND last_used_at < ?`, | |
| 100 | fmtTime(now), fmtTime(now.Add(WebSessionIdle)), hash, fmtTime(now.Add(-time.Minute))) | |
| 86 | 101 | return s.UserByID(userID) |
| 87 | 102 | } |
| 88 | 103 | |
| @@ -95,14 +110,15 @@ func (s *Store) DeleteWebSession(hash string) error { | ||
| 95 | 110 | // twelve hex digits of the stored token hash: enough to name it, and a |
| 96 | 111 | // hash of the cookie rather than the cookie. |
| 97 | 112 | type WebSession struct { |
| 98 | ID string `json:"id"` | |
| 99 | CreatedAt string `json:"created_at"` | |
| 100 | ExpiresAt string `json:"expires_at"` | |
| 113 | ID string `json:"id"` | |
| 114 | CreatedAt string `json:"created_at"` | |
| 115 | ExpiresAt string `json:"expires_at"` | |
| 116 | LastUsedAt string `json:"last_used_at"` | |
| 101 | 117 | } |
| 102 | 118 | |
| 103 | 119 | // ListWebSessions lists the user's unexpired browser sessions, newest first. |
| 104 | 120 | func (s *Store) ListWebSessions(userID int64) ([]WebSession, error) { |
| 105 | rows, err := s.DB.Query(`SELECT substr(token_hash, 1, 12), created_at, expires_at | |
| 121 | rows, err := s.DB.Query(`SELECT substr(token_hash, 1, 12), created_at, expires_at, COALESCE(last_used_at, created_at) | |
| 106 | 122 | FROM web_sessions WHERE user_id = ? AND expires_at > ? ORDER BY created_at DESC`, |
| 107 | 123 | userID, fmtTime(time.Now())) |
| 108 | 124 | if err != nil { |
| @@ -112,7 +128,7 @@ func (s *Store) ListWebSessions(userID int64) ([]WebSession, error) { | ||
| 112 | 128 | var out []WebSession |
| 113 | 129 | for rows.Next() { |
| 114 | 130 | var ws WebSession |
| 115 | if err := rows.Scan(&ws.ID, &ws.CreatedAt, &ws.ExpiresAt); err != nil { | |
| 131 | if err := rows.Scan(&ws.ID, &ws.CreatedAt, &ws.ExpiresAt, &ws.LastUsedAt); err != nil { | |
| 116 | 132 | return nil, err |
| 117 | 133 | } |
| 118 | 134 | out = append(out, ws) |
internal/store/sessions_test.go +73
| @@ -1,6 +1,7 @@ | ||
| 1 | 1 | package store |
| 2 | 2 | |
| 3 | 3 | import ( |
| 4 | "database/sql" | |
| 4 | 5 | "testing" |
| 5 | 6 | "time" |
| 6 | 7 | ) |
| @@ -44,3 +45,75 @@ func TestCountLoginTokensSince(t *testing.T) { | ||
| 44 | 45 | t.Fatalf("other account count = %d, %v; want 0", n, err) |
| 45 | 46 | } |
| 46 | 47 | } |
| 48 | ||
| 49 | func sessionFixture(t *testing.T) (*Store, int64) { | |
| 50 | t.Helper() | |
| 51 | s := open(t) | |
| 52 | if err := s.MigrateUp(); err != nil { | |
| 53 | t.Fatal(err) | |
| 54 | } | |
| 55 | uid, err := s.CreateUser("cmc", false) | |
| 56 | if err != nil { | |
| 57 | t.Fatal(err) | |
| 58 | } | |
| 59 | return s, uid | |
| 60 | } | |
| 61 | ||
| 62 | func sessionTimes(t *testing.T, s *Store, hash string) (expires, absolute time.Time) { | |
| 63 | t.Helper() | |
| 64 | var e, a string | |
| 65 | if err := s.DB.QueryRow("SELECT expires_at, absolute_expires_at FROM web_sessions WHERE token_hash = ?", hash).Scan(&e, &a); err != nil { | |
| 66 | t.Fatal(err) | |
| 67 | } | |
| 68 | return *parseTime(sql.NullString{String: e, Valid: true}), *parseTime(sql.NullString{String: a, Valid: true}) | |
| 69 | } | |
| 70 | ||
| 71 | func TestWebSessionIdleExpiry(t *testing.T) { | |
| 72 | s, uid := sessionFixture(t) | |
| 73 | if err := s.CreateWebSession("h", uid, 7*24*time.Hour); err != nil { | |
| 74 | t.Fatal(err) | |
| 75 | } | |
| 76 | exp, abs := sessionTimes(t, s, "h") | |
| 77 | if d := time.Until(exp); d < WebSessionIdle-time.Minute || d > WebSessionIdle { | |
| 78 | t.Fatalf("a new session expires in %s, want %s", d, WebSessionIdle) | |
| 79 | } | |
| 80 | if d := time.Until(abs); d < 7*24*time.Hour-time.Minute { | |
| 81 | t.Fatalf("absolute cap in %s", d) | |
| 82 | } | |
| 83 | // Idle past the window: gone. | |
| 84 | old := fmtTime(time.Now().Add(-time.Second)) | |
| 85 | s.DB.Exec("UPDATE web_sessions SET expires_at = ? WHERE token_hash = 'h'", old) | |
| 86 | if _, err := s.WebSessionUser("h"); err != ErrNotFound { | |
| 87 | t.Fatalf("idle session: %v", err) | |
| 88 | } | |
| 89 | } | |
| 90 | ||
| 91 | func TestWebSessionRenewsUpToTheCap(t *testing.T) { | |
| 92 | s, uid := sessionFixture(t) | |
| 93 | if err := s.CreateWebSession("h", uid, 7*24*time.Hour); err != nil { | |
| 94 | t.Fatal(err) | |
| 95 | } | |
| 96 | // Last used two minutes ago, one minute left: a request renews it. | |
| 97 | s.DB.Exec("UPDATE web_sessions SET last_used_at = ?, expires_at = ? WHERE token_hash = 'h'", | |
| 98 | fmtTime(time.Now().Add(-2*time.Minute)), fmtTime(time.Now().Add(time.Minute))) | |
| 99 | if _, err := s.WebSessionUser("h"); err != nil { | |
| 100 | t.Fatal(err) | |
| 101 | } | |
| 102 | if exp, _ := sessionTimes(t, s, "h"); time.Until(exp) < WebSessionIdle-time.Minute { | |
| 103 | t.Fatalf("not renewed: expires in %s", time.Until(exp)) | |
| 104 | } | |
| 105 | // Near the cap, renewal stops at it. | |
| 106 | capAt := time.Now().Add(time.Hour) | |
| 107 | s.DB.Exec("UPDATE web_sessions SET last_used_at = ?, absolute_expires_at = ? WHERE token_hash = 'h'", | |
| 108 | fmtTime(time.Now().Add(-2*time.Minute)), fmtTime(capAt)) | |
| 109 | if _, err := s.WebSessionUser("h"); err != nil { | |
| 110 | t.Fatal(err) | |
| 111 | } | |
| 112 | if exp, _ := sessionTimes(t, s, "h"); exp.After(capAt) { | |
| 113 | t.Fatalf("renewed past the cap: %s > %s", exp, capAt) | |
| 114 | } | |
| 115 | list, err := s.ListWebSessions(uid) | |
| 116 | if err != nil || len(list) != 1 || list[0].LastUsedAt == "" { | |
| 117 | t.Fatalf("list: %+v %v", list, err) | |
| 118 | } | |
| 119 | } | |