package apikey import ( "context" "encoding/json" "errors" "fmt" "time" platformid "aigateway.local/core/internal/platform/id" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) type Repository struct{ pool *pgxpool.Pool } func NewRepository(pool *pgxpool.Pool) *Repository { return &Repository{pool: pool} } func (r *Repository) List(ctx context.Context) ([]Record, error) { if r.pool == nil { return nil, ErrStore } rows, err := r.pool.Query(ctx, ` SELECT id::text, tenant_id::text, name, key_prefix, scopes, enabled, requests_per_minute, monthly_request_quota, monthly_token_quota, expires_at, last_used_at, created_at FROM gateway.api_keys ORDER BY created_at DESC`) if err != nil { return nil, fmt.Errorf("%w: %v", ErrStore, err) } defer rows.Close() var records []Record for rows.Next() { var record Record if err := rows.Scan(&record.ID, &record.TenantID, &record.Name, &record.KeyPrefix, &record.Scopes, &record.Enabled, &record.RequestsPerMinute, &record.MonthlyRequestQuota, &record.MonthlyTokenQuota, &record.ExpiresAt, &record.LastUsedAt, &record.CreatedAt); err != nil { return nil, fmt.Errorf("%w: %v", ErrStore, err) } records = append(records, record) } return records, rows.Err() } func (r *Repository) Create(ctx context.Context, name string, scopes []string, requestsPerMinute int, monthlyRequestQuota, monthlyTokenQuota int64, expiresAt *time.Time, actorID string) (Record, string, error) { if r.pool == nil { return Record{}, "", ErrStore } secret, prefix, hash, err := Generate() if err != nil { return Record{}, "", err } id, err := platformid.NewUUID() if err != nil { return Record{}, "", err } eventID, err := platformid.NewUUID() if err != nil { return Record{}, "", err } tx, err := r.pool.Begin(ctx) if err != nil { return Record{}, "", fmt.Errorf("%w: %v", ErrStore, err) } defer func() { _ = tx.Rollback(ctx) }() record := Record{ID: id, Name: name, KeyPrefix: prefix, KeyHash: hash, Scopes: scopes, Enabled: true, RequestsPerMinute: requestsPerMinute, MonthlyRequestQuota: monthlyRequestQuota, MonthlyTokenQuota: monthlyTokenQuota, ExpiresAt: expiresAt} err = tx.QueryRow(ctx, ` INSERT INTO gateway.api_keys (id, name, key_prefix, key_hash, scopes, requests_per_minute, monthly_request_quota, monthly_token_quota, expires_at, created_by) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, nullif($10,'')::uuid) RETURNING created_at`, id, name, prefix, hash, scopes, requestsPerMinute, monthlyRequestQuota, monthlyTokenQuota, expiresAt, actorID).Scan(&record.CreatedAt) if err != nil { return Record{}, "", fmt.Errorf("%w: %v", ErrStore, err) } payload, _ := json.Marshal(map[string]any{"api_key_id": id, "name": name, "key_prefix": prefix, "requests_per_minute": requestsPerMinute, "monthly_request_quota": monthlyRequestQuota, "monthly_token_quota": monthlyTokenQuota}) if _, err := tx.Exec(ctx, ` INSERT INTO gateway.outbox_events (event_id, event_type, event_version, aggregate_type, aggregate_id, payload) VALUES ($1, 'api_key.created', 1, 'api_key', $2, $3)`, eventID, id, payload); err != nil { return Record{}, "", fmt.Errorf("%w: %v", ErrStore, err) } if err := tx.Commit(ctx); err != nil { return Record{}, "", fmt.Errorf("%w: %v", ErrStore, err) } return record, secret, nil } func (r *Repository) Validate(ctx context.Context, hash []byte) (Principal, error) { if r.pool == nil { return Principal{}, ErrStore } var principal Principal err := r.pool.QueryRow(ctx, ` UPDATE gateway.api_keys SET last_used_at = CASE WHEN last_used_at IS NULL OR last_used_at < clock_timestamp() - interval '5 minutes' THEN clock_timestamp() ELSE last_used_at END WHERE key_hash = $1 AND enabled AND (expires_at IS NULL OR expires_at > clock_timestamp()) RETURNING id::text, tenant_id::text, scopes, requests_per_minute, monthly_request_quota, monthly_token_quota`, hash).Scan(&principal.APIKeyID, &principal.TenantID, &principal.Scopes, &principal.RequestsPerMinute, &principal.MonthlyRequestQuota, &principal.MonthlyTokenQuota) if errors.Is(err, pgx.ErrNoRows) { return Principal{}, ErrInvalid } if err != nil { return Principal{}, fmt.Errorf("%w: %v", ErrStore, err) } return principal, nil } func (r *Repository) Revoke(ctx context.Context, id, actorID string) ([]byte, error) { if r.pool == nil { return nil, ErrStore } eventID, err := platformid.NewUUID() if err != nil { return nil, err } tx, err := r.pool.Begin(ctx) if err != nil { return nil, fmt.Errorf("%w: %v", ErrStore, err) } defer func() { _ = tx.Rollback(ctx) }() var hash []byte err = tx.QueryRow(ctx, ` UPDATE gateway.api_keys SET enabled = false, updated_at = clock_timestamp() WHERE id = $1 AND enabled RETURNING key_hash`, id).Scan(&hash) if errors.Is(err, pgx.ErrNoRows) { return nil, ErrInvalid } if err != nil { return nil, fmt.Errorf("%w: %v", ErrStore, err) } payload, _ := json.Marshal(map[string]any{"api_key_id": id, "actor_id": actorID}) if _, err := tx.Exec(ctx, ` INSERT INTO gateway.outbox_events (event_id, event_type, event_version, aggregate_type, aggregate_id, payload) VALUES ($1, 'api_key.revoked', 1, 'api_key', $2, $3)`, eventID, id, payload); err != nil { return nil, fmt.Errorf("%w: %v", ErrStore, err) } if err := tx.Commit(ctx); err != nil { return nil, fmt.Errorf("%w: %v", ErrStore, err) } return hash, nil } func (r *Repository) UpdateLimits(ctx context.Context, id string, requestsPerMinute int, monthlyRequestQuota, monthlyTokenQuota int64, actorID string) (Record, []byte, error) { if r.pool == nil { return Record{}, nil, ErrStore } eventID, err := platformid.NewUUID() if err != nil { return Record{}, nil, err } tx, err := r.pool.Begin(ctx) if err != nil { return Record{}, nil, fmt.Errorf("%w: %v", ErrStore, err) } defer func() { _ = tx.Rollback(ctx) }() var record Record err = tx.QueryRow(ctx, ` UPDATE gateway.api_keys SET requests_per_minute = $2, monthly_request_quota = $3, monthly_token_quota = $4, updated_at = clock_timestamp() WHERE id = $1 AND enabled RETURNING id::text, tenant_id::text, name, key_prefix, key_hash, scopes, enabled, requests_per_minute, monthly_request_quota, monthly_token_quota, expires_at, last_used_at, created_at`, id, requestsPerMinute, monthlyRequestQuota, monthlyTokenQuota, ).Scan(&record.ID, &record.TenantID, &record.Name, &record.KeyPrefix, &record.KeyHash, &record.Scopes, &record.Enabled, &record.RequestsPerMinute, &record.MonthlyRequestQuota, &record.MonthlyTokenQuota, &record.ExpiresAt, &record.LastUsedAt, &record.CreatedAt) if errors.Is(err, pgx.ErrNoRows) { return Record{}, nil, ErrInvalid } if err != nil { return Record{}, nil, fmt.Errorf("%w: %v", ErrStore, err) } payload, _ := json.Marshal(map[string]any{ "api_key_id": id, "actor_id": actorID, "requests_per_minute": requestsPerMinute, "monthly_request_quota": monthlyRequestQuota, "monthly_token_quota": monthlyTokenQuota, }) if _, err := tx.Exec(ctx, ` INSERT INTO gateway.outbox_events (event_id, event_type, event_version, aggregate_type, aggregate_id, payload) VALUES ($1, 'api_key.limits_updated', 1, 'api_key', $2, $3)`, eventID, id, payload); err != nil { return Record{}, nil, fmt.Errorf("%w: %v", ErrStore, err) } if err := tx.Commit(ctx); err != nil { return Record{}, nil, fmt.Errorf("%w: %v", ErrStore, err) } return record, record.KeyHash, nil }