internal/store/tokens.go

e6cd75b5f28bacf51620bb531320c30fd4e66bfd
gitbay/internal/store/tokens.go history · blame · raw

191 lines · 5563 bytes

9 symbols in this file
  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}