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:] }