internal/store/tokens.go
178 lines · 5141 bytes
1package store
2
3import (
4 "database/sql"
5 "errors"
6 "fmt"
7 "strings"
8 "time"
9)
10
11type APIToken struct {
12 ID int64
13 Name string
14 Scope string
15 CreatedAt string
16 ExpiresAt *time.Time
17 LastUsedAt *time.Time
18 CreatedBy string // name of the token that created this one; "" for none
19}
20
21// nullID stores 0 as NULL.
22func nullID(id int64) any {
23 if id == 0 {
24 return nil
25 }
26 return id
27}
28
29// CreateAPIToken stores a token hash; expires nil means no expiry,
30// createdByToken 0 means it was not created through a token.
31func (s *Store) CreateAPIToken(userID int64, name, tokenHash, scope string, expires *time.Time, createdByToken int64) error {
32 var exp any
33 if expires != nil {
34 exp = fmtTime(*expires)
35 }
36 _, err := s.DB.Exec(
37 "INSERT INTO api_tokens (user_id, name, token_hash, scope, expires_at, created_by_token) VALUES (?, ?, ?, ?, ?, ?)",
38 userID, name, tokenHash, scope, exp, nullID(createdByToken))
39 if isUniqueErr(err) {
40 return fmt.Errorf("you already have a token named %q", name)
41 }
42 return err
43}
44
45// APITokenUser resolves a presented token to its user and the token;
46// expired and unknown tokens fail identically.
47func (s *Store) APITokenUser(tokenHash string) (User, APIToken, error) {
48 var userID int64
49 var t APIToken
50 var exp sql.NullString
51 err := s.DB.QueryRow(`
52 SELECT user_id, id, name, scope, expires_at FROM api_tokens
53 WHERE token_hash = ? AND (expires_at IS NULL OR expires_at > ?)`,
54 tokenHash, fmtTime(time.Now())).Scan(&userID, &t.ID, &t.Name, &t.Scope, &exp)
55 if errors.Is(err, sql.ErrNoRows) {
56 return User{}, APIToken{}, ErrNotFound
57 }
58 if err != nil {
59 return User{}, APIToken{}, err
60 }
61 t.ExpiresAt = parseTime(exp)
62 s.DB.Exec("UPDATE api_tokens SET last_used_at = strftime('%Y-%m-%dT%H:%M:%fZ','now') WHERE token_hash = ?", tokenHash)
63 u, err := s.UserByID(userID)
64 return u, t, err
65}
66
67func (s *Store) ListAPITokens(userID int64) ([]APIToken, error) {
68 rows, err := s.DB.Query(`
69 SELECT t.id, t.name, t.scope, t.created_at, t.expires_at, t.last_used_at, COALESCE(p.name, '')
70 FROM api_tokens t LEFT JOIN api_tokens p ON p.id = t.created_by_token
71 WHERE t.user_id = ? ORDER BY t.name`, userID)
72 if err != nil {
73 return nil, err
74 }
75 defer rows.Close()
76 var out []APIToken
77 for rows.Next() {
78 var t APIToken
79 var exp, used sql.NullString
80 if err := rows.Scan(&t.ID, &t.Name, &t.Scope, &t.CreatedAt, &exp, &used, &t.CreatedBy); err != nil {
81 return nil, err
82 }
83 t.ExpiresAt = parseTime(exp)
84 t.LastUsedAt = parseTime(used)
85 out = append(out, t)
86 }
87 return out, rows.Err()
88}
89
90// Created is what a token made, directly or through tokens it made:
91// token names and SSH key fingerprints.
92type Created struct {
93 Tokens []string `json:"tokens"`
94 Keys []string `json:"keys"`
95}
96
97// chainCTE selects the token named by the first argument and every
98// token created from it, at any depth.
99const chainCTE = `WITH RECURSIVE chain(id) AS (
100 SELECT ? UNION SELECT t.id FROM api_tokens t JOIN chain ON t.created_by_token = chain.id)`
101
102// RevokeAPIToken deletes the user's token by name and returns what it
103// created. withCreated deletes those too; otherwise they stay and lose
104// the link to the revoked token.
105func (s *Store) RevokeAPIToken(userID int64, name string, withCreated bool) (Created, error) {
106 tx, err := s.DB.Begin()
107 if err != nil {
108 return Created{}, err
109 }
110 defer tx.Rollback()
111 var id int64
112 err = tx.QueryRow("SELECT id FROM api_tokens WHERE user_id = ? AND name = ?", userID, name).Scan(&id)
113 if errors.Is(err, sql.ErrNoRows) {
114 return Created{}, ErrNotFound
115 }
116 if err != nil {
117 return Created{}, err
118 }
119 var c Created
120 rows, err := tx.Query(chainCTE+` SELECT name FROM api_tokens WHERE id IN (SELECT id FROM chain) AND id != ? ORDER BY name`, id, id)
121 if err != nil {
122 return Created{}, err
123 }
124 for rows.Next() {
125 var n string
126 if err := rows.Scan(&n); err != nil {
127 rows.Close()
128 return Created{}, err
129 }
130 c.Tokens = append(c.Tokens, n)
131 }
132 rows.Close()
133 var keyIDs []int64
134 rows, err = tx.Query(chainCTE+` SELECT id, fingerprint FROM ssh_keys WHERE created_by_token IN (SELECT id FROM chain) ORDER BY id`, id)
135 if err != nil {
136 return Created{}, err
137 }
138 for rows.Next() {
139 var kid int64
140 var fp string
141 if err := rows.Scan(&kid, &fp); err != nil {
142 rows.Close()
143 return Created{}, err
144 }
145 keyIDs = append(keyIDs, kid)
146 c.Keys = append(c.Keys, fp)
147 }
148 rows.Close()
149
150 if !withCreated {
151 if _, err := tx.Exec("DELETE FROM api_tokens WHERE id = ?", id); err != nil {
152 return Created{}, err
153 }
154 return c, tx.Commit()
155 }
156 if len(keyIDs) > 0 {
157 args := make([]any, len(keyIDs))
158 for i, k := range keyIDs {
159 args[i] = k
160 }
161 if _, err := tx.Exec("DELETE FROM ssh_keys WHERE id IN (?"+strings.Repeat(", ?", len(keyIDs)-1)+")", args...); err != nil {
162 return Created{}, err
163 }
164 if err := bumpKeyEpoch(tx); err != nil {
165 return Created{}, err
166 }
167 }
168 if _, err := tx.Exec(chainCTE+` DELETE FROM api_tokens WHERE id IN (SELECT id FROM chain)`, id); err != nil {
169 return Created{}, err
170 }
171 if err := tx.Commit(); err != nil {
172 return Created{}, err
173 }
174 if len(keyIDs) > 0 {
175 s.announce(Revoked{KeyIDs: keyIDs})
176 }
177 return c, nil
178}