package store import ( "context" "database/sql" "errors" "fmt" "gitea.stevedudenhoeffer.com/steve/pansy/internal/domain" ) // changeSetColumns lists change_sets columns in the order scanChangeSet expects. const changeSetColumns = `id, garden_id, actor_id, source, summary, agent_run_id, reverts_id, created_at` func scanChangeSet(s scanner) (*domain.ChangeSet, error) { var cs domain.ChangeSet if err := s.Scan( &cs.ID, &cs.GardenID, &cs.ActorID, &cs.Source, &cs.Summary, &cs.AgentRunID, &cs.RevertsID, &cs.CreatedAt, ); err != nil { return nil, err } return &cs, nil } // revisionColumns lists revisions columns in the order scanRevision expects. const revisionColumns = `id, change_set_id, seq, entity_type, entity_id, op, before, after` func scanRevision(s scanner) (*domain.Revision, error) { var r domain.Revision if err := s.Scan( &r.ID, &r.ChangeSetID, &r.Seq, &r.EntityType, &r.EntityID, &r.Op, &r.Before, &r.After, ); err != nil { return nil, err } return &r, nil } // WriteChangeSet inserts a change set and all of its revisions in one // transaction, so history never records half an operation. The service buffers // revisions while the operation runs and calls this once it has succeeded — // which is also why an operation that fails partway leaves no change set behind. // seq is assigned from the slice order. Returns the stored change set. func (d *DB) WriteChangeSet(ctx context.Context, cs *domain.ChangeSet, revs []domain.Revision) (*domain.ChangeSet, error) { tx, err := d.sql.BeginTx(ctx, nil) if err != nil { return nil, fmt.Errorf("store: begin change set: %w", err) } defer func() { _ = tx.Rollback() }() created, err := scanChangeSet(tx.QueryRowContext(ctx, `INSERT INTO change_sets (garden_id, actor_id, source, summary, agent_run_id, reverts_id) VALUES (?, ?, ?, ?, ?, ?) RETURNING `+changeSetColumns, cs.GardenID, cs.ActorID, cs.Source, cs.Summary, cs.AgentRunID, cs.RevertsID)) if err != nil { return nil, fmt.Errorf("store: insert change set: %w", err) } for i := range revs { r := &revs[i] if _, err := tx.ExecContext(ctx, `INSERT INTO revisions (change_set_id, seq, entity_type, entity_id, op, before, after) VALUES (?, ?, ?, ?, ?, ?, ?)`, created.ID, int64(i+1), r.EntityType, r.EntityID, r.Op, r.Before, r.After, ); err != nil { return nil, fmt.Errorf("store: insert revision: %w", err) } } if err := tx.Commit(); err != nil { return nil, fmt.Errorf("store: commit change set: %w", err) } return created, nil } // ListChangeSets returns a garden's change sets newest-first, one page at a time. // Each carries its actor's display name, the id of the change set that reverted // it (if any), and per-(entity,op) counts — everything the history list renders, // without a second round-trip per row. Always a non-nil slice. func (d *DB) ListChangeSets(ctx context.Context, gardenID int64, limit, offset int) ([]domain.ChangeSet, error) { // The "was this reverted?" lookup is a scalar subquery, NOT a LEFT JOIN: a // change set can be reverted more than once (undo, redo, undo again), and a // join would then emit one duplicate row per revert and silently corrupt the // page. MIN(id) names the first revert, which is the one worth showing. rows, err := d.sql.QueryContext(ctx, `SELECT `+qualifyColumns("cs", changeSetColumns)+`, u.display_name, (SELECT MIN(r.id) FROM change_sets r WHERE r.reverts_id = cs.id) FROM change_sets cs JOIN users u ON u.id = cs.actor_id WHERE cs.garden_id = ? ORDER BY cs.id DESC LIMIT ? OFFSET ?`, gardenID, limit, offset) if err != nil { return nil, fmt.Errorf("store: list change sets: %w", err) } defer rows.Close() sets := []domain.ChangeSet{} ids := []any{} byID := map[int64]int{} for rows.Next() { var cs domain.ChangeSet if err := rows.Scan( &cs.ID, &cs.GardenID, &cs.ActorID, &cs.Source, &cs.Summary, &cs.AgentRunID, &cs.RevertsID, &cs.CreatedAt, &cs.ActorName, &cs.RevertedByID, ); err != nil { return nil, fmt.Errorf("store: scan change set: %w", err) } byID[cs.ID] = len(sets) ids = append(ids, cs.ID) sets = append(sets, cs) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("store: iterate change sets: %w", err) } if len(sets) == 0 { return sets, nil } // One grouped query for the whole page rather than N per-row counts. countRows, err := d.sql.QueryContext(ctx, `SELECT change_set_id, entity_type, op, COUNT(*) FROM revisions WHERE change_set_id IN (`+placeholders(len(ids))+`) GROUP BY change_set_id, entity_type, op ORDER BY change_set_id, entity_type, op`, ids...) if err != nil { return nil, fmt.Errorf("store: count revisions: %w", err) } defer countRows.Close() for countRows.Next() { var csID int64 var c domain.ChangeCount if err := countRows.Scan(&csID, &c.EntityType, &c.Op, &c.N); err != nil { return nil, fmt.Errorf("store: scan revision count: %w", err) } if i, ok := byID[csID]; ok { sets[i].Counts = append(sets[i].Counts, c) } } if err := countRows.Err(); err != nil { return nil, fmt.Errorf("store: iterate revision counts: %w", err) } return sets, nil } // GetChangeSet returns one change set with its revisions loaded in seq order, or // domain.ErrNotFound. func (d *DB) GetChangeSet(ctx context.Context, id int64) (*domain.ChangeSet, error) { cs, err := scanChangeSet(d.sql.QueryRowContext(ctx, `SELECT `+changeSetColumns+` FROM change_sets WHERE id = ?`, id)) if errors.Is(err, sql.ErrNoRows) { return nil, domain.ErrNotFound } if err != nil { return nil, fmt.Errorf("store: get change set: %w", err) } rows, err := d.sql.QueryContext(ctx, `SELECT `+revisionColumns+` FROM revisions WHERE change_set_id = ? ORDER BY seq`, id) if err != nil { return nil, fmt.Errorf("store: list revisions: %w", err) } defer rows.Close() cs.Revisions = []domain.Revision{} for rows.Next() { r, err := scanRevision(rows) if err != nil { return nil, fmt.Errorf("store: scan revision: %w", err) } cs.Revisions = append(cs.Revisions, *r) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("store: iterate revisions: %w", err) } return cs, nil } // placeholders returns "?, ?, …" for an IN clause of n values. func placeholders(n int) string { if n <= 0 { return "NULL" } b := make([]byte, 0, n*3) for i := 0; i < n; i++ { if i > 0 { b = append(b, ',', ' ') } b = append(b, '?') } return string(b) }