Commit cab50f5004

cab50f5004df82a3ff699f71bb31f049f8deca75

parent: e2a32d5f8d

Verified · cmc

cmc <hello@cleberg.net> · 2026-09-28 07:24 UTC

store: ssh_keys.expires_at; expired keys are not live

Ref #277

Layout: unified · split

internal/store/keyexpiry_test.go added +48
@@ -0,0 +1,48 @@
1package store
2
3import (
4 "testing"
5 "time"
6)
7
8func TestKeyExpiry(t *testing.T) {
9 s, uid, _ := revokeFixture(t)
10 past, future := time.Now().Add(-time.Minute), time.Now().Add(time.Hour)
11 for fp, exp := range map[string]*time.Time{"SHA256:old": &past, "SHA256:new": &future, "SHA256:ever": nil} {
12 if err := s.AddSSHKeyFrom(uid, fp, "ssh-ed25519", []byte(fp), "full", "", KeyOrigin{ExpiresAt: exp}); err != nil {
13 t.Fatal(err)
14 }
15 }
16 now := time.Now()
17 ids := map[string]int64{}
18 for _, fp := range []string{"SHA256:old", "SHA256:new", "SHA256:ever"} {
19 k, err := s.SSHKeyByFingerprint(fp)
20 if err != nil {
21 t.Fatal(err)
22 }
23 ids[fp] = k.ID
24 byID, err := s.SSHKeyByID(k.ID)
25 if err != nil || (byID.ExpiresAt == nil) != (k.ExpiresAt == nil) {
26 t.Fatalf("%s by id: %+v %v", fp, byID, err)
27 }
28 if got, want := k.Expired(now), fp == "SHA256:old"; got != want {
29 t.Errorf("%s Expired = %v, want %v", fp, got, want)
30 }
31 }
32 live, err := s.LiveSSHKeys([]int64{ids["SHA256:old"], ids["SHA256:new"], ids["SHA256:ever"]})
33 if err != nil {
34 t.Fatal(err)
35 }
36 if live[ids["SHA256:old"]] || !live[ids["SHA256:new"]] || !live[ids["SHA256:ever"]] {
37 t.Fatalf("live = %v", live)
38 }
39 keys, err := s.ListSSHKeys(uid)
40 if err != nil || len(keys) != 3 {
41 t.Fatalf("list: %+v %v", keys, err)
42 }
43 for _, k := range keys {
44 if k.Fingerprint == "SHA256:ever" && k.ExpiresAt != nil {
45 t.Fatalf("list: %s should have nil ExpiresAt: %+v", k.Fingerprint, k)
46 }
47 }
48}
internal/store/migrations/0061_ssh_key_expiry.down.sql added +1
@@ -0,0 +1 @@
1ALTER TABLE ssh_keys DROP COLUMN expires_at;
internal/store/migrations/0061_ssh_key_expiry.up.sql added +2
@@ -0,0 +1,2 @@
1-- When the key stops authenticating; NULL for never.
2ALTER TABLE ssh_keys ADD COLUMN expires_at TEXT;
internal/store/revoke.go +8 −6
@@ -3,6 +3,7 @@ package store
3import ( 3import (
4 "slices" 4 "slices"
5 "strings" 5 "strings"
6 "time"
6) 7)
7 8
8// Revoked names SSH keys that stopped being valid: by id, or every key 9// Revoked names SSH keys that stopped being valid: by id, or every key
@@ -32,19 +33,20 @@ func (s *Store) announce(r Revoked) {
32 } 33 }
33} 34}
34 35
35// LiveSSHKeys reports which of ids still name a registered key on an 36// LiveSSHKeys reports which of ids still name a registered, unexpired
36// account that is not disabled. 37// key on an account that is not disabled.
37func (s *Store) LiveSSHKeys(ids []int64) (map[int64]bool, error) { 38func (s *Store) LiveSSHKeys(ids []int64) (map[int64]bool, error) {
38 live := map[int64]bool{} 39 live := map[int64]bool{}
39 if len(ids) == 0 { 40 if len(ids) == 0 {
40 return live, nil 41 return live, nil
41 } 42 }
42 args := make([]any, len(ids)) 43 args := []any{fmtTime(time.Now())}
43 for i, id := range ids { 44 for _, id := range ids {
44 args[i] = id 45 args = append(args, id)
45 } 46 }
46 rows, err := s.DB.Query(`SELECT k.id FROM ssh_keys k JOIN users u ON u.id = k.user_id 47 rows, err := s.DB.Query(`SELECT k.id FROM ssh_keys k JOIN users u ON u.id = k.user_id
47 WHERE u.disabled = 0 AND k.id IN (?`+strings.Repeat(", ?", len(ids)-1)+`)`, args...) 48 WHERE u.disabled = 0 AND (k.expires_at IS NULL OR k.expires_at > ?)
49 AND k.id IN (?`+strings.Repeat(", ?", len(ids)-1)+`)`, args...)
48 if err != nil { 50 if err != nil {
49 return nil, err 51 return nil, err
50 } 52 }
internal/store/users.go +34 −13
@@ -5,6 +5,7 @@ import (
5 "errors" 5 "errors"
6 "fmt" 6 "fmt"
7 "strings" 7 "strings"
8 "time"
8) 9)
9 10
10type User struct { 11type User struct {
@@ -24,8 +25,14 @@ type SSHKey struct {
24 Scope string 25 Scope string
25 Label string // "" when the key was added with no name 26 Label string // "" when the key was added with no name
26 CreatedAt string 27 CreatedAt string
27 LastUsedAt string // "" when the key has never authenticated 28 LastUsedAt string // "" when the key has never authenticated
28 CreatedBy string // name of the API token that added the key; "" for none. ListSSHKeys only. 29 CreatedBy string // name of the API token that added the key; "" for none. ListSSHKeys only.
30 ExpiresAt *time.Time // nil when the key never expires
31}
32
33// Expired reports whether the key has lapsed at now.
34func (k SSHKey) Expired(now time.Time) bool {
35 return k.ExpiresAt != nil && !k.ExpiresAt.After(now)
29} 36}
30 37
31// ErrDuplicateKey carries the exact user-facing message from the spec. It 38// ErrDuplicateKey carries the exact user-facing message from the spec. It
@@ -273,7 +280,8 @@ func (s *Store) UserByID(id int64) (User, error) {
273 280
274// KeyOrigin is how a key came to be. 281// KeyOrigin is how a key came to be.
275type KeyOrigin struct { 282type KeyOrigin struct {
276 CreatedByToken int64 // the API token that added it; 0 for none 283 CreatedByToken int64 // the API token that added it; 0 for none
284 ExpiresAt *time.Time // when it stops authenticating; nil for never
277} 285}
278 286
279// AddSSHKey registers a key and bumps the key epoch in one transaction. 287// AddSSHKey registers a key and bumps the key epoch in one transaction.
@@ -288,9 +296,13 @@ func (s *Store) AddSSHKeyFrom(userID int64, fingerprint, algo string, blob []byt
288 return err 296 return err
289 } 297 }
290 defer tx.Rollback() 298 defer tx.Rollback()
299 var exp any
300 if o.ExpiresAt != nil {
301 exp = fmtTime(*o.ExpiresAt)
302 }
291 if _, err := tx.Exec( 303 if _, err := tx.Exec(
292 "INSERT INTO ssh_keys (user_id, fingerprint, algo, blob, scope, label, created_by_token) VALUES (?, ?, ?, ?, ?, ?, ?)", 304 "INSERT INTO ssh_keys (user_id, fingerprint, algo, blob, scope, label, created_by_token, expires_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
293 userID, fingerprint, algo, blob, scope, label, nullID(o.CreatedByToken)); err != nil { 305 userID, fingerprint, algo, blob, scope, label, nullID(o.CreatedByToken), exp); err != nil {
294 if isUniqueErr(err) { 306 if isUniqueErr(err) {
295 return ErrDuplicateKey 307 return ErrDuplicateKey
296 } 308 }
@@ -343,19 +355,21 @@ func (s *Store) SetSSHKeyLabel(userID int64, fingerprint, label string) error {
343 355
344func (s *Store) SSHKeyByFingerprint(fingerprint string) (SSHKey, error) { 356func (s *Store) SSHKeyByFingerprint(fingerprint string) (SSHKey, error) {
345 var k SSHKey 357 var k SSHKey
358 var exp sql.NullString
346 err := s.DB.QueryRow( 359 err := s.DB.QueryRow(
347 "SELECT id, user_id, fingerprint, algo, blob, scope, label FROM ssh_keys WHERE fingerprint = ?", 360 "SELECT id, user_id, fingerprint, algo, blob, scope, label, expires_at FROM ssh_keys WHERE fingerprint = ?",
348 fingerprint).Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label) 361 fingerprint).Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &exp)
349 if errors.Is(err, sql.ErrNoRows) { 362 if errors.Is(err, sql.ErrNoRows) {
350 return k, ErrNotFound 363 return k, ErrNotFound
351 } 364 }
365 k.ExpiresAt = parseTime(exp)
352 return k, err 366 return k, err
353} 367}
354 368
355func (s *Store) ListSSHKeys(userID int64) ([]SSHKey, error) { 369func (s *Store) ListSSHKeys(userID int64) ([]SSHKey, error) {
356 rows, err := s.DB.Query( 370 rows, err := s.DB.Query(
357 `SELECT k.id, k.user_id, k.fingerprint, k.algo, k.blob, k.scope, k.label, k.created_at, 371 `SELECT k.id, k.user_id, k.fingerprint, k.algo, k.blob, k.scope, k.label, k.created_at,
358 COALESCE(k.last_used_at, ''), COALESCE(t.name, '') 372 COALESCE(k.last_used_at, ''), COALESCE(t.name, ''), k.expires_at
359 FROM ssh_keys k LEFT JOIN api_tokens t ON t.id = k.created_by_token 373 FROM ssh_keys k LEFT JOIN api_tokens t ON t.id = k.created_by_token
360 WHERE k.user_id = ? ORDER BY k.id`, 374 WHERE k.user_id = ? ORDER BY k.id`,
361 userID) 375 userID)
@@ -366,9 +380,11 @@ func (s *Store) ListSSHKeys(userID int64) ([]SSHKey, error) {
366 var keys []SSHKey 380 var keys []SSHKey
367 for rows.Next() { 381 for rows.Next() {
368 var k SSHKey 382 var k SSHKey
369 if err := rows.Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &k.CreatedAt, &k.LastUsedAt, &k.CreatedBy); err != nil { 383 var exp sql.NullString
384 if err := rows.Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &k.CreatedAt, &k.LastUsedAt, &k.CreatedBy, &exp); err != nil {
370 return nil, err 385 return nil, err
371 } 386 }
387 k.ExpiresAt = parseTime(exp)
372 keys = append(keys, k) 388 keys = append(keys, k)
373 } 389 }
374 return keys, rows.Err() 390 return keys, rows.Err()
@@ -523,19 +539,22 @@ func isUniqueErr(err error) bool {
523 539
524func (s *Store) SSHKeyByID(id int64) (SSHKey, error) { 540func (s *Store) SSHKeyByID(id int64) (SSHKey, error) {
525 var k SSHKey 541 var k SSHKey
542 var exp sql.NullString
526 err := s.DB.QueryRow( 543 err := s.DB.QueryRow(
527 "SELECT id, user_id, fingerprint, algo, blob, scope, label FROM ssh_keys WHERE id = ?", 544 "SELECT id, user_id, fingerprint, algo, blob, scope, label, expires_at FROM ssh_keys WHERE id = ?",
528 id).Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label) 545 id).Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &exp)
529 if errors.Is(err, sql.ErrNoRows) { 546 if errors.Is(err, sql.ErrNoRows) {
530 return k, ErrNotFound 547 return k, ErrNotFound
531 } 548 }
549 k.ExpiresAt = parseTime(exp)
532 return k, err 550 return k, err
533} 551}
534 552
535// ListDeployKeys returns the deploy keys bound to a repository. 553// ListDeployKeys returns the deploy keys bound to a repository.
536func (s *Store) ListDeployKeys(repoID int64) ([]SSHKey, error) { 554func (s *Store) ListDeployKeys(repoID int64) ([]SSHKey, error) {
537 rows, err := s.DB.Query( 555 rows, err := s.DB.Query(
538 "SELECT id, user_id, fingerprint, algo, blob, scope, label FROM ssh_keys WHERE scope LIKE 'deploy:' || ? || ':%' ORDER BY id", 556 `SELECT id, user_id, fingerprint, algo, blob, scope, label, COALESCE(last_used_at, ''), expires_at
557 FROM ssh_keys WHERE scope LIKE 'deploy:' || ? || ':%' ORDER BY id`,
539 repoID) 558 repoID)
540 if err != nil { 559 if err != nil {
541 return nil, err 560 return nil, err
@@ -544,9 +563,11 @@ func (s *Store) ListDeployKeys(repoID int64) ([]SSHKey, error) {
544 var keys []SSHKey 563 var keys []SSHKey
545 for rows.Next() { 564 for rows.Next() {
546 var k SSHKey 565 var k SSHKey
547 if err := rows.Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label); err != nil { 566 var exp sql.NullString
567 if err := rows.Scan(&k.ID, &k.UserID, &k.Fingerprint, &k.Algo, &k.Blob, &k.Scope, &k.Label, &k.LastUsedAt, &exp); err != nil {
548 return nil, err 568 return nil, err
549 } 569 }
570 k.ExpiresAt = parseTime(exp)
550 keys = append(keys, k) 571 keys = append(keys, k)
551 } 572 }
552 return keys, rows.Err() 573 return keys, rows.Err()