internal/store/tokens.go

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

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}