213 lines
6.9 KiB
Go
213 lines
6.9 KiB
Go
package deviceauth
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"database/sql"
|
|
"encoding/hex"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
"unicode"
|
|
)
|
|
|
|
const (
|
|
StatusActive = "ACTIVE"
|
|
StatusRevoked = "REVOKED"
|
|
)
|
|
|
|
var (
|
|
ErrInvalidCredential = errors.New("invalid device credential input")
|
|
ErrCredentialNotFound = errors.New("device credential not found")
|
|
)
|
|
|
|
type Credential struct {
|
|
DeviceID string `json:"device_id"`
|
|
DisplayName string `json:"display_name"`
|
|
Status string `json:"status"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
RevokedAt *time.Time `json:"revoked_at,omitempty"`
|
|
}
|
|
|
|
// IssuedCredential is the only value that can carry the plaintext token. It is returned only
|
|
// after SQLite has committed the hash and is intended for the management CLI's one stdout write.
|
|
type IssuedCredential struct {
|
|
Credential
|
|
Token string `json:"-"`
|
|
}
|
|
|
|
type CredentialStore struct {
|
|
database *sql.DB
|
|
now func() time.Time
|
|
random io.Reader
|
|
randomMu sync.Mutex
|
|
}
|
|
|
|
func NewCredentialStore(database *sql.DB) (*CredentialStore, error) {
|
|
if database == nil {
|
|
return nil, errors.New("device credential database is required")
|
|
}
|
|
if _, err := database.Exec("SELECT device_id FROM device_credentials LIMIT 1"); err != nil {
|
|
return nil, errors.New("device credential migration is not available")
|
|
}
|
|
return &CredentialStore{database: database, now: time.Now, random: rand.Reader}, nil
|
|
}
|
|
|
|
func (store *CredentialStore) Issue(ctx context.Context, displayName string) (IssuedCredential, error) {
|
|
if !ValidDisplayName(displayName) {
|
|
return IssuedCredential{}, ErrInvalidCredential
|
|
}
|
|
randomBytes := make([]byte, 16+32)
|
|
store.randomMu.Lock()
|
|
_, randomErr := io.ReadFull(store.random, randomBytes)
|
|
store.randomMu.Unlock()
|
|
if randomErr != nil {
|
|
return IssuedCredential{}, fmt.Errorf("generate device credential: %w", randomErr)
|
|
}
|
|
deviceID := formatUUIDv4(randomBytes[:16])
|
|
token := hex.EncodeToString(randomBytes[16:])
|
|
tokenHash := sha256.Sum256(randomBytes[16:])
|
|
createdAt := store.now().UTC()
|
|
if createdAt.IsZero() {
|
|
return IssuedCredential{}, errors.New("device credential clock is invalid")
|
|
}
|
|
_, err := store.database.ExecContext(ctx, `INSERT INTO device_credentials
|
|
(device_id, display_name, token_sha256, status, created_at, revoked_at)
|
|
VALUES (?, ?, ?, ?, ?, NULL)`,
|
|
deviceID, displayName, tokenHash[:], StatusActive, createdAt.Format(time.RFC3339Nano))
|
|
if err != nil {
|
|
return IssuedCredential{}, fmt.Errorf("persist device credential: %w", err)
|
|
}
|
|
return IssuedCredential{Credential: Credential{
|
|
DeviceID: deviceID, DisplayName: displayName, Status: StatusActive, CreatedAt: createdAt,
|
|
}, Token: token}, nil
|
|
}
|
|
|
|
func (store *CredentialStore) List(ctx context.Context) ([]Credential, error) {
|
|
rows, err := store.database.QueryContext(ctx, `SELECT device_id, display_name, status, created_at, revoked_at
|
|
FROM device_credentials ORDER BY created_at, device_id`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("list device credentials: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
credentials := make([]Credential, 0)
|
|
for rows.Next() {
|
|
credential, err := scanCredential(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
credentials = append(credentials, credential)
|
|
}
|
|
if err := rows.Err(); err != nil {
|
|
return nil, fmt.Errorf("list device credentials: %w", err)
|
|
}
|
|
return credentials, nil
|
|
}
|
|
|
|
func (store *CredentialStore) Revoke(ctx context.Context, deviceID string) (Credential, bool, error) {
|
|
if !ValidDeviceID(deviceID) {
|
|
return Credential{}, false, ErrInvalidCredential
|
|
}
|
|
revokedAt := store.now().UTC()
|
|
if revokedAt.IsZero() {
|
|
return Credential{}, false, errors.New("device credential clock is invalid")
|
|
}
|
|
transaction, err := store.database.BeginTx(ctx, nil)
|
|
if err != nil {
|
|
return Credential{}, false, fmt.Errorf("begin device credential revocation: %w", err)
|
|
}
|
|
defer transaction.Rollback()
|
|
result, err := transaction.ExecContext(ctx, `UPDATE device_credentials
|
|
SET status = ?, revoked_at = ? WHERE device_id = ? AND status = ?`,
|
|
StatusRevoked, revokedAt.Format(time.RFC3339Nano), deviceID, StatusActive)
|
|
if err != nil {
|
|
return Credential{}, false, fmt.Errorf("revoke device credential: %w", err)
|
|
}
|
|
changedRows, err := result.RowsAffected()
|
|
if err != nil {
|
|
return Credential{}, false, fmt.Errorf("inspect device credential revocation: %w", err)
|
|
}
|
|
credential, err := scanCredential(transaction.QueryRowContext(ctx, `SELECT device_id, display_name, status, created_at, revoked_at
|
|
FROM device_credentials WHERE device_id = ?`, deviceID))
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return Credential{}, false, ErrCredentialNotFound
|
|
}
|
|
if err != nil {
|
|
return Credential{}, false, err
|
|
}
|
|
if err := transaction.Commit(); err != nil {
|
|
return Credential{}, false, fmt.Errorf("commit device credential revocation: %w", err)
|
|
}
|
|
return credential, changedRows == 1, nil
|
|
}
|
|
|
|
type rowScanner interface {
|
|
Scan(...any) error
|
|
}
|
|
|
|
func scanCredential(row rowScanner) (Credential, error) {
|
|
var credential Credential
|
|
var created string
|
|
var revoked sql.NullString
|
|
if err := row.Scan(&credential.DeviceID, &credential.DisplayName, &credential.Status, &created, &revoked); err != nil {
|
|
return Credential{}, err
|
|
}
|
|
if !ValidDeviceID(credential.DeviceID) || !ValidDisplayName(credential.DisplayName) || (credential.Status != StatusActive && credential.Status != StatusRevoked) {
|
|
return Credential{}, errors.New("stored device credential metadata is invalid")
|
|
}
|
|
createdAt, err := parseStoredTime(created)
|
|
if err != nil {
|
|
return Credential{}, err
|
|
}
|
|
credential.CreatedAt = createdAt
|
|
if revoked.Valid {
|
|
revokedAt, err := parseStoredTime(revoked.String)
|
|
if err != nil {
|
|
return Credential{}, err
|
|
}
|
|
if revokedAt.Before(createdAt) {
|
|
return Credential{}, errors.New("stored device credential status is invalid")
|
|
}
|
|
credential.RevokedAt = &revokedAt
|
|
}
|
|
if (credential.Status == StatusActive) != (credential.RevokedAt == nil) {
|
|
return Credential{}, errors.New("stored device credential status is invalid")
|
|
}
|
|
return credential, nil
|
|
}
|
|
|
|
func ValidDisplayName(value string) bool {
|
|
if value == "" || len([]rune(value)) > 128 || strings.TrimSpace(value) != value {
|
|
return false
|
|
}
|
|
for _, character := range value {
|
|
if unicode.IsControl(character) {
|
|
return false
|
|
}
|
|
}
|
|
return true
|
|
}
|
|
|
|
func parseStoredTime(value string) (time.Time, error) {
|
|
if strings.TrimSpace(value) != value || !strings.HasSuffix(value, "Z") {
|
|
return time.Time{}, errors.New("stored device credential time is invalid")
|
|
}
|
|
parsed, err := time.Parse(time.RFC3339Nano, value)
|
|
if err != nil || parsed.Location() != time.UTC {
|
|
return time.Time{}, errors.New("stored device credential time is invalid")
|
|
}
|
|
return parsed, nil
|
|
}
|
|
|
|
func formatUUIDv4(bytes []byte) string {
|
|
copyBytes := append([]byte(nil), bytes...)
|
|
copyBytes[6] = (copyBytes[6] & 0x0f) | 0x40
|
|
copyBytes[8] = (copyBytes[8] & 0x3f) | 0x80
|
|
encoded := hex.EncodeToString(copyBytes)
|
|
return encoded[:8] + "-" + encoded[8:12] + "-" + encoded[12:16] + "-" + encoded[16:20] + "-" + encoded[20:]
|
|
}
|