Files
silo-server/internal/auth/api_key_repository.go

212 lines
6.0 KiB
Go

package auth
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/Silo-Server/silo-server/internal/models"
)
// Sentinel errors for API key operations.
var (
ErrAPIKeyNotFound = errors.New("api key not found")
)
const apiKeyColumns = `id, user_id, label, api_key, rate_tier, created_at, last_used_at`
// APIKeyRepository provides CRUD operations for the api_keys table.
type APIKeyRepository struct {
pool *pgxpool.Pool
}
// NewAPIKeyRepository creates a new APIKeyRepository backed by the given pool.
func NewAPIKeyRepository(pool *pgxpool.Pool) *APIKeyRepository {
return &APIKeyRepository{pool: pool}
}
// scanAPIKey scans a single row into a *models.APIKey.
func scanAPIKey(row pgx.Row) (*models.APIKey, error) {
var k models.APIKey
err := row.Scan(
&k.ID,
&k.UserID,
&k.Label,
&k.Key,
&k.RateTier,
&k.CreatedAt,
&k.LastUsedAt,
)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, ErrAPIKeyNotFound
}
return nil, fmt.Errorf("scanning api key: %w", err)
}
return &k, nil
}
// scanAPIKeys scans multiple rows into a []*models.APIKey slice.
func scanAPIKeys(rows pgx.Rows) ([]*models.APIKey, error) {
var keys []*models.APIKey
for rows.Next() {
var k models.APIKey
err := rows.Scan(
&k.ID,
&k.UserID,
&k.Label,
&k.Key,
&k.RateTier,
&k.CreatedAt,
&k.LastUsedAt,
)
if err != nil {
return nil, fmt.Errorf("scanning api key row: %w", err)
}
keys = append(keys, &k)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("iterating api key rows: %w", err)
}
return keys, nil
}
// generateAPIKey creates a cryptographically random API key with the "sa_" prefix.
func generateAPIKey() (string, error) {
b := make([]byte, 32)
if _, err := rand.Read(b); err != nil {
return "", fmt.Errorf("generating random bytes: %w", err)
}
return "sa_" + hex.EncodeToString(b), nil
}
// Create generates a new API key for the given user and returns the full record.
func (r *APIKeyRepository) Create(ctx context.Context, userID int, label string) (*models.APIKey, error) {
key, err := generateAPIKey()
if err != nil {
return nil, err
}
query := `INSERT INTO api_keys (user_id, label, api_key)
VALUES ($1, $2, $3)
RETURNING ` + apiKeyColumns
row := r.pool.QueryRow(ctx, query, userID, label, key)
return scanAPIKey(row)
}
// ListByUser returns all API keys belonging to the given user, ordered by creation time.
func (r *APIKeyRepository) ListByUser(ctx context.Context, userID int) ([]*models.APIKey, error) {
query := `SELECT ` + apiKeyColumns + ` FROM api_keys WHERE user_id = $1 ORDER BY created_at DESC`
rows, err := r.pool.Query(ctx, query, userID)
if err != nil {
return nil, fmt.Errorf("listing api keys: %w", err)
}
defer rows.Close()
return scanAPIKeys(rows)
}
// GetByKey looks up an API key by its full key string (including "sa_" prefix).
func (r *APIKeyRepository) GetByKey(ctx context.Context, key string) (*models.APIKey, error) {
query := `SELECT ` + apiKeyColumns + ` FROM api_keys WHERE api_key = $1`
return scanAPIKey(r.pool.QueryRow(ctx, query, key))
}
// Delete removes an API key owned by the given user.
func (r *APIKeyRepository) Delete(ctx context.Context, id int64, userID int) error {
tag, err := r.pool.Exec(ctx, "DELETE FROM api_keys WHERE id = $1 AND user_id = $2", id, userID)
if err != nil {
return fmt.Errorf("deleting api key: %w", err)
}
if tag.RowsAffected() == 0 {
return ErrAPIKeyNotFound
}
return nil
}
// DeleteByAdmin removes any API key by its ID regardless of owner.
func (r *APIKeyRepository) DeleteByAdmin(ctx context.Context, id int64) error {
tag, err := r.pool.Exec(ctx, "DELETE FROM api_keys WHERE id = $1", id)
if err != nil {
return fmt.Errorf("deleting api key: %w", err)
}
if tag.RowsAffected() == 0 {
return ErrAPIKeyNotFound
}
return nil
}
// ListByUserAdmin returns all API keys for a specific user (admin view).
func (r *APIKeyRepository) ListByUserAdmin(ctx context.Context, userID int) ([]*models.APIKey, error) {
return r.ListByUser(ctx, userID)
}
// ListAll returns all API keys across all users, ordered by creation time descending.
// Each entry includes the owning user's username.
func (r *APIKeyRepository) ListAll(ctx context.Context) ([]*models.APIKeyWithUser, error) {
query := `SELECT ak.id, ak.user_id, u.username, ak.label, ak.api_key, ak.rate_tier, ak.created_at, ak.last_used_at
FROM api_keys ak
JOIN users u ON u.id = ak.user_id
ORDER BY ak.created_at DESC`
rows, err := r.pool.Query(ctx, query)
if err != nil {
return nil, fmt.Errorf("listing all api keys: %w", err)
}
defer rows.Close()
var keys []*models.APIKeyWithUser
for rows.Next() {
var k models.APIKeyWithUser
err := rows.Scan(
&k.ID,
&k.UserID,
&k.Username,
&k.Label,
&k.Key,
&k.RateTier,
&k.CreatedAt,
&k.LastUsedAt,
)
if err != nil {
return nil, fmt.Errorf("scanning api key with user row: %w", err)
}
keys = append(keys, &k)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("iterating api key with user rows: %w", err)
}
return keys, nil
}
// GetByID looks up an API key by its ID.
func (r *APIKeyRepository) GetByID(ctx context.Context, id int64) (*models.APIKey, error) {
query := `SELECT ` + apiKeyColumns + ` FROM api_keys WHERE id = $1`
return scanAPIKey(r.pool.QueryRow(ctx, query, id))
}
// UpdateTier updates the rate tier for an API key.
func (r *APIKeyRepository) UpdateTier(ctx context.Context, id int64, tier string) error {
result, err := r.pool.Exec(ctx, `UPDATE api_keys SET rate_tier = $1 WHERE id = $2`, tier, id)
if err != nil {
return fmt.Errorf("update api key tier: %w", err)
}
if result.RowsAffected() == 0 {
return ErrAPIKeyNotFound
}
return nil
}
// UpdateLastUsed sets last_used_at to now. Intended to be called asynchronously.
func (r *APIKeyRepository) UpdateLastUsed(ctx context.Context, id int64) error {
_, err := r.pool.Exec(ctx, "UPDATE api_keys SET last_used_at = NOW() WHERE id = $1", id)
if err != nil {
return fmt.Errorf("updating api key last_used_at: %w", err)
}
return nil
}