diff --git a/CHANGELOG.md b/CHANGELOG.md index 2ba3f79..0df852f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -15,6 +15,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Deepened audit-log signature verification behind a single shared core used by both single-log and batch verification. Batch `verify-audit-logs` now reports KEK-missing logs in their own `kek_missing_count`/`kek_missing_logs` bucket instead of lumping them into `invalid_count`, so batch agrees with single-log verification (`ErrKekNotFoundForLog`). - Internal refactors with no API change: consolidated the six retention-sweep CLI commands (`purge-secrets`, `purge-transit-keys`, `purge-tokenization-keys`, `clean-expired-tokens`, `clean-audit-logs`, `purge-auth-tokens`) behind a single `RunRetentionSweep` module, and relocated the interactive policy-prompt helpers from `internal/ui` into the CLI commands package (#141). +- Collapsed the hand-rolled latest-version cursor pagination and dry-run hard-delete SQL shared by the `secrets`, `transit`, and `tokenization` key repositories behind two helpers in `internal/database` (`ListLatestCursor`, `HardDeleteOlderThan`), and removed the unused offset-pagination `List` on the secrets repository. - Folded the master-key/KMS lifecycle into the `keyring` deep module. KEK loading (`bootstrapWith`) and the KMS decrypt path are now unit-testable without a database or live KMS; `KMSKeeper` gained an explicit `Encrypt` operation so the create/rotate-master-key CLI commands no longer type-assert the concrete keeper; added an in-memory `keyring.FakeKMSService` test double and removed the unused `internal/tokenization/testing` helper. ### Fixed diff --git a/internal/database/versioned.go b/internal/database/versioned.go new file mode 100644 index 0000000..ad76015 --- /dev/null +++ b/internal/database/versioned.go @@ -0,0 +1,111 @@ +package database + +import ( + "context" + "database/sql" + "time" +) + +// ListLatestCursor returns, for each logical key, its latest non-deleted row, +// ordered by key ascending with cursor-based pagination. The two-branch +// latest-version JOIN, row iteration, and empty-slice normalisation live here +// so feature repositories don't hand-copy them. +// +// selectCols must be the SELECT column list for the OUTER query, each column +// prefixed with the alias t (e.g. "t.id, t.name, ..."). table and keyCol are +// the physical table name and its logical key column. scan maps one row to T. +// +// A nil afterKey returns the first page; otherwise rows with key > afterKey. +// Returns a non-nil empty slice when nothing matches. +func ListLatestCursor[T any]( + ctx context.Context, + q Querier, + table, keyCol, selectCols string, + afterKey *string, + limit int, + scan func(*sql.Rows) (T, error), +) ([]T, error) { + var query string + var args []any + + if afterKey == nil { + //nolint:gosec // table/keyCol/selectCols are compile-time constants from callers, not user input. + query = ` + SELECT ` + selectCols + ` + FROM ` + table + ` t + INNER JOIN ( + SELECT ` + keyCol + `, MAX(version) as max_version + FROM ` + table + ` + WHERE deleted_at IS NULL + GROUP BY ` + keyCol + ` + ORDER BY ` + keyCol + ` ASC + LIMIT $1 + ) latest ON t.` + keyCol + ` = latest.` + keyCol + ` AND t.version = latest.max_version + ORDER BY t.` + keyCol + ` ASC` + args = []any{limit} + } else { + //nolint:gosec // same as above: constant fragments, not user input. + query = ` + SELECT ` + selectCols + ` + FROM ` + table + ` t + INNER JOIN ( + SELECT ` + keyCol + `, MAX(version) as max_version + FROM ` + table + ` + WHERE deleted_at IS NULL AND ` + keyCol + ` > $1 + GROUP BY ` + keyCol + ` + ORDER BY ` + keyCol + ` ASC + LIMIT $2 + ) latest ON t.` + keyCol + ` = latest.` + keyCol + ` AND t.version = latest.max_version + ORDER BY t.` + keyCol + ` ASC` + args = []any{*afterKey, limit} + } + + rows, err := q.QueryContext(ctx, query, args...) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + out := make([]T, 0, limit) + for rows.Next() { + v, err := scan(rows) + if err != nil { + return nil, err + } + out = append(out, v) + } + if err := rows.Err(); err != nil { + return nil, err + } + return out, nil +} + +// HardDeleteOlderThan permanently removes soft-deleted rows older than the +// cutoff. In dryRun mode it returns the count without deleting. +func HardDeleteOlderThan( + ctx context.Context, + q Querier, + table string, + olderThan time.Time, + dryRun bool, +) (int64, error) { + if dryRun { + query := `SELECT COUNT(*) FROM ` + table + ` WHERE deleted_at IS NOT NULL AND deleted_at < $1` + var count int64 + if err := q.QueryRowContext(ctx, query, olderThan).Scan(&count); err != nil { + return 0, err + } + return count, nil + } + + query := `DELETE FROM ` + table + ` WHERE deleted_at IS NOT NULL AND deleted_at < $1` + result, err := q.ExecContext(ctx, query, olderThan) + if err != nil { + return 0, err + } + count, err := result.RowsAffected() + if err != nil { + return 0, err + } + return count, nil +} diff --git a/internal/secrets/repository/secret_repository.go b/internal/secrets/repository/secret_repository.go index 9cffd8c..dbdeb67 100644 --- a/internal/secrets/repository/secret_repository.go +++ b/internal/secrets/repository/secret_repository.go @@ -134,143 +134,41 @@ func (p *SecretRepository) Delete(ctx context.Context, path string) error { return nil } -// List retrieves secrets ordered by path ascending with pagination. -func (p *SecretRepository) List( - ctx context.Context, - offset, limit int, -) ([]*secretsDomain.Secret, error) { - querier := database.GetTx(ctx, p.db) - - query := ` - SELECT s.id, s.path, s.version, s.dek_id, s.ciphertext, s.nonce, s.created_at, s.deleted_at - FROM secrets s - INNER JOIN ( - SELECT path, MAX(version) as max_version - FROM secrets - WHERE deleted_at IS NULL - GROUP BY path - ORDER BY path ASC - LIMIT $1 OFFSET $2 - ) latest ON s.path = latest.path AND s.version = latest.max_version - ORDER BY s.path ASC` - - rows, err := querier.QueryContext(ctx, query, limit, offset) - if err != nil { - return nil, apperrors.Wrap(err, "failed to list secrets") - } - defer func() { - _ = rows.Close() - }() - - var secrets []*secretsDomain.Secret - for rows.Next() { - var secret secretsDomain.Secret - err := rows.Scan( - &secret.ID, - &secret.Path, - &secret.Version, - &secret.DekID, - &secret.Ciphertext, - &secret.Nonce, - &secret.CreatedAt, - &secret.DeletedAt, - ) - if err != nil { - return nil, apperrors.Wrap(err, "failed to scan secret") - } - secrets = append(secrets, &secret) - } - - if err := rows.Err(); err != nil { - return nil, apperrors.Wrap(err, "error iterating secrets") - } - - if secrets == nil { - secrets = make([]*secretsDomain.Secret, 0) - } - - return secrets, nil -} - // ListCursor retrieves secrets ordered by path ascending using cursor-based pagination. func (p *SecretRepository) ListCursor( ctx context.Context, afterPath *string, limit int, ) ([]*secretsDomain.Secret, error) { - querier := database.GetTx(ctx, p.db) - - var query string - var args []interface{} - - if afterPath == nil { - // First page: no cursor - query = ` - SELECT s.id, s.path, s.version, s.dek_id, s.ciphertext, s.nonce, s.created_at, s.deleted_at - FROM secrets s - INNER JOIN ( - SELECT path, MAX(version) as max_version - FROM secrets - WHERE deleted_at IS NULL - GROUP BY path - ORDER BY path ASC - LIMIT $1 - ) latest ON s.path = latest.path AND s.version = latest.max_version - ORDER BY s.path ASC` - args = []interface{}{limit} - } else { - // Subsequent pages: use cursor (path > afterPath) - query = ` - SELECT s.id, s.path, s.version, s.dek_id, s.ciphertext, s.nonce, s.created_at, s.deleted_at - FROM secrets s - INNER JOIN ( - SELECT path, MAX(version) as max_version - FROM secrets - WHERE deleted_at IS NULL AND path > $1 - GROUP BY path - ORDER BY path ASC - LIMIT $2 - ) latest ON s.path = latest.path AND s.version = latest.max_version - ORDER BY s.path ASC` - args = []interface{}{*afterPath, limit} - } - - rows, err := querier.QueryContext(ctx, query, args...) + records, err := database.ListLatestCursor( + ctx, + database.GetTx(ctx, p.db), + "secrets", + "path", + "t.id, t.path, t.version, t.dek_id, t.ciphertext, t.nonce, t.created_at, t.deleted_at", + afterPath, + limit, + func(rows *sql.Rows) (*secretsDomain.Secret, error) { + var secret secretsDomain.Secret + if err := rows.Scan( + &secret.ID, + &secret.Path, + &secret.Version, + &secret.DekID, + &secret.Ciphertext, + &secret.Nonce, + &secret.CreatedAt, + &secret.DeletedAt, + ); err != nil { + return nil, err + } + return &secret, nil + }, + ) if err != nil { return nil, apperrors.Wrap(err, "failed to list secrets with cursor") } - defer func() { - _ = rows.Close() - }() - - var secrets []*secretsDomain.Secret - for rows.Next() { - var secret secretsDomain.Secret - err := rows.Scan( - &secret.ID, - &secret.Path, - &secret.Version, - &secret.DekID, - &secret.Ciphertext, - &secret.Nonce, - &secret.CreatedAt, - &secret.DeletedAt, - ) - if err != nil { - return nil, apperrors.Wrap(err, "failed to scan secret") - } - secrets = append(secrets, &secret) - } - - if err := rows.Err(); err != nil { - return nil, apperrors.Wrap(err, "error iterating secrets") - } - - if secrets == nil { - secrets = make([]*secretsDomain.Secret, 0) - } - - return secrets, nil + return records, nil } // HardDelete permanently removes soft-deleted secrets older than the specified time. @@ -282,29 +180,10 @@ func (p *SecretRepository) HardDelete( olderThan time.Time, dryRun bool, ) (int64, error) { - querier := database.GetTx(ctx, p.db) - - if dryRun { - query := `SELECT COUNT(*) FROM secrets WHERE deleted_at IS NOT NULL AND deleted_at < $1` - var count int64 - err := querier.QueryRowContext(ctx, query, olderThan).Scan(&count) - if err != nil { - return 0, apperrors.Wrap(err, "failed to count secrets for deletion") - } - return count, nil - } - - query := `DELETE FROM secrets WHERE deleted_at IS NOT NULL AND deleted_at < $1` - result, err := querier.ExecContext(ctx, query, olderThan) + count, err := database.HardDeleteOlderThan(ctx, database.GetTx(ctx, p.db), "secrets", olderThan, dryRun) if err != nil { return 0, apperrors.Wrap(err, "failed to hard delete secrets") } - - count, err := result.RowsAffected() - if err != nil { - return 0, apperrors.Wrap(err, "failed to get affected rows count") - } - return count, nil } diff --git a/internal/tokenization/repository/repository.go b/internal/tokenization/repository/repository.go index 03a1154..886f279 100644 --- a/internal/tokenization/repository/repository.go +++ b/internal/tokenization/repository/repository.go @@ -183,84 +183,38 @@ func (p *TokenizationKeyRepository) ListCursor( afterName *string, limit int, ) ([]*tokenizationDomain.TokenizationKey, error) { - querier := database.GetTx(ctx, p.db) - - var query string - var args []interface{} - - if afterName == nil { - // First page: no cursor - query = ` - SELECT tk.id, tk.name, tk.version, tk.format_type, tk.is_deterministic, tk.salt, tk.dek_id, tk.created_at, tk.deleted_at - FROM tokenization_keys tk - INNER JOIN ( - SELECT name, MAX(version) as max_version - FROM tokenization_keys - WHERE deleted_at IS NULL - GROUP BY name - ORDER BY name ASC - LIMIT $1 - ) latest ON tk.name = latest.name AND tk.version = latest.max_version - ORDER BY tk.name ASC` - args = []interface{}{limit} - } else { - // Subsequent pages: use cursor (name > afterName) - query = ` - SELECT tk.id, tk.name, tk.version, tk.format_type, tk.is_deterministic, tk.salt, tk.dek_id, tk.created_at, tk.deleted_at - FROM tokenization_keys tk - INNER JOIN ( - SELECT name, MAX(version) as max_version - FROM tokenization_keys - WHERE deleted_at IS NULL AND name > $1 - GROUP BY name - ORDER BY name ASC - LIMIT $2 - ) latest ON tk.name = latest.name AND tk.version = latest.max_version - ORDER BY tk.name ASC` - args = []interface{}{*afterName, limit} - } - - rows, err := querier.QueryContext(ctx, query, args...) + records, err := database.ListLatestCursor( + ctx, + database.GetTx(ctx, p.db), + "tokenization_keys", + "name", + "t.id, t.name, t.version, t.format_type, t.is_deterministic, t.salt, t.dek_id, t.created_at, t.deleted_at", + afterName, + limit, + func(rows *sql.Rows) (*tokenizationDomain.TokenizationKey, error) { + var key tokenizationDomain.TokenizationKey + var formatType string + if err := rows.Scan( + &key.ID, + &key.Name, + &key.Version, + &formatType, + &key.IsDeterministic, + &key.Salt, + &key.DekID, + &key.CreatedAt, + &key.DeletedAt, + ); err != nil { + return nil, err + } + key.FormatType = tokenizationDomain.FormatType(formatType) + return &key, nil + }, + ) if err != nil { return nil, apperrors.Wrap(err, "failed to list tokenization keys with cursor") } - defer func() { - _ = rows.Close() - }() - - var keys []*tokenizationDomain.TokenizationKey - for rows.Next() { - var key tokenizationDomain.TokenizationKey - var formatType string - - err := rows.Scan( - &key.ID, - &key.Name, - &key.Version, - &formatType, - &key.IsDeterministic, - &key.Salt, - &key.DekID, - &key.CreatedAt, - &key.DeletedAt, - ) - if err != nil { - return nil, apperrors.Wrap(err, "failed to scan tokenization key") - } - - key.FormatType = tokenizationDomain.FormatType(formatType) - keys = append(keys, &key) - } - - if err := rows.Err(); err != nil { - return nil, apperrors.Wrap(err, "error iterating tokenization keys") - } - - if keys == nil { - keys = make([]*tokenizationDomain.TokenizationKey, 0) - } - - return keys, nil + return records, nil } // HardDelete permanently removes soft-deleted tokenization keys and their associated tokens. diff --git a/internal/transit/repository/transit_key_repository.go b/internal/transit/repository/transit_key_repository.go index 72ec216..7a4f39e 100644 --- a/internal/transit/repository/transit_key_repository.go +++ b/internal/transit/repository/transit_key_repository.go @@ -177,77 +177,33 @@ func (p *TransitKeyRepository) ListCursor( afterName *string, limit int, ) ([]*transitDomain.TransitKey, error) { - querier := database.GetTx(ctx, p.db) - - var query string - var args []interface{} - - if afterName == nil { - // First page: no cursor - query = ` - SELECT tk.id, tk.name, tk.version, tk.dek_id, tk.created_at, tk.deleted_at - FROM transit_keys tk - INNER JOIN ( - SELECT name, MAX(version) as max_version - FROM transit_keys - WHERE deleted_at IS NULL - GROUP BY name - ORDER BY name ASC - LIMIT $1 - ) latest ON tk.name = latest.name AND tk.version = latest.max_version - ORDER BY tk.name ASC` - args = []interface{}{limit} - } else { - // Subsequent pages: use cursor (name > afterName) - query = ` - SELECT tk.id, tk.name, tk.version, tk.dek_id, tk.created_at, tk.deleted_at - FROM transit_keys tk - INNER JOIN ( - SELECT name, MAX(version) as max_version - FROM transit_keys - WHERE deleted_at IS NULL AND name > $1 - GROUP BY name - ORDER BY name ASC - LIMIT $2 - ) latest ON tk.name = latest.name AND tk.version = latest.max_version - ORDER BY tk.name ASC` - args = []interface{}{*afterName, limit} - } - - rows, err := querier.QueryContext(ctx, query, args...) + records, err := database.ListLatestCursor( + ctx, + database.GetTx(ctx, p.db), + "transit_keys", + "name", + "t.id, t.name, t.version, t.dek_id, t.created_at, t.deleted_at", + afterName, + limit, + func(rows *sql.Rows) (*transitDomain.TransitKey, error) { + var key transitDomain.TransitKey + if err := rows.Scan( + &key.ID, + &key.Name, + &key.Version, + &key.DekID, + &key.CreatedAt, + &key.DeletedAt, + ); err != nil { + return nil, err + } + return &key, nil + }, + ) if err != nil { return nil, apperrors.Wrap(err, "failed to list transit keys with cursor") } - defer func() { - _ = rows.Close() - }() - - var transitKeys []*transitDomain.TransitKey - for rows.Next() { - var transitKey transitDomain.TransitKey - err := rows.Scan( - &transitKey.ID, - &transitKey.Name, - &transitKey.Version, - &transitKey.DekID, - &transitKey.CreatedAt, - &transitKey.DeletedAt, - ) - if err != nil { - return nil, apperrors.Wrap(err, "failed to scan transit key") - } - transitKeys = append(transitKeys, &transitKey) - } - - if err := rows.Err(); err != nil { - return nil, apperrors.Wrap(err, "error iterating transit keys") - } - - if transitKeys == nil { - transitKeys = make([]*transitDomain.TransitKey, 0) - } - - return transitKeys, nil + return records, nil } // HardDelete permanently removes soft-deleted transit keys older than the specified time. @@ -256,29 +212,16 @@ func (p *TransitKeyRepository) HardDelete( olderThan time.Time, dryRun bool, ) (int64, error) { - querier := database.GetTx(ctx, p.db) - - if dryRun { - query := `SELECT COUNT(*) FROM transit_keys WHERE deleted_at IS NOT NULL AND deleted_at < $1` - var count int64 - err := querier.QueryRowContext(ctx, query, olderThan).Scan(&count) - if err != nil { - return 0, apperrors.Wrap(err, "failed to count transit keys for hard delete") - } - return count, nil - } - - query := `DELETE FROM transit_keys WHERE deleted_at IS NOT NULL AND deleted_at < $1` - result, err := querier.ExecContext(ctx, query, olderThan) + count, err := database.HardDeleteOlderThan( + ctx, + database.GetTx(ctx, p.db), + "transit_keys", + olderThan, + dryRun, + ) if err != nil { return 0, apperrors.Wrap(err, "failed to hard delete transit keys") } - - count, err := result.RowsAffected() - if err != nil { - return 0, apperrors.Wrap(err, "failed to get rows affected for hard delete") - } - return count, nil }