internal/store/tokens.go
191 lines · 5563 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
67// APITokenByID loads a token by id, expired or not.
68func (s *Store) APITokenByID(id int64) (APIToken, error) {
69 var t APIToken
70 var exp sql.NullString
71 err := s.DB.QueryRow("SELECT id, name, scope, created_at, expires_at FROM api_tokens WHERE id = ?", id).
72 Scan(&t.ID, &t.Name, &t.Scope, &t.CreatedAt, &exp)
73 if errors.Is(err, sql.ErrNoRows) {
74 return t, ErrNotFound
75 }
76 t.ExpiresAt = parseTime(exp)
77 return t, err
78}
79
80func (s *Store) ListAPITokens(userID int64) ([]APIToken, error) {
81 rows, err := s.DB.Query(`
82 SELECT t.id, t.name, t.scope, t.created_at, t.expires_at, t.last_used_at, COALESCE(p.name, '')
83 FROM api_tokens t LEFT JOIN api_tokens p ON p.id = t.created_by_token
84 WHERE t.user_id = ? ORDER BY t.name`, userID)
85 if err != nil {
86 return nil, err
87 }
88 defer rows.Close()
89 var out []APIToken
90 for rows.Next() {
91 var t APIToken
92 var exp, used sql.NullString
93 if err := rows.Scan(&t.ID, &t.Name, &t.Scope, &t.CreatedAt, &exp, &used, &t.CreatedBy); err != nil {
94 return nil, err
95 }
96 t.ExpiresAt = parseTime(exp)
97 t.LastUsedAt = parseTime(used)
98 out = append(out, t)
99 }
100 return out, rows.Err()
101}
102
103// Created is what a token made, directly or through tokens it made:
104// token names and SSH key fingerprints.
105type Created struct {
106 Tokens []string `json:"tokens"`
107 Keys []string `json:"keys"`
108}
109
110// chainCTE selects the token named by the first argument and every
111// token created from it, at any depth.
112const chainCTE = `WITH RECURSIVE chain(id) AS (
113 SELECT ? UNION SELECT t.id FROM api_tokens t JOIN chain ON t.created_by_token = chain.id)`
114
115// RevokeAPIToken deletes the user's token by name and returns what it
116// created. withCreated deletes those too; otherwise they stay and lose
117// the link to the revoked token.
118func (s *Store) RevokeAPIToken(userID int64, name string, withCreated bool) (Created, error) {
119 tx, err := s.DB.Begin()
120 if err != nil {
121 return Created{}, err
122 }
123 defer tx.Rollback()
124 var id int64
125 err = tx.QueryRow("SELECT id FROM api_tokens WHERE user_id = ? AND name = ?", userID, name).Scan(&id)
126 if errors.Is(err, sql.ErrNoRows) {
127 return Created{}, ErrNotFound
128 }
129 if err != nil {
130 return Created{}, err
131 }
132 var c Created
133 rows, err := tx.Query(chainCTE+` SELECT name FROM api_tokens WHERE id IN (SELECT id FROM chain) AND id != ? ORDER BY name`, id, id)
134 if err != nil {
135 return Created{}, err
136 }
137 for rows.Next() {
138 var n string
139 if err := rows.Scan(&n); err != nil {
140 rows.Close()
141 return Created{}, err
142 }
143 c.Tokens = append(c.Tokens, n)
144 }
145 rows.Close()
146 var keyIDs []int64
147 rows, err = tx.Query(chainCTE+` SELECT id, fingerprint FROM ssh_keys WHERE created_by_token IN (SELECT id FROM chain) ORDER BY id`, id)
148 if err != nil {
149 return Created{}, err
150 }
151 for rows.Next() {
152 var kid int64
153 var fp string
154 if err := rows.Scan(&kid, &fp); err != nil {
155 rows.Close()
156 return Created{}, err
157 }
158 keyIDs = append(keyIDs, kid)
159 c.Keys = append(c.Keys, fp)
160 }
161 rows.Close()
162
163 if !withCreated {
164 if _, err := tx.Exec("DELETE FROM api_tokens WHERE id = ?", id); err != nil {
165 return Created{}, err
166 }
167 return c, tx.Commit()
168 }
169 if len(keyIDs) > 0 {
170 args := make([]any, len(keyIDs))
171 for i, k := range keyIDs {
172 args[i] = k
173 }
174 if _, err := tx.Exec("DELETE FROM ssh_keys WHERE id IN (?"+strings.Repeat(", ?", len(keyIDs)-1)+")", args...); err != nil {
175 return Created{}, err
176 }
177 if err := bumpKeyEpoch(tx); err != nil {
178 return Created{}, err
179 }
180 }
181 if _, err := tx.Exec(chainCTE+` DELETE FROM api_tokens WHERE id IN (SELECT id FROM chain)`, id); err != nil {
182 return Created{}, err
183 }
184 if err := tx.Commit(); err != nil {
185 return Created{}, err
186 }
187 if len(keyIDs) > 0 {
188 s.announce(Revoked{KeyIDs: keyIDs})
189 }
190 return c, nil
191}