From ea9d2cb7bc58254c69b82232610bbaa985308e77 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 13:11:09 -0400 Subject: [PATCH 01/33] feat(psql): make derived select queries immutable --- dialect/psql/immutable_select.go | 587 ++++++++++++++++++++++++++ dialect/psql/immutable_select_test.go | 168 ++++++++ dialect/psql/view.go | 2 + 3 files changed, 757 insertions(+) create mode 100644 dialect/psql/immutable_select.go create mode 100644 dialect/psql/immutable_select_test.go diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go new file mode 100644 index 00000000..584af6c4 --- /dev/null +++ b/dialect/psql/immutable_select.go @@ -0,0 +1,587 @@ +package psql + +import ( + "context" + "database/sql" + "fmt" + "io" + "strconv" + "strings" + + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/clause" + psqldialect "github.com/stephenafamo/bob/dialect/psql/dialect" + "github.com/stephenafamo/bob/mods" + "github.com/stephenafamo/bob/orm" + "github.com/stephenafamo/scan" +) + +type ImmutableSelectQuery struct { + state immutableSelectState +} + +type immutableSelectState struct { + With clause.With + SelectColumns []any + PreloadColumns []any + Distinct psqldialect.Distinct + TableRef clause.TableRef + Where clause.Where + GroupBy clause.GroupBy + Having clause.Having + Windows clause.Windows + Combines clause.Combines + OrderBy clause.OrderBy + Limit clause.Limit + Offset clause.Offset + Fetch clause.Fetch + Locks clause.Locks + + CombinedOrder clause.OrderBy + CombinedLimit clause.Limit + CombinedFetch clause.Fetch + CombinedOffset clause.Offset +} + +func newImmutableSelect(queryMods ...bob.Mod[*psqldialect.SelectQuery]) ImmutableSelectQuery { + mutable := Select(queryMods...) + return ImmutableSelectQuery{state: immutableStateFromMutable(mutable.Expression)} +} + +func asImmutable(q bob.BaseQuery[*psqldialect.SelectQuery]) ImmutableSelectQuery { + return ImmutableSelectQuery{state: immutableStateFromMutable(q.Expression)} +} + +func (q ImmutableSelectQuery) Type() bob.QueryType { + return bob.QueryTypeSelect +} + +func (q ImmutableSelectQuery) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) ImmutableSelectQuery { + next, ok := q.state.withMods(queryMods...) + if ok { + return ImmutableSelectQuery{state: next} + } + + mutable := q.state.toMutable() + for _, mod := range queryMods { + mod.Apply(&mutable) + } + + return ImmutableSelectQuery{state: immutableStateFromMutable(&mutable)} +} + +func (q ImmutableSelectQuery) AsCount() ImmutableSelectQuery { + next := q.state + next.SelectColumns = []any{"count(1)"} + next.PreloadColumns = nil + next.OrderBy.Expressions = nil + next.GroupBy.Groups = nil + next.GroupBy.With = "" + next.GroupBy.Distinct = false + next.Offset.Count = nil + next.Limit.Count = 1 + + return ImmutableSelectQuery{state: next} +} + +func (q ImmutableSelectQuery) Build(ctx context.Context) (string, []any, error) { + return q.BuildN(ctx, 1) +} + +func (q ImmutableSelectQuery) BuildN(ctx context.Context, start int) (string, []any, error) { + var sb strings.Builder + args, err := q.WriteQuery(ctx, &sb, start) + if err != nil { + return "", nil, err + } + + return sb.String(), args, nil +} + +func (q ImmutableSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { + writer := immutableSelectWriter{ + ctx: ctx, + w: w, + start: start, + } + + if err := writer.writeQuery(q.state); err != nil { + return nil, err + } + + return writer.args, nil +} + +func (q ImmutableSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { + w.WriteString("(") + args, err := q.WriteQuery(ctx, w, start) + if err != nil { + return nil, err + } + w.WriteString(")") + return args, nil +} + +type ImmutableViewQuery[T any, Ts ~[]T] struct { + Query ImmutableSelectQuery + Scanner scan.Mapper[T] + Hooks *bob.Hooks[*psqldialect.SelectQuery, bob.SkipQueryHooksKey] +} + +func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) ImmutableViewQuery[T, Ts] { + state := immutableStateFromMutable(q.BaseQuery.Expression) + if len(state.SelectColumns) == 0 && q.defaultSelect != nil { + state.SelectColumns = append(state.SelectColumns, q.defaultSelect) + } + + return ImmutableViewQuery[T, Ts]{ + Query: ImmutableSelectQuery{state: state}.With(queryMods...), + Scanner: q.Scanner, + Hooks: q.Hooks, + } +} + +func (q ImmutableViewQuery[T, Ts]) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) ImmutableViewQuery[T, Ts] { + q.Query = q.Query.With(queryMods...) + return q +} + +func (q ImmutableViewQuery[T, Ts]) One(ctx context.Context, exec bob.Executor) (T, error) { + return q.mutable().One(ctx, exec) +} + +func (q ImmutableViewQuery[T, Ts]) All(ctx context.Context, exec bob.Executor) (Ts, error) { + return q.mutable().All(ctx, exec) +} + +func (q ImmutableViewQuery[T, Ts]) Cursor(ctx context.Context, exec bob.Executor) (scan.ICursor[T], error) { + return q.mutable().Cursor(ctx, exec) +} + +func (q ImmutableViewQuery[T, Ts]) Each(ctx context.Context, exec bob.Executor) (func(func(T, error) bool), error) { + return q.mutable().Each(ctx, exec) +} + +func (q ImmutableViewQuery[T, Ts]) Count(ctx context.Context, exec bob.Executor) (int64, error) { + mq := q.mutable() + ctx, err := mq.RunHooks(ctx, exec) + if err != nil { + return 0, err + } + return bob.One(ctx, exec, asCountQuery(mq.BaseQuery), scan.SingleColumnMapper[int64]) +} + +func (q ImmutableViewQuery[T, Ts]) Exists(ctx context.Context, exec bob.Executor) (bool, error) { + count, err := q.Count(ctx, exec) + return count > 0, err +} + +func (q ImmutableViewQuery[T, Ts]) CountQuery() ImmutableSelectQuery { + return q.Query.AsCount() +} + +func (q ImmutableViewQuery[T, Ts]) Build(ctx context.Context) (string, []any, error) { + return q.Query.Build(ctx) +} + +func (q ImmutableViewQuery[T, Ts]) mutable() orm.Query[*psqldialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] { + mutable := q.Query.state.toMutable() + return orm.Query[*psqldialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]]{ + ExecQuery: orm.ExecQuery[*psqldialect.SelectQuery]{ + BaseQuery: bob.BaseQuery[*psqldialect.SelectQuery]{ + Expression: &mutable, + Dialect: psqldialect.Dialect, + QueryType: bob.QueryTypeSelect, + }, + Hooks: q.Hooks, + }, + Scanner: q.Scanner, + } +} + +func immutableStateFromMutable(q *psqldialect.SelectQuery) immutableSelectState { + return immutableSelectState{ + With: clause.With{ + Recursive: q.With.Recursive, + CTEs: append([]bob.Expression(nil), q.With.CTEs...), + }, + SelectColumns: append([]any(nil), q.SelectList.Columns...), + PreloadColumns: append([]any(nil), q.SelectList.PreloadColumns...), + Distinct: psqldialect.Distinct{On: append([]any(nil), q.Distinct.On...)}, + TableRef: cloneTableRef(q.TableRef), + Where: clause.Where{Conditions: append([]any(nil), q.Where.Conditions...)}, + GroupBy: clause.GroupBy{ + Groups: append([]any(nil), q.GroupBy.Groups...), + Distinct: q.GroupBy.Distinct, + With: q.GroupBy.With, + }, + Having: clause.Having{ + Conditions: append([]any(nil), q.Having.Conditions...), + }, + Windows: clause.Windows{ + Windows: append([]bob.Expression(nil), q.Windows.Windows...), + }, + Combines: clause.Combines{ + Queries: append([]clause.Combine(nil), q.Combines.Queries...), + }, + OrderBy: clause.OrderBy{ + Expressions: append([]bob.Expression(nil), q.OrderBy.Expressions...), + }, + Limit: clause.Limit{Count: q.Limit.Count}, + Offset: clause.Offset{ + Count: q.Offset.Count, + }, + Fetch: clause.Fetch{ + Count: q.Fetch.Count, + WithTies: q.Fetch.WithTies, + }, + Locks: clause.Locks{ + Locks: append([]bob.Expression(nil), q.Locks.Locks...), + }, + CombinedOrder: clause.OrderBy{ + Expressions: append([]bob.Expression(nil), q.CombinedOrder.Expressions...), + }, + CombinedLimit: clause.Limit{Count: q.CombinedLimit.Count}, + CombinedFetch: clause.Fetch{ + Count: q.CombinedFetch.Count, + WithTies: q.CombinedFetch.WithTies, + }, + CombinedOffset: clause.Offset{Count: q.CombinedOffset.Count}, + } +} + +func (s immutableSelectState) toMutable() psqldialect.SelectQuery { + return psqldialect.SelectQuery{ + With: s.With, + SelectList: clause.SelectList{Columns: s.SelectColumns, PreloadColumns: s.PreloadColumns}, + Distinct: s.Distinct, + TableRef: s.TableRef, + Where: s.Where, + GroupBy: s.GroupBy, + Having: s.Having, + Windows: s.Windows, + Combines: s.Combines, + OrderBy: s.OrderBy, + Limit: s.Limit, + Offset: s.Offset, + Fetch: s.Fetch, + Locks: s.Locks, + CombinedOrder: s.CombinedOrder, + CombinedLimit: s.CombinedLimit, + CombinedFetch: s.CombinedFetch, + CombinedOffset: s.CombinedOffset, + } +} + +func (s immutableSelectState) withMods(queryMods ...bob.Mod[*psqldialect.SelectQuery]) (immutableSelectState, bool) { + next := s + var cloneSelect, cloneWhere, cloneGroup, cloneHaving, cloneOrder, cloneWindows, cloneLocks bool + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Select[*psqldialect.SelectQuery]: + if !cloneSelect { + next.SelectColumns = append([]any(nil), s.SelectColumns...) + cloneSelect = true + } + next.SelectColumns = append(next.SelectColumns, []any(m)...) + case mods.Where[*psqldialect.SelectQuery]: + if !cloneWhere { + next.Where.Conditions = append([]any(nil), s.Where.Conditions...) + cloneWhere = true + } + next.Where.Conditions = append(next.Where.Conditions, m.E) + case mods.GroupBy[*psqldialect.SelectQuery]: + if !cloneGroup { + next.GroupBy.Groups = append([]any(nil), s.GroupBy.Groups...) + cloneGroup = true + } + next.GroupBy.Groups = append(next.GroupBy.Groups, m.E) + case mods.Having[*psqldialect.SelectQuery]: + if !cloneHaving { + next.Having.Conditions = append([]any(nil), s.Having.Conditions...) + cloneHaving = true + } + next.Having.Conditions = append(next.Having.Conditions, []any(m)...) + case mods.Limit[*psqldialect.SelectQuery]: + next.Limit.Count = m.Count + case mods.Offset[*psqldialect.SelectQuery]: + next.Offset.Count = m.Count + case mods.Fetch[*psqldialect.SelectQuery]: + next.Fetch = clause.Fetch(m) + case psqldialect.OrderBy[*psqldialect.SelectQuery]: + if !cloneOrder { + next.OrderBy.Expressions = append([]bob.Expression(nil), s.OrderBy.Expressions...) + cloneOrder = true + } + next.OrderBy.Expressions = append(next.OrderBy.Expressions, m()) + case mods.NamedWindow[*psqldialect.SelectQuery]: + if !cloneWindows { + next.Windows.Windows = append([]bob.Expression(nil), s.Windows.Windows...) + cloneWindows = true + } + next.Windows.Windows = append(next.Windows.Windows, clause.NamedWindow(m)) + case psqldialect.LockChain[*psqldialect.SelectQuery]: + if !cloneLocks { + next.Locks.Locks = append([]bob.Expression(nil), s.Locks.Locks...) + cloneLocks = true + } + next.Locks.Locks = append(next.Locks.Locks, m()) + case psqldialect.FromChain[*psqldialect.SelectQuery]: + next.TableRef = cloneTableRef(m()) + default: + return next, false + } + } + + return next, true +} + +func cloneTableRef(from clause.TableRef) clause.TableRef { + from.Columns = append([]string(nil), from.Columns...) + from.Partitions = append([]string(nil), from.Partitions...) + from.IndexHints = append([]clause.IndexHint(nil), from.IndexHints...) + from.Joins = append([]clause.Join(nil), from.Joins...) + for i := range from.Joins { + from.Joins[i].On = append([]bob.Expression(nil), from.Joins[i].On...) + from.Joins[i].Using = append([]string(nil), from.Joins[i].Using...) + from.Joins[i].To = cloneTableRef(from.Joins[i].To) + } + return from +} + +type immutableSelectWriter struct { + ctx context.Context + w io.StringWriter + args []any + start int +} + +func (w *immutableSelectWriter) writeQuery(q immutableSelectState) error { + if len(q.With.CTEs) > 0 { + if _, err := q.With.WriteSQL(w.ctx, w.w, psqldialect.Dialect, w.argPos()); err != nil { + return err + } + w.w.WriteString("\n") + } + + w.w.WriteString("SELECT ") + + if q.Distinct.On != nil { + w.w.WriteString("DISTINCT") + if len(q.Distinct.On) > 0 { + w.w.WriteString(" ON (") + if err := w.writeSliceAny(q.Distinct.On, ", "); err != nil { + return err + } + w.w.WriteString(")") + } + w.w.WriteString(" ") + } + + w.w.WriteString("\n") + if len(q.SelectColumns) == 0 && len(q.PreloadColumns) == 0 { + w.w.WriteString("*") + } else { + allCols := append([]any(nil), q.SelectColumns...) + allCols = append(allCols, q.PreloadColumns...) + if err := w.writeSliceAny(allCols, ", "); err != nil { + return err + } + } + + if q.TableRef.Expression != nil { + w.w.WriteString("\nFROM ") + args, err := q.TableRef.WriteSQL(w.ctx, w.w, psqldialect.Dialect, w.argPos()) + if err != nil { + return err + } + w.args = append(w.args, args...) + } + + if len(q.Where.Conditions) > 0 { + w.w.WriteString("\nWHERE ") + if err := w.writeSliceAny(q.Where.Conditions, " AND "); err != nil { + return err + } + } + + if len(q.GroupBy.Groups) > 0 { + w.w.WriteString("\nGROUP BY ") + if q.GroupBy.Distinct { + w.w.WriteString("DISTINCT ") + } + if err := w.writeSliceAny(q.GroupBy.Groups, ", "); err != nil { + return err + } + if q.GroupBy.With != "" { + w.w.WriteString(" WITH ") + w.w.WriteString(q.GroupBy.With) + } + } + + if len(q.Having.Conditions) > 0 { + w.w.WriteString("\nHAVING ") + if err := w.writeSliceAny(q.Having.Conditions, " AND "); err != nil { + return err + } + } + + if len(q.Windows.Windows) > 0 { + w.w.WriteString("\nWINDOW ") + if err := w.writeSliceExpr(q.Windows.Windows, ", "); err != nil { + return err + } + } + + if len(q.OrderBy.Expressions) > 0 { + w.w.WriteString("\nORDER BY ") + if err := w.writeOrderExprs(q.OrderBy.Expressions); err != nil { + return err + } + } + + if q.Limit.Count != nil { + w.w.WriteString("\nLIMIT ") + if err := w.writeAny(q.Limit.Count); err != nil { + return err + } + } + + if q.Offset.Count != nil { + w.w.WriteString("\nOFFSET ") + if err := w.writeAny(q.Offset.Count); err != nil { + return err + } + } + + if q.Fetch.Count != nil { + w.w.WriteString("\nFETCH NEXT ") + if err := w.writeAny(q.Fetch.Count); err != nil { + return err + } + if q.Fetch.WithTies { + w.w.WriteString(" ROWS WITH TIES") + } else { + w.w.WriteString(" ROWS ONLY") + } + } + + for _, lock := range q.Locks.Locks { + w.w.WriteString("\n") + if err := w.writeAny(lock); err != nil { + return err + } + } + + w.w.WriteString("\n") + return nil +} + +func (w *immutableSelectWriter) argPos() int { + return w.start + len(w.args) +} + +func (w *immutableSelectWriter) writeSliceAny(values []any, sep string) error { + for i, value := range values { + if i > 0 { + w.w.WriteString(sep) + } + if err := w.writeAny(value); err != nil { + return err + } + } + return nil +} + +func (w *immutableSelectWriter) writeSliceExpr(values []bob.Expression, sep string) error { + for i, value := range values { + if i > 0 { + w.w.WriteString(sep) + } + if err := w.writeExpression(value); err != nil { + return err + } + } + return nil +} + +func (w *immutableSelectWriter) writeOrderExprs(values []bob.Expression) error { + for i, value := range values { + if i > 0 { + w.w.WriteString(", ") + } + + switch order := value.(type) { + case clause.OrderDef: + if err := w.writeAny(order.Expression); err != nil { + return err + } + if order.Collation != "" { + w.w.WriteString(" COLLATE ") + psqldialect.Dialect.WriteQuoted(w.w, order.Collation) + } + if order.Direction != "" { + w.w.WriteString(" ") + w.w.WriteString(order.Direction) + } + if order.Nulls != "" { + w.w.WriteString(" NULLS ") + w.w.WriteString(order.Nulls) + } + default: + if err := w.writeExpression(value); err != nil { + return err + } + } + } + return nil +} + +func (w *immutableSelectWriter) writeExpression(value bob.Expression) error { + args, err := value.WriteSQL(w.ctx, w.w, psqldialect.Dialect, w.argPos()) + if err != nil { + return err + } + w.args = append(w.args, args...) + return nil +} + +func (w *immutableSelectWriter) writeAny(value any) error { + switch v := value.(type) { + case nil: + w.w.WriteString("NULL") + case string: + w.w.WriteString(v) + case []byte: + w.w.WriteString(string(v)) + case int: + w.w.WriteString(strconv.Itoa(v)) + case int8: + w.w.WriteString(strconv.FormatInt(int64(v), 10)) + case int16: + w.w.WriteString(strconv.FormatInt(int64(v), 10)) + case int32: + w.w.WriteString(strconv.FormatInt(int64(v), 10)) + case int64: + w.w.WriteString(strconv.FormatInt(v, 10)) + case uint: + w.w.WriteString(strconv.FormatUint(uint64(v), 10)) + case uint8: + w.w.WriteString(strconv.FormatUint(uint64(v), 10)) + case uint16: + w.w.WriteString(strconv.FormatUint(uint64(v), 10)) + case uint32: + w.w.WriteString(strconv.FormatUint(uint64(v), 10)) + case uint64: + w.w.WriteString(strconv.FormatUint(v, 10)) + case sql.NamedArg: + return fmt.Errorf("named args are not supported by psql dialect") + case bob.Expression: + return w.writeExpression(v) + default: + w.w.WriteString(fmt.Sprint(v)) + } + + return nil +} diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go new file mode 100644 index 00000000..92c65c66 --- /dev/null +++ b/dialect/psql/immutable_select_test.go @@ -0,0 +1,168 @@ +package psql + +import ( + "context" + "testing" + + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/dialect/psql/sm" +) + +func TestImmutableSelectQueryWithDoesNotMutateOriginal(t *testing.T) { + base := newImmutableSelect( + sm.Columns("id"), + sm.From("users"), + ) + + derived := base.With( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + if derived.Type() != bob.QueryTypeSelect { + t.Fatalf("expected derived query type %q, got %q", bob.QueryTypeSelect, derived.Type()) + } + + baseSQL, _, err := base.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "SELECT \nid\nFROM users\n" { + t.Fatalf("base query changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "SELECT \nid\nFROM users\nORDER BY id DESC\nLIMIT 10\nOFFSET 20\n" { + t.Fatalf("derived query mismatch: %#v", derivedSQL) + } +} + +func TestImmutableViewQueryWithDoesNotMutateOriginal(t *testing.T) { + base := someStructView.Query( + sm.Where(Quote("id").GT(Arg(0))), + ).With() + + derived := base.With( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + if derived.Scanner == nil { + t.Fatal("expected derived view query to preserve scanner") + } + + baseSQL, _, err := base.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "SELECT \n\"some_struct\".\"id\" AS \"id\", \"some_struct\".\"name\" AS \"name\", \"some_struct\".\"email\" AS \"email\"\nFROM \"public\".\"some_struct\" AS \"public.some_struct\"\nWHERE (\"id\" > $1)\n" { + t.Fatalf("base view query changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "SELECT \n\"some_struct\".\"id\" AS \"id\", \"some_struct\".\"name\" AS \"name\", \"some_struct\".\"email\" AS \"email\"\nFROM \"public\".\"some_struct\" AS \"public.some_struct\"\nWHERE (\"id\" > $1)\nORDER BY id DESC\nLIMIT 10\nOFFSET 20\n" { + t.Fatalf("derived view query mismatch: %#v", derivedSQL) + } +} + +func BenchmarkBaseQueryApplyMain(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Select( + sm.Columns("id", "name"), + sm.From("users"), + sm.Where(Quote("tenant_id").EQ(Arg(42))), + ) + q.Apply( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + if _, _, err := q.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkBaseQueryImmutableNativeHotPath(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := asImmutable(Select( + sm.Columns("id", "name"), + sm.From("users"), + sm.Where(Quote("tenant_id").EQ(Arg(42))), + )) + derived := q.With( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + if _, _, err := derived.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := someStructView.Query( + sm.Where(Quote("id").GT(Arg(0))), + ) + + if _, _, err := asCountQuery(q.BaseQuery).Build(ctx); err != nil { + b.Fatal(err) + } + + q.Apply( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + if _, _, err := q.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := someStructView.Query( + sm.Where(Quote("id").GT(Arg(0))), + ) + + if _, _, err := asCountQuery(q.BaseQuery).Build(ctx); err != nil { + b.Fatal(err) + } + + derived := q.With( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + if _, _, err := derived.Build(ctx); err != nil { + b.Fatal(err) + } + } +} diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 104ca0de..07c3ee25 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -87,6 +87,7 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * }, Scanner: v.scanner, }, + defaultSelect: v.Columns, } q.Expression.AppendContextualModFunc( @@ -105,6 +106,7 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * type ViewQuery[T any, Ts ~[]T] struct { orm.Query[*dialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] + defaultSelect bob.Expression } // Count the number of matching rows From bb27f9ffea97e4f9d34cea5b2cf5ad3989ca6065 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 13:12:49 -0400 Subject: [PATCH 02/33] refactor(psql): route select with through immutable queries --- dialect/psql/immutable_select.go | 5 ----- dialect/psql/immutable_select_test.go | 8 ++++---- dialect/psql/select.go | 20 +++++++++++++++----- dialect/psql/view.go | 2 +- dialect/psql/view_test.go | 3 +-- 5 files changed, 21 insertions(+), 17 deletions(-) diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index 584af6c4..391fa046 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -43,11 +43,6 @@ type immutableSelectState struct { CombinedOffset clause.Offset } -func newImmutableSelect(queryMods ...bob.Mod[*psqldialect.SelectQuery]) ImmutableSelectQuery { - mutable := Select(queryMods...) - return ImmutableSelectQuery{state: immutableStateFromMutable(mutable.Expression)} -} - func asImmutable(q bob.BaseQuery[*psqldialect.SelectQuery]) ImmutableSelectQuery { return ImmutableSelectQuery{state: immutableStateFromMutable(q.Expression)} } diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 92c65c66..92717138 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -9,10 +9,10 @@ import ( ) func TestImmutableSelectQueryWithDoesNotMutateOriginal(t *testing.T) { - base := newImmutableSelect( + base := Select( sm.Columns("id"), sm.From("users"), - ) + ).With() derived := base.With( sm.OrderBy("id").Desc(), @@ -100,11 +100,11 @@ func BenchmarkBaseQueryImmutableNativeHotPath(b *testing.B) { b.ReportAllocs() for i := 0; i < b.N; i++ { - q := asImmutable(Select( + q := Select( sm.Columns("id", "name"), sm.From("users"), sm.Where(Quote("tenant_id").EQ(Arg(42))), - )) + ) derived := q.With( sm.OrderBy("id").Desc(), sm.Limit(10), diff --git a/dialect/psql/select.go b/dialect/psql/select.go index e177663d..54923be6 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -5,15 +5,25 @@ import ( "github.com/stephenafamo/bob/dialect/psql/dialect" ) -func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) bob.BaseQuery[*dialect.SelectQuery] { +type SelectQuery struct { + bob.BaseQuery[*dialect.SelectQuery] +} + +func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) ImmutableSelectQuery { + return asImmutable(q.BaseQuery).With(queryMods...) +} + +func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { q := &dialect.SelectQuery{} for _, mod := range queryMods { mod.Apply(q) } - return bob.BaseQuery[*dialect.SelectQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeSelect, + return SelectQuery{ + BaseQuery: bob.BaseQuery[*dialect.SelectQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeSelect, + }, } } diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 07c3ee25..79cd6c5b 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -82,7 +82,7 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * q := &ViewQuery[T, Tslice]{ Query: orm.Query[*dialect.SelectQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ ExecQuery: orm.ExecQuery[*dialect.SelectQuery]{ - BaseQuery: Select(sm.From(v.NameAs())), + BaseQuery: Select(sm.From(v.NameAs())).BaseQuery, Hooks: &v.SelectQueryHooks, }, Scanner: v.scanner, diff --git a/dialect/psql/view_test.go b/dialect/psql/view_test.go index 9e68a9d3..05c2bb5f 100644 --- a/dialect/psql/view_test.go +++ b/dialect/psql/view_test.go @@ -7,7 +7,6 @@ import ( _ "github.com/lib/pq" "github.com/stephenafamo/bob" - "github.com/stephenafamo/bob/dialect/psql/dialect" "github.com/stephenafamo/bob/dialect/psql/sm" "github.com/stephenafamo/bob/expr" ) @@ -56,7 +55,7 @@ func TestSomeViewQuery(t *testing.T) { } } -func selectToString(t *testing.T, query bob.BaseQuery[*dialect.SelectQuery], argsLen int) string { +func selectToString(t *testing.T, query bob.Query, argsLen int) string { t.Helper() ctx := context.Background() buf := new(bytes.Buffer) From 2eb36dfd27a0eda3d6d6ad6ca47c20bbdf8be5ad Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 13:26:02 -0400 Subject: [PATCH 03/33] refactor(psql): hide immutable query internals --- dialect/psql/immutable_select.go | 58 +++++++++++++++----------------- dialect/psql/select.go | 2 +- 2 files changed, 28 insertions(+), 32 deletions(-) diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index 391fa046..73d62f12 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -16,7 +16,7 @@ import ( "github.com/stephenafamo/scan" ) -type ImmutableSelectQuery struct { +type derivedSelectQuery struct { state immutableSelectState } @@ -43,18 +43,18 @@ type immutableSelectState struct { CombinedOffset clause.Offset } -func asImmutable(q bob.BaseQuery[*psqldialect.SelectQuery]) ImmutableSelectQuery { - return ImmutableSelectQuery{state: immutableStateFromMutable(q.Expression)} +func asImmutable(q bob.BaseQuery[*psqldialect.SelectQuery]) derivedSelectQuery { + return derivedSelectQuery{state: immutableStateFromMutable(q.Expression)} } -func (q ImmutableSelectQuery) Type() bob.QueryType { +func (q derivedSelectQuery) Type() bob.QueryType { return bob.QueryTypeSelect } -func (q ImmutableSelectQuery) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) ImmutableSelectQuery { +func (q derivedSelectQuery) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) derivedSelectQuery { next, ok := q.state.withMods(queryMods...) if ok { - return ImmutableSelectQuery{state: next} + return derivedSelectQuery{state: next} } mutable := q.state.toMutable() @@ -62,10 +62,10 @@ func (q ImmutableSelectQuery) With(queryMods ...bob.Mod[*psqldialect.SelectQuery mod.Apply(&mutable) } - return ImmutableSelectQuery{state: immutableStateFromMutable(&mutable)} + return derivedSelectQuery{state: immutableStateFromMutable(&mutable)} } -func (q ImmutableSelectQuery) AsCount() ImmutableSelectQuery { +func (q derivedSelectQuery) AsCount() derivedSelectQuery { next := q.state next.SelectColumns = []any{"count(1)"} next.PreloadColumns = nil @@ -76,14 +76,14 @@ func (q ImmutableSelectQuery) AsCount() ImmutableSelectQuery { next.Offset.Count = nil next.Limit.Count = 1 - return ImmutableSelectQuery{state: next} + return derivedSelectQuery{state: next} } -func (q ImmutableSelectQuery) Build(ctx context.Context) (string, []any, error) { +func (q derivedSelectQuery) Build(ctx context.Context) (string, []any, error) { return q.BuildN(ctx, 1) } -func (q ImmutableSelectQuery) BuildN(ctx context.Context, start int) (string, []any, error) { +func (q derivedSelectQuery) BuildN(ctx context.Context, start int) (string, []any, error) { var sb strings.Builder args, err := q.WriteQuery(ctx, &sb, start) if err != nil { @@ -93,7 +93,7 @@ func (q ImmutableSelectQuery) BuildN(ctx context.Context, start int) (string, [] return sb.String(), args, nil } -func (q ImmutableSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { +func (q derivedSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { writer := immutableSelectWriter{ ctx: ctx, w: w, @@ -107,7 +107,7 @@ func (q ImmutableSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, return writer.args, nil } -func (q ImmutableSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { +func (q derivedSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { w.WriteString("(") args, err := q.WriteQuery(ctx, w, start) if err != nil { @@ -117,47 +117,47 @@ func (q ImmutableSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ return args, nil } -type ImmutableViewQuery[T any, Ts ~[]T] struct { - Query ImmutableSelectQuery +type derivedViewQuery[T any, Ts ~[]T] struct { + Query derivedSelectQuery Scanner scan.Mapper[T] Hooks *bob.Hooks[*psqldialect.SelectQuery, bob.SkipQueryHooksKey] } -func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) ImmutableViewQuery[T, Ts] { +func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) derivedViewQuery[T, Ts] { state := immutableStateFromMutable(q.BaseQuery.Expression) if len(state.SelectColumns) == 0 && q.defaultSelect != nil { state.SelectColumns = append(state.SelectColumns, q.defaultSelect) } - return ImmutableViewQuery[T, Ts]{ - Query: ImmutableSelectQuery{state: state}.With(queryMods...), + return derivedViewQuery[T, Ts]{ + Query: derivedSelectQuery{state: state}.With(queryMods...), Scanner: q.Scanner, Hooks: q.Hooks, } } -func (q ImmutableViewQuery[T, Ts]) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) ImmutableViewQuery[T, Ts] { +func (q derivedViewQuery[T, Ts]) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) derivedViewQuery[T, Ts] { q.Query = q.Query.With(queryMods...) return q } -func (q ImmutableViewQuery[T, Ts]) One(ctx context.Context, exec bob.Executor) (T, error) { +func (q derivedViewQuery[T, Ts]) One(ctx context.Context, exec bob.Executor) (T, error) { return q.mutable().One(ctx, exec) } -func (q ImmutableViewQuery[T, Ts]) All(ctx context.Context, exec bob.Executor) (Ts, error) { +func (q derivedViewQuery[T, Ts]) All(ctx context.Context, exec bob.Executor) (Ts, error) { return q.mutable().All(ctx, exec) } -func (q ImmutableViewQuery[T, Ts]) Cursor(ctx context.Context, exec bob.Executor) (scan.ICursor[T], error) { +func (q derivedViewQuery[T, Ts]) Cursor(ctx context.Context, exec bob.Executor) (scan.ICursor[T], error) { return q.mutable().Cursor(ctx, exec) } -func (q ImmutableViewQuery[T, Ts]) Each(ctx context.Context, exec bob.Executor) (func(func(T, error) bool), error) { +func (q derivedViewQuery[T, Ts]) Each(ctx context.Context, exec bob.Executor) (func(func(T, error) bool), error) { return q.mutable().Each(ctx, exec) } -func (q ImmutableViewQuery[T, Ts]) Count(ctx context.Context, exec bob.Executor) (int64, error) { +func (q derivedViewQuery[T, Ts]) Count(ctx context.Context, exec bob.Executor) (int64, error) { mq := q.mutable() ctx, err := mq.RunHooks(ctx, exec) if err != nil { @@ -166,20 +166,16 @@ func (q ImmutableViewQuery[T, Ts]) Count(ctx context.Context, exec bob.Executor) return bob.One(ctx, exec, asCountQuery(mq.BaseQuery), scan.SingleColumnMapper[int64]) } -func (q ImmutableViewQuery[T, Ts]) Exists(ctx context.Context, exec bob.Executor) (bool, error) { +func (q derivedViewQuery[T, Ts]) Exists(ctx context.Context, exec bob.Executor) (bool, error) { count, err := q.Count(ctx, exec) return count > 0, err } -func (q ImmutableViewQuery[T, Ts]) CountQuery() ImmutableSelectQuery { - return q.Query.AsCount() -} - -func (q ImmutableViewQuery[T, Ts]) Build(ctx context.Context) (string, []any, error) { +func (q derivedViewQuery[T, Ts]) Build(ctx context.Context) (string, []any, error) { return q.Query.Build(ctx) } -func (q ImmutableViewQuery[T, Ts]) mutable() orm.Query[*psqldialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] { +func (q derivedViewQuery[T, Ts]) mutable() orm.Query[*psqldialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] { mutable := q.Query.state.toMutable() return orm.Query[*psqldialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]]{ ExecQuery: orm.ExecQuery[*psqldialect.SelectQuery]{ diff --git a/dialect/psql/select.go b/dialect/psql/select.go index 54923be6..36fafadf 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -9,7 +9,7 @@ type SelectQuery struct { bob.BaseQuery[*dialect.SelectQuery] } -func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) ImmutableSelectQuery { +func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) derivedSelectQuery { return asImmutable(q.BaseQuery).With(queryMods...) } From 352c6512d5b6e1002dedc89fa6ff320e9e296ae7 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 13:57:23 -0400 Subject: [PATCH 04/33] feat(psql): add immutable with paths for write queries Benchmarks:\n- BenchmarkUpdateQueryApplyMain: 3038 ns/op, 2319 B/op, 47 allocs/op\n- BenchmarkUpdateQueryImmutableNativeHotPath: 1974 ns/op, 1360 B/op, 41 allocs/op\n- BenchmarkDeleteQueryApplyMain: 1926 ns/op, 1956 B/op, 33 allocs/op\n- BenchmarkDeleteQueryImmutableNativeHotPath: 1345 ns/op, 1027 B/op, 26 allocs/op\n- BenchmarkInsertQueryApplyMain: 1473 ns/op, 1576 B/op, 24 allocs/op\n- BenchmarkInsertQueryImmutableNativeHotPath: 1143 ns/op, 920 B/op, 20 allocs/op --- dialect/psql/delete.go | 20 +- dialect/psql/immutable_write.go | 618 +++++++++++++++++++++++++++ dialect/psql/immutable_write_test.go | 208 +++++++++ dialect/psql/insert.go | 20 +- dialect/psql/table.go | 6 +- dialect/psql/update.go | 20 +- 6 files changed, 874 insertions(+), 18 deletions(-) create mode 100644 dialect/psql/immutable_write.go create mode 100644 dialect/psql/immutable_write_test.go diff --git a/dialect/psql/delete.go b/dialect/psql/delete.go index cdbdd214..9b023db7 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -5,15 +5,25 @@ import ( "github.com/stephenafamo/bob/dialect/psql/dialect" ) -func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) bob.BaseQuery[*dialect.DeleteQuery] { +type DeleteQuery struct { + bob.BaseQuery[*dialect.DeleteQuery] +} + +func (q DeleteQuery) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) derivedDeleteQuery { + return asImmutableDelete(q.BaseQuery).With(queryMods...) +} + +func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { q := &dialect.DeleteQuery{} for _, mod := range queryMods { mod.Apply(q) } - return bob.BaseQuery[*dialect.DeleteQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeDelete, + return DeleteQuery{ + BaseQuery: bob.BaseQuery[*dialect.DeleteQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeDelete, + }, } } diff --git a/dialect/psql/immutable_write.go b/dialect/psql/immutable_write.go new file mode 100644 index 00000000..f0be166f --- /dev/null +++ b/dialect/psql/immutable_write.go @@ -0,0 +1,618 @@ +package psql + +import ( + "context" + "database/sql" + "io" + "strings" + + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/clause" + psqldialect "github.com/stephenafamo/bob/dialect/psql/dialect" + "github.com/stephenafamo/bob/mods" + "github.com/stephenafamo/scan" +) + +type derivedUpdateQuery struct { + state immutableUpdateState + load bob.Load + hooks bob.EmbeddedHook + contextualMods []bob.ContextualMod[*psqldialect.UpdateQuery] +} + +type immutableUpdateState struct { + With clause.With + Only bool + Table clause.TableRef + From clause.TableRef + Set clause.Set + Where clause.Where + Returning clause.Returning +} + +func asImmutableUpdate(q bob.BaseQuery[*psqldialect.UpdateQuery]) derivedUpdateQuery { + return derivedUpdateQuery{ + state: immutableUpdateState{ + With: clause.With{ + Recursive: q.Expression.With.Recursive, + CTEs: append([]bob.Expression(nil), q.Expression.With.CTEs...), + }, + Only: q.Expression.Only, + Table: cloneTableRef(q.Expression.Table), + From: cloneTableRef(q.Expression.TableRef), + Set: clause.Set{ + Set: append([]any(nil), q.Expression.Set.Set...), + }, + Where: clause.Where{ + Conditions: append([]any(nil), q.Expression.Where.Conditions...), + }, + Returning: clause.Returning{ + Expressions: append([]any(nil), q.Expression.Returning.Expressions...), + }, + }, + load: q.Expression.Load, + hooks: q.Expression.EmbeddedHook, + contextualMods: append([]bob.ContextualMod[*psqldialect.UpdateQuery](nil), q.Expression.ContextualModdable.Mods...), + } +} + +func (q derivedUpdateQuery) Type() bob.QueryType { return bob.QueryTypeUpdate } + +func (q derivedUpdateQuery) With(queryMods ...bob.Mod[*psqldialect.UpdateQuery]) derivedUpdateQuery { + next, ok := q.state.withMods(queryMods...) + if ok { + q.state = next + return q + } + + base := q.mutableBase() + mutable := base.Expression + for _, mod := range queryMods { + mod.Apply(mutable) + } + + return asImmutableUpdate(base) +} + +func (q derivedUpdateQuery) Exec(ctx context.Context, exec bob.Executor) (sql.Result, error) { + return bob.Exec(ctx, exec, q) +} + +func (q derivedUpdateQuery) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { + return q.hooks.RunHooks(ctx, exec) +} + +func (q derivedUpdateQuery) GetLoaders() []bob.Loader { + return q.load.GetLoaders() +} + +func (q derivedUpdateQuery) GetMapperMods() []scan.MapperMod { + return q.load.GetMapperMods() +} + +func (q derivedUpdateQuery) Build(ctx context.Context) (string, []any, error) { + return q.BuildN(ctx, 1) +} + +func (q derivedUpdateQuery) BuildN(ctx context.Context, start int) (string, []any, error) { + var sb strings.Builder + args, err := q.WriteQuery(ctx, &sb, start) + if err != nil { + return "", nil, err + } + return sb.String(), args, nil +} + +func (q derivedUpdateQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { + if len(q.contextualMods) > 0 { + return q.mutableBase().WriteQuery(ctx, w, start) + } + + var args []any + + if len(q.state.With.CTEs) > 0 { + withArgs, err := q.state.With.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, withArgs...) + w.WriteString("\n") + } + + w.WriteString("UPDATE ") + if q.state.Only { + w.WriteString("ONLY ") + } + + tableArgs, err := q.state.Table.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, tableArgs...) + + w.WriteString(" SET\n") + setArgs, err := q.state.Set.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, setArgs...) + + if q.state.From.Expression != nil { + w.WriteString("\nFROM ") + fromArgs, err := q.state.From.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, fromArgs...) + } + + if len(q.state.Where.Conditions) > 0 { + w.WriteString("\n") + whereArgs, err := q.state.Where.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, whereArgs...) + } + + if len(q.state.Returning.Expressions) > 0 { + w.WriteString("\n") + retArgs, err := q.state.Returning.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, retArgs...) + } + + return args, nil +} + +func (q derivedUpdateQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { + return q.WriteQuery(ctx, w, start) +} + +func (q derivedUpdateQuery) mutableBase() bob.BaseQuery[*psqldialect.UpdateQuery] { + mutable := &psqldialect.UpdateQuery{ + With: q.state.With, + Only: q.state.Only, + Table: q.state.Table, + Set: q.state.Set, + TableRef: q.state.From, + Where: q.state.Where, + Returning: q.state.Returning, + Load: q.load, + EmbeddedHook: q.hooks, + ContextualModdable: bob.ContextualModdable[*psqldialect.UpdateQuery]{ + Mods: append([]bob.ContextualMod[*psqldialect.UpdateQuery](nil), q.contextualMods...), + }, + } + + return bob.BaseQuery[*psqldialect.UpdateQuery]{ + Expression: mutable, + Dialect: psqldialect.Dialect, + QueryType: bob.QueryTypeUpdate, + } +} + +func (s immutableUpdateState) withMods(queryMods ...bob.Mod[*psqldialect.UpdateQuery]) (immutableUpdateState, bool) { + next := s + var cloneWhere, cloneReturning bool + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Where[*psqldialect.UpdateQuery]: + if !cloneWhere { + next.Where.Conditions = append([]any(nil), s.Where.Conditions...) + cloneWhere = true + } + next.Where.Conditions = append(next.Where.Conditions, m.E) + case mods.Returning[*psqldialect.UpdateQuery]: + if !cloneReturning { + next.Returning.Expressions = append([]any(nil), s.Returning.Expressions...) + cloneReturning = true + } + next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + case psqldialect.FromChain[*psqldialect.UpdateQuery]: + next.From = cloneTableRef(m()) + default: + return next, false + } + } + + return next, true +} + +type derivedDeleteQuery struct { + state immutableDeleteState + load bob.Load + hooks bob.EmbeddedHook + contextualMods []bob.ContextualMod[*psqldialect.DeleteQuery] +} + +type immutableDeleteState struct { + With clause.With + Only bool + Table clause.TableRef + Using clause.TableRef + Where clause.Where + Returning clause.Returning +} + +func asImmutableDelete(q bob.BaseQuery[*psqldialect.DeleteQuery]) derivedDeleteQuery { + return derivedDeleteQuery{ + state: immutableDeleteState{ + With: clause.With{ + Recursive: q.Expression.With.Recursive, + CTEs: append([]bob.Expression(nil), q.Expression.With.CTEs...), + }, + Only: q.Expression.Only, + Table: cloneTableRef(q.Expression.Table), + Using: cloneTableRef(q.Expression.TableRef), + Where: clause.Where{Conditions: append([]any(nil), q.Expression.Where.Conditions...)}, + Returning: clause.Returning{ + Expressions: append([]any(nil), q.Expression.Returning.Expressions...), + }, + }, + load: q.Expression.Load, + hooks: q.Expression.EmbeddedHook, + contextualMods: append([]bob.ContextualMod[*psqldialect.DeleteQuery](nil), q.Expression.ContextualModdable.Mods...), + } +} + +func (q derivedDeleteQuery) Type() bob.QueryType { return bob.QueryTypeDelete } + +func (q derivedDeleteQuery) With(queryMods ...bob.Mod[*psqldialect.DeleteQuery]) derivedDeleteQuery { + next, ok := q.state.withMods(queryMods...) + if ok { + q.state = next + return q + } + + base := q.mutableBase() + mutable := base.Expression + for _, mod := range queryMods { + mod.Apply(mutable) + } + + return asImmutableDelete(base) +} + +func (q derivedDeleteQuery) Exec(ctx context.Context, exec bob.Executor) (sql.Result, error) { + return bob.Exec(ctx, exec, q) +} + +func (q derivedDeleteQuery) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { + return q.hooks.RunHooks(ctx, exec) +} + +func (q derivedDeleteQuery) GetLoaders() []bob.Loader { return q.load.GetLoaders() } + +func (q derivedDeleteQuery) GetMapperMods() []scan.MapperMod { return q.load.GetMapperMods() } + +func (q derivedDeleteQuery) Build(ctx context.Context) (string, []any, error) { + return q.BuildN(ctx, 1) +} + +func (q derivedDeleteQuery) BuildN(ctx context.Context, start int) (string, []any, error) { + var sb strings.Builder + args, err := q.WriteQuery(ctx, &sb, start) + if err != nil { + return "", nil, err + } + return sb.String(), args, nil +} + +func (q derivedDeleteQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { + if len(q.contextualMods) > 0 { + return q.mutableBase().WriteQuery(ctx, w, start) + } + + var args []any + if len(q.state.With.CTEs) > 0 { + withArgs, err := q.state.With.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, withArgs...) + w.WriteString("\n") + } + + w.WriteString("DELETE FROM ") + if q.state.Only { + w.WriteString("ONLY ") + } + + tableArgs, err := q.state.Table.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, tableArgs...) + + if q.state.Using.Expression != nil { + w.WriteString("\nUSING ") + usingArgs, err := q.state.Using.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, usingArgs...) + } + + if len(q.state.Where.Conditions) > 0 { + w.WriteString("\n") + whereArgs, err := q.state.Where.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, whereArgs...) + } + + if len(q.state.Returning.Expressions) > 0 { + w.WriteString("\n") + retArgs, err := q.state.Returning.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, retArgs...) + } + + return args, nil +} + +func (q derivedDeleteQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { + return q.WriteQuery(ctx, w, start) +} + +func (q derivedDeleteQuery) mutableBase() bob.BaseQuery[*psqldialect.DeleteQuery] { + mutable := &psqldialect.DeleteQuery{ + With: q.state.With, + Only: q.state.Only, + Table: q.state.Table, + TableRef: q.state.Using, + Where: q.state.Where, + Returning: q.state.Returning, + Load: q.load, + EmbeddedHook: q.hooks, + ContextualModdable: bob.ContextualModdable[*psqldialect.DeleteQuery]{ + Mods: append([]bob.ContextualMod[*psqldialect.DeleteQuery](nil), q.contextualMods...), + }, + } + + return bob.BaseQuery[*psqldialect.DeleteQuery]{ + Expression: mutable, + Dialect: psqldialect.Dialect, + QueryType: bob.QueryTypeDelete, + } +} + +func (s immutableDeleteState) withMods(queryMods ...bob.Mod[*psqldialect.DeleteQuery]) (immutableDeleteState, bool) { + next := s + var cloneWhere, cloneReturning bool + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Where[*psqldialect.DeleteQuery]: + if !cloneWhere { + next.Where.Conditions = append([]any(nil), s.Where.Conditions...) + cloneWhere = true + } + next.Where.Conditions = append(next.Where.Conditions, m.E) + case mods.Returning[*psqldialect.DeleteQuery]: + if !cloneReturning { + next.Returning.Expressions = append([]any(nil), s.Returning.Expressions...) + cloneReturning = true + } + next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + case psqldialect.FromChain[*psqldialect.DeleteQuery]: + next.Using = cloneTableRef(m()) + default: + return next, false + } + } + + return next, true +} + +type derivedInsertQuery struct { + state immutableInsertState + load bob.Load + hooks bob.EmbeddedHook + contextualMods []bob.ContextualMod[*psqldialect.InsertQuery] +} + +type immutableInsertState struct { + With clause.With + Overriding string + Table clause.TableRef + Values clause.Values + Conflict clause.Conflict + Returning clause.Returning +} + +func asImmutableInsert(q bob.BaseQuery[*psqldialect.InsertQuery]) derivedInsertQuery { + return derivedInsertQuery{ + state: immutableInsertState{ + With: clause.With{ + Recursive: q.Expression.With.Recursive, + CTEs: append([]bob.Expression(nil), q.Expression.With.CTEs...), + }, + Overriding: q.Expression.Overriding, + Table: cloneTableRef(q.Expression.TableRef), + Values: clause.Values{ + Query: q.Expression.Values.Query, + Vals: append([]clause.Value(nil), q.Expression.Values.Vals...), + }, + Conflict: clause.Conflict{Expression: q.Expression.Conflict.Expression}, + Returning: clause.Returning{ + Expressions: append([]any(nil), q.Expression.Returning.Expressions...), + }, + }, + load: q.Expression.Load, + hooks: q.Expression.EmbeddedHook, + contextualMods: append([]bob.ContextualMod[*psqldialect.InsertQuery](nil), q.Expression.ContextualModdable.Mods...), + } +} + +func (q derivedInsertQuery) Type() bob.QueryType { return bob.QueryTypeInsert } + +func (q derivedInsertQuery) With(queryMods ...bob.Mod[*psqldialect.InsertQuery]) derivedInsertQuery { + next, ok := q.state.withMods(queryMods...) + if ok { + q.state = next + return q + } + + base := q.mutableBase() + mutable := base.Expression + for _, mod := range queryMods { + mod.Apply(mutable) + } + + return asImmutableInsert(base) +} + +func (q derivedInsertQuery) Exec(ctx context.Context, exec bob.Executor) (sql.Result, error) { + return bob.Exec(ctx, exec, q) +} + +func (q derivedInsertQuery) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { + return q.hooks.RunHooks(ctx, exec) +} + +func (q derivedInsertQuery) GetLoaders() []bob.Loader { return q.load.GetLoaders() } + +func (q derivedInsertQuery) GetMapperMods() []scan.MapperMod { return q.load.GetMapperMods() } + +func (q derivedInsertQuery) Build(ctx context.Context) (string, []any, error) { + return q.BuildN(ctx, 1) +} + +func (q derivedInsertQuery) BuildN(ctx context.Context, start int) (string, []any, error) { + var sb strings.Builder + args, err := q.WriteQuery(ctx, &sb, start) + if err != nil { + return "", nil, err + } + return sb.String(), args, nil +} + +func (q derivedInsertQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { + if len(q.contextualMods) > 0 { + return q.mutableBase().WriteQuery(ctx, w, start) + } + + var args []any + if len(q.state.With.CTEs) > 0 { + withArgs, err := q.state.With.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, withArgs...) + w.WriteString("\n") + } + + w.WriteString("INSERT INTO ") + tableArgs, err := q.state.Table.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, tableArgs...) + + if q.state.Overriding != "" { + w.WriteString("\nOVERRIDING ") + w.WriteString(q.state.Overriding) + w.WriteString(" VALUE") + } + + w.WriteString("\n") + valArgs, err := q.state.Values.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, valArgs...) + + if q.state.Conflict.Expression != nil { + w.WriteString("\n") + conflictArgs, err := q.state.Conflict.Expression.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, conflictArgs...) + } + + if len(q.state.Returning.Expressions) > 0 { + w.WriteString("\n") + retArgs, err := q.state.Returning.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) + if err != nil { + return nil, err + } + args = append(args, retArgs...) + } + + w.WriteString("\n") + return args, nil +} + +func (q derivedInsertQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { + return q.WriteQuery(ctx, w, start) +} + +func (q derivedInsertQuery) mutableBase() bob.BaseQuery[*psqldialect.InsertQuery] { + values := clause.Values{ + Query: q.state.Values.Query, + Vals: append([]clause.Value(nil), q.state.Values.Vals...), + } + + mutable := &psqldialect.InsertQuery{ + With: q.state.With, + Overriding: q.state.Overriding, + TableRef: q.state.Table, + Values: values, + Conflict: q.state.Conflict, + Returning: q.state.Returning, + Load: q.load, + EmbeddedHook: q.hooks, + ContextualModdable: bob.ContextualModdable[*psqldialect.InsertQuery]{ + Mods: append([]bob.ContextualMod[*psqldialect.InsertQuery](nil), q.contextualMods...), + }, + } + + return bob.BaseQuery[*psqldialect.InsertQuery]{ + Expression: mutable, + Dialect: psqldialect.Dialect, + QueryType: bob.QueryTypeInsert, + } +} + +func (s immutableInsertState) withMods(queryMods ...bob.Mod[*psqldialect.InsertQuery]) (immutableInsertState, bool) { + next := s + var cloneReturning, cloneVals bool + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Returning[*psqldialect.InsertQuery]: + if !cloneReturning { + next.Returning.Expressions = append([]any(nil), s.Returning.Expressions...) + cloneReturning = true + } + next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + case mods.Values[*psqldialect.InsertQuery]: + if !cloneVals { + next.Values.Vals = append([]clause.Value(nil), s.Values.Vals...) + cloneVals = true + } + next.Values.Vals = append(next.Values.Vals, clause.Value(m)) + case mods.Rows[*psqldialect.InsertQuery]: + if !cloneVals { + next.Values.Vals = append([]clause.Value(nil), s.Values.Vals...) + cloneVals = true + } + for _, row := range m { + next.Values.Vals = append(next.Values.Vals, clause.Value(row)) + } + default: + return next, false + } + } + + return next, true +} diff --git a/dialect/psql/immutable_write_test.go b/dialect/psql/immutable_write_test.go new file mode 100644 index 00000000..a4b763b8 --- /dev/null +++ b/dialect/psql/immutable_write_test.go @@ -0,0 +1,208 @@ +package psql + +import ( + "context" + "testing" + + "github.com/stephenafamo/bob/dialect/psql/dm" + "github.com/stephenafamo/bob/dialect/psql/im" + "github.com/stephenafamo/bob/dialect/psql/um" +) + +func TestUpdateWithDoesNotMutateOriginal(t *testing.T) { + base := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + ) + + derived := base.With( + um.Where(Quote("kind").EQ(Arg("Drama"))), + um.Returning("id"), + ) + + baseSQL, _, err := base.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "UPDATE films SET\n\"kind\" = $1" { + t.Fatalf("base update changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "UPDATE films SET\n\"kind\" = $1\nWHERE (\"kind\" = $2)\nRETURNING id" { + t.Fatalf("derived update mismatch: %#v", derivedSQL) + } +} + +func TestDeleteWithDoesNotMutateOriginal(t *testing.T) { + base := Delete( + dm.From("films"), + ) + + derived := base.With( + dm.Where(Quote("kind").EQ(Arg("Drama"))), + dm.Returning("id"), + ) + + baseSQL, _, err := base.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "DELETE FROM films" { + t.Fatalf("base delete changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "DELETE FROM films\nWHERE (\"kind\" = $1)\nRETURNING id" { + t.Fatalf("derived delete mismatch: %#v", derivedSQL) + } +} + +func TestInsertWithDoesNotMutateOriginal(t *testing.T) { + base := Insert( + im.Into("films"), + im.Values(Arg("UA502", "Bananas")), + ) + + derived := base.With( + im.Returning("id"), + ) + + baseSQL, _, err := base.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "INSERT INTO films\nVALUES ($1, $2)\n" { + t.Fatalf("base insert changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(context.Background()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "INSERT INTO films\nVALUES ($1, $2)\nRETURNING id\n" { + t.Fatalf("derived insert mismatch: %#v", derivedSQL) + } +} + +func BenchmarkUpdateQueryApplyMain(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + ) + q.Apply( + um.Where(Quote("kind").EQ(Arg("Drama"))), + um.Returning("id"), + ) + + if _, _, err := q.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkUpdateQueryImmutableNativeHotPath(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + ) + derived := q.With( + um.Where(Quote("kind").EQ(Arg("Drama"))), + um.Returning("id"), + ) + + if _, _, err := derived.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkDeleteQueryApplyMain(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Delete( + dm.From("films"), + ) + q.Apply( + dm.Where(Quote("kind").EQ(Arg("Drama"))), + dm.Returning("id"), + ) + + if _, _, err := q.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkDeleteQueryImmutableNativeHotPath(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Delete( + dm.From("films"), + ) + derived := q.With( + dm.Where(Quote("kind").EQ(Arg("Drama"))), + dm.Returning("id"), + ) + + if _, _, err := derived.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkInsertQueryApplyMain(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Insert( + im.Into("films"), + im.Values(Arg("UA502", "Bananas")), + ) + q.Apply( + im.Returning("id"), + ) + + if _, _, err := q.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkInsertQueryImmutableNativeHotPath(b *testing.B) { + ctx := context.Background() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Insert( + im.Into("films"), + im.Values(Arg("UA502", "Bananas")), + ) + derived := q.With( + im.Returning("id"), + ) + + if _, _, err := derived.Build(ctx); err != nil { + b.Fatal(err) + } + } +} diff --git a/dialect/psql/insert.go b/dialect/psql/insert.go index 8e7a1fd5..e2b9a1c3 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -5,15 +5,25 @@ import ( "github.com/stephenafamo/bob/dialect/psql/dialect" ) -func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) bob.BaseQuery[*dialect.InsertQuery] { +type InsertQuery struct { + bob.BaseQuery[*dialect.InsertQuery] +} + +func (q InsertQuery) With(queryMods ...bob.Mod[*dialect.InsertQuery]) derivedInsertQuery { + return asImmutableInsert(q.BaseQuery).With(queryMods...) +} + +func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { q := &dialect.InsertQuery{} for _, mod := range queryMods { mod.Apply(q) } - return bob.BaseQuery[*dialect.InsertQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeInsert, + return InsertQuery{ + BaseQuery: bob.BaseQuery[*dialect.InsertQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeInsert, + }, } } diff --git a/dialect/psql/table.go b/dialect/psql/table.go index 63cc5f03..5c48fd1e 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -76,7 +76,7 @@ func (t *Table[T, Tslice, Tset, C]) PrimaryKey() expr.ColumnsExpr { func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { q := &ormInsertQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ - BaseQuery: Insert(im.Into(t.NameAs(), t.nonGeneratedCols...)), + BaseQuery: Insert(im.Into(t.NameAs(), t.nonGeneratedCols...)).BaseQuery, Hooks: &t.InsertQueryHooks, }, Scanner: t.scanner, @@ -100,7 +100,7 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { q := &ormUpdateQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ - BaseQuery: Update(um.Table(t.NameAs())), + BaseQuery: Update(um.Table(t.NameAs())).BaseQuery, Hooks: &t.UpdateQueryHooks, }, Scanner: t.scanner, @@ -124,7 +124,7 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { q := &ormDeleteQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ - BaseQuery: Delete(dm.From(t.NameAs())), + BaseQuery: Delete(dm.From(t.NameAs())).BaseQuery, Hooks: &t.DeleteQueryHooks, }, Scanner: t.scanner, diff --git a/dialect/psql/update.go b/dialect/psql/update.go index 8df715df..9d20d934 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -5,15 +5,25 @@ import ( "github.com/stephenafamo/bob/dialect/psql/dialect" ) -func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) bob.BaseQuery[*dialect.UpdateQuery] { +type UpdateQuery struct { + bob.BaseQuery[*dialect.UpdateQuery] +} + +func (q UpdateQuery) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) derivedUpdateQuery { + return asImmutableUpdate(q.BaseQuery).With(queryMods...) +} + +func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { q := &dialect.UpdateQuery{} for _, mod := range queryMods { mod.Apply(q) } - return bob.BaseQuery[*dialect.UpdateQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeUpdate, + return UpdateQuery{ + BaseQuery: bob.BaseQuery[*dialect.UpdateQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeUpdate, + }, } } From 230c162826abf79d60948cdc5587350eb8ccfd71 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 14:03:18 -0400 Subject: [PATCH 05/33] test(psql): add with regression coverage --- dialect/psql/with_regression_test.go | 227 +++++++++++++++++++++++++++ 1 file changed, 227 insertions(+) create mode 100644 dialect/psql/with_regression_test.go diff --git a/dialect/psql/with_regression_test.go b/dialect/psql/with_regression_test.go new file mode 100644 index 00000000..8c2933e9 --- /dev/null +++ b/dialect/psql/with_regression_test.go @@ -0,0 +1,227 @@ +package psql_test + +import ( + "context" + "testing" + + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/dialect/psql" + "github.com/stephenafamo/bob/dialect/psql/dm" + "github.com/stephenafamo/bob/dialect/psql/im" + "github.com/stephenafamo/bob/dialect/psql/sm" + "github.com/stephenafamo/bob/dialect/psql/um" + "github.com/stephenafamo/bob/expr" + testutils "github.com/stephenafamo/bob/test/utils" +) + +type withTestStruct struct { + ID int64 `db:"id,pk"` + Name string `db:"name"` + Email string `db:"email"` +} + +var withTestStructView = psql.NewView[*withTestStruct, bob.Expression]( + "public", + "with_test_struct", + expr.ColsForStruct[withTestStruct]("with_test_struct"), +) + +func TestSelectWithRegression(t *testing.T) { + t.Run("native path matches direct construction", func(t *testing.T) { + base := psql.Select( + sm.Columns("id", "name"), + sm.From("users"), + sm.Where(psql.Quote("tenant_id").EQ(psql.Arg(42))), + ) + + derived := base.With( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + assertQueriesEqual(t, base, psql.Select( + sm.Columns("id", "name"), + sm.From("users"), + sm.Where(psql.Quote("tenant_id").EQ(psql.Arg(42))), + )) + assertQueriesEqual(t, derived, psql.Select( + sm.Columns("id", "name"), + sm.From("users"), + sm.Where(psql.Quote("tenant_id").EQ(psql.Arg(42))), + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + )) + }) + + t.Run("fallback path matches direct construction", func(t *testing.T) { + base := psql.Select( + sm.Columns("users.id", "users.name"), + sm.From("users"), + ) + + derived := base.With( + sm.LeftJoin("teams").Using("id"), + sm.ForUpdate("users").SkipLocked(), + ) + + assertQueriesEqual(t, base, psql.Select( + sm.Columns("users.id", "users.name"), + sm.From("users"), + )) + assertQueriesEqual(t, derived, psql.Select( + sm.Columns("users.id", "users.name"), + sm.From("users"), + sm.LeftJoin("teams").Using("id"), + sm.ForUpdate("users").SkipLocked(), + )) + }) +} + +func TestViewQueryWithRegression(t *testing.T) { + base := withTestStructView.Query( + sm.Where(psql.Quote("id").GT(psql.Arg(0))), + ) + + derived := base.With( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + assertQueriesEqual(t, base, withTestStructView.Query( + sm.Where(psql.Quote("id").GT(psql.Arg(0))), + )) + assertQueriesEqual(t, derived.Query, withTestStructView.Query( + sm.Where(psql.Quote("id").GT(psql.Arg(0))), + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + )) +} + +func TestUpdateWithRegression(t *testing.T) { + base := psql.Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + ) + + derived := base.With( + um.SetCol("updated_at").To("NOW()"), + um.Where(psql.Quote("kind").EQ(psql.Arg("Drama"))), + um.Returning("id"), + ) + + assertQueriesEqual(t, base, psql.Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + )) + assertQueriesEqual(t, derived, psql.Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + um.SetCol("updated_at").To("NOW()"), + um.Where(psql.Quote("kind").EQ(psql.Arg("Drama"))), + um.Returning("id"), + )) +} + +func TestDeleteWithRegression(t *testing.T) { + base := psql.Delete( + dm.From("employees"), + ) + + derived := base.With( + dm.Using("accounts"), + dm.Where(psql.Quote("accounts", "name").EQ(psql.Arg("Acme Corporation"))), + dm.Where(psql.Quote("employees", "id").EQ(psql.Quote("accounts", "sales_person"))), + dm.Returning("id"), + ) + + assertQueriesEqual(t, base, psql.Delete( + dm.From("employees"), + )) + assertQueriesEqual(t, derived, psql.Delete( + dm.From("employees"), + dm.Using("accounts"), + dm.Where(psql.Quote("accounts", "name").EQ(psql.Arg("Acme Corporation"))), + dm.Where(psql.Quote("employees", "id").EQ(psql.Quote("accounts", "sales_person"))), + dm.Returning("id"), + )) +} + +func TestInsertWithRegression(t *testing.T) { + t.Run("native path matches direct construction", func(t *testing.T) { + base := psql.Insert( + im.Into("films"), + im.Values(psql.Arg("UA502", "Bananas")), + ) + + derived := base.With( + im.Returning("id"), + ) + + assertQueriesEqual(t, base, psql.Insert( + im.Into("films"), + im.Values(psql.Arg("UA502", "Bananas")), + )) + assertQueriesEqual(t, derived, psql.Insert( + im.Into("films"), + im.Values(psql.Arg("UA502", "Bananas")), + im.Returning("id"), + )) + }) + + t.Run("fallback path matches direct construction", func(t *testing.T) { + base := psql.Insert( + im.IntoAs("distributors", "d", "did", "dname"), + im.Values(psql.Arg(8, "Anvil Distribution")), + ) + + derived := base.With( + im.OnConflict("did").DoUpdate( + im.SetExcluded("dname"), + im.Where(psql.Quote("d", "zipcode").NE(psql.S("21201"))), + ), + ) + + assertQueriesEqual(t, base, psql.Insert( + im.IntoAs("distributors", "d", "did", "dname"), + im.Values(psql.Arg(8, "Anvil Distribution")), + )) + assertQueriesEqual(t, derived, psql.Insert( + im.IntoAs("distributors", "d", "did", "dname"), + im.Values(psql.Arg(8, "Anvil Distribution")), + im.OnConflict("did").DoUpdate( + im.SetExcluded("dname"), + im.Where(psql.Quote("d", "zipcode").NE(psql.S("21201"))), + ), + )) + }) +} + +func assertQueriesEqual(t *testing.T, got bob.Query, want bob.Query) { + t.Helper() + + gotSQL, gotArgs, err := bob.Build(context.Background(), got) + if err != nil { + t.Fatalf("build got: %v", err) + } + + wantSQL, wantArgs, err := bob.Build(context.Background(), want) + if err != nil { + t.Fatalf("build want: %v", err) + } + + diff, err := testutils.QueryDiff(wantSQL, gotSQL, formatter) + if err != nil { + t.Fatalf("query diff error: %v", err) + } + if diff != "" { + t.Fatalf("sql diff: %s", diff) + } + + if diff := testutils.ArgsDiff(wantArgs, gotArgs); diff != "" { + t.Fatalf("args diff: %s", diff) + } +} From 835cb441f95d4f43eeb0c7506435bee31cad950b Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 16:25:55 -0400 Subject: [PATCH 06/33] test(psql): use test contexts in query tests --- dialect/psql/immutable_select_test.go | 17 ++++++++--------- dialect/psql/immutable_write_test.go | 25 ++++++++++++------------- dialect/psql/view_test.go | 3 +-- dialect/psql/with_regression_test.go | 5 ++--- 4 files changed, 23 insertions(+), 27 deletions(-) diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 92717138..80d7cf9f 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -1,7 +1,6 @@ package psql import ( - "context" "testing" "github.com/stephenafamo/bob" @@ -24,7 +23,7 @@ func TestImmutableSelectQueryWithDoesNotMutateOriginal(t *testing.T) { t.Fatalf("expected derived query type %q, got %q", bob.QueryTypeSelect, derived.Type()) } - baseSQL, _, err := base.Build(context.Background()) + baseSQL, _, err := base.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -32,7 +31,7 @@ func TestImmutableSelectQueryWithDoesNotMutateOriginal(t *testing.T) { t.Fatalf("base query changed unexpectedly: %#v", baseSQL) } - derivedSQL, _, err := derived.Build(context.Background()) + derivedSQL, _, err := derived.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -56,7 +55,7 @@ func TestImmutableViewQueryWithDoesNotMutateOriginal(t *testing.T) { t.Fatal("expected derived view query to preserve scanner") } - baseSQL, _, err := base.Build(context.Background()) + baseSQL, _, err := base.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -64,7 +63,7 @@ func TestImmutableViewQueryWithDoesNotMutateOriginal(t *testing.T) { t.Fatalf("base view query changed unexpectedly: %#v", baseSQL) } - derivedSQL, _, err := derived.Build(context.Background()) + derivedSQL, _, err := derived.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -74,7 +73,7 @@ func TestImmutableViewQueryWithDoesNotMutateOriginal(t *testing.T) { } func BenchmarkBaseQueryApplyMain(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -96,7 +95,7 @@ func BenchmarkBaseQueryApplyMain(b *testing.B) { } func BenchmarkBaseQueryImmutableNativeHotPath(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -118,7 +117,7 @@ func BenchmarkBaseQueryImmutableNativeHotPath(b *testing.B) { } func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -143,7 +142,7 @@ func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { } func BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { diff --git a/dialect/psql/immutable_write_test.go b/dialect/psql/immutable_write_test.go index a4b763b8..cfe7702b 100644 --- a/dialect/psql/immutable_write_test.go +++ b/dialect/psql/immutable_write_test.go @@ -1,7 +1,6 @@ package psql import ( - "context" "testing" "github.com/stephenafamo/bob/dialect/psql/dm" @@ -20,7 +19,7 @@ func TestUpdateWithDoesNotMutateOriginal(t *testing.T) { um.Returning("id"), ) - baseSQL, _, err := base.Build(context.Background()) + baseSQL, _, err := base.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -28,7 +27,7 @@ func TestUpdateWithDoesNotMutateOriginal(t *testing.T) { t.Fatalf("base update changed unexpectedly: %#v", baseSQL) } - derivedSQL, _, err := derived.Build(context.Background()) + derivedSQL, _, err := derived.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -47,7 +46,7 @@ func TestDeleteWithDoesNotMutateOriginal(t *testing.T) { dm.Returning("id"), ) - baseSQL, _, err := base.Build(context.Background()) + baseSQL, _, err := base.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -55,7 +54,7 @@ func TestDeleteWithDoesNotMutateOriginal(t *testing.T) { t.Fatalf("base delete changed unexpectedly: %#v", baseSQL) } - derivedSQL, _, err := derived.Build(context.Background()) + derivedSQL, _, err := derived.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -74,7 +73,7 @@ func TestInsertWithDoesNotMutateOriginal(t *testing.T) { im.Returning("id"), ) - baseSQL, _, err := base.Build(context.Background()) + baseSQL, _, err := base.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -82,7 +81,7 @@ func TestInsertWithDoesNotMutateOriginal(t *testing.T) { t.Fatalf("base insert changed unexpectedly: %#v", baseSQL) } - derivedSQL, _, err := derived.Build(context.Background()) + derivedSQL, _, err := derived.Build(t.Context()) if err != nil { t.Fatal(err) } @@ -92,7 +91,7 @@ func TestInsertWithDoesNotMutateOriginal(t *testing.T) { } func BenchmarkUpdateQueryApplyMain(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -112,7 +111,7 @@ func BenchmarkUpdateQueryApplyMain(b *testing.B) { } func BenchmarkUpdateQueryImmutableNativeHotPath(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -132,7 +131,7 @@ func BenchmarkUpdateQueryImmutableNativeHotPath(b *testing.B) { } func BenchmarkDeleteQueryApplyMain(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -151,7 +150,7 @@ func BenchmarkDeleteQueryApplyMain(b *testing.B) { } func BenchmarkDeleteQueryImmutableNativeHotPath(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -170,7 +169,7 @@ func BenchmarkDeleteQueryImmutableNativeHotPath(b *testing.B) { } func BenchmarkInsertQueryApplyMain(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { @@ -189,7 +188,7 @@ func BenchmarkInsertQueryApplyMain(b *testing.B) { } func BenchmarkInsertQueryImmutableNativeHotPath(b *testing.B) { - ctx := context.Background() + ctx := b.Context() b.ReportAllocs() for i := 0; i < b.N; i++ { diff --git a/dialect/psql/view_test.go b/dialect/psql/view_test.go index 05c2bb5f..468c82ce 100644 --- a/dialect/psql/view_test.go +++ b/dialect/psql/view_test.go @@ -2,7 +2,6 @@ package psql import ( "bytes" - "context" "testing" _ "github.com/lib/pq" @@ -57,7 +56,7 @@ func TestSomeViewQuery(t *testing.T) { func selectToString(t *testing.T, query bob.Query, argsLen int) string { t.Helper() - ctx := context.Background() + ctx := t.Context() buf := new(bytes.Buffer) args, err := query.WriteQuery(ctx, buf, 0) if err != nil { diff --git a/dialect/psql/with_regression_test.go b/dialect/psql/with_regression_test.go index 8c2933e9..d743dc18 100644 --- a/dialect/psql/with_regression_test.go +++ b/dialect/psql/with_regression_test.go @@ -1,7 +1,6 @@ package psql_test import ( - "context" "testing" "github.com/stephenafamo/bob" @@ -203,12 +202,12 @@ func TestInsertWithRegression(t *testing.T) { func assertQueriesEqual(t *testing.T, got bob.Query, want bob.Query) { t.Helper() - gotSQL, gotArgs, err := bob.Build(context.Background(), got) + gotSQL, gotArgs, err := bob.Build(t.Context(), got) if err != nil { t.Fatalf("build got: %v", err) } - wantSQL, wantArgs, err := bob.Build(context.Background(), want) + wantSQL, wantArgs, err := bob.Build(t.Context(), want) if err != nil { t.Fatalf("build want: %v", err) } From e8aafe0ac943df5cac992354724319327f9117af Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 22:28:07 -0400 Subject: [PATCH 07/33] refactor(psql): make queries immutable by default Apply now matches With for psql select/view/insert/update/delete wrappers, so derivation no longer mutates the source query. Added regression coverage for immutable Apply semantics and preserved fallback SQL behavior for combined/distinct select shapes. Benchmarks versus stored upstream/main baselines: - BaseQuery immutable: 2404 ns/op vs 3039 ns/op upstream Apply - View count+paginate immutable: 5906 ns/op vs 10070 ns/op upstream Apply - Update immutable: 2147 ns/op vs 3038 ns/op upstream Apply - Delete immutable: 1587 ns/op vs 1926 ns/op upstream Apply - Insert immutable: 1295 ns/op vs 1473 ns/op upstream Apply --- dialect/psql/delete.go | 33 ++++-- dialect/psql/immutable_select.go | 140 ++++++++++++++------------ dialect/psql/immutable_select_test.go | 114 ++++++++++++++++++++- dialect/psql/immutable_write_test.go | 88 +++++++++++++++- dialect/psql/insert.go | 33 ++++-- dialect/psql/select.go | 33 ++++-- dialect/psql/table.go | 6 +- dialect/psql/update.go | 33 ++++-- dialect/psql/view.go | 111 ++++++++++++++++---- dialect/psql/view_test.go | 2 +- 10 files changed, 462 insertions(+), 131 deletions(-) diff --git a/dialect/psql/delete.go b/dialect/psql/delete.go index 9b023db7..8c656580 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -6,11 +6,25 @@ import ( ) type DeleteQuery struct { - bob.BaseQuery[*dialect.DeleteQuery] + derivedDeleteQuery + materialized *bob.BaseQuery[*dialect.DeleteQuery] } -func (q DeleteQuery) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) derivedDeleteQuery { - return asImmutableDelete(q.BaseQuery).With(queryMods...) +func (q DeleteQuery) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { + q.derivedDeleteQuery = q.derivedDeleteQuery.With(queryMods...) + q.materialized = nil + return q +} + +func (q DeleteQuery) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { + return q.With(queryMods...) +} + +func (q DeleteQuery) baseQuery() bob.BaseQuery[*dialect.DeleteQuery] { + if q.materialized != nil { + return *q.materialized + } + return q.derivedDeleteQuery.mutableBase() } func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { @@ -19,11 +33,14 @@ func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { mod.Apply(q) } + base := bob.BaseQuery[*dialect.DeleteQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeDelete, + } + return DeleteQuery{ - BaseQuery: bob.BaseQuery[*dialect.DeleteQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeDelete, - }, + derivedDeleteQuery: asImmutableDelete(base), + materialized: &base, } } diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index 73d62f12..e8d1ed65 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -12,12 +12,14 @@ import ( "github.com/stephenafamo/bob/clause" psqldialect "github.com/stephenafamo/bob/dialect/psql/dialect" "github.com/stephenafamo/bob/mods" - "github.com/stephenafamo/bob/orm" "github.com/stephenafamo/scan" ) type derivedSelectQuery struct { - state immutableSelectState + state immutableSelectState + load bob.Load + hooks bob.EmbeddedHook + contextualMods []bob.ContextualMod[*psqldialect.SelectQuery] } type immutableSelectState struct { @@ -44,7 +46,13 @@ type immutableSelectState struct { } func asImmutable(q bob.BaseQuery[*psqldialect.SelectQuery]) derivedSelectQuery { - return derivedSelectQuery{state: immutableStateFromMutable(q.Expression)} + return derivedSelectQuery{ + state: immutableStateFromMutable(q.Expression), + load: q.Expression.Load, + hooks: q.Expression.EmbeddedHook, + contextualMods: append([]bob.ContextualMod[*psqldialect.SelectQuery](nil), + q.Expression.ContextualModdable.Mods...), + } } func (q derivedSelectQuery) Type() bob.QueryType { @@ -54,15 +62,17 @@ func (q derivedSelectQuery) Type() bob.QueryType { func (q derivedSelectQuery) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) derivedSelectQuery { next, ok := q.state.withMods(queryMods...) if ok { - return derivedSelectQuery{state: next} + q.state = next + return q } - mutable := q.state.toMutable() + base := q.mutableBase() + mutable := base.Expression for _, mod := range queryMods { - mod.Apply(&mutable) + mod.Apply(mutable) } - return derivedSelectQuery{state: immutableStateFromMutable(&mutable)} + return asImmutable(base) } func (q derivedSelectQuery) AsCount() derivedSelectQuery { @@ -94,6 +104,10 @@ func (q derivedSelectQuery) BuildN(ctx context.Context, start int) (string, []an } func (q derivedSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { + if len(q.contextualMods) > 0 || !q.state.supportsNativeWrite() { + return q.mutableBase().WriteQuery(ctx, w, start) + } + writer := immutableSelectWriter{ ctx: ctx, w: w, @@ -117,76 +131,55 @@ func (q derivedSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ b return args, nil } -type derivedViewQuery[T any, Ts ~[]T] struct { - Query derivedSelectQuery - Scanner scan.Mapper[T] - Hooks *bob.Hooks[*psqldialect.SelectQuery, bob.SkipQueryHooksKey] -} - -func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) derivedViewQuery[T, Ts] { - state := immutableStateFromMutable(q.BaseQuery.Expression) - if len(state.SelectColumns) == 0 && q.defaultSelect != nil { - state.SelectColumns = append(state.SelectColumns, q.defaultSelect) - } - - return derivedViewQuery[T, Ts]{ - Query: derivedSelectQuery{state: state}.With(queryMods...), - Scanner: q.Scanner, - Hooks: q.Hooks, - } +func (q derivedSelectQuery) Exec(ctx context.Context, exec bob.Executor) (sql.Result, error) { + return bob.Exec(ctx, exec, q) } -func (q derivedViewQuery[T, Ts]) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) derivedViewQuery[T, Ts] { - q.Query = q.Query.With(queryMods...) - return q +func (q derivedSelectQuery) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { + return q.hooks.RunHooks(ctx, exec) } -func (q derivedViewQuery[T, Ts]) One(ctx context.Context, exec bob.Executor) (T, error) { - return q.mutable().One(ctx, exec) +func (q derivedSelectQuery) GetLoaders() []bob.Loader { + return q.load.GetLoaders() } -func (q derivedViewQuery[T, Ts]) All(ctx context.Context, exec bob.Executor) (Ts, error) { - return q.mutable().All(ctx, exec) +func (q derivedSelectQuery) GetMapperMods() []scan.MapperMod { + return q.load.GetMapperMods() } -func (q derivedViewQuery[T, Ts]) Cursor(ctx context.Context, exec bob.Executor) (scan.ICursor[T], error) { - return q.mutable().Cursor(ctx, exec) -} +func (q derivedSelectQuery) mutableBase() bob.BaseQuery[*psqldialect.SelectQuery] { + mutable := &psqldialect.SelectQuery{ + With: q.state.With, + SelectList: clause.SelectList{Columns: q.state.SelectColumns, PreloadColumns: q.state.PreloadColumns}, + Distinct: q.state.Distinct, + TableRef: q.state.TableRef, + Where: q.state.Where, + GroupBy: q.state.GroupBy, + Having: q.state.Having, + Windows: q.state.Windows, + Combines: q.state.Combines, + OrderBy: q.state.OrderBy, + Limit: q.state.Limit, + Offset: q.state.Offset, + Fetch: q.state.Fetch, + Locks: q.state.Locks, -func (q derivedViewQuery[T, Ts]) Each(ctx context.Context, exec bob.Executor) (func(func(T, error) bool), error) { - return q.mutable().Each(ctx, exec) -} + Load: q.load, + EmbeddedHook: q.hooks, + ContextualModdable: bob.ContextualModdable[*psqldialect.SelectQuery]{ + Mods: append([]bob.ContextualMod[*psqldialect.SelectQuery](nil), q.contextualMods...), + }, -func (q derivedViewQuery[T, Ts]) Count(ctx context.Context, exec bob.Executor) (int64, error) { - mq := q.mutable() - ctx, err := mq.RunHooks(ctx, exec) - if err != nil { - return 0, err + CombinedOrder: q.state.CombinedOrder, + CombinedLimit: q.state.CombinedLimit, + CombinedFetch: q.state.CombinedFetch, + CombinedOffset: q.state.CombinedOffset, } - return bob.One(ctx, exec, asCountQuery(mq.BaseQuery), scan.SingleColumnMapper[int64]) -} - -func (q derivedViewQuery[T, Ts]) Exists(ctx context.Context, exec bob.Executor) (bool, error) { - count, err := q.Count(ctx, exec) - return count > 0, err -} -func (q derivedViewQuery[T, Ts]) Build(ctx context.Context) (string, []any, error) { - return q.Query.Build(ctx) -} - -func (q derivedViewQuery[T, Ts]) mutable() orm.Query[*psqldialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] { - mutable := q.Query.state.toMutable() - return orm.Query[*psqldialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]]{ - ExecQuery: orm.ExecQuery[*psqldialect.SelectQuery]{ - BaseQuery: bob.BaseQuery[*psqldialect.SelectQuery]{ - Expression: &mutable, - Dialect: psqldialect.Dialect, - QueryType: bob.QueryTypeSelect, - }, - Hooks: q.Hooks, - }, - Scanner: q.Scanner, + return bob.BaseQuery[*psqldialect.SelectQuery]{ + Expression: mutable, + Dialect: psqldialect.Dialect, + QueryType: bob.QueryTypeSelect, } } @@ -198,7 +191,7 @@ func immutableStateFromMutable(q *psqldialect.SelectQuery) immutableSelectState }, SelectColumns: append([]any(nil), q.SelectList.Columns...), PreloadColumns: append([]any(nil), q.SelectList.PreloadColumns...), - Distinct: psqldialect.Distinct{On: append([]any(nil), q.Distinct.On...)}, + Distinct: psqldialect.Distinct{On: cloneAnySlice(q.Distinct.On)}, TableRef: cloneTableRef(q.TableRef), Where: clause.Where{Conditions: append([]any(nil), q.Where.Conditions...)}, GroupBy: clause.GroupBy{ @@ -241,6 +234,21 @@ func immutableStateFromMutable(q *psqldialect.SelectQuery) immutableSelectState } } +func (s immutableSelectState) supportsNativeWrite() bool { + return len(s.Combines.Queries) == 0 && + len(s.CombinedOrder.Expressions) == 0 && + s.CombinedLimit.Count == nil && + s.CombinedOffset.Count == nil && + s.CombinedFetch.Count == nil +} + +func cloneAnySlice(values []any) []any { + if values == nil { + return nil + } + return append(make([]any, 0, len(values)), values...) +} + func (s immutableSelectState) toMutable() psqldialect.SelectQuery { return psqldialect.SelectQuery{ With: s.With, diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 80d7cf9f..20478c1a 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -72,6 +72,112 @@ func TestImmutableViewQueryWithDoesNotMutateOriginal(t *testing.T) { } } +func TestImmutableSelectQueryApplyDoesNotMutateOriginal(t *testing.T) { + base := Select( + sm.Columns("id"), + sm.From("users"), + ) + + derived := base.Apply( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "SELECT \nid\nFROM users\n" { + t.Fatalf("base query changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "SELECT \nid\nFROM users\nORDER BY id DESC\nLIMIT 10\nOFFSET 20\n" { + t.Fatalf("derived query mismatch: %#v", derivedSQL) + } +} + +func TestImmutableSelectQueryApplyFallbackDoesNotMutateOriginal(t *testing.T) { + base := Select( + sm.Columns("id", "name"), + sm.From("users"), + ) + + derived := base.Apply( + sm.Distinct(), + sm.Union(Select( + sm.Columns("id", "name"), + sm.From("admins"), + )), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "SELECT \nid, name\nFROM users\n" { + t.Fatalf("base query changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, derivedArgs, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expected := Select( + sm.Columns("id", "name"), + sm.From("users"), + sm.Distinct(), + sm.Union(Select( + sm.Columns("id", "name"), + sm.From("admins"), + )), + ) + + expectedSQL, expectedArgs, err := expected.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != expectedSQL { + t.Fatalf("derived fallback query mismatch: got %#v want %#v", derivedSQL, expectedSQL) + } + if len(derivedArgs) != len(expectedArgs) { + t.Fatalf("derived fallback args mismatch: got %d want %d", len(derivedArgs), len(expectedArgs)) + } +} + +func TestImmutableViewQueryApplyDoesNotMutateOriginal(t *testing.T) { + base := someStructView.Query( + sm.Where(Quote("id").GT(Arg(0))), + ) + + derived := base.Apply( + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "SELECT \n\"some_struct\".\"id\" AS \"id\", \"some_struct\".\"name\" AS \"name\", \"some_struct\".\"email\" AS \"email\"\nFROM \"public\".\"some_struct\" AS \"public.some_struct\"\nWHERE (\"id\" > $1)\n" { + t.Fatalf("base view query changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "SELECT \n\"some_struct\".\"id\" AS \"id\", \"some_struct\".\"name\" AS \"name\", \"some_struct\".\"email\" AS \"email\"\nFROM \"public\".\"some_struct\" AS \"public.some_struct\"\nWHERE (\"id\" > $1)\nORDER BY id DESC\nLIMIT 10\nOFFSET 20\n" { + t.Fatalf("derived view query mismatch: %#v", derivedSQL) + } +} + func BenchmarkBaseQueryApplyMain(b *testing.B) { ctx := b.Context() @@ -82,7 +188,7 @@ func BenchmarkBaseQueryApplyMain(b *testing.B) { sm.From("users"), sm.Where(Quote("tenant_id").EQ(Arg(42))), ) - q.Apply( + q = q.Apply( sm.OrderBy("id").Desc(), sm.Limit(10), sm.Offset(20), @@ -125,11 +231,11 @@ func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := asCountQuery(q.BaseQuery).Build(ctx); err != nil { + if _, _, err := asCountQuery(q.baseQuery()).Build(ctx); err != nil { b.Fatal(err) } - q.Apply( + q = q.Apply( sm.OrderBy("id").Desc(), sm.Limit(10), sm.Offset(20), @@ -150,7 +256,7 @@ func BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := asCountQuery(q.BaseQuery).Build(ctx); err != nil { + if _, _, err := q.Query.derivedSelectQuery.AsCount().Build(ctx); err != nil { b.Fatal(err) } diff --git a/dialect/psql/immutable_write_test.go b/dialect/psql/immutable_write_test.go index cfe7702b..4032f798 100644 --- a/dialect/psql/immutable_write_test.go +++ b/dialect/psql/immutable_write_test.go @@ -90,6 +90,88 @@ func TestInsertWithDoesNotMutateOriginal(t *testing.T) { } } +func TestUpdateApplyDoesNotMutateOriginal(t *testing.T) { + base := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + ) + + derived := base.Apply( + um.Where(Quote("kind").EQ(Arg("Drama"))), + um.Returning("id"), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "UPDATE films SET\n\"kind\" = $1" { + t.Fatalf("base update changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "UPDATE films SET\n\"kind\" = $1\nWHERE (\"kind\" = $2)\nRETURNING id" { + t.Fatalf("derived update mismatch: %#v", derivedSQL) + } +} + +func TestDeleteApplyDoesNotMutateOriginal(t *testing.T) { + base := Delete( + dm.From("films"), + ) + + derived := base.Apply( + dm.Where(Quote("kind").EQ(Arg("Drama"))), + dm.Returning("id"), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "DELETE FROM films" { + t.Fatalf("base delete changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "DELETE FROM films\nWHERE (\"kind\" = $1)\nRETURNING id" { + t.Fatalf("derived delete mismatch: %#v", derivedSQL) + } +} + +func TestInsertApplyDoesNotMutateOriginal(t *testing.T) { + base := Insert( + im.Into("films"), + im.Values(Arg("UA502", "Bananas")), + ) + + derived := base.Apply( + im.Returning("id"), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "INSERT INTO films\nVALUES ($1, $2)\n" { + t.Fatalf("base insert changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != "INSERT INTO films\nVALUES ($1, $2)\nRETURNING id\n" { + t.Fatalf("derived insert mismatch: %#v", derivedSQL) + } +} + func BenchmarkUpdateQueryApplyMain(b *testing.B) { ctx := b.Context() @@ -99,7 +181,7 @@ func BenchmarkUpdateQueryApplyMain(b *testing.B) { um.Table("films"), um.SetCol("kind").ToArg("Dramatic"), ) - q.Apply( + q = q.Apply( um.Where(Quote("kind").EQ(Arg("Drama"))), um.Returning("id"), ) @@ -138,7 +220,7 @@ func BenchmarkDeleteQueryApplyMain(b *testing.B) { q := Delete( dm.From("films"), ) - q.Apply( + q = q.Apply( dm.Where(Quote("kind").EQ(Arg("Drama"))), dm.Returning("id"), ) @@ -177,7 +259,7 @@ func BenchmarkInsertQueryApplyMain(b *testing.B) { im.Into("films"), im.Values(Arg("UA502", "Bananas")), ) - q.Apply( + q = q.Apply( im.Returning("id"), ) diff --git a/dialect/psql/insert.go b/dialect/psql/insert.go index e2b9a1c3..10b84992 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -6,11 +6,25 @@ import ( ) type InsertQuery struct { - bob.BaseQuery[*dialect.InsertQuery] + derivedInsertQuery + materialized *bob.BaseQuery[*dialect.InsertQuery] } -func (q InsertQuery) With(queryMods ...bob.Mod[*dialect.InsertQuery]) derivedInsertQuery { - return asImmutableInsert(q.BaseQuery).With(queryMods...) +func (q InsertQuery) With(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { + q.derivedInsertQuery = q.derivedInsertQuery.With(queryMods...) + q.materialized = nil + return q +} + +func (q InsertQuery) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { + return q.With(queryMods...) +} + +func (q InsertQuery) baseQuery() bob.BaseQuery[*dialect.InsertQuery] { + if q.materialized != nil { + return *q.materialized + } + return q.derivedInsertQuery.mutableBase() } func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { @@ -19,11 +33,14 @@ func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { mod.Apply(q) } + base := bob.BaseQuery[*dialect.InsertQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeInsert, + } + return InsertQuery{ - BaseQuery: bob.BaseQuery[*dialect.InsertQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeInsert, - }, + derivedInsertQuery: asImmutableInsert(base), + materialized: &base, } } diff --git a/dialect/psql/select.go b/dialect/psql/select.go index 36fafadf..87f87932 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -6,11 +6,25 @@ import ( ) type SelectQuery struct { - bob.BaseQuery[*dialect.SelectQuery] + derivedSelectQuery + materialized *bob.BaseQuery[*dialect.SelectQuery] } -func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) derivedSelectQuery { - return asImmutable(q.BaseQuery).With(queryMods...) +func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { + q.derivedSelectQuery = q.derivedSelectQuery.With(queryMods...) + q.materialized = nil + return q +} + +func (q SelectQuery) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { + return q.With(queryMods...) +} + +func (q SelectQuery) baseQuery() bob.BaseQuery[*dialect.SelectQuery] { + if q.materialized != nil { + return *q.materialized + } + return q.derivedSelectQuery.mutableBase() } func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { @@ -19,11 +33,14 @@ func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { mod.Apply(q) } + base := bob.BaseQuery[*dialect.SelectQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeSelect, + } + return SelectQuery{ - BaseQuery: bob.BaseQuery[*dialect.SelectQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeSelect, - }, + derivedSelectQuery: asImmutable(base), + materialized: &base, } } diff --git a/dialect/psql/table.go b/dialect/psql/table.go index 5c48fd1e..a2f71e08 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -76,7 +76,7 @@ func (t *Table[T, Tslice, Tset, C]) PrimaryKey() expr.ColumnsExpr { func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { q := &ormInsertQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ - BaseQuery: Insert(im.Into(t.NameAs(), t.nonGeneratedCols...)).BaseQuery, + BaseQuery: Insert(im.Into(t.NameAs(), t.nonGeneratedCols...)).baseQuery(), Hooks: &t.InsertQueryHooks, }, Scanner: t.scanner, @@ -100,7 +100,7 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { q := &ormUpdateQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ - BaseQuery: Update(um.Table(t.NameAs())).BaseQuery, + BaseQuery: Update(um.Table(t.NameAs())).baseQuery(), Hooks: &t.UpdateQueryHooks, }, Scanner: t.scanner, @@ -124,7 +124,7 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { q := &ormDeleteQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ - BaseQuery: Delete(dm.From(t.NameAs())).BaseQuery, + BaseQuery: Delete(dm.From(t.NameAs())).baseQuery(), Hooks: &t.DeleteQueryHooks, }, Scanner: t.scanner, diff --git a/dialect/psql/update.go b/dialect/psql/update.go index 9d20d934..5e56d09d 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -6,11 +6,25 @@ import ( ) type UpdateQuery struct { - bob.BaseQuery[*dialect.UpdateQuery] + derivedUpdateQuery + materialized *bob.BaseQuery[*dialect.UpdateQuery] } -func (q UpdateQuery) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) derivedUpdateQuery { - return asImmutableUpdate(q.BaseQuery).With(queryMods...) +func (q UpdateQuery) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { + q.derivedUpdateQuery = q.derivedUpdateQuery.With(queryMods...) + q.materialized = nil + return q +} + +func (q UpdateQuery) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { + return q.With(queryMods...) +} + +func (q UpdateQuery) baseQuery() bob.BaseQuery[*dialect.UpdateQuery] { + if q.materialized != nil { + return *q.materialized + } + return q.derivedUpdateQuery.mutableBase() } func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { @@ -19,11 +33,14 @@ func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { mod.Apply(q) } + base := bob.BaseQuery[*dialect.UpdateQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeUpdate, + } + return UpdateQuery{ - BaseQuery: bob.BaseQuery[*dialect.UpdateQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeUpdate, - }, + derivedUpdateQuery: asImmutableUpdate(base), + materialized: &base, } } diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 79cd6c5b..02e8dacf 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -3,6 +3,7 @@ package psql import ( "context" "fmt" + "io" "reflect" "github.com/stephenafamo/bob" @@ -10,6 +11,7 @@ import ( "github.com/stephenafamo/bob/dialect/psql/sm" "github.com/stephenafamo/bob/expr" "github.com/stephenafamo/bob/internal/mappings" + bobmods "github.com/stephenafamo/bob/mods" "github.com/stephenafamo/bob/orm" "github.com/stephenafamo/scan" ) @@ -80,42 +82,68 @@ func (v *View[T, Tslice, C]) Alias() string { // Starts a select query func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Tslice] { q := &ViewQuery[T, Tslice]{ - Query: orm.Query[*dialect.SelectQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ - ExecQuery: orm.ExecQuery[*dialect.SelectQuery]{ - BaseQuery: Select(sm.From(v.NameAs())).BaseQuery, - Hooks: &v.SelectQueryHooks, - }, - Scanner: v.scanner, - }, + Query: Select(sm.Columns(v.Columns), sm.From(v.NameAs())), + Scanner: v.scanner, + Hooks: &v.SelectQueryHooks, defaultSelect: v.Columns, + defaulted: true, + } + if len(queryMods) == 0 { + return q } + return q.Apply(queryMods...) +} - q.Expression.AppendContextualModFunc( - func(ctx context.Context, q *dialect.SelectQuery) (context.Context, error) { - if len(q.SelectList.Columns) == 0 { - q.AppendSelect(v.Columns) - } - return ctx, nil - }, - ) +type ViewQuery[T any, Ts ~[]T] struct { + Query SelectQuery + Scanner scan.Mapper[T] + Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] + defaultSelect bob.Expression + defaulted bool +} + +func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { + next := *q + if hasSelectMod(queryMods) && next.defaulted { + next.Query.derivedSelectQuery.state.SelectColumns = nil + next.defaulted = false + } + next.Query = next.Query.Apply(queryMods...) + return &next +} - q.Apply(queryMods...) +func (q *ViewQuery[T, Ts]) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { + return q.With(queryMods...) +} - return q +func (q *ViewQuery[T, Ts]) Type() bob.QueryType { + return q.Query.Type() } -type ViewQuery[T any, Ts ~[]T] struct { - orm.Query[*dialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] - defaultSelect bob.Expression +func (q *ViewQuery[T, Ts]) Build(ctx context.Context) (string, []any, error) { + return q.Query.Build(ctx) +} + +func (q *ViewQuery[T, Ts]) BuildN(ctx context.Context, start int) (string, []any, error) { + return q.Query.BuildN(ctx, start) +} + +func (q *ViewQuery[T, Ts]) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { + return q.Query.WriteQuery(ctx, w, start) +} + +func (q *ViewQuery[T, Ts]) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { + return q.Query.WriteSQL(ctx, w, d, start) } // Count the number of matching rows func (v *ViewQuery[T, Tslice]) Count(ctx context.Context, exec bob.Executor) (int64, error) { - ctx, err := v.RunHooks(ctx, exec) + mq := v.mutable() + ctx, err := mq.RunHooks(ctx, exec) if err != nil { return 0, err } - return bob.One(ctx, exec, asCountQuery(v.BaseQuery), scan.SingleColumnMapper[int64]) + return bob.One(ctx, exec, asCountQuery(mq.BaseQuery), scan.SingleColumnMapper[int64]) } // Exists checks if there is any matching row @@ -124,6 +152,45 @@ func (v *ViewQuery[T, Tslice]) Exists(ctx context.Context, exec bob.Executor) (b return count > 0, err } +func (q *ViewQuery[T, Ts]) One(ctx context.Context, exec bob.Executor) (T, error) { + return q.mutable().One(ctx, exec) +} + +func (q *ViewQuery[T, Ts]) All(ctx context.Context, exec bob.Executor) (Ts, error) { + return q.mutable().All(ctx, exec) +} + +func (q *ViewQuery[T, Ts]) Cursor(ctx context.Context, exec bob.Executor) (scan.ICursor[T], error) { + return q.mutable().Cursor(ctx, exec) +} + +func (q *ViewQuery[T, Ts]) Each(ctx context.Context, exec bob.Executor) (func(func(T, error) bool), error) { + return q.mutable().Each(ctx, exec) +} + +func hasSelectMod(mods []bob.Mod[*dialect.SelectQuery]) bool { + for _, mod := range mods { + if _, ok := mod.(bobmods.Select[*dialect.SelectQuery]); ok { + return true + } + } + return false +} + +func (q *ViewQuery[T, Ts]) baseQuery() bob.BaseQuery[*dialect.SelectQuery] { + return q.Query.baseQuery() +} + +func (q *ViewQuery[T, Ts]) mutable() orm.Query[*dialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] { + return orm.Query[*dialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]]{ + ExecQuery: orm.ExecQuery[*dialect.SelectQuery]{ + BaseQuery: q.baseQuery(), + Hooks: q.Hooks, + }, + Scanner: q.Scanner, + } +} + // asCountQuery clones and rewrites an existing query to a count query func asCountQuery(query bob.BaseQuery[*dialect.SelectQuery]) bob.BaseQuery[*dialect.SelectQuery] { // clone the original query, so it's not being modified silently diff --git a/dialect/psql/view_test.go b/dialect/psql/view_test.go index 468c82ce..064dd2fc 100644 --- a/dialect/psql/view_test.go +++ b/dialect/psql/view_test.go @@ -72,5 +72,5 @@ func selectToString(t *testing.T, query bob.Query, argsLen int) string { func viewToString(t *testing.T, query *ViewQuery[*someStruct, []*someStruct]) string { t.Helper() - return selectToString(t, query.BaseQuery, 3) + return selectToString(t, query, 3) } From d81d2a7107f3a0c4a41e345bf78c70e872db88f1 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 23:04:10 -0400 Subject: [PATCH 08/33] refactor(psql): drop materialized query cache Remove the wrapper-level cached mutable BaseQuery compatibility form from psql select/update/delete/insert wrappers. Verification: - go test ./dialect/psql - go test ./dialect/psql -run '^$' -bench '^Benchmark(BaseQuery|ViewQueryCountThenPaginate|UpdateQuery|DeleteQuery|InsertQuery)(ApplyMain|ImmutableNativeHotPath)$' -benchmem Benchmark snapshot after cleanup: - BaseQuery immutable: 2357 ns/op, 1705 B/op, 36 allocs/op - View count+paginate immutable: 5811 ns/op, 4731 B/op, 60 allocs/op - Update immutable: 2111 ns/op, 1360 B/op, 41 allocs/op - Delete immutable: 1555 ns/op, 1027 B/op, 26 allocs/op - Insert immutable: 1253 ns/op, 920 B/op, 20 allocs/op --- dialect/psql/delete.go | 18 +++++------------- dialect/psql/insert.go | 18 +++++------------- dialect/psql/select.go | 18 +++++------------- dialect/psql/update.go | 18 +++++------------- 4 files changed, 20 insertions(+), 52 deletions(-) diff --git a/dialect/psql/delete.go b/dialect/psql/delete.go index 8c656580..15776dc8 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -7,12 +7,10 @@ import ( type DeleteQuery struct { derivedDeleteQuery - materialized *bob.BaseQuery[*dialect.DeleteQuery] } func (q DeleteQuery) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { q.derivedDeleteQuery = q.derivedDeleteQuery.With(queryMods...) - q.materialized = nil return q } @@ -21,9 +19,6 @@ func (q DeleteQuery) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQue } func (q DeleteQuery) baseQuery() bob.BaseQuery[*dialect.DeleteQuery] { - if q.materialized != nil { - return *q.materialized - } return q.derivedDeleteQuery.mutableBase() } @@ -33,14 +28,11 @@ func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { mod.Apply(q) } - base := bob.BaseQuery[*dialect.DeleteQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeDelete, - } - return DeleteQuery{ - derivedDeleteQuery: asImmutableDelete(base), - materialized: &base, + derivedDeleteQuery: asImmutableDelete(bob.BaseQuery[*dialect.DeleteQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeDelete, + }), } } diff --git a/dialect/psql/insert.go b/dialect/psql/insert.go index 10b84992..b44f71f8 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -7,12 +7,10 @@ import ( type InsertQuery struct { derivedInsertQuery - materialized *bob.BaseQuery[*dialect.InsertQuery] } func (q InsertQuery) With(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { q.derivedInsertQuery = q.derivedInsertQuery.With(queryMods...) - q.materialized = nil return q } @@ -21,9 +19,6 @@ func (q InsertQuery) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQue } func (q InsertQuery) baseQuery() bob.BaseQuery[*dialect.InsertQuery] { - if q.materialized != nil { - return *q.materialized - } return q.derivedInsertQuery.mutableBase() } @@ -33,14 +28,11 @@ func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { mod.Apply(q) } - base := bob.BaseQuery[*dialect.InsertQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeInsert, - } - return InsertQuery{ - derivedInsertQuery: asImmutableInsert(base), - materialized: &base, + derivedInsertQuery: asImmutableInsert(bob.BaseQuery[*dialect.InsertQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeInsert, + }), } } diff --git a/dialect/psql/select.go b/dialect/psql/select.go index 87f87932..dfc885d1 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -7,12 +7,10 @@ import ( type SelectQuery struct { derivedSelectQuery - materialized *bob.BaseQuery[*dialect.SelectQuery] } func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { q.derivedSelectQuery = q.derivedSelectQuery.With(queryMods...) - q.materialized = nil return q } @@ -21,9 +19,6 @@ func (q SelectQuery) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQue } func (q SelectQuery) baseQuery() bob.BaseQuery[*dialect.SelectQuery] { - if q.materialized != nil { - return *q.materialized - } return q.derivedSelectQuery.mutableBase() } @@ -33,14 +28,11 @@ func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { mod.Apply(q) } - base := bob.BaseQuery[*dialect.SelectQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeSelect, - } - return SelectQuery{ - derivedSelectQuery: asImmutable(base), - materialized: &base, + derivedSelectQuery: asImmutable(bob.BaseQuery[*dialect.SelectQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeSelect, + }), } } diff --git a/dialect/psql/update.go b/dialect/psql/update.go index 5e56d09d..ec10e0df 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -7,12 +7,10 @@ import ( type UpdateQuery struct { derivedUpdateQuery - materialized *bob.BaseQuery[*dialect.UpdateQuery] } func (q UpdateQuery) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { q.derivedUpdateQuery = q.derivedUpdateQuery.With(queryMods...) - q.materialized = nil return q } @@ -21,9 +19,6 @@ func (q UpdateQuery) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQue } func (q UpdateQuery) baseQuery() bob.BaseQuery[*dialect.UpdateQuery] { - if q.materialized != nil { - return *q.materialized - } return q.derivedUpdateQuery.mutableBase() } @@ -33,14 +28,11 @@ func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { mod.Apply(q) } - base := bob.BaseQuery[*dialect.UpdateQuery]{ - Expression: q, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeUpdate, - } - return UpdateQuery{ - derivedUpdateQuery: asImmutableUpdate(base), - materialized: &base, + derivedUpdateQuery: asImmutableUpdate(bob.BaseQuery[*dialect.UpdateQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeUpdate, + }), } } From d07fed297e6d6acea0acbdcce8d08f321c3e037a Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 23:05:21 -0400 Subject: [PATCH 09/33] refactor(psql): simplify immutable view state Remove the unused defaultSelect field from ViewQuery, keep only the one bit of state that matters for overriding the auto-generated select list, and name the mutable writer fallback path explicitly. Verification: - go test ./dialect/psql - go test ./dialect/psql -run '^$' -bench '^Benchmark(BaseQuery|ViewQueryCountThenPaginate|UpdateQuery|DeleteQuery|InsertQuery)(ApplyMain|ImmutableNativeHotPath)$' -benchmem Benchmark snapshot after cleanup: - BaseQuery immutable: 2410 ns/op, 1704 B/op, 36 allocs/op - View count+paginate immutable: 5655 ns/op, 4538 B/op, 60 allocs/op - Update immutable: 2092 ns/op, 1360 B/op, 41 allocs/op - Delete immutable: 1558 ns/op, 1027 B/op, 26 allocs/op - Insert immutable: 1283 ns/op, 920 B/op, 20 allocs/op --- dialect/psql/immutable_select.go | 6 +++++- dialect/psql/view.go | 24 +++++++++++------------- 2 files changed, 16 insertions(+), 14 deletions(-) diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index e8d1ed65..133b3698 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -104,7 +104,7 @@ func (q derivedSelectQuery) BuildN(ctx context.Context, start int) (string, []an } func (q derivedSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - if len(q.contextualMods) > 0 || !q.state.supportsNativeWrite() { + if q.requiresMutableWrite() { return q.mutableBase().WriteQuery(ctx, w, start) } @@ -121,6 +121,10 @@ func (q derivedSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, s return writer.args, nil } +func (q derivedSelectQuery) requiresMutableWrite() bool { + return len(q.contextualMods) > 0 || !q.state.supportsNativeWrite() +} + func (q derivedSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { w.WriteString("(") args, err := q.WriteQuery(ctx, w, start) diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 02e8dacf..557a50e3 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -82,11 +82,10 @@ func (v *View[T, Tslice, C]) Alias() string { // Starts a select query func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Tslice] { q := &ViewQuery[T, Tslice]{ - Query: Select(sm.Columns(v.Columns), sm.From(v.NameAs())), - Scanner: v.scanner, - Hooks: &v.SelectQueryHooks, - defaultSelect: v.Columns, - defaulted: true, + Query: Select(sm.Columns(v.Columns), sm.From(v.NameAs())), + Scanner: v.scanner, + Hooks: &v.SelectQueryHooks, + usesDefaultSelect: true, } if len(queryMods) == 0 { return q @@ -95,18 +94,17 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * } type ViewQuery[T any, Ts ~[]T] struct { - Query SelectQuery - Scanner scan.Mapper[T] - Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] - defaultSelect bob.Expression - defaulted bool + Query SelectQuery + Scanner scan.Mapper[T] + Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] + usesDefaultSelect bool } func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { next := *q - if hasSelectMod(queryMods) && next.defaulted { + if next.usesDefaultSelect && overridesDefaultSelect(queryMods) { next.Query.derivedSelectQuery.state.SelectColumns = nil - next.defaulted = false + next.usesDefaultSelect = false } next.Query = next.Query.Apply(queryMods...) return &next @@ -168,7 +166,7 @@ func (q *ViewQuery[T, Ts]) Each(ctx context.Context, exec bob.Executor) (func(fu return q.mutable().Each(ctx, exec) } -func hasSelectMod(mods []bob.Mod[*dialect.SelectQuery]) bool { +func overridesDefaultSelect(mods []bob.Mod[*dialect.SelectQuery]) bool { for _, mod := range mods { if _, ok := mod.(bobmods.Select[*dialect.SelectQuery]); ok { return true From ac889e8e6d5243c5b9cd67988763fd7d65389bfb Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 23:11:33 -0400 Subject: [PATCH 10/33] refactor(psql): remove view wrapper fallback state Move default view-column handling into immutable select state and route ViewQuery execution directly through bob.One/All/Cursor/Each instead of wrapper-local baseQuery/mutable helpers. Added a regression test to ensure an explicit view Select mod replaces the generated default columns. Verification: - go test ./dialect/psql ./orm ./mods ./clause ./dialect/mysql ./dialect/sqlite - go test ./dialect/psql -run '^$' -bench '^Benchmark(BaseQuery|ViewQueryCountThenPaginate|UpdateQuery|DeleteQuery|InsertQuery)(ApplyMain|ImmutableNativeHotPath)$' -benchmem Benchmark snapshot after cleanup: - BaseQuery immutable: 2468 ns/op, 1705 B/op, 36 allocs/op - View count+paginate immutable: 5982 ns/op, 4674 B/op, 57 allocs/op - Update immutable: 2209 ns/op, 1360 B/op, 41 allocs/op - Delete immutable: 1606 ns/op, 1027 B/op, 26 allocs/op - Insert immutable: 1253 ns/op, 920 B/op, 20 allocs/op --- dialect/psql/immutable_select.go | 52 +++++++++++++-------- dialect/psql/immutable_select_test.go | 2 +- dialect/psql/view.go | 67 +++++++++++++-------------- dialect/psql/view_test.go | 14 ++++++ 4 files changed, 78 insertions(+), 57 deletions(-) diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index 133b3698..a6237184 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -23,21 +23,22 @@ type derivedSelectQuery struct { } type immutableSelectState struct { - With clause.With - SelectColumns []any - PreloadColumns []any - Distinct psqldialect.Distinct - TableRef clause.TableRef - Where clause.Where - GroupBy clause.GroupBy - Having clause.Having - Windows clause.Windows - Combines clause.Combines - OrderBy clause.OrderBy - Limit clause.Limit - Offset clause.Offset - Fetch clause.Fetch - Locks clause.Locks + DefaultSelectColumns []any + With clause.With + SelectColumns []any + PreloadColumns []any + Distinct psqldialect.Distinct + TableRef clause.TableRef + Where clause.Where + GroupBy clause.GroupBy + Having clause.Having + Windows clause.Windows + Combines clause.Combines + OrderBy clause.OrderBy + Limit clause.Limit + Offset clause.Offset + Fetch clause.Fetch + Locks clause.Locks CombinedOrder clause.OrderBy CombinedLimit clause.Limit @@ -78,6 +79,7 @@ func (q derivedSelectQuery) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) func (q derivedSelectQuery) AsCount() derivedSelectQuery { next := q.state next.SelectColumns = []any{"count(1)"} + next.DefaultSelectColumns = nil next.PreloadColumns = nil next.OrderBy.Expressions = nil next.GroupBy.Groups = nil @@ -86,7 +88,8 @@ func (q derivedSelectQuery) AsCount() derivedSelectQuery { next.Offset.Count = nil next.Limit.Count = 1 - return derivedSelectQuery{state: next} + q.state = next + return q } func (q derivedSelectQuery) Build(ctx context.Context) (string, []any, error) { @@ -154,7 +157,7 @@ func (q derivedSelectQuery) GetMapperMods() []scan.MapperMod { func (q derivedSelectQuery) mutableBase() bob.BaseQuery[*psqldialect.SelectQuery] { mutable := &psqldialect.SelectQuery{ With: q.state.With, - SelectList: clause.SelectList{Columns: q.state.SelectColumns, PreloadColumns: q.state.PreloadColumns}, + SelectList: clause.SelectList{Columns: q.state.selectColumns(), PreloadColumns: q.state.PreloadColumns}, Distinct: q.state.Distinct, TableRef: q.state.TableRef, Where: q.state.Where, @@ -189,6 +192,7 @@ func (q derivedSelectQuery) mutableBase() bob.BaseQuery[*psqldialect.SelectQuery func immutableStateFromMutable(q *psqldialect.SelectQuery) immutableSelectState { return immutableSelectState{ + DefaultSelectColumns: nil, With: clause.With{ Recursive: q.With.Recursive, CTEs: append([]bob.Expression(nil), q.With.CTEs...), @@ -246,6 +250,13 @@ func (s immutableSelectState) supportsNativeWrite() bool { s.CombinedFetch.Count == nil } +func (s immutableSelectState) selectColumns() []any { + if len(s.SelectColumns) > 0 { + return s.SelectColumns + } + return s.DefaultSelectColumns +} + func cloneAnySlice(values []any) []any { if values == nil { return nil @@ -256,7 +267,7 @@ func cloneAnySlice(values []any) []any { func (s immutableSelectState) toMutable() psqldialect.SelectQuery { return psqldialect.SelectQuery{ With: s.With, - SelectList: clause.SelectList{Columns: s.SelectColumns, PreloadColumns: s.PreloadColumns}, + SelectList: clause.SelectList{Columns: s.selectColumns(), PreloadColumns: s.PreloadColumns}, Distinct: s.Distinct, TableRef: s.TableRef, Where: s.Where, @@ -383,10 +394,11 @@ func (w *immutableSelectWriter) writeQuery(q immutableSelectState) error { } w.w.WriteString("\n") - if len(q.SelectColumns) == 0 && len(q.PreloadColumns) == 0 { + selectColumns := q.selectColumns() + if len(selectColumns) == 0 && len(q.PreloadColumns) == 0 { w.w.WriteString("*") } else { - allCols := append([]any(nil), q.SelectColumns...) + allCols := append([]any(nil), selectColumns...) allCols = append(allCols, q.PreloadColumns...) if err := w.writeSliceAny(allCols, ", "); err != nil { return err diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 20478c1a..6d121500 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -231,7 +231,7 @@ func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := asCountQuery(q.baseQuery()).Build(ctx); err != nil { + if _, _, err := asCountQuery(q.Query.baseQuery()).Build(ctx); err != nil { b.Fatal(err) } diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 557a50e3..7c9c5633 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -11,7 +11,6 @@ import ( "github.com/stephenafamo/bob/dialect/psql/sm" "github.com/stephenafamo/bob/expr" "github.com/stephenafamo/bob/internal/mappings" - bobmods "github.com/stephenafamo/bob/mods" "github.com/stephenafamo/bob/orm" "github.com/stephenafamo/scan" ) @@ -82,11 +81,11 @@ func (v *View[T, Tslice, C]) Alias() string { // Starts a select query func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Tslice] { q := &ViewQuery[T, Tslice]{ - Query: Select(sm.Columns(v.Columns), sm.From(v.NameAs())), - Scanner: v.scanner, - Hooks: &v.SelectQueryHooks, - usesDefaultSelect: true, + Query: Select(sm.From(v.NameAs())), + Scanner: v.scanner, + Hooks: &v.SelectQueryHooks, } + q.Query.derivedSelectQuery.state.DefaultSelectColumns = []any{v.Columns} if len(queryMods) == 0 { return q } @@ -94,18 +93,13 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * } type ViewQuery[T any, Ts ~[]T] struct { - Query SelectQuery - Scanner scan.Mapper[T] - Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] - usesDefaultSelect bool + Query SelectQuery + Scanner scan.Mapper[T] + Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] } func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { next := *q - if next.usesDefaultSelect && overridesDefaultSelect(queryMods) { - next.Query.derivedSelectQuery.state.SelectColumns = nil - next.usesDefaultSelect = false - } next.Query = next.Query.Apply(queryMods...) return &next } @@ -136,12 +130,15 @@ func (q *ViewQuery[T, Ts]) WriteSQL(ctx context.Context, w io.StringWriter, d bo // Count the number of matching rows func (v *ViewQuery[T, Tslice]) Count(ctx context.Context, exec bob.Executor) (int64, error) { - mq := v.mutable() - ctx, err := mq.RunHooks(ctx, exec) + ctx, err := v.RunHooks(ctx, exec) if err != nil { return 0, err } - return bob.One(ctx, exec, asCountQuery(mq.BaseQuery), scan.SingleColumnMapper[int64]) + sql, args, err := v.Query.AsCount().Build(ctx) + if err != nil { + return 0, err + } + return scan.One(ctx, exec, scan.SingleColumnMapper[int64], sql, args...) } // Exists checks if there is any matching row @@ -151,42 +148,40 @@ func (v *ViewQuery[T, Tslice]) Exists(ctx context.Context, exec bob.Executor) (b } func (q *ViewQuery[T, Ts]) One(ctx context.Context, exec bob.Executor) (T, error) { - return q.mutable().One(ctx, exec) + return bob.One(ctx, exec, q, q.Scanner) } func (q *ViewQuery[T, Ts]) All(ctx context.Context, exec bob.Executor) (Ts, error) { - return q.mutable().All(ctx, exec) + return bob.Allx[bob.SliceTransformer[T, Ts]](ctx, exec, q, q.Scanner) } func (q *ViewQuery[T, Ts]) Cursor(ctx context.Context, exec bob.Executor) (scan.ICursor[T], error) { - return q.mutable().Cursor(ctx, exec) + return bob.Cursor(ctx, exec, q, q.Scanner) } func (q *ViewQuery[T, Ts]) Each(ctx context.Context, exec bob.Executor) (func(func(T, error) bool), error) { - return q.mutable().Each(ctx, exec) + return bob.Each(ctx, exec, q, q.Scanner) } -func overridesDefaultSelect(mods []bob.Mod[*dialect.SelectQuery]) bool { - for _, mod := range mods { - if _, ok := mod.(bobmods.Select[*dialect.SelectQuery]); ok { - return true - } +func (q *ViewQuery[T, Ts]) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { + ctx, err := q.Query.RunHooks(ctx, exec) + if err != nil { + return ctx, err } - return false + + if q.Hooks == nil { + return ctx, nil + } + + return q.Hooks.RunHooks(ctx, exec, q.Query.baseQuery().Expression) } -func (q *ViewQuery[T, Ts]) baseQuery() bob.BaseQuery[*dialect.SelectQuery] { - return q.Query.baseQuery() +func (q *ViewQuery[T, Ts]) GetLoaders() []bob.Loader { + return q.Query.GetLoaders() } -func (q *ViewQuery[T, Ts]) mutable() orm.Query[*dialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] { - return orm.Query[*dialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]]{ - ExecQuery: orm.ExecQuery[*dialect.SelectQuery]{ - BaseQuery: q.baseQuery(), - Hooks: q.Hooks, - }, - Scanner: q.Scanner, - } +func (q *ViewQuery[T, Ts]) GetMapperMods() []scan.MapperMod { + return q.Query.GetMapperMods() } // asCountQuery clones and rewrites an existing query to a count query diff --git a/dialect/psql/view_test.go b/dialect/psql/view_test.go index 064dd2fc..af4e21f5 100644 --- a/dialect/psql/view_test.go +++ b/dialect/psql/view_test.go @@ -54,6 +54,20 @@ func TestSomeViewQuery(t *testing.T) { } } +func TestSomeViewQueryExplicitSelectReplacesDefaultColumns(t *testing.T) { + q := someStructView.Query().Apply( + sm.Columns("id"), + sm.Where(Quote("id").EQ(Arg(1))), + ) + + query := selectToString(t, q, 1) + expected := "SELECT \nid\nFROM \"public\".\"some_struct\" AS \"public.some_struct\"\nWHERE (\"id\" = $0)\n" + + if query != expected { + t.Errorf("Expected '%#v' but got '%#v'", expected, query) + } +} + func selectToString(t *testing.T, query bob.Query, argsLen int) string { t.Helper() ctx := t.Context() From c967589c8de5a46eb1b09ae2f49c104536ed055d Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 23:15:36 -0400 Subject: [PATCH 11/33] refactor(psql): drop wrapper base query helpers Remove the public wrapper-level baseQuery compatibility methods from psql select/update/delete/insert types and update internal callers to use the immutable internals directly. Verification: - go test ./dialect/psql ./orm ./mods ./clause ./dialect/mysql ./dialect/sqlite - go test ./dialect/psql -run '^$' -bench '^Benchmark(BaseQuery|ViewQueryCountThenPaginate|UpdateQuery|DeleteQuery|InsertQuery)(ApplyMain|ImmutableNativeHotPath)$' -benchmem Benchmark snapshot after cleanup: - BaseQuery immutable: 2369 ns/op, 1705 B/op, 36 allocs/op - View count+paginate immutable: 5711 ns/op, 4674 B/op, 57 allocs/op - Update immutable: 2097 ns/op, 1360 B/op, 41 allocs/op - Delete immutable: 1568 ns/op, 1027 B/op, 26 allocs/op - Insert immutable: 1291 ns/op, 920 B/op, 20 allocs/op --- dialect/psql/delete.go | 4 ---- dialect/psql/immutable_select_test.go | 2 +- dialect/psql/insert.go | 4 ---- dialect/psql/select.go | 4 ---- dialect/psql/table.go | 6 +++--- dialect/psql/update.go | 4 ---- dialect/psql/view.go | 2 +- 7 files changed, 5 insertions(+), 21 deletions(-) diff --git a/dialect/psql/delete.go b/dialect/psql/delete.go index 15776dc8..c40e94c2 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -18,10 +18,6 @@ func (q DeleteQuery) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQue return q.With(queryMods...) } -func (q DeleteQuery) baseQuery() bob.BaseQuery[*dialect.DeleteQuery] { - return q.derivedDeleteQuery.mutableBase() -} - func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { q := &dialect.DeleteQuery{} for _, mod := range queryMods { diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 6d121500..f0564965 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -231,7 +231,7 @@ func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := asCountQuery(q.Query.baseQuery()).Build(ctx); err != nil { + if _, _, err := asCountQuery(q.Query.derivedSelectQuery.mutableBase()).Build(ctx); err != nil { b.Fatal(err) } diff --git a/dialect/psql/insert.go b/dialect/psql/insert.go index b44f71f8..16ad95f7 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -18,10 +18,6 @@ func (q InsertQuery) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQue return q.With(queryMods...) } -func (q InsertQuery) baseQuery() bob.BaseQuery[*dialect.InsertQuery] { - return q.derivedInsertQuery.mutableBase() -} - func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { q := &dialect.InsertQuery{} for _, mod := range queryMods { diff --git a/dialect/psql/select.go b/dialect/psql/select.go index dfc885d1..e99102bc 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -18,10 +18,6 @@ func (q SelectQuery) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQue return q.With(queryMods...) } -func (q SelectQuery) baseQuery() bob.BaseQuery[*dialect.SelectQuery] { - return q.derivedSelectQuery.mutableBase() -} - func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { q := &dialect.SelectQuery{} for _, mod := range queryMods { diff --git a/dialect/psql/table.go b/dialect/psql/table.go index a2f71e08..a10a07b1 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -76,7 +76,7 @@ func (t *Table[T, Tslice, Tset, C]) PrimaryKey() expr.ColumnsExpr { func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { q := &ormInsertQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ - BaseQuery: Insert(im.Into(t.NameAs(), t.nonGeneratedCols...)).baseQuery(), + BaseQuery: Insert(im.Into(t.NameAs(), t.nonGeneratedCols...)).derivedInsertQuery.mutableBase(), Hooks: &t.InsertQueryHooks, }, Scanner: t.scanner, @@ -100,7 +100,7 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { q := &ormUpdateQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ - BaseQuery: Update(um.Table(t.NameAs())).baseQuery(), + BaseQuery: Update(um.Table(t.NameAs())).derivedUpdateQuery.mutableBase(), Hooks: &t.UpdateQueryHooks, }, Scanner: t.scanner, @@ -124,7 +124,7 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { q := &ormDeleteQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ - BaseQuery: Delete(dm.From(t.NameAs())).baseQuery(), + BaseQuery: Delete(dm.From(t.NameAs())).derivedDeleteQuery.mutableBase(), Hooks: &t.DeleteQueryHooks, }, Scanner: t.scanner, diff --git a/dialect/psql/update.go b/dialect/psql/update.go index ec10e0df..6cbe2053 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -18,10 +18,6 @@ func (q UpdateQuery) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQue return q.With(queryMods...) } -func (q UpdateQuery) baseQuery() bob.BaseQuery[*dialect.UpdateQuery] { - return q.derivedUpdateQuery.mutableBase() -} - func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { q := &dialect.UpdateQuery{} for _, mod := range queryMods { diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 7c9c5633..da466029 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -173,7 +173,7 @@ func (q *ViewQuery[T, Ts]) RunHooks(ctx context.Context, exec bob.Executor) (con return ctx, nil } - return q.Hooks.RunHooks(ctx, exec, q.Query.baseQuery().Expression) + return q.Hooks.RunHooks(ctx, exec, q.Query.derivedSelectQuery.mutableBase().Expression) } func (q *ViewQuery[T, Ts]) GetLoaders() []bob.Loader { From 6737d42374efdd1d32e1d7f43cb846b474300b5c Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Mon, 20 Apr 2026 23:26:40 -0400 Subject: [PATCH 12/33] refactor(psql): remove table returning contextual mods Replace the psql table query contextual mods that injected default RETURNING columns with constructor-time defaults, while preserving explicit RETURNING overrides. Added regression coverage for update/insert/delete table queries to assert both the default RETURNING-all-columns behavior and explicit RETURNING override behavior. Verification: - go test ./dialect/psql - go test ./dialect/psql -run '^$' -bench '^Benchmark(BaseQuery|ViewQueryCountThenPaginate|UpdateQuery|DeleteQuery|InsertQuery)(ApplyMain|ImmutableNativeHotPath)$' -benchmem Benchmark snapshot after cleanup: - BaseQuery immutable: 2329 ns/op, 1705 B/op, 36 allocs/op - View count+paginate immutable: 5759 ns/op, 4674 B/op, 57 allocs/op - Update immutable: 2138 ns/op, 1360 B/op, 41 allocs/op - Delete immutable: 1577 ns/op, 1027 B/op, 26 allocs/op - Insert immutable: 1214 ns/op, 920 B/op, 20 allocs/op --- dialect/psql/table.go | 85 +++++++++++++++++--------- dialect/psql/table_test.go | 121 +++++++++++++++++++++++++++++++++++++ 2 files changed, 176 insertions(+), 30 deletions(-) diff --git a/dialect/psql/table.go b/dialect/psql/table.go index a10a07b1..a93c3a17 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -13,6 +13,7 @@ import ( "github.com/stephenafamo/bob/expr" "github.com/stephenafamo/bob/internal" "github.com/stephenafamo/bob/internal/mappings" + bobmods "github.com/stephenafamo/bob/mods" "github.com/stephenafamo/bob/orm" ) @@ -76,21 +77,12 @@ func (t *Table[T, Tslice, Tset, C]) PrimaryKey() expr.ColumnsExpr { func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { q := &ormInsertQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ - BaseQuery: Insert(im.Into(t.NameAs(), t.nonGeneratedCols...)).derivedInsertQuery.mutableBase(), + BaseQuery: insertTableBaseQuery(t.NameAs(), t.nonGeneratedCols, t.Columns, queryMods), Hooks: &t.InsertQueryHooks, }, Scanner: t.scanner, } - q.Expression.AppendContextualModFunc( - func(ctx context.Context, q *dialect.InsertQuery) (context.Context, error) { - if !q.HasReturning() { - q.AppendReturning(t.Columns) - } - return ctx, nil - }, - ) - q.Apply(queryMods...) return q @@ -100,21 +92,12 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { q := &ormUpdateQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ - BaseQuery: Update(um.Table(t.NameAs())).derivedUpdateQuery.mutableBase(), + BaseQuery: updateTableBaseQuery(t.NameAs(), t.Columns, queryMods), Hooks: &t.UpdateQueryHooks, }, Scanner: t.scanner, } - q.Expression.AppendContextualModFunc( - func(ctx context.Context, q *dialect.UpdateQuery) (context.Context, error) { - if !q.HasReturning() { - q.AppendReturning(t.Columns) - } - return ctx, nil - }, - ) - q.Apply(queryMods...) return q @@ -124,21 +107,12 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { q := &ormDeleteQuery[T, Tslice]{ ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ - BaseQuery: Delete(dm.From(t.NameAs())).derivedDeleteQuery.mutableBase(), + BaseQuery: deleteTableBaseQuery(t.NameAs(), t.Columns, queryMods), Hooks: &t.DeleteQueryHooks, }, Scanner: t.scanner, } - q.Expression.AppendContextualModFunc( - func(ctx context.Context, q *dialect.DeleteQuery) (context.Context, error) { - if !q.HasReturning() { - q.AppendReturning(t.Columns) - } - return ctx, nil - }, - ) - q.Apply(queryMods...) return q @@ -172,3 +146,54 @@ func (t *Table[T, Tslice, Tset, C]) Merge(queryMods ...bob.Mod[*dialect.MergeQue return q } + +func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Expression, queryMods []bob.Mod[*dialect.InsertQuery]) bob.BaseQuery[*dialect.InsertQuery] { + base := Insert(im.Into(name, nonGeneratedCols...)).derivedInsertQuery.mutableBase() + if !hasInsertReturning(queryMods) { + base.Expression.AppendReturning(returning) + } + return base +} + +func updateTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mod[*dialect.UpdateQuery]) bob.BaseQuery[*dialect.UpdateQuery] { + base := Update(um.Table(name)).derivedUpdateQuery.mutableBase() + if !hasUpdateReturning(queryMods) { + base.Expression.AppendReturning(returning) + } + return base +} + +func deleteTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mod[*dialect.DeleteQuery]) bob.BaseQuery[*dialect.DeleteQuery] { + base := Delete(dm.From(name)).derivedDeleteQuery.mutableBase() + if !hasDeleteReturning(queryMods) { + base.Expression.AppendReturning(returning) + } + return base +} + +func hasInsertReturning(mods []bob.Mod[*dialect.InsertQuery]) bool { + for _, mod := range mods { + if _, ok := mod.(bobmods.Returning[*dialect.InsertQuery]); ok { + return true + } + } + return false +} + +func hasUpdateReturning(mods []bob.Mod[*dialect.UpdateQuery]) bool { + for _, mod := range mods { + if _, ok := mod.(bobmods.Returning[*dialect.UpdateQuery]); ok { + return true + } + } + return false +} + +func hasDeleteReturning(mods []bob.Mod[*dialect.DeleteQuery]) bool { + for _, mod := range mods { + if _, ok := mod.(bobmods.Returning[*dialect.DeleteQuery]); ok { + return true + } + } + return false +} diff --git a/dialect/psql/table_test.go b/dialect/psql/table_test.go index c7303733..324082d0 100644 --- a/dialect/psql/table_test.go +++ b/dialect/psql/table_test.go @@ -14,6 +14,8 @@ import ( _ "github.com/lib/pq" "github.com/stephenafamo/bob" "github.com/stephenafamo/bob/dialect/psql/dialect" + "github.com/stephenafamo/bob/dialect/psql/dm" + "github.com/stephenafamo/bob/dialect/psql/im" "github.com/stephenafamo/bob/dialect/psql/mm" "github.com/stephenafamo/bob/dialect/psql/um" "github.com/stephenafamo/bob/expr" @@ -131,6 +133,125 @@ func (s UserSetter) Expressions(prefix ...string) []bob.Expression { var userTable = NewTable[*User, *UserSetter, bob.Expression]("", "users", expr.ColsForStruct[User]("users")) +func TestTableUpdateDefaultsReturningAllColumns(t *testing.T) { + q := userTable.Update( + um.SetCol("name").ToArg("Stephen"), + um.Where(Quote("id").EQ(Arg(1))), + ) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "UPDATE \"users\" AS \"users\" SET\n\"name\" = $1\nWHERE (\"id\" = $2)\nRETURNING \"users\".\"id\" AS \"id\", \"users\".\"name\" AS \"name\", \"users\".\"email\" AS \"email\"" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 2 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + +func TestTableUpdateExplicitReturningOverridesDefault(t *testing.T) { + q := userTable.Update( + um.SetCol("name").ToArg("Stephen"), + um.Where(Quote("id").EQ(Arg(1))), + um.Returning("id"), + ) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "UPDATE \"users\" AS \"users\" SET\n\"name\" = $1\nWHERE (\"id\" = $2)\nRETURNING id" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 2 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + +func TestTableInsertDefaultsReturningAllColumns(t *testing.T) { + q := userTable.Insert( + im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), + ) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "INSERT INTO \"users\" AS \"users\"(\"id\", \"name\", \"email\")\nVALUES ($1, $2, $3)\nRETURNING \"users\".\"id\" AS \"id\", \"users\".\"name\" AS \"name\", \"users\".\"email\" AS \"email\"\n" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 3 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + +func TestTableInsertExplicitReturningOverridesDefault(t *testing.T) { + q := userTable.Insert( + im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), + im.Returning("id"), + ) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "INSERT INTO \"users\" AS \"users\"(\"id\", \"name\", \"email\")\nVALUES ($1, $2, $3)\nRETURNING id\n" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 3 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + +func TestTableDeleteDefaultsReturningAllColumns(t *testing.T) { + q := userTable.Delete( + dm.Where(Quote("id").EQ(Arg(1))), + ) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "DELETE FROM \"users\" AS \"users\"\nWHERE (\"id\" = $1)\nRETURNING \"users\".\"id\" AS \"id\", \"users\".\"name\" AS \"name\", \"users\".\"email\" AS \"email\"" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 1 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + +func TestTableDeleteExplicitReturningOverridesDefault(t *testing.T) { + q := userTable.Delete( + dm.Where(Quote("id").EQ(Arg(1))), + dm.Returning("id"), + ) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "DELETE FROM \"users\" AS \"users\"\nWHERE (\"id\" = $1)\nRETURNING id" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 1 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + func TestUpdate(t *testing.T) { ctx := t.Context() From 5345cd6c1646b924a6e9be25e081070157ee65e1 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 09:51:57 -0400 Subject: [PATCH 13/33] feat(psql): support native immutable combined selects --- dialect/psql/immutable_select.go | 63 ++++++++++++++++++++++++--- dialect/psql/immutable_select_test.go | 63 +++++++++++++++++++++++++++ 2 files changed, 121 insertions(+), 5 deletions(-) diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index a6237184..7b84bb6e 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -243,11 +243,7 @@ func immutableStateFromMutable(q *psqldialect.SelectQuery) immutableSelectState } func (s immutableSelectState) supportsNativeWrite() bool { - return len(s.Combines.Queries) == 0 && - len(s.CombinedOrder.Expressions) == 0 && - s.CombinedLimit.Count == nil && - s.CombinedOffset.Count == nil && - s.CombinedFetch.Count == nil + return true } func (s immutableSelectState) selectColumns() []any { @@ -379,6 +375,17 @@ func (w *immutableSelectWriter) writeQuery(q immutableSelectState) error { w.w.WriteString("\n") } + needsParens := len(q.Combines.Queries) > 0 && + (len(q.OrderBy.Expressions) > 0 || + q.Limit.Count != nil || + q.Offset.Count != nil || + q.Fetch.Count != nil || + len(q.Locks.Locks) > 0) + + if needsParens { + w.w.WriteString("(") + } + w.w.WriteString("SELECT ") if q.Distinct.On != nil { @@ -489,6 +496,52 @@ func (w *immutableSelectWriter) writeQuery(q immutableSelectState) error { } } + if needsParens { + w.w.WriteString(")") + } + + for _, combine := range q.Combines.Queries { + w.w.WriteString("\n") + args, err := combine.WriteSQL(w.ctx, w.w, psqldialect.Dialect, w.argPos()) + if err != nil { + return err + } + w.args = append(w.args, args...) + } + + if len(q.CombinedOrder.Expressions) > 0 { + w.w.WriteString("\nORDER BY ") + if err := w.writeOrderExprs(q.CombinedOrder.Expressions); err != nil { + return err + } + } + + if q.CombinedLimit.Count != nil { + w.w.WriteString("\nLIMIT ") + if err := w.writeAny(q.CombinedLimit.Count); err != nil { + return err + } + } + + if q.CombinedOffset.Count != nil { + w.w.WriteString("\nOFFSET ") + if err := w.writeAny(q.CombinedOffset.Count); err != nil { + return err + } + } + + if q.CombinedFetch.Count != nil { + w.w.WriteString("\nFETCH NEXT ") + if err := w.writeAny(q.CombinedFetch.Count); err != nil { + return err + } + if q.CombinedFetch.WithTies { + w.w.WriteString(" ROWS WITH TIES") + } else { + w.w.WriteString(" ROWS ONLY") + } + } + w.w.WriteString("\n") return nil } diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index f0564965..6a767d17 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -150,6 +150,69 @@ func TestImmutableSelectQueryApplyFallbackDoesNotMutateOriginal(t *testing.T) { } } +func TestImmutableSelectQueryApplyCombinedDoesNotMutateOriginal(t *testing.T) { + base := Select( + sm.Columns("id", "name"), + sm.From("users"), + sm.Limit(100), + sm.OrderBy("id"), + ) + + derived := base.Apply( + sm.Union(Select( + sm.Columns("id", "name"), + sm.From("admins"), + sm.Limit(10), + sm.OrderBy("id"), + )), + sm.OrderCombined("id"), + sm.LimitCombined(1000), + ) + + if derived.derivedSelectQuery.requiresMutableWrite() { + t.Fatal("expected combined select shape to use native immutable writer") + } + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "SELECT \nid, name\nFROM users\nORDER BY id\nLIMIT 100\n" { + t.Fatalf("base query changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, derivedArgs, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expected := Select( + sm.Columns("id", "name"), + sm.From("users"), + sm.Limit(100), + sm.OrderBy("id"), + sm.Union(Select( + sm.Columns("id", "name"), + sm.From("admins"), + sm.Limit(10), + sm.OrderBy("id"), + )), + sm.OrderCombined("id"), + sm.LimitCombined(1000), + ) + + expectedSQL, expectedArgs, err := expected.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != expectedSQL { + t.Fatalf("derived combined query mismatch: got %#v want %#v", derivedSQL, expectedSQL) + } + if len(derivedArgs) != len(expectedArgs) { + t.Fatalf("derived combined args mismatch: got %d want %d", len(derivedArgs), len(expectedArgs)) + } +} + func TestImmutableViewQueryApplyDoesNotMutateOriginal(t *testing.T) { base := someStructView.Query( sm.Where(Quote("id").GT(Arg(0))), From 9c7f39189417efb730d9d0752056e94d2a88b20e Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 09:57:44 -0400 Subject: [PATCH 14/33] refactor(orm): make table queries immutable by default --- clause/returning.go | 8 ++++ dialect/mysql/table.go | 27 +++++++----- dialect/psql/table.go | 18 +++----- dialect/psql/table_test.go | 88 ++++++++++++++++++++++++++++++++++++++ dialect/sqlite/table.go | 12 ++---- mods/mods.go | 4 ++ orm/default_returning.go | 24 +++++++++++ orm/query.go | 65 ++++++++++++++++++++++++++++ query.go | 2 + 9 files changed, 217 insertions(+), 31 deletions(-) create mode 100644 orm/default_returning.go diff --git a/clause/returning.go b/clause/returning.go index e8cdb1d5..f1ee5f35 100644 --- a/clause/returning.go +++ b/clause/returning.go @@ -15,6 +15,14 @@ func (r *Returning) HasReturning() bool { return len(r.Expressions) > 0 } +func (r Returning) ReturningExpressions() []any { + return r.Expressions +} + +func (r *Returning) SetReturning(columns ...any) { + r.Expressions = append(r.Expressions[:0], columns...) +} + func (r *Returning) AppendReturning(columns ...any) { r.Expressions = append(r.Expressions, columns...) } diff --git a/dialect/mysql/table.go b/dialect/mysql/table.go index 3d32a296..63fea022 100644 --- a/dialect/mysql/table.go +++ b/dialect/mysql/table.go @@ -93,9 +93,7 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ table: t, } - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } // Starts an update query for this table @@ -104,9 +102,7 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ BaseQuery: Update(um.Table(t.NameAs())), Hooks: &t.UpdateQueryHooks, } - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } // Starts a delete query for this table @@ -116,9 +112,7 @@ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQ Hooks: &t.DeleteQueryHooks, } - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } type insertQuery[T any, Ts ~[]T, Tset setter[T], C bob.Expression] struct { @@ -126,6 +120,20 @@ type insertQuery[T any, Ts ~[]T, Tset setter[T], C bob.Expression] struct { table *Table[T, Ts, Tset, C] } +func (t *insertQuery[T, Ts, Tset, C]) With(queryMods ...bob.Mod[*dialect.InsertQuery]) *insertQuery[T, Ts, Tset, C] { + if t == nil { + return nil + } + + next := *t + next.ExecQuery = *t.ExecQuery.With(queryMods...) + return &next +} + +func (t *insertQuery[T, Ts, Tset, C]) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) *insertQuery[T, Ts, Tset, C] { + return t.With(queryMods...) +} + // Insert One Row // NOTE: Because MySQL does not support RETURNING, this will insert the row and then run a SELECT query // to retrieve the row. @@ -288,7 +296,6 @@ func (t *insertQuery[T, Tslice, Tset, C]) getInserted(vals []clause.Value, resul } query.Apply(sm.Where(Or(filters...))) - return query, nil } diff --git a/dialect/psql/table.go b/dialect/psql/table.go index a93c3a17..f6666eae 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -83,9 +83,7 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ Scanner: t.scanner, } - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } // Starts an Update query for this table @@ -98,9 +96,7 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ Scanner: t.scanner, } - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } // Starts a Delete query for this table @@ -113,9 +109,7 @@ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQ Scanner: t.scanner, } - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } // Starts a Merge query for this table @@ -150,7 +144,7 @@ func (t *Table[T, Tslice, Tset, C]) Merge(queryMods ...bob.Mod[*dialect.MergeQue func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Expression, queryMods []bob.Mod[*dialect.InsertQuery]) bob.BaseQuery[*dialect.InsertQuery] { base := Insert(im.Into(name, nonGeneratedCols...)).derivedInsertQuery.mutableBase() if !hasInsertReturning(queryMods) { - base.Expression.AppendReturning(returning) + base.Expression.AppendReturning(orm.DefaultReturning(returning)) } return base } @@ -158,7 +152,7 @@ func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Exp func updateTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mod[*dialect.UpdateQuery]) bob.BaseQuery[*dialect.UpdateQuery] { base := Update(um.Table(name)).derivedUpdateQuery.mutableBase() if !hasUpdateReturning(queryMods) { - base.Expression.AppendReturning(returning) + base.Expression.AppendReturning(orm.DefaultReturning(returning)) } return base } @@ -166,7 +160,7 @@ func updateTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mo func deleteTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mod[*dialect.DeleteQuery]) bob.BaseQuery[*dialect.DeleteQuery] { base := Delete(dm.From(name)).derivedDeleteQuery.mutableBase() if !hasDeleteReturning(queryMods) { - base.Expression.AppendReturning(returning) + base.Expression.AppendReturning(orm.DefaultReturning(returning)) } return base } diff --git a/dialect/psql/table_test.go b/dialect/psql/table_test.go index 324082d0..f951f821 100644 --- a/dialect/psql/table_test.go +++ b/dialect/psql/table_test.go @@ -252,6 +252,94 @@ func TestTableDeleteExplicitReturningOverridesDefault(t *testing.T) { } } +func TestTableUpdateApplyDoesNotMutateOriginal(t *testing.T) { + base := userTable.Update( + um.SetCol("name").ToArg("Stephen"), + ) + + derived := base.Apply( + um.Where(Quote("id").EQ(Arg(1))), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + derivedSQL, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedBase := "UPDATE \"users\" AS \"users\" SET\n\"name\" = $1\nRETURNING \"users\".\"id\" AS \"id\", \"users\".\"name\" AS \"name\", \"users\".\"email\" AS \"email\"" + if baseSQL != expectedBase { + t.Fatalf("unexpected base SQL: %#v", baseSQL) + } + + expectedDerived := "UPDATE \"users\" AS \"users\" SET\n\"name\" = $1\nWHERE (\"id\" = $2)\nRETURNING \"users\".\"id\" AS \"id\", \"users\".\"name\" AS \"name\", \"users\".\"email\" AS \"email\"" + if derivedSQL != expectedDerived { + t.Fatalf("unexpected derived SQL: %#v", derivedSQL) + } +} + +func TestTableInsertWithDoesNotMutateOriginal(t *testing.T) { + base := userTable.Insert( + im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), + ) + + derived := base.With(im.Returning("id")) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + derivedSQL, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedBase := "INSERT INTO \"users\" AS \"users\"(\"id\", \"name\", \"email\")\nVALUES ($1, $2, $3)\nRETURNING \"users\".\"id\" AS \"id\", \"users\".\"name\" AS \"name\", \"users\".\"email\" AS \"email\"\n" + if baseSQL != expectedBase { + t.Fatalf("unexpected base SQL: %#v", baseSQL) + } + + expectedDerived := "INSERT INTO \"users\" AS \"users\"(\"id\", \"name\", \"email\")\nVALUES ($1, $2, $3)\nRETURNING id\n" + if derivedSQL != expectedDerived { + t.Fatalf("unexpected derived SQL: %#v", derivedSQL) + } +} + +func TestTableDeleteApplyDoesNotMutateOriginal(t *testing.T) { + base := userTable.Delete( + dm.Where(Quote("email").EQ(Arg("stephen@example.com"))), + ) + + derived := base.Apply( + dm.Returning("id"), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + derivedSQL, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedBase := "DELETE FROM \"users\" AS \"users\"\nWHERE (\"email\" = $1)\nRETURNING \"users\".\"id\" AS \"id\", \"users\".\"name\" AS \"name\", \"users\".\"email\" AS \"email\"" + if baseSQL != expectedBase { + t.Fatalf("unexpected base SQL: %#v", baseSQL) + } + + expectedDerived := "DELETE FROM \"users\" AS \"users\"\nWHERE (\"email\" = $1)\nRETURNING id" + if derivedSQL != expectedDerived { + t.Fatalf("unexpected derived SQL: %#v", derivedSQL) + } +} + func TestUpdate(t *testing.T) { ctx := t.Context() diff --git a/dialect/sqlite/table.go b/dialect/sqlite/table.go index 2a66a11f..6f56b63a 100644 --- a/dialect/sqlite/table.go +++ b/dialect/sqlite/table.go @@ -82,9 +82,7 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ }, ) - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } // Starts an Update query for this table @@ -106,9 +104,7 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ }, ) - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } // Starts a Delete query for this table @@ -130,7 +126,5 @@ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQ }, ) - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } diff --git a/mods/mods.go b/mods/mods.go index f870a2c3..f95c2434 100644 --- a/mods/mods.go +++ b/mods/mods.go @@ -161,6 +161,10 @@ func (s Returning[Q]) Apply(q Q) { q.AppendReturning(s...) } +func (s Returning[Q]) ReturningValues() []any { + return []any(s) +} + type Set[Q interface{ AppendSet(clauses ...any) }] []string func (s Set[Q]) To(to any) bob.Mod[Q] { diff --git a/orm/default_returning.go b/orm/default_returning.go new file mode 100644 index 00000000..8ec1bb32 --- /dev/null +++ b/orm/default_returning.go @@ -0,0 +1,24 @@ +package orm + +import ( + "context" + "io" + + "github.com/stephenafamo/bob" +) + +type defaultReturning struct { + expr bob.Expression +} + +func DefaultReturning(expr bob.Expression) bob.Expression { + return defaultReturning{expr: expr} +} + +func (d defaultReturning) IsDefaultReturning() bool { + return true +} + +func (d defaultReturning) WriteSQL(ctx context.Context, w io.StringWriter, dl bob.Dialect, start int) ([]any, error) { + return d.expr.WriteSQL(ctx, w, dl, start) +} diff --git a/orm/query.go b/orm/query.go index 5ed7729d..ef3e28d0 100644 --- a/orm/query.go +++ b/orm/query.go @@ -10,6 +10,11 @@ import ( "github.com/stephenafamo/scan" ) +type returningAware interface { + ReturningExpressions() []any + SetReturning(...any) +} + type ExecQuery[Q bob.Expression] struct { bob.BaseQuery[Q] Hooks *bob.Hooks[Q, bob.SkipQueryHooksKey] @@ -22,6 +27,20 @@ func (q ExecQuery[Q]) Clone() ExecQuery[Q] { } } +func (q *ExecQuery[Q]) With(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { + if q == nil { + return nil + } + + next := q.Clone() + applyQueryMods(next.BaseQuery.Expression, queryMods...) + return &next +} + +func (q *ExecQuery[Q]) Apply(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { + return q.With(queryMods...) +} + func (q ExecQuery[Q]) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { var err error @@ -55,9 +74,24 @@ type Query[Q bob.Expression, T, Ts any, Tr bob.Transformer[T, Ts]] struct { func (q Query[Q, T, Ts, Tr]) Clone() Query[Q, T, Ts, Tr] { return Query[Q, T, Ts, Tr]{ ExecQuery: q.ExecQuery.Clone(), + Scanner: q.Scanner, } } +func (q *Query[Q, T, Ts, Tr]) With(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] { + if q == nil { + return nil + } + + next := q.Clone() + applyQueryMods(next.BaseQuery.Expression, queryMods...) + return &next +} + +func (q *Query[Q, T, Ts, Tr]) Apply(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] { + return q.With(queryMods...) +} + // First matching row func (q Query[Q, T, Ts, Tr]) One(ctx context.Context, exec bob.Executor) (T, error) { return bob.One(ctx, exec, q, q.Scanner) @@ -96,6 +130,37 @@ func (q ModQuery[Q, E, T, Ts, Tr]) Apply(e Q) { q.Mod.Apply(e) } +func applyQueryMods[Q any](query Q, queryMods ...bob.Mod[Q]) { + replacedDefaultReturning := false + + for _, mod := range queryMods { + if returning, ok := any(mod).(interface{ ReturningValues() []any }); ok && !replacedDefaultReturning { + if returningClause, ok := any(query).(returningAware); ok && hasOnlyDefaultReturning(returningClause.ReturningExpressions()) { + returningClause.SetReturning(returning.ReturningValues()...) + replacedDefaultReturning = true + continue + } + } + + mod.Apply(query) + } +} + +func hasOnlyDefaultReturning(expressions []any) bool { + if len(expressions) == 0 { + return false + } + + for _, expression := range expressions { + marker, ok := expression.(interface{ IsDefaultReturning() bool }) + if !ok || !marker.IsDefaultReturning() { + return false + } + } + + return true +} + func ArgsToExpression(querySQL string, from, to int, argIter iter.Seq[ArgWithPosition]) bob.Expression { return bob.ExpressionFunc(func(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { args := []any{} diff --git a/query.go b/query.go index eb227f8c..73eca1af 100644 --- a/query.go +++ b/query.go @@ -74,12 +74,14 @@ func (b BaseQuery[E]) Clone() BaseQuery[E] { return BaseQuery[E]{ Expression: c.Clone(), Dialect: b.Dialect, + QueryType: b.QueryType, } } return BaseQuery[E]{ Expression: reprint.This(b.Expression).(E), Dialect: b.Dialect, + QueryType: b.QueryType, } } From f69f1220984b8743b64f6754a93a715d360cb5a7 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:02:39 -0400 Subject: [PATCH 15/33] refactor(psql): run view hooks on immutable queries --- dialect/psql/immutable_select_test.go | 35 +++++++++++++++++++++++++++ dialect/psql/view.go | 6 ++--- 2 files changed, 38 insertions(+), 3 deletions(-) diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 6a767d17..1c977d78 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -1,10 +1,12 @@ package psql import ( + "context" "testing" "github.com/stephenafamo/bob" "github.com/stephenafamo/bob/dialect/psql/sm" + "github.com/stephenafamo/bob/expr" ) func TestImmutableSelectQueryWithDoesNotMutateOriginal(t *testing.T) { @@ -241,6 +243,39 @@ func TestImmutableViewQueryApplyDoesNotMutateOriginal(t *testing.T) { } } +func TestViewSelectQueryHooksUseImmutableSelectQuery(t *testing.T) { + view := NewView[*someStruct, bob.Expression]("public", "some_struct", expr.ColsForStruct[someStruct]("some_struct")) + + var hookSQL string + view.SelectQueryHooks.AppendHooks(func(ctx context.Context, exec bob.Executor, q *SelectQuery) (context.Context, error) { + sql, _, err := q.Build(ctx) + if err != nil { + return ctx, err + } + hookSQL = sql + return context.WithValue(ctx, "view-hook-ran", true), nil + }) + + query := view.Query( + sm.Where(Quote("id").EQ(Arg(1))), + sm.OrderBy("name"), + ) + + ctx, err := query.RunHooks(t.Context(), nil) + if err != nil { + t.Fatal(err) + } + + if got := ctx.Value("view-hook-ran"); got != true { + t.Fatalf("expected hook marker in context, got %#v", got) + } + + expected := "SELECT \n\"some_struct\".\"id\" AS \"id\", \"some_struct\".\"name\" AS \"name\", \"some_struct\".\"email\" AS \"email\"\nFROM \"public\".\"some_struct\" AS \"public.some_struct\"\nWHERE (\"id\" = $1)\nORDER BY name\n" + if hookSQL != expected { + t.Fatalf("unexpected hook SQL: %#v", hookSQL) + } +} + func BenchmarkBaseQueryApplyMain(b *testing.B) { ctx := b.Context() diff --git a/dialect/psql/view.go b/dialect/psql/view.go index da466029..f1bb3112 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -58,7 +58,7 @@ type View[T any, Tslice ~[]T, C bob.Expression] struct { Columns C AfterSelectHooks bob.Hooks[Tslice, bob.SkipModelHooksKey] - SelectQueryHooks bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] + SelectQueryHooks bob.Hooks[*SelectQuery, bob.SkipQueryHooksKey] } func (v *View[T, Tslice, C]) Name() Expression { @@ -95,7 +95,7 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * type ViewQuery[T any, Ts ~[]T] struct { Query SelectQuery Scanner scan.Mapper[T] - Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] + Hooks *bob.Hooks[*SelectQuery, bob.SkipQueryHooksKey] } func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { @@ -173,7 +173,7 @@ func (q *ViewQuery[T, Ts]) RunHooks(ctx context.Context, exec bob.Executor) (con return ctx, nil } - return q.Hooks.RunHooks(ctx, exec, q.Query.derivedSelectQuery.mutableBase().Expression) + return q.Hooks.RunHooks(ctx, exec, &q.Query) } func (q *ViewQuery[T, Ts]) GetLoaders() []bob.Loader { From cb9440ef3c5fbe966e3b0800f3ff3c00e85844e2 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:04:20 -0400 Subject: [PATCH 16/33] refactor(psql): drop dead immutable contextual state --- dialect/psql/immutable_select.go | 14 +++----- dialect/psql/immutable_write.go | 57 +++++++++----------------------- 2 files changed, 19 insertions(+), 52 deletions(-) diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index 7b84bb6e..78f051d0 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -16,10 +16,9 @@ import ( ) type derivedSelectQuery struct { - state immutableSelectState - load bob.Load - hooks bob.EmbeddedHook - contextualMods []bob.ContextualMod[*psqldialect.SelectQuery] + state immutableSelectState + load bob.Load + hooks bob.EmbeddedHook } type immutableSelectState struct { @@ -51,8 +50,6 @@ func asImmutable(q bob.BaseQuery[*psqldialect.SelectQuery]) derivedSelectQuery { state: immutableStateFromMutable(q.Expression), load: q.Expression.Load, hooks: q.Expression.EmbeddedHook, - contextualMods: append([]bob.ContextualMod[*psqldialect.SelectQuery](nil), - q.Expression.ContextualModdable.Mods...), } } @@ -125,7 +122,7 @@ func (q derivedSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, s } func (q derivedSelectQuery) requiresMutableWrite() bool { - return len(q.contextualMods) > 0 || !q.state.supportsNativeWrite() + return !q.state.supportsNativeWrite() } func (q derivedSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { @@ -173,9 +170,6 @@ func (q derivedSelectQuery) mutableBase() bob.BaseQuery[*psqldialect.SelectQuery Load: q.load, EmbeddedHook: q.hooks, - ContextualModdable: bob.ContextualModdable[*psqldialect.SelectQuery]{ - Mods: append([]bob.ContextualMod[*psqldialect.SelectQuery](nil), q.contextualMods...), - }, CombinedOrder: q.state.CombinedOrder, CombinedLimit: q.state.CombinedLimit, diff --git a/dialect/psql/immutable_write.go b/dialect/psql/immutable_write.go index f0be166f..6fe22119 100644 --- a/dialect/psql/immutable_write.go +++ b/dialect/psql/immutable_write.go @@ -14,10 +14,9 @@ import ( ) type derivedUpdateQuery struct { - state immutableUpdateState - load bob.Load - hooks bob.EmbeddedHook - contextualMods []bob.ContextualMod[*psqldialect.UpdateQuery] + state immutableUpdateState + load bob.Load + hooks bob.EmbeddedHook } type immutableUpdateState struct { @@ -50,9 +49,8 @@ func asImmutableUpdate(q bob.BaseQuery[*psqldialect.UpdateQuery]) derivedUpdateQ Expressions: append([]any(nil), q.Expression.Returning.Expressions...), }, }, - load: q.Expression.Load, - hooks: q.Expression.EmbeddedHook, - contextualMods: append([]bob.ContextualMod[*psqldialect.UpdateQuery](nil), q.Expression.ContextualModdable.Mods...), + load: q.Expression.Load, + hooks: q.Expression.EmbeddedHook, } } @@ -104,10 +102,6 @@ func (q derivedUpdateQuery) BuildN(ctx context.Context, start int) (string, []an } func (q derivedUpdateQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - if len(q.contextualMods) > 0 { - return q.mutableBase().WriteQuery(ctx, w, start) - } - var args []any if len(q.state.With.CTEs) > 0 { @@ -182,9 +176,6 @@ func (q derivedUpdateQuery) mutableBase() bob.BaseQuery[*psqldialect.UpdateQuery Returning: q.state.Returning, Load: q.load, EmbeddedHook: q.hooks, - ContextualModdable: bob.ContextualModdable[*psqldialect.UpdateQuery]{ - Mods: append([]bob.ContextualMod[*psqldialect.UpdateQuery](nil), q.contextualMods...), - }, } return bob.BaseQuery[*psqldialect.UpdateQuery]{ @@ -223,10 +214,9 @@ func (s immutableUpdateState) withMods(queryMods ...bob.Mod[*psqldialect.UpdateQ } type derivedDeleteQuery struct { - state immutableDeleteState - load bob.Load - hooks bob.EmbeddedHook - contextualMods []bob.ContextualMod[*psqldialect.DeleteQuery] + state immutableDeleteState + load bob.Load + hooks bob.EmbeddedHook } type immutableDeleteState struct { @@ -253,9 +243,8 @@ func asImmutableDelete(q bob.BaseQuery[*psqldialect.DeleteQuery]) derivedDeleteQ Expressions: append([]any(nil), q.Expression.Returning.Expressions...), }, }, - load: q.Expression.Load, - hooks: q.Expression.EmbeddedHook, - contextualMods: append([]bob.ContextualMod[*psqldialect.DeleteQuery](nil), q.Expression.ContextualModdable.Mods...), + load: q.Expression.Load, + hooks: q.Expression.EmbeddedHook, } } @@ -303,10 +292,6 @@ func (q derivedDeleteQuery) BuildN(ctx context.Context, start int) (string, []an } func (q derivedDeleteQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - if len(q.contextualMods) > 0 { - return q.mutableBase().WriteQuery(ctx, w, start) - } - var args []any if len(q.state.With.CTEs) > 0 { withArgs, err := q.state.With.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) @@ -372,9 +357,6 @@ func (q derivedDeleteQuery) mutableBase() bob.BaseQuery[*psqldialect.DeleteQuery Returning: q.state.Returning, Load: q.load, EmbeddedHook: q.hooks, - ContextualModdable: bob.ContextualModdable[*psqldialect.DeleteQuery]{ - Mods: append([]bob.ContextualMod[*psqldialect.DeleteQuery](nil), q.contextualMods...), - }, } return bob.BaseQuery[*psqldialect.DeleteQuery]{ @@ -413,10 +395,9 @@ func (s immutableDeleteState) withMods(queryMods ...bob.Mod[*psqldialect.DeleteQ } type derivedInsertQuery struct { - state immutableInsertState - load bob.Load - hooks bob.EmbeddedHook - contextualMods []bob.ContextualMod[*psqldialect.InsertQuery] + state immutableInsertState + load bob.Load + hooks bob.EmbeddedHook } type immutableInsertState struct { @@ -446,9 +427,8 @@ func asImmutableInsert(q bob.BaseQuery[*psqldialect.InsertQuery]) derivedInsertQ Expressions: append([]any(nil), q.Expression.Returning.Expressions...), }, }, - load: q.Expression.Load, - hooks: q.Expression.EmbeddedHook, - contextualMods: append([]bob.ContextualMod[*psqldialect.InsertQuery](nil), q.Expression.ContextualModdable.Mods...), + load: q.Expression.Load, + hooks: q.Expression.EmbeddedHook, } } @@ -496,10 +476,6 @@ func (q derivedInsertQuery) BuildN(ctx context.Context, start int) (string, []an } func (q derivedInsertQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - if len(q.contextualMods) > 0 { - return q.mutableBase().WriteQuery(ctx, w, start) - } - var args []any if len(q.state.With.CTEs) > 0 { withArgs, err := q.state.With.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) @@ -571,9 +547,6 @@ func (q derivedInsertQuery) mutableBase() bob.BaseQuery[*psqldialect.InsertQuery Returning: q.state.Returning, Load: q.load, EmbeddedHook: q.hooks, - ContextualModdable: bob.ContextualModdable[*psqldialect.InsertQuery]{ - Mods: append([]bob.ContextualMod[*psqldialect.InsertQuery](nil), q.contextualMods...), - }, } return bob.BaseQuery[*psqldialect.InsertQuery]{ From 29f2602ef7a39e28f9fb4aec6f07fc5453a7cfd3 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:05:39 -0400 Subject: [PATCH 17/33] refactor(psql): remove dead immutable select fallback --- dialect/psql/immutable_select.go | 12 ------------ dialect/psql/immutable_select_test.go | 6 +----- 2 files changed, 1 insertion(+), 17 deletions(-) diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index 78f051d0..6f48c5ba 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -104,10 +104,6 @@ func (q derivedSelectQuery) BuildN(ctx context.Context, start int) (string, []an } func (q derivedSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - if q.requiresMutableWrite() { - return q.mutableBase().WriteQuery(ctx, w, start) - } - writer := immutableSelectWriter{ ctx: ctx, w: w, @@ -121,10 +117,6 @@ func (q derivedSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, s return writer.args, nil } -func (q derivedSelectQuery) requiresMutableWrite() bool { - return !q.state.supportsNativeWrite() -} - func (q derivedSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { w.WriteString("(") args, err := q.WriteQuery(ctx, w, start) @@ -236,10 +228,6 @@ func immutableStateFromMutable(q *psqldialect.SelectQuery) immutableSelectState } } -func (s immutableSelectState) supportsNativeWrite() bool { - return true -} - func (s immutableSelectState) selectColumns() []any { if len(s.SelectColumns) > 0 { return s.SelectColumns diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 1c977d78..e398c2ca 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -171,10 +171,6 @@ func TestImmutableSelectQueryApplyCombinedDoesNotMutateOriginal(t *testing.T) { sm.LimitCombined(1000), ) - if derived.derivedSelectQuery.requiresMutableWrite() { - t.Fatal("expected combined select shape to use native immutable writer") - } - baseSQL, _, err := base.Build(t.Context()) if err != nil { t.Fatal(err) @@ -329,7 +325,7 @@ func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := asCountQuery(q.Query.derivedSelectQuery.mutableBase()).Build(ctx); err != nil { + if _, _, err := q.Query.derivedSelectQuery.AsCount().Build(ctx); err != nil { b.Fatal(err) } From d16c25a34575abc19d36768faa0406224f0c1a39 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:09:28 -0400 Subject: [PATCH 18/33] refactor(psql): support common immutable select mods --- dialect/psql/dialect/mods.go | 34 ++++++++++++ dialect/psql/immutable_select.go | 53 ++++++++++++++++++- dialect/psql/immutable_select_test.go | 76 +++++++++++++++++++++++++++ dialect/psql/sm/qm.go | 17 ++---- 4 files changed, 166 insertions(+), 14 deletions(-) diff --git a/dialect/psql/dialect/mods.go b/dialect/psql/dialect/mods.go index 223dccae..ebf7cac9 100644 --- a/dialect/psql/dialect/mods.go +++ b/dialect/psql/dialect/mods.go @@ -19,6 +19,14 @@ func (di Distinct) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialec return bob.ExpressSlice(ctx, w, d, start, di.On, " ON (", ", ", ")") } +type DistinctMod struct { + On []any +} + +func (d DistinctMod) Apply(q *SelectQuery) { + q.Distinct.On = d.On +} + func With[Q interface{ AppendCTE(bob.Expression) }](name string, columns ...string) CTEChain[Q] { return CTEChain[Q](func() clause.CTE { return clause.CTE{ @@ -223,6 +231,32 @@ func (o OrderCombined) Apply(q *SelectQuery) { q.CombinedOrder.AppendOrder(o()) } +type LimitCombined struct { + Count any +} + +func (l LimitCombined) Apply(q *SelectQuery) { + q.CombinedLimit.SetLimit(l.Count) +} + +type OffsetCombined struct { + Count any +} + +func (o OffsetCombined) Apply(q *SelectQuery) { + q.CombinedOffset.SetOffset(o.Count) +} + +type FetchCombined struct { + Count any + WithTies bool +} + +func (f FetchCombined) Apply(q *SelectQuery) { + q.CombinedFetch.Count = f.Count + q.CombinedFetch.WithTies = f.WithTies +} + type OrderBy[Q interface{ AppendOrder(bob.Expression) }] func() clause.OrderDef func (s OrderBy[Q]) Apply(q Q) { diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go index 6f48c5ba..4f0b4675 100644 --- a/dialect/psql/immutable_select.go +++ b/dialect/psql/immutable_select.go @@ -267,16 +267,32 @@ func (s immutableSelectState) toMutable() psqldialect.SelectQuery { func (s immutableSelectState) withMods(queryMods ...bob.Mod[*psqldialect.SelectQuery]) (immutableSelectState, bool) { next := s - var cloneSelect, cloneWhere, cloneGroup, cloneHaving, cloneOrder, cloneWindows, cloneLocks bool + var cloneWith, cloneSelect, cloneWhere, cloneGroup, cloneHaving, cloneOrder, cloneWindows, cloneLocks, cloneJoins, cloneCombines, cloneCombinedOrder, clonePreload bool for _, mod := range queryMods { switch m := mod.(type) { + case mods.Recursive[*psqldialect.SelectQuery]: + next.With.Recursive = bool(m) + case psqldialect.CTEChain[*psqldialect.SelectQuery]: + if !cloneWith { + next.With.CTEs = append([]bob.Expression(nil), s.With.CTEs...) + cloneWith = true + } + next.With.CTEs = append(next.With.CTEs, m()) + case psqldialect.DistinctMod: + next.Distinct.On = cloneAnySlice(m.On) case mods.Select[*psqldialect.SelectQuery]: if !cloneSelect { next.SelectColumns = append([]any(nil), s.SelectColumns...) cloneSelect = true } next.SelectColumns = append(next.SelectColumns, []any(m)...) + case mods.Preload[*psqldialect.SelectQuery]: + if !clonePreload { + next.PreloadColumns = append([]any(nil), s.PreloadColumns...) + clonePreload = true + } + next.PreloadColumns = append(next.PreloadColumns, []any(m)...) case mods.Where[*psqldialect.SelectQuery]: if !cloneWhere { next.Where.Conditions = append([]any(nil), s.Where.Conditions...) @@ -289,6 +305,10 @@ func (s immutableSelectState) withMods(queryMods ...bob.Mod[*psqldialect.SelectQ cloneGroup = true } next.GroupBy.Groups = append(next.GroupBy.Groups, m.E) + case mods.GroupByDistinct[*psqldialect.SelectQuery]: + next.GroupBy.Distinct = bool(m) + case mods.GroupWith[*psqldialect.SelectQuery]: + next.GroupBy.With = string(m) case mods.Having[*psqldialect.SelectQuery]: if !cloneHaving { next.Having.Conditions = append([]any(nil), s.Having.Conditions...) @@ -307,6 +327,18 @@ func (s immutableSelectState) withMods(queryMods ...bob.Mod[*psqldialect.SelectQ cloneOrder = true } next.OrderBy.Expressions = append(next.OrderBy.Expressions, m()) + case mods.Join[*psqldialect.SelectQuery]: + if !cloneJoins { + next.TableRef.Joins = append([]clause.Join(nil), s.TableRef.Joins...) + cloneJoins = true + } + next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) + case psqldialect.CrossJoinChain[*psqldialect.SelectQuery]: + if !cloneJoins { + next.TableRef.Joins = append([]clause.Join(nil), s.TableRef.Joins...) + cloneJoins = true + } + next.TableRef.Joins = append(next.TableRef.Joins, m()) case mods.NamedWindow[*psqldialect.SelectQuery]: if !cloneWindows { next.Windows.Windows = append([]bob.Expression(nil), s.Windows.Windows...) @@ -319,6 +351,25 @@ func (s immutableSelectState) withMods(queryMods ...bob.Mod[*psqldialect.SelectQ cloneLocks = true } next.Locks.Locks = append(next.Locks.Locks, m()) + case mods.Combine[*psqldialect.SelectQuery]: + if !cloneCombines { + next.Combines.Queries = append([]clause.Combine(nil), s.Combines.Queries...) + cloneCombines = true + } + next.Combines.Queries = append(next.Combines.Queries, clause.Combine(m)) + case psqldialect.OrderCombined: + if !cloneCombinedOrder { + next.CombinedOrder.Expressions = append([]bob.Expression(nil), s.CombinedOrder.Expressions...) + cloneCombinedOrder = true + } + next.CombinedOrder.Expressions = append(next.CombinedOrder.Expressions, m()) + case psqldialect.LimitCombined: + next.CombinedLimit.Count = m.Count + case psqldialect.OffsetCombined: + next.CombinedOffset.Count = m.Count + case psqldialect.FetchCombined: + next.CombinedFetch.Count = m.Count + next.CombinedFetch.WithTies = m.WithTies case psqldialect.FromChain[*psqldialect.SelectQuery]: next.TableRef = cloneTableRef(m()) default: diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index e398c2ca..c379fae1 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -272,6 +272,82 @@ func TestViewSelectQueryHooksUseImmutableSelectQuery(t *testing.T) { } } +func TestImmutableSelectQueryApplySupportsCommonDerivedMods(t *testing.T) { + base := Select( + sm.Columns("users.id", "users.name"), + sm.From("users"), + sm.Where(Quote("users", "active").EQ(Arg(true))), + ) + + derived := base.Apply( + sm.With("admins", "id").As(Select( + sm.Columns("id"), + sm.From("admins"), + )), + sm.Recursive(true), + sm.Distinct("users.id"), + sm.LeftJoin("admins").As("a").OnEQ(Quote("a", "id"), Quote("users", "id")), + sm.GroupBy("users.id"), + sm.GroupByDistinct(true), + sm.UnionAll(Select( + sm.Columns("users.id", "users.name"), + sm.From("archived_users"), + )), + sm.OrderCombined("users.id"), + sm.LimitCombined(10), + sm.OffsetCombined(5), + sm.FetchCombined(3, false), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "SELECT \nusers.id, users.name\nFROM users\nWHERE (\"users\".\"active\" = $1)\n" { + t.Fatalf("base query changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, derivedArgs, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expected := Select( + sm.Columns("users.id", "users.name"), + sm.From("users"), + sm.Where(Quote("users", "active").EQ(Arg(true))), + sm.With("admins", "id").As(Select( + sm.Columns("id"), + sm.From("admins"), + )), + sm.Recursive(true), + sm.Distinct("users.id"), + sm.LeftJoin("admins").As("a").OnEQ(Quote("a", "id"), Quote("users", "id")), + sm.GroupBy("users.id"), + sm.GroupByDistinct(true), + sm.UnionAll(Select( + sm.Columns("users.id", "users.name"), + sm.From("archived_users"), + )), + sm.OrderCombined("users.id"), + sm.LimitCombined(10), + sm.OffsetCombined(5), + sm.FetchCombined(3, false), + ) + + expectedSQL, expectedArgs, err := expected.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + if derivedSQL != expectedSQL { + t.Fatalf("derived query mismatch: got %#v want %#v", derivedSQL, expectedSQL) + } + if len(derivedArgs) != len(expectedArgs) { + t.Fatalf("derived args mismatch: got %d want %d", len(derivedArgs), len(expectedArgs)) + } +} + func BenchmarkBaseQueryApplyMain(b *testing.B) { ctx := b.Context() diff --git a/dialect/psql/sm/qm.go b/dialect/psql/sm/qm.go index 81a81dd9..8fcaf5f1 100644 --- a/dialect/psql/sm/qm.go +++ b/dialect/psql/sm/qm.go @@ -20,9 +20,7 @@ func Distinct(on ...any) bob.Mod[*dialect.SelectQuery] { on = []any{} // nil means no distinct } - return bob.ModFunc[*dialect.SelectQuery](func(q *dialect.SelectQuery) { - q.Distinct.On = on - }) + return dialect.DistinctMod{On: on} } func Columns(clauses ...any) bob.Mod[*dialect.SelectQuery] { @@ -219,22 +217,15 @@ func OrderCombined(e any) dialect.OrderCombined { // To apply limit to the result of a UNION, INTERSECT, or EXCEPT query func LimitCombined(count any) bob.Mod[*dialect.SelectQuery] { - return bob.ModFunc[*dialect.SelectQuery](func(q *dialect.SelectQuery) { - q.CombinedLimit.SetLimit(count) - }) + return dialect.LimitCombined{Count: count} } // To apply offset to the result of a UNION, INTERSECT, or EXCEPT query func OffsetCombined(count any) bob.Mod[*dialect.SelectQuery] { - return bob.ModFunc[*dialect.SelectQuery](func(q *dialect.SelectQuery) { - q.CombinedOffset.SetOffset(count) - }) + return dialect.OffsetCombined{Count: count} } // To apply fetch to the result of a UNION, INTERSECT, or EXCEPT query func FetchCombined(count any, withTies bool) bob.Mod[*dialect.SelectQuery] { - return bob.ModFunc[*dialect.SelectQuery](func(q *dialect.SelectQuery) { - q.CombinedFetch.Count = count - q.CombinedFetch.WithTies = withTies - }) + return dialect.FetchCombined{Count: count, WithTies: withTies} } From 0a91ca3429c3258d3c55cf1de69143cbd47eda75 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:19:35 -0400 Subject: [PATCH 19/33] refactor(psql): support common immutable write mods --- dialect/psql/dialect/mods.go | 50 +++++++ dialect/psql/dm/qm.go | 23 ++-- dialect/psql/im/qm.go | 34 ++--- dialect/psql/immutable_write.go | 76 ++++++++++- dialect/psql/immutable_write_test.go | 191 +++++++++++++++++++++++++++ dialect/psql/um/qm.go | 27 ++-- 6 files changed, 343 insertions(+), 58 deletions(-) diff --git a/dialect/psql/dialect/mods.go b/dialect/psql/dialect/mods.go index ebf7cac9..e573815b 100644 --- a/dialect/psql/dialect/mods.go +++ b/dialect/psql/dialect/mods.go @@ -27,6 +27,56 @@ func (d DistinctMod) Apply(q *SelectQuery) { q.Distinct.On = d.On } +type UpdateOnly bool + +func (o UpdateOnly) Apply(q *UpdateQuery) { + q.Only = bool(o) +} + +type DeleteOnly bool + +func (o DeleteOnly) Apply(q *DeleteQuery) { + q.Only = bool(o) +} + +type UpdateTable clause.TableRef + +func (t UpdateTable) Apply(q *UpdateQuery) { + q.Table = clause.TableRef(t) +} + +type DeleteTable clause.TableRef + +func (t DeleteTable) Apply(q *DeleteQuery) { + q.Table = clause.TableRef(t) +} + +type InsertTable clause.TableRef + +func (t InsertTable) Apply(q *InsertQuery) { + q.TableRef = clause.TableRef(t) +} + +type UpdateSet []any + +func (s UpdateSet) Apply(q *UpdateQuery) { + q.Set.Set = append(q.Set.Set, []any(s)...) +} + +type InsertOverriding string + +func (o InsertOverriding) Apply(q *InsertQuery) { + q.Overriding = string(o) +} + +type InsertQuerySource struct { + Query bob.Query +} + +func (s InsertQuerySource) Apply(q *InsertQuery) { + q.Query = s.Query +} + func With[Q interface{ AppendCTE(bob.Expression) }](name string, columns ...string) CTEChain[Q] { return CTEChain[Q](func() clause.CTE { return clause.CTE{ diff --git a/dialect/psql/dm/qm.go b/dialect/psql/dm/qm.go index b282f86e..4635356d 100644 --- a/dialect/psql/dm/qm.go +++ b/dialect/psql/dm/qm.go @@ -2,7 +2,6 @@ package dm import ( "github.com/stephenafamo/bob" - "github.com/stephenafamo/bob/clause" "github.com/stephenafamo/bob/dialect/psql/dialect" "github.com/stephenafamo/bob/mods" ) @@ -16,26 +15,20 @@ func Recursive(r bool) bob.Mod[*dialect.DeleteQuery] { } func Only() bob.Mod[*dialect.DeleteQuery] { - return bob.ModFunc[*dialect.DeleteQuery](func(d *dialect.DeleteQuery) { - d.Only = true - }) + return dialect.DeleteOnly(true) } func From(name any) bob.Mod[*dialect.DeleteQuery] { - return bob.ModFunc[*dialect.DeleteQuery](func(u *dialect.DeleteQuery) { - u.Table = clause.TableRef{ - Expression: name, - } - }) + return dialect.DeleteTable{ + Expression: name, + } } func FromAs(name any, alias string) bob.Mod[*dialect.DeleteQuery] { - return bob.ModFunc[*dialect.DeleteQuery](func(u *dialect.DeleteQuery) { - u.Table = clause.TableRef{ - Expression: name, - Alias: alias, - } - }) + return dialect.DeleteTable{ + Expression: name, + Alias: alias, + } } func Using(table any) dialect.FromChain[*dialect.DeleteQuery] { diff --git a/dialect/psql/im/qm.go b/dialect/psql/im/qm.go index ca8464be..bb6c835b 100644 --- a/dialect/psql/im/qm.go +++ b/dialect/psql/im/qm.go @@ -18,34 +18,26 @@ func Recursive(r bool) bob.Mod[*dialect.InsertQuery] { } func Into(name any, columns ...string) bob.Mod[*dialect.InsertQuery] { - return bob.ModFunc[*dialect.InsertQuery](func(i *dialect.InsertQuery) { - i.TableRef = clause.TableRef{ - Expression: name, - Columns: columns, - } - }) + return dialect.InsertTable{ + Expression: name, + Columns: columns, + } } func IntoAs(name any, alias string, columns ...string) bob.Mod[*dialect.InsertQuery] { - return bob.ModFunc[*dialect.InsertQuery](func(i *dialect.InsertQuery) { - i.TableRef = clause.TableRef{ - Expression: name, - Alias: alias, - Columns: columns, - } - }) + return dialect.InsertTable{ + Expression: name, + Alias: alias, + Columns: columns, + } } func OverridingSystem() bob.Mod[*dialect.InsertQuery] { - return bob.ModFunc[*dialect.InsertQuery](func(i *dialect.InsertQuery) { - i.Overriding = dialect.OverridingSystem - }) + return dialect.InsertOverriding("SYSTEM") } func OverridingUser() bob.Mod[*dialect.InsertQuery] { - return bob.ModFunc[*dialect.InsertQuery](func(i *dialect.InsertQuery) { - i.Overriding = dialect.OverridingUser - }) + return dialect.InsertOverriding("USER") } func Values(clauses ...bob.Expression) bob.Mod[*dialect.InsertQuery] { @@ -58,9 +50,7 @@ func Rows(rows ...[]bob.Expression) bob.Mod[*dialect.InsertQuery] { // Insert from a query func Query(q bob.Query) bob.Mod[*dialect.InsertQuery] { - return bob.ModFunc[*dialect.InsertQuery](func(i *dialect.InsertQuery) { - i.Query = q - }) + return dialect.InsertQuerySource{Query: q} } // The column to target. Will auto add brackets diff --git a/dialect/psql/immutable_write.go b/dialect/psql/immutable_write.go index 6fe22119..17237263 100644 --- a/dialect/psql/immutable_write.go +++ b/dialect/psql/immutable_write.go @@ -187,10 +187,28 @@ func (q derivedUpdateQuery) mutableBase() bob.BaseQuery[*psqldialect.UpdateQuery func (s immutableUpdateState) withMods(queryMods ...bob.Mod[*psqldialect.UpdateQuery]) (immutableUpdateState, bool) { next := s - var cloneWhere, cloneReturning bool + var cloneWith, cloneSet, cloneWhere, cloneReturning, cloneJoins bool for _, mod := range queryMods { switch m := mod.(type) { + case mods.Recursive[*psqldialect.UpdateQuery]: + next.With.Recursive = bool(m) + case psqldialect.CTEChain[*psqldialect.UpdateQuery]: + if !cloneWith { + next.With.CTEs = append([]bob.Expression(nil), s.With.CTEs...) + cloneWith = true + } + next.With.CTEs = append(next.With.CTEs, m()) + case psqldialect.UpdateOnly: + next.Only = bool(m) + case psqldialect.UpdateTable: + next.Table = cloneTableRef(clause.TableRef(m)) + case psqldialect.UpdateSet: + if !cloneSet { + next.Set.Set = append([]any(nil), s.Set.Set...) + cloneSet = true + } + next.Set.Set = append(next.Set.Set, []any(m)...) case mods.Where[*psqldialect.UpdateQuery]: if !cloneWhere { next.Where.Conditions = append([]any(nil), s.Where.Conditions...) @@ -205,6 +223,18 @@ func (s immutableUpdateState) withMods(queryMods ...bob.Mod[*psqldialect.UpdateQ next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) case psqldialect.FromChain[*psqldialect.UpdateQuery]: next.From = cloneTableRef(m()) + case mods.Join[*psqldialect.UpdateQuery]: + if !cloneJoins { + next.From.Joins = append([]clause.Join(nil), s.From.Joins...) + cloneJoins = true + } + next.From.Joins = append(next.From.Joins, clause.Join(m)) + case psqldialect.CrossJoinChain[*psqldialect.UpdateQuery]: + if !cloneJoins { + next.From.Joins = append([]clause.Join(nil), s.From.Joins...) + cloneJoins = true + } + next.From.Joins = append(next.From.Joins, m()) default: return next, false } @@ -368,10 +398,22 @@ func (q derivedDeleteQuery) mutableBase() bob.BaseQuery[*psqldialect.DeleteQuery func (s immutableDeleteState) withMods(queryMods ...bob.Mod[*psqldialect.DeleteQuery]) (immutableDeleteState, bool) { next := s - var cloneWhere, cloneReturning bool + var cloneWith, cloneWhere, cloneReturning, cloneJoins bool for _, mod := range queryMods { switch m := mod.(type) { + case mods.Recursive[*psqldialect.DeleteQuery]: + next.With.Recursive = bool(m) + case psqldialect.CTEChain[*psqldialect.DeleteQuery]: + if !cloneWith { + next.With.CTEs = append([]bob.Expression(nil), s.With.CTEs...) + cloneWith = true + } + next.With.CTEs = append(next.With.CTEs, m()) + case psqldialect.DeleteOnly: + next.Only = bool(m) + case psqldialect.DeleteTable: + next.Table = cloneTableRef(clause.TableRef(m)) case mods.Where[*psqldialect.DeleteQuery]: if !cloneWhere { next.Where.Conditions = append([]any(nil), s.Where.Conditions...) @@ -386,6 +428,18 @@ func (s immutableDeleteState) withMods(queryMods ...bob.Mod[*psqldialect.DeleteQ next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) case psqldialect.FromChain[*psqldialect.DeleteQuery]: next.Using = cloneTableRef(m()) + case mods.Join[*psqldialect.DeleteQuery]: + if !cloneJoins { + next.Using.Joins = append([]clause.Join(nil), s.Using.Joins...) + cloneJoins = true + } + next.Using.Joins = append(next.Using.Joins, clause.Join(m)) + case psqldialect.CrossJoinChain[*psqldialect.DeleteQuery]: + if !cloneJoins { + next.Using.Joins = append([]clause.Join(nil), s.Using.Joins...) + cloneJoins = true + } + next.Using.Joins = append(next.Using.Joins, m()) default: return next, false } @@ -558,10 +612,24 @@ func (q derivedInsertQuery) mutableBase() bob.BaseQuery[*psqldialect.InsertQuery func (s immutableInsertState) withMods(queryMods ...bob.Mod[*psqldialect.InsertQuery]) (immutableInsertState, bool) { next := s - var cloneReturning, cloneVals bool + var cloneWith, cloneReturning, cloneVals bool for _, mod := range queryMods { switch m := mod.(type) { + case mods.Recursive[*psqldialect.InsertQuery]: + next.With.Recursive = bool(m) + case psqldialect.CTEChain[*psqldialect.InsertQuery]: + if !cloneWith { + next.With.CTEs = append([]bob.Expression(nil), s.With.CTEs...) + cloneWith = true + } + next.With.CTEs = append(next.With.CTEs, m()) + case psqldialect.InsertTable: + next.Table = cloneTableRef(clause.TableRef(m)) + case psqldialect.InsertOverriding: + next.Overriding = string(m) + case psqldialect.InsertQuerySource: + next.Values.Query = m.Query case mods.Returning[*psqldialect.InsertQuery]: if !cloneReturning { next.Returning.Expressions = append([]any(nil), s.Returning.Expressions...) @@ -582,6 +650,8 @@ func (s immutableInsertState) withMods(queryMods ...bob.Mod[*psqldialect.InsertQ for _, row := range m { next.Values.Vals = append(next.Values.Vals, clause.Value(row)) } + case mods.Conflict[*psqldialect.InsertQuery]: + next.Conflict.Expression = m() default: return next, false } diff --git a/dialect/psql/immutable_write_test.go b/dialect/psql/immutable_write_test.go index 4032f798..5e8afa7c 100644 --- a/dialect/psql/immutable_write_test.go +++ b/dialect/psql/immutable_write_test.go @@ -3,8 +3,10 @@ package psql import ( "testing" + "github.com/stephenafamo/bob" "github.com/stephenafamo/bob/dialect/psql/dm" "github.com/stephenafamo/bob/dialect/psql/im" + "github.com/stephenafamo/bob/dialect/psql/sm" "github.com/stephenafamo/bob/dialect/psql/um" ) @@ -172,6 +174,195 @@ func TestInsertApplyDoesNotMutateOriginal(t *testing.T) { } } +func TestUpdateApplySupportsCommonDerivedMods(t *testing.T) { + base := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Drama"), + ) + + derived := base.Apply( + um.With("recent").As(Select( + sm.Columns("id"), + sm.From("recent_films"), + )), + um.Recursive(true), + um.Only(), + um.TableAs("films", "f"), + um.Set(Quote("rating").EQ(Arg("PG"))), + um.From("producers"), + um.LeftJoin("studios").As("s").OnEQ(Quote("s", "id"), Quote("f", "studio_id")), + um.Where(Quote("f", "id").EQ(Arg(1))), + um.Returning("f.id"), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "UPDATE films SET\n\"kind\" = $1" { + t.Fatalf("base update changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, derivedArgs, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expected := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Drama"), + um.With("recent").As(Select( + sm.Columns("id"), + sm.From("recent_films"), + )), + um.Recursive(true), + um.Only(), + um.TableAs("films", "f"), + um.Set(Quote("rating").EQ(Arg("PG"))), + um.From("producers"), + um.LeftJoin("studios").As("s").OnEQ(Quote("s", "id"), Quote("f", "studio_id")), + um.Where(Quote("f", "id").EQ(Arg(1))), + um.Returning("f.id"), + ) + + expectedSQL, expectedArgs, err := expected.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != expectedSQL { + t.Fatalf("derived update mismatch: got %#v want %#v", derivedSQL, expectedSQL) + } + if len(derivedArgs) != len(expectedArgs) { + t.Fatalf("derived update args mismatch: got %d want %d", len(derivedArgs), len(expectedArgs)) + } +} + +func TestDeleteApplySupportsCommonDerivedMods(t *testing.T) { + base := Delete( + dm.From("films"), + ) + + derived := base.Apply( + dm.With("recent").As(Select( + sm.Columns("id"), + sm.From("recent_films"), + )), + dm.Recursive(true), + dm.Only(), + dm.FromAs("films", "f"), + dm.Using("producers"), + dm.LeftJoin("studios").As("s").OnEQ(Quote("s", "id"), Quote("f", "studio_id")), + dm.Where(Quote("f", "id").EQ(Arg(1))), + dm.Returning("f.id"), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "DELETE FROM films" { + t.Fatalf("base delete changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, derivedArgs, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expected := Delete( + dm.From("films"), + dm.With("recent").As(Select( + sm.Columns("id"), + sm.From("recent_films"), + )), + dm.Recursive(true), + dm.Only(), + dm.FromAs("films", "f"), + dm.Using("producers"), + dm.LeftJoin("studios").As("s").OnEQ(Quote("s", "id"), Quote("f", "studio_id")), + dm.Where(Quote("f", "id").EQ(Arg(1))), + dm.Returning("f.id"), + ) + + expectedSQL, expectedArgs, err := expected.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != expectedSQL { + t.Fatalf("derived delete mismatch: got %#v want %#v", derivedSQL, expectedSQL) + } + if len(derivedArgs) != len(expectedArgs) { + t.Fatalf("derived delete args mismatch: got %d want %d", len(derivedArgs), len(expectedArgs)) + } +} + +func TestInsertApplySupportsCommonDerivedMods(t *testing.T) { + base := Insert( + im.Into("films"), + ) + + derived := base.Apply( + im.With("recent").As(Select( + sm.Columns("id"), + sm.From("recent_films"), + )), + im.Recursive(true), + im.IntoAs("films", "f", "code", "title"), + im.OverridingUser(), + im.Rows( + []bob.Expression{Arg("UA502"), Arg("Bananas")}, + []bob.Expression{Arg("UA503"), Arg("Grapes")}, + ), + im.OnConflict("code").DoUpdate( + im.SetExcluded("title"), + ), + im.Returning("f.id"), + ) + + baseSQL, _, err := base.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if baseSQL != "INSERT INTO films\nDEFAULT VALUES\n" { + t.Fatalf("base insert changed unexpectedly: %#v", baseSQL) + } + + derivedSQL, derivedArgs, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expected := Insert( + im.Into("films"), + im.With("recent").As(Select( + sm.Columns("id"), + sm.From("recent_films"), + )), + im.Recursive(true), + im.IntoAs("films", "f", "code", "title"), + im.OverridingUser(), + im.Rows( + []bob.Expression{Arg("UA502"), Arg("Bananas")}, + []bob.Expression{Arg("UA503"), Arg("Grapes")}, + ), + im.OnConflict("code").DoUpdate( + im.SetExcluded("title"), + ), + im.Returning("f.id"), + ) + + expectedSQL, expectedArgs, err := expected.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if derivedSQL != expectedSQL { + t.Fatalf("derived insert mismatch: got %#v want %#v", derivedSQL, expectedSQL) + } + if len(derivedArgs) != len(expectedArgs) { + t.Fatalf("derived insert args mismatch: got %d want %d", len(derivedArgs), len(expectedArgs)) + } +} + func BenchmarkUpdateQueryApplyMain(b *testing.B) { ctx := b.Context() diff --git a/dialect/psql/um/qm.go b/dialect/psql/um/qm.go index 85abc22b..c93f0302 100644 --- a/dialect/psql/um/qm.go +++ b/dialect/psql/um/qm.go @@ -2,7 +2,6 @@ package um import ( "github.com/stephenafamo/bob" - "github.com/stephenafamo/bob/clause" "github.com/stephenafamo/bob/dialect/psql/dialect" "github.com/stephenafamo/bob/internal" "github.com/stephenafamo/bob/mods" @@ -17,32 +16,24 @@ func Recursive(r bool) bob.Mod[*dialect.UpdateQuery] { } func Only() bob.Mod[*dialect.UpdateQuery] { - return bob.ModFunc[*dialect.UpdateQuery](func(u *dialect.UpdateQuery) { - u.Only = true - }) + return dialect.UpdateOnly(true) } func Table(name any) bob.Mod[*dialect.UpdateQuery] { - return bob.ModFunc[*dialect.UpdateQuery](func(u *dialect.UpdateQuery) { - u.Table = clause.TableRef{ - Expression: name, - } - }) + return dialect.UpdateTable{ + Expression: name, + } } func TableAs(name any, alias string) bob.Mod[*dialect.UpdateQuery] { - return bob.ModFunc[*dialect.UpdateQuery](func(u *dialect.UpdateQuery) { - u.Table = clause.TableRef{ - Expression: name, - Alias: alias, - } - }) + return dialect.UpdateTable{ + Expression: name, + Alias: alias, + } } func Set(sets ...bob.Expression) bob.Mod[*dialect.UpdateQuery] { - return bob.ModFunc[*dialect.UpdateQuery](func(q *dialect.UpdateQuery) { - q.Set.Set = append(q.Set.Set, internal.ToAnySlice(sets)...) - }) + return dialect.UpdateSet(internal.ToAnySlice(sets)) } func SetCol(from string) mods.Set[*dialect.UpdateQuery] { From faf981fdcd7602ce8e46dadc79e5b1cab3222109 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:20:48 -0400 Subject: [PATCH 20/33] refactor(psql): build table queries directly --- dialect/psql/table.go | 35 +++++++++++++++++++++++++++++------ 1 file changed, 29 insertions(+), 6 deletions(-) diff --git a/dialect/psql/table.go b/dialect/psql/table.go index f6666eae..8930b305 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -5,11 +5,9 @@ import ( "reflect" "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/clause" "github.com/stephenafamo/bob/dialect/psql/dialect" - "github.com/stephenafamo/bob/dialect/psql/dm" - "github.com/stephenafamo/bob/dialect/psql/im" "github.com/stephenafamo/bob/dialect/psql/mm" - "github.com/stephenafamo/bob/dialect/psql/um" "github.com/stephenafamo/bob/expr" "github.com/stephenafamo/bob/internal" "github.com/stephenafamo/bob/internal/mappings" @@ -142,7 +140,16 @@ func (t *Table[T, Tslice, Tset, C]) Merge(queryMods ...bob.Mod[*dialect.MergeQue } func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Expression, queryMods []bob.Mod[*dialect.InsertQuery]) bob.BaseQuery[*dialect.InsertQuery] { - base := Insert(im.Into(name, nonGeneratedCols...)).derivedInsertQuery.mutableBase() + base := bob.BaseQuery[*dialect.InsertQuery]{ + Expression: &dialect.InsertQuery{ + TableRef: clause.TableRef{ + Expression: name, + Columns: nonGeneratedCols, + }, + }, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeInsert, + } if !hasInsertReturning(queryMods) { base.Expression.AppendReturning(orm.DefaultReturning(returning)) } @@ -150,7 +157,15 @@ func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Exp } func updateTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mod[*dialect.UpdateQuery]) bob.BaseQuery[*dialect.UpdateQuery] { - base := Update(um.Table(name)).derivedUpdateQuery.mutableBase() + base := bob.BaseQuery[*dialect.UpdateQuery]{ + Expression: &dialect.UpdateQuery{ + Table: clause.TableRef{ + Expression: name, + }, + }, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeUpdate, + } if !hasUpdateReturning(queryMods) { base.Expression.AppendReturning(orm.DefaultReturning(returning)) } @@ -158,7 +173,15 @@ func updateTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mo } func deleteTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mod[*dialect.DeleteQuery]) bob.BaseQuery[*dialect.DeleteQuery] { - base := Delete(dm.From(name)).derivedDeleteQuery.mutableBase() + base := bob.BaseQuery[*dialect.DeleteQuery]{ + Expression: &dialect.DeleteQuery{ + Table: clause.TableRef{ + Expression: name, + }, + }, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeDelete, + } if !hasDeleteReturning(queryMods) { base.Expression.AppendReturning(orm.DefaultReturning(returning)) } From 8f279d9c690fc17229258db9eca3bd5336ae66b3 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:23:23 -0400 Subject: [PATCH 21/33] refactor(psql): localize table returning behavior --- clause/returning.go | 4 -- dialect/psql/table.go | 136 +++++++++++++++++++++++++++++++++------ orm/default_returning.go | 24 ------- orm/query.go | 40 +----------- 4 files changed, 118 insertions(+), 86 deletions(-) delete mode 100644 orm/default_returning.go diff --git a/clause/returning.go b/clause/returning.go index f1ee5f35..ad0a48d3 100644 --- a/clause/returning.go +++ b/clause/returning.go @@ -15,10 +15,6 @@ func (r *Returning) HasReturning() bool { return len(r.Expressions) > 0 } -func (r Returning) ReturningExpressions() []any { - return r.Expressions -} - func (r *Returning) SetReturning(columns ...any) { r.Expressions = append(r.Expressions[:0], columns...) } diff --git a/dialect/psql/table.go b/dialect/psql/table.go index 8930b305..97c33dc2 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -16,13 +16,100 @@ import ( ) type ( - setter[T any] = orm.Setter[T, *dialect.InsertQuery, *dialect.UpdateQuery] - ormInsertQuery[T any, Tslice ~[]T] = orm.Query[*dialect.InsertQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - ormUpdateQuery[T any, Tslice ~[]T] = orm.Query[*dialect.UpdateQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - ormDeleteQuery[T any, Tslice ~[]T] = orm.Query[*dialect.DeleteQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - ormMergeQuery[T any, Tslice ~[]T] = orm.Query[*dialect.MergeQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] + setter[T any] = orm.Setter[T, *dialect.InsertQuery, *dialect.UpdateQuery] + ormMergeQuery[T any, Tslice ~[]T] = orm.Query[*dialect.MergeQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] ) +type ormInsertQuery[T any, Tslice ~[]T] struct { + orm.Query[*dialect.InsertQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] + hasDefaultReturning bool +} + +type ormUpdateQuery[T any, Tslice ~[]T] struct { + orm.Query[*dialect.UpdateQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] + hasDefaultReturning bool +} + +type ormDeleteQuery[T any, Tslice ~[]T] struct { + orm.Query[*dialect.DeleteQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] + hasDefaultReturning bool +} + +func (q ormInsertQuery[T, Tslice]) clone() ormInsertQuery[T, Tslice] { + return ormInsertQuery[T, Tslice]{ + Query: q.Query.Clone(), + hasDefaultReturning: q.hasDefaultReturning, + } +} + +func (q *ormInsertQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { + if q == nil { + return nil + } + + next := q.clone() + applyTableQueryMods(next.Expression, &next.hasDefaultReturning, queryMods...) + return &next +} + +func (q *ormInsertQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { + return q.With(queryMods...) +} + +func (q ormUpdateQuery[T, Tslice]) clone() ormUpdateQuery[T, Tslice] { + return ormUpdateQuery[T, Tslice]{ + Query: q.Query.Clone(), + hasDefaultReturning: q.hasDefaultReturning, + } +} + +func (q *ormUpdateQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { + if q == nil { + return nil + } + + next := q.clone() + applyTableQueryMods(next.Expression, &next.hasDefaultReturning, queryMods...) + return &next +} + +func (q *ormUpdateQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { + return q.With(queryMods...) +} + +func (q ormDeleteQuery[T, Tslice]) clone() ormDeleteQuery[T, Tslice] { + return ormDeleteQuery[T, Tslice]{ + Query: q.Query.Clone(), + hasDefaultReturning: q.hasDefaultReturning, + } +} + +func (q *ormDeleteQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { + if q == nil { + return nil + } + + next := q.clone() + applyTableQueryMods(next.Expression, &next.hasDefaultReturning, queryMods...) + return &next +} + +func (q *ormDeleteQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { + return q.With(queryMods...) +} + +func applyTableQueryMods[Q interface{ SetReturning(...any) }](query Q, hasDefaultReturning *bool, queryMods ...bob.Mod[Q]) { + for _, mod := range queryMods { + if returning, ok := any(mod).(interface{ ReturningValues() []any }); ok && *hasDefaultReturning { + query.SetReturning(returning.ReturningValues()...) + *hasDefaultReturning = false + continue + } + + mod.Apply(query) + } +} + func NewTable[T any, Tset setter[T], C bob.Expression](schema, tableName string, columns C) *Table[T, []T, Tset, C] { return NewTablex[T, []T, Tset](schema, tableName, columns) } @@ -74,11 +161,14 @@ func (t *Table[T, Tslice, Tset, C]) PrimaryKey() expr.ColumnsExpr { // Starts an insert query for this table func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { q := &ormInsertQuery[T, Tslice]{ - ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ - BaseQuery: insertTableBaseQuery(t.NameAs(), t.nonGeneratedCols, t.Columns, queryMods), - Hooks: &t.InsertQueryHooks, + Query: orm.Query[*dialect.InsertQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ + ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ + BaseQuery: insertTableBaseQuery(t.NameAs(), t.nonGeneratedCols, t.Columns, queryMods), + Hooks: &t.InsertQueryHooks, + }, + Scanner: t.scanner, }, - Scanner: t.scanner, + hasDefaultReturning: !hasInsertReturning(queryMods), } return q.Apply(queryMods...) @@ -87,11 +177,14 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ // Starts an Update query for this table func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { q := &ormUpdateQuery[T, Tslice]{ - ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ - BaseQuery: updateTableBaseQuery(t.NameAs(), t.Columns, queryMods), - Hooks: &t.UpdateQueryHooks, + Query: orm.Query[*dialect.UpdateQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ + ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ + BaseQuery: updateTableBaseQuery(t.NameAs(), t.Columns, queryMods), + Hooks: &t.UpdateQueryHooks, + }, + Scanner: t.scanner, }, - Scanner: t.scanner, + hasDefaultReturning: !hasUpdateReturning(queryMods), } return q.Apply(queryMods...) @@ -100,11 +193,14 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ // Starts a Delete query for this table func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { q := &ormDeleteQuery[T, Tslice]{ - ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ - BaseQuery: deleteTableBaseQuery(t.NameAs(), t.Columns, queryMods), - Hooks: &t.DeleteQueryHooks, + Query: orm.Query[*dialect.DeleteQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ + ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ + BaseQuery: deleteTableBaseQuery(t.NameAs(), t.Columns, queryMods), + Hooks: &t.DeleteQueryHooks, + }, + Scanner: t.scanner, }, - Scanner: t.scanner, + hasDefaultReturning: !hasDeleteReturning(queryMods), } return q.Apply(queryMods...) @@ -151,7 +247,7 @@ func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Exp QueryType: bob.QueryTypeInsert, } if !hasInsertReturning(queryMods) { - base.Expression.AppendReturning(orm.DefaultReturning(returning)) + base.Expression.AppendReturning(returning) } return base } @@ -167,7 +263,7 @@ func updateTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mo QueryType: bob.QueryTypeUpdate, } if !hasUpdateReturning(queryMods) { - base.Expression.AppendReturning(orm.DefaultReturning(returning)) + base.Expression.AppendReturning(returning) } return base } @@ -183,7 +279,7 @@ func deleteTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mo QueryType: bob.QueryTypeDelete, } if !hasDeleteReturning(queryMods) { - base.Expression.AppendReturning(orm.DefaultReturning(returning)) + base.Expression.AppendReturning(returning) } return base } diff --git a/orm/default_returning.go b/orm/default_returning.go deleted file mode 100644 index 8ec1bb32..00000000 --- a/orm/default_returning.go +++ /dev/null @@ -1,24 +0,0 @@ -package orm - -import ( - "context" - "io" - - "github.com/stephenafamo/bob" -) - -type defaultReturning struct { - expr bob.Expression -} - -func DefaultReturning(expr bob.Expression) bob.Expression { - return defaultReturning{expr: expr} -} - -func (d defaultReturning) IsDefaultReturning() bool { - return true -} - -func (d defaultReturning) WriteSQL(ctx context.Context, w io.StringWriter, dl bob.Dialect, start int) ([]any, error) { - return d.expr.WriteSQL(ctx, w, dl, start) -} diff --git a/orm/query.go b/orm/query.go index ef3e28d0..d7dfe83a 100644 --- a/orm/query.go +++ b/orm/query.go @@ -10,11 +10,6 @@ import ( "github.com/stephenafamo/scan" ) -type returningAware interface { - ReturningExpressions() []any - SetReturning(...any) -} - type ExecQuery[Q bob.Expression] struct { bob.BaseQuery[Q] Hooks *bob.Hooks[Q, bob.SkipQueryHooksKey] @@ -33,7 +28,7 @@ func (q *ExecQuery[Q]) With(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { } next := q.Clone() - applyQueryMods(next.BaseQuery.Expression, queryMods...) + next.BaseQuery.Apply(queryMods...) return &next } @@ -84,7 +79,7 @@ func (q *Query[Q, T, Ts, Tr]) With(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] } next := q.Clone() - applyQueryMods(next.BaseQuery.Expression, queryMods...) + next.BaseQuery.Apply(queryMods...) return &next } @@ -130,37 +125,6 @@ func (q ModQuery[Q, E, T, Ts, Tr]) Apply(e Q) { q.Mod.Apply(e) } -func applyQueryMods[Q any](query Q, queryMods ...bob.Mod[Q]) { - replacedDefaultReturning := false - - for _, mod := range queryMods { - if returning, ok := any(mod).(interface{ ReturningValues() []any }); ok && !replacedDefaultReturning { - if returningClause, ok := any(query).(returningAware); ok && hasOnlyDefaultReturning(returningClause.ReturningExpressions()) { - returningClause.SetReturning(returning.ReturningValues()...) - replacedDefaultReturning = true - continue - } - } - - mod.Apply(query) - } -} - -func hasOnlyDefaultReturning(expressions []any) bool { - if len(expressions) == 0 { - return false - } - - for _, expression := range expressions { - marker, ok := expression.(interface{ IsDefaultReturning() bool }) - if !ok || !marker.IsDefaultReturning() { - return false - } - } - - return true -} - func ArgsToExpression(querySQL string, from, to int, argIter iter.Seq[ArgWithPosition]) bob.Expression { return bob.ExpressionFunc(func(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { args := []any{} From 5e0b3589bb3ecaefc4a5f66c7a4a3a07408c2f16 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:25:00 -0400 Subject: [PATCH 22/33] perf(psql): build common queries immutably --- dialect/psql/delete.go | 9 +++++++++ dialect/psql/insert.go | 9 +++++++++ dialect/psql/select.go | 9 +++++++++ dialect/psql/update.go | 9 +++++++++ 4 files changed, 36 insertions(+) diff --git a/dialect/psql/delete.go b/dialect/psql/delete.go index c40e94c2..f8eb758d 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -19,6 +19,15 @@ func (q DeleteQuery) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQue } func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { + state, ok := (immutableDeleteState{}).withMods(queryMods...) + if ok { + return DeleteQuery{ + derivedDeleteQuery: derivedDeleteQuery{ + state: state, + }, + } + } + q := &dialect.DeleteQuery{} for _, mod := range queryMods { mod.Apply(q) diff --git a/dialect/psql/insert.go b/dialect/psql/insert.go index 16ad95f7..b5e82651 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -19,6 +19,15 @@ func (q InsertQuery) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQue } func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { + state, ok := (immutableInsertState{}).withMods(queryMods...) + if ok { + return InsertQuery{ + derivedInsertQuery: derivedInsertQuery{ + state: state, + }, + } + } + q := &dialect.InsertQuery{} for _, mod := range queryMods { mod.Apply(q) diff --git a/dialect/psql/select.go b/dialect/psql/select.go index e99102bc..2530a27e 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -19,6 +19,15 @@ func (q SelectQuery) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQue } func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { + state, ok := (immutableSelectState{}).withMods(queryMods...) + if ok { + return SelectQuery{ + derivedSelectQuery: derivedSelectQuery{ + state: state, + }, + } + } + q := &dialect.SelectQuery{} for _, mod := range queryMods { mod.Apply(q) diff --git a/dialect/psql/update.go b/dialect/psql/update.go index 6cbe2053..6109763c 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -19,6 +19,15 @@ func (q UpdateQuery) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQue } func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { + state, ok := (immutableUpdateState{}).withMods(queryMods...) + if ok { + return UpdateQuery{ + derivedUpdateQuery: derivedUpdateQuery{ + state: state, + }, + } + } + q := &dialect.UpdateQuery{} for _, mod := range queryMods { mod.Apply(q) From 831ca5097261a393b5e37fbfcfb143ebe791241c Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:35:58 -0400 Subject: [PATCH 23/33] refactor(query): make base queries immutable by default --- dialect/mysql/table.go | 2 +- dialect/mysql/view.go | 6 ++-- dialect/sqlite/view.go | 6 ++-- orm/query.go | 4 +-- query.go | 10 +++++-- query_immutable_test.go | 62 +++++++++++++++++++++++++++++++++++++++++ 6 files changed, 79 insertions(+), 11 deletions(-) create mode 100644 query_immutable_test.go diff --git a/dialect/mysql/table.go b/dialect/mysql/table.go index 63fea022..b6861987 100644 --- a/dialect/mysql/table.go +++ b/dialect/mysql/table.go @@ -295,7 +295,7 @@ func (t *insertQuery[T, Tslice, Tset, C]) getInserted(vals []clause.Value, resul filters = append(filters, Group(t.table.uniqueColNames(i)...).In(args...)) } - query.Apply(sm.Where(Or(filters...))) + query = query.Apply(sm.Where(Or(filters...))) return query, nil } diff --git a/dialect/mysql/view.go b/dialect/mysql/view.go index fe1e4524..8126865e 100644 --- a/dialect/mysql/view.go +++ b/dialect/mysql/view.go @@ -86,9 +86,9 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * }, ) - q.Apply(queryMods...) - - return q + next := *q + next.Query = *q.Query.Apply(queryMods...) + return &next } type ViewQuery[T any, Ts ~[]T] struct { diff --git a/dialect/sqlite/view.go b/dialect/sqlite/view.go index edcc537a..5e022770 100644 --- a/dialect/sqlite/view.go +++ b/dialect/sqlite/view.go @@ -108,9 +108,9 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * }, ) - q.Apply(queryMods...) - - return q + next := *q + next.Query = *q.Query.Apply(queryMods...) + return &next } type ViewQuery[T any, Ts ~[]T] struct { diff --git a/orm/query.go b/orm/query.go index d7dfe83a..c6ab7f80 100644 --- a/orm/query.go +++ b/orm/query.go @@ -28,7 +28,7 @@ func (q *ExecQuery[Q]) With(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { } next := q.Clone() - next.BaseQuery.Apply(queryMods...) + next.BaseQuery = next.BaseQuery.Apply(queryMods...) return &next } @@ -79,7 +79,7 @@ func (q *Query[Q, T, Ts, Tr]) With(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] } next := q.Clone() - next.BaseQuery.Apply(queryMods...) + next.BaseQuery = next.BaseQuery.Apply(queryMods...) return &next } diff --git a/query.go b/query.go index 73eca1af..767d3df4 100644 --- a/query.go +++ b/query.go @@ -117,10 +117,16 @@ func (b BaseQuery[E]) GetMapperMods() []scan.MapperMod { return nil } -func (b BaseQuery[E]) Apply(mods ...Mod[E]) { +func (b BaseQuery[E]) With(mods ...Mod[E]) BaseQuery[E] { + next := b.Clone() for _, mod := range mods { - mod.Apply(b.Expression) + mod.Apply(next.Expression) } + return next +} + +func (b BaseQuery[E]) Apply(mods ...Mod[E]) BaseQuery[E] { + return b.With(mods...) } func (b BaseQuery[E]) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { diff --git a/query_immutable_test.go b/query_immutable_test.go new file mode 100644 index 00000000..44acac0f --- /dev/null +++ b/query_immutable_test.go @@ -0,0 +1,62 @@ +package bob + +import ( + "context" + "io" + "slices" + "testing" +) + +type cloneableExpr struct { + parts []string +} + +func (e *cloneableExpr) Clone() *cloneableExpr { + return &cloneableExpr{ + parts: append([]string(nil), e.parts...), + } +} + +func (e *cloneableExpr) WriteSQL(context.Context, io.StringWriter, Dialect, int) ([]any, error) { + return nil, nil +} + +type appendExprMod string + +func (m appendExprMod) Apply(e *cloneableExpr) { + e.parts = append(e.parts, string(m)) +} + +func TestBaseQueryWithDoesNotMutateOriginal(t *testing.T) { + base := BaseQuery[*cloneableExpr]{ + Expression: &cloneableExpr{parts: []string{"base"}}, + QueryType: QueryTypeSelect, + } + + derived := base.With(appendExprMod("derived")) + + if !slices.Equal(base.Expression.parts, []string{"base"}) { + t.Fatalf("base query changed unexpectedly: %#v", base.Expression.parts) + } + + if !slices.Equal(derived.Expression.parts, []string{"base", "derived"}) { + t.Fatalf("derived query mismatch: %#v", derived.Expression.parts) + } +} + +func TestBaseQueryApplyDoesNotMutateOriginal(t *testing.T) { + base := BaseQuery[*cloneableExpr]{ + Expression: &cloneableExpr{parts: []string{"base"}}, + QueryType: QueryTypeSelect, + } + + derived := base.Apply(appendExprMod("derived")) + + if !slices.Equal(base.Expression.parts, []string{"base"}) { + t.Fatalf("base query changed unexpectedly: %#v", base.Expression.parts) + } + + if !slices.Equal(derived.Expression.parts, []string{"base", "derived"}) { + t.Fatalf("derived query mismatch: %#v", derived.Expression.parts) + } +} From aa059fe17c828c581fbb56a26fe633780a9f72d9 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 10:59:01 -0400 Subject: [PATCH 24/33] refactor(psql): add clone support for dialect queries --- dialect/psql/dialect/clone.go | 273 +++++++++++++++++++++++++++++ dialect/psql/dialect/clone_test.go | 92 ++++++++++ 2 files changed, 365 insertions(+) create mode 100644 dialect/psql/dialect/clone.go create mode 100644 dialect/psql/dialect/clone_test.go diff --git a/dialect/psql/dialect/clone.go b/dialect/psql/dialect/clone.go new file mode 100644 index 00000000..a04155d0 --- /dev/null +++ b/dialect/psql/dialect/clone.go @@ -0,0 +1,273 @@ +package dialect + +import ( + "context" + + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/clause" +) + +func cloneAnySlice(values []any) []any { + if values == nil { + return nil + } + return append([]any(nil), values...) +} + +func cloneStringSlice(values []string) []string { + if values == nil { + return nil + } + return append([]string(nil), values...) +} + +func cloneExpressionSlice(values []bob.Expression) []bob.Expression { + if values == nil { + return nil + } + return append([]bob.Expression(nil), values...) +} + +func cloneWith(with clause.With) clause.With { + return clause.With{ + Recursive: with.Recursive, + CTEs: cloneExpressionSlice(with.CTEs), + } +} + +func cloneSelectList(list clause.SelectList) clause.SelectList { + return clause.SelectList{ + Columns: cloneAnySlice(list.Columns), + PreloadColumns: cloneAnySlice(list.PreloadColumns), + } +} + +func cloneWhere(where clause.Where) clause.Where { + return clause.Where{Conditions: cloneAnySlice(where.Conditions)} +} + +func cloneGroupBy(groupBy clause.GroupBy) clause.GroupBy { + return clause.GroupBy{ + Groups: cloneAnySlice(groupBy.Groups), + Distinct: groupBy.Distinct, + With: groupBy.With, + } +} + +func cloneHaving(having clause.Having) clause.Having { + return clause.Having{Conditions: cloneAnySlice(having.Conditions)} +} + +func cloneWindows(windows clause.Windows) clause.Windows { + return clause.Windows{Windows: cloneExpressionSlice(windows.Windows)} +} + +func cloneOrderBy(orderBy clause.OrderBy) clause.OrderBy { + return clause.OrderBy{Expressions: cloneExpressionSlice(orderBy.Expressions)} +} + +func cloneLocks(locks clause.Locks) clause.Locks { + return clause.Locks{Locks: cloneExpressionSlice(locks.Locks)} +} + +func cloneLimit(limit clause.Limit) clause.Limit { + return clause.Limit{Count: limit.Count} +} + +func cloneOffset(offset clause.Offset) clause.Offset { + return clause.Offset{Count: offset.Count} +} + +func cloneFetch(fetch clause.Fetch) clause.Fetch { + return clause.Fetch{ + Count: fetch.Count, + WithTies: fetch.WithTies, + } +} + +func cloneReturning(returning clause.Returning) clause.Returning { + return clause.Returning{Expressions: cloneAnySlice(returning.Expressions)} +} + +func cloneSet(set clause.Set) clause.Set { + return clause.Set{Set: cloneAnySlice(set.Set)} +} + +func cloneConflict(conflict clause.Conflict) clause.Conflict { + return clause.Conflict{Expression: conflict.Expression} +} + +func cloneValues(values clause.Values) clause.Values { + cloned := clause.Values{ + Query: values.Query, + Vals: make([]clause.Value, 0, len(values.Vals)), + } + + for _, row := range values.Vals { + cloned.Vals = append(cloned.Vals, append(clause.Value(nil), row...)) + } + + return cloned +} + +func cloneCombines(combines clause.Combines) clause.Combines { + if combines.Queries == nil { + return clause.Combines{} + } + + queries := make([]clause.Combine, 0, len(combines.Queries)) + for _, combine := range combines.Queries { + queries = append(queries, clause.Combine{ + Strategy: combine.Strategy, + Query: combine.Query, + All: combine.All, + }) + } + + return clause.Combines{Queries: queries} +} + +func cloneTableRef(ref clause.TableRef) clause.TableRef { + var indexedBy *string + if ref.IndexedBy != nil { + indexed := *ref.IndexedBy + indexedBy = &indexed + } + + indexHints := make([]clause.IndexHint, 0, len(ref.IndexHints)) + for _, hint := range ref.IndexHints { + indexHints = append(indexHints, clause.IndexHint{ + Type: hint.Type, + Indexes: cloneStringSlice(hint.Indexes), + For: hint.For, + }) + } + + joins := make([]clause.Join, 0, len(ref.Joins)) + for _, join := range ref.Joins { + joins = append(joins, clause.Join{ + Type: join.Type, + Natural: join.Natural, + To: cloneTableRef(join.To), + On: cloneExpressionSlice(join.On), + Using: cloneStringSlice(join.Using), + }) + } + + return clause.TableRef{ + Expression: ref.Expression, + Alias: ref.Alias, + Columns: cloneStringSlice(ref.Columns), + Only: ref.Only, + Lateral: ref.Lateral, + WithOrdinality: ref.WithOrdinality, + IndexedBy: indexedBy, + Partitions: cloneStringSlice(ref.Partitions), + IndexHints: indexHints, + Joins: joins, + } +} + +func cloneLoad(load bob.Load) bob.Load { + var cloned bob.Load + cloned.SetLoaders(load.GetLoaders()...) + cloned.SetMapperMods(load.GetMapperMods()...) + return cloned +} + +func cloneEmbeddedHook(hook bob.EmbeddedHook) bob.EmbeddedHook { + return bob.EmbeddedHook{ + Hooks: append([]func(context.Context, bob.Executor) (context.Context, error){}, hook.Hooks...), + } +} + +func cloneContextualModdable[T any](mods bob.ContextualModdable[T]) bob.ContextualModdable[T] { + return bob.ContextualModdable[T]{ + Mods: append([]bob.ContextualMod[T](nil), mods.Mods...), + } +} + +func (s *SelectQuery) Clone() *SelectQuery { + if s == nil { + return nil + } + + return &SelectQuery{ + With: cloneWith(s.With), + SelectList: cloneSelectList(s.SelectList), + Distinct: Distinct{On: cloneAnySlice(s.Distinct.On)}, + TableRef: cloneTableRef(s.TableRef), + Where: cloneWhere(s.Where), + GroupBy: cloneGroupBy(s.GroupBy), + Having: cloneHaving(s.Having), + Windows: cloneWindows(s.Windows), + Combines: cloneCombines(s.Combines), + OrderBy: cloneOrderBy(s.OrderBy), + Limit: cloneLimit(s.Limit), + Offset: cloneOffset(s.Offset), + Fetch: cloneFetch(s.Fetch), + Locks: cloneLocks(s.Locks), + Load: cloneLoad(s.Load), + EmbeddedHook: cloneEmbeddedHook(s.EmbeddedHook), + ContextualModdable: cloneContextualModdable(s.ContextualModdable), + CombinedOrder: cloneOrderBy(s.CombinedOrder), + CombinedLimit: cloneLimit(s.CombinedLimit), + CombinedFetch: cloneFetch(s.CombinedFetch), + CombinedOffset: cloneOffset(s.CombinedOffset), + } +} + +func (u *UpdateQuery) Clone() *UpdateQuery { + if u == nil { + return nil + } + + return &UpdateQuery{ + With: cloneWith(u.With), + Only: u.Only, + Table: cloneTableRef(u.Table), + Set: cloneSet(u.Set), + TableRef: cloneTableRef(u.TableRef), + Where: cloneWhere(u.Where), + Returning: cloneReturning(u.Returning), + Load: cloneLoad(u.Load), + EmbeddedHook: cloneEmbeddedHook(u.EmbeddedHook), + ContextualModdable: cloneContextualModdable(u.ContextualModdable), + } +} + +func (d *DeleteQuery) Clone() *DeleteQuery { + if d == nil { + return nil + } + + return &DeleteQuery{ + With: cloneWith(d.With), + Only: d.Only, + Table: cloneTableRef(d.Table), + TableRef: cloneTableRef(d.TableRef), + Where: cloneWhere(d.Where), + Returning: cloneReturning(d.Returning), + Load: cloneLoad(d.Load), + EmbeddedHook: cloneEmbeddedHook(d.EmbeddedHook), + ContextualModdable: cloneContextualModdable(d.ContextualModdable), + } +} + +func (i *InsertQuery) Clone() *InsertQuery { + if i == nil { + return nil + } + + return &InsertQuery{ + With: cloneWith(i.With), + Overriding: i.Overriding, + TableRef: cloneTableRef(i.TableRef), + Values: cloneValues(i.Values), + Conflict: cloneConflict(i.Conflict), + Returning: cloneReturning(i.Returning), + Load: cloneLoad(i.Load), + EmbeddedHook: cloneEmbeddedHook(i.EmbeddedHook), + ContextualModdable: cloneContextualModdable(i.ContextualModdable), + } +} diff --git a/dialect/psql/dialect/clone_test.go b/dialect/psql/dialect/clone_test.go new file mode 100644 index 00000000..4cec6660 --- /dev/null +++ b/dialect/psql/dialect/clone_test.go @@ -0,0 +1,92 @@ +package dialect + +import ( + "context" + "io" + "testing" + + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/clause" +) + +func rawExpression(s string) bob.Expression { + return bob.ExpressionFunc(func(context.Context, io.StringWriter, bob.Dialect, int) ([]any, error) { + return nil, nil + }) +} + +func TestSelectQueryCloneDoesNotShareMutableState(t *testing.T) { + original := &SelectQuery{ + With: clause.With{ + CTEs: []bob.Expression{rawExpression("cte")}, + }, + SelectList: clause.SelectList{ + Columns: []any{"id"}, + PreloadColumns: []any{"email"}, + }, + Distinct: Distinct{ + On: []any{"id"}, + }, + TableRef: clause.TableRef{ + Expression: "users", + Joins: []clause.Join{{ + Type: clause.LeftJoin, + To: clause.TableRef{ + Expression: "profiles", + Alias: "p", + }, + }}, + }, + Where: clause.Where{ + Conditions: []any{"tenant_id = 1"}, + }, + GroupBy: clause.GroupBy{ + Groups: []any{"id"}, + }, + OrderBy: clause.OrderBy{ + Expressions: []bob.Expression{rawExpression("id DESC")}, + }, + Load: bob.Load{}, + } + original.AppendLoader(bob.LoaderFunc(func(_ context.Context, _ bob.Executor, _ any) error { return nil })) + + cloned := original.Clone() + + cloned.With.CTEs = append(cloned.With.CTEs, rawExpression("other_cte")) + cloned.SelectList.Columns = append(cloned.SelectList.Columns, "name") + cloned.SelectList.PreloadColumns = append(cloned.SelectList.PreloadColumns, "phone") + cloned.Distinct.On = append(cloned.Distinct.On, "name") + cloned.TableRef.Joins[0].To.Alias = "profiles_alias" + cloned.Where.Conditions = append(cloned.Where.Conditions, "active = true") + cloned.GroupBy.Groups = append(cloned.GroupBy.Groups, "name") + cloned.OrderBy.Expressions = append(cloned.OrderBy.Expressions, rawExpression("name ASC")) + cloned.AppendLoader(bob.LoaderFunc(func(_ context.Context, _ bob.Executor, _ any) error { return nil })) + + if len(original.With.CTEs) != 1 { + t.Fatalf("original with changed unexpectedly: %#v", original.With.CTEs) + } + if len(original.SelectList.Columns) != 1 { + t.Fatalf("original select columns changed unexpectedly: %#v", original.SelectList.Columns) + } + if len(original.SelectList.PreloadColumns) != 1 { + t.Fatalf("original preload columns changed unexpectedly: %#v", original.SelectList.PreloadColumns) + } + if len(original.Distinct.On) != 1 { + t.Fatalf("original distinct changed unexpectedly: %#v", original.Distinct.On) + } + if original.TableRef.Joins[0].To.Alias != "p" { + t.Fatalf("original join alias changed unexpectedly: %#v", original.TableRef.Joins[0].To.Alias) + } + if len(original.Where.Conditions) != 1 { + t.Fatalf("original where changed unexpectedly: %#v", original.Where.Conditions) + } + if len(original.GroupBy.Groups) != 1 { + t.Fatalf("original group by changed unexpectedly: %#v", original.GroupBy.Groups) + } + if len(original.OrderBy.Expressions) != 1 { + t.Fatalf("original order by changed unexpectedly: %#v", original.OrderBy.Expressions) + } + if len(original.GetLoaders()) != 1 { + t.Fatalf("original loaders changed unexpectedly: %d", len(original.GetLoaders())) + } +} From 2710d2b9f34de39059a26f3438e14eba5f9be168 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 11:10:24 -0400 Subject: [PATCH 25/33] refactor(psql): use a single immutable query core --- dialect/psql/delete.go | 21 +- dialect/psql/derive.go | 310 ++++++++++++ dialect/psql/dialect/clone.go | 2 +- dialect/psql/dialect/delete.go | 61 ++- dialect/psql/dialect/insert.go | 95 ++-- dialect/psql/dialect/select.go | 223 +++++---- dialect/psql/dialect/update.go | 70 +-- dialect/psql/dialect/writer.go | 144 ++++++ dialect/psql/immutable_select.go | 688 -------------------------- dialect/psql/immutable_select_test.go | 4 +- dialect/psql/immutable_write.go | 661 ------------------------- dialect/psql/insert.go | 21 +- dialect/psql/select.go | 34 +- dialect/psql/update.go | 21 +- dialect/psql/view.go | 48 +- 15 files changed, 788 insertions(+), 1615 deletions(-) create mode 100644 dialect/psql/derive.go create mode 100644 dialect/psql/dialect/writer.go delete mode 100644 dialect/psql/immutable_select.go delete mode 100644 dialect/psql/immutable_write.go diff --git a/dialect/psql/delete.go b/dialect/psql/delete.go index f8eb758d..135589df 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -6,11 +6,15 @@ import ( ) type DeleteQuery struct { - derivedDeleteQuery + bob.BaseQuery[*dialect.DeleteQuery] } func (q DeleteQuery) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { - q.derivedDeleteQuery = q.derivedDeleteQuery.With(queryMods...) + if next, ok := deriveDelete(q.Expression, queryMods...); ok { + q.Expression = next + return q + } + q.BaseQuery = q.BaseQuery.Apply(queryMods...) return q } @@ -19,25 +23,16 @@ func (q DeleteQuery) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQue } func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { - state, ok := (immutableDeleteState{}).withMods(queryMods...) - if ok { - return DeleteQuery{ - derivedDeleteQuery: derivedDeleteQuery{ - state: state, - }, - } - } - q := &dialect.DeleteQuery{} for _, mod := range queryMods { mod.Apply(q) } return DeleteQuery{ - derivedDeleteQuery: asImmutableDelete(bob.BaseQuery[*dialect.DeleteQuery]{ + BaseQuery: bob.BaseQuery[*dialect.DeleteQuery]{ Expression: q, Dialect: dialect.Dialect, QueryType: bob.QueryTypeDelete, - }), + }, } } diff --git a/dialect/psql/derive.go b/dialect/psql/derive.go new file mode 100644 index 00000000..d0f94ccb --- /dev/null +++ b/dialect/psql/derive.go @@ -0,0 +1,310 @@ +package psql + +import ( + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/clause" + "github.com/stephenafamo/bob/dialect/psql/dialect" + "github.com/stephenafamo/bob/mods" +) + +func copyAnySlice(values []any) []any { + if values == nil { + return nil + } + return append([]any(nil), values...) +} + +func copyExpressionSlice(values []bob.Expression) []bob.Expression { + if values == nil { + return nil + } + return append([]bob.Expression(nil), values...) +} + +func copyTableRef(from clause.TableRef) clause.TableRef { + from.Columns = append([]string(nil), from.Columns...) + from.Partitions = append([]string(nil), from.Partitions...) + from.IndexHints = append([]clause.IndexHint(nil), from.IndexHints...) + from.Joins = append([]clause.Join(nil), from.Joins...) + for i := range from.Joins { + from.Joins[i].On = append([]bob.Expression(nil), from.Joins[i].On...) + from.Joins[i].Using = append([]string(nil), from.Joins[i].Using...) + from.Joins[i].To = copyTableRef(from.Joins[i].To) + } + return from +} + +func deriveSelect(base *dialect.SelectQuery, queryMods ...bob.Mod[*dialect.SelectQuery]) (*dialect.SelectQuery, bool) { + next := *base + var cloneWith, cloneSelect, clonePreload, cloneWhere, cloneGroup, cloneHaving, cloneOrder, cloneWindows, cloneLocks, cloneJoins, cloneCombines, cloneCombinedOrder bool + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Recursive[*dialect.SelectQuery]: + next.With.Recursive = bool(m) + case dialect.CTEChain[*dialect.SelectQuery]: + if !cloneWith { + next.With.CTEs = copyExpressionSlice(base.With.CTEs) + cloneWith = true + } + next.With.CTEs = append(next.With.CTEs, m()) + case dialect.DistinctMod: + next.Distinct.On = append(make([]any, 0, len(m.On)), m.On...) + case mods.Select[*dialect.SelectQuery]: + if !cloneSelect { + next.SelectList.Columns = copyAnySlice(base.SelectList.Columns) + cloneSelect = true + } + next.SelectList.Columns = append(next.SelectList.Columns, []any(m)...) + case mods.Preload[*dialect.SelectQuery]: + if !clonePreload { + next.SelectList.PreloadColumns = copyAnySlice(base.SelectList.PreloadColumns) + clonePreload = true + } + next.SelectList.PreloadColumns = append(next.SelectList.PreloadColumns, []any(m)...) + case mods.Where[*dialect.SelectQuery]: + if !cloneWhere { + next.Where.Conditions = copyAnySlice(base.Where.Conditions) + cloneWhere = true + } + next.Where.Conditions = append(next.Where.Conditions, m.E) + case mods.GroupBy[*dialect.SelectQuery]: + if !cloneGroup { + next.GroupBy.Groups = copyAnySlice(base.GroupBy.Groups) + cloneGroup = true + } + next.GroupBy.Groups = append(next.GroupBy.Groups, m.E) + case mods.GroupByDistinct[*dialect.SelectQuery]: + next.GroupBy.Distinct = bool(m) + case mods.GroupWith[*dialect.SelectQuery]: + next.GroupBy.With = string(m) + case mods.Having[*dialect.SelectQuery]: + if !cloneHaving { + next.Having.Conditions = copyAnySlice(base.Having.Conditions) + cloneHaving = true + } + next.Having.Conditions = append(next.Having.Conditions, []any(m)...) + case mods.Limit[*dialect.SelectQuery]: + next.Limit.Count = m.Count + case mods.Offset[*dialect.SelectQuery]: + next.Offset.Count = m.Count + case mods.Fetch[*dialect.SelectQuery]: + next.Fetch = clause.Fetch(m) + case dialect.OrderBy[*dialect.SelectQuery]: + if !cloneOrder { + next.OrderBy.Expressions = copyExpressionSlice(base.OrderBy.Expressions) + cloneOrder = true + } + next.OrderBy.Expressions = append(next.OrderBy.Expressions, m()) + case mods.Join[*dialect.SelectQuery]: + if !cloneJoins { + next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) + cloneJoins = true + } + next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) + case dialect.CrossJoinChain[*dialect.SelectQuery]: + if !cloneJoins { + next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) + cloneJoins = true + } + next.TableRef.Joins = append(next.TableRef.Joins, m()) + case mods.NamedWindow[*dialect.SelectQuery]: + if !cloneWindows { + next.Windows.Windows = copyExpressionSlice(base.Windows.Windows) + cloneWindows = true + } + next.Windows.Windows = append(next.Windows.Windows, clause.NamedWindow(m)) + case dialect.LockChain[*dialect.SelectQuery]: + if !cloneLocks { + next.Locks.Locks = copyExpressionSlice(base.Locks.Locks) + cloneLocks = true + } + next.Locks.Locks = append(next.Locks.Locks, m()) + case mods.Combine[*dialect.SelectQuery]: + if !cloneCombines { + next.Combines.Queries = append([]clause.Combine(nil), base.Combines.Queries...) + cloneCombines = true + } + next.Combines.Queries = append(next.Combines.Queries, clause.Combine(m)) + case dialect.OrderCombined: + if !cloneCombinedOrder { + next.CombinedOrder.Expressions = copyExpressionSlice(base.CombinedOrder.Expressions) + cloneCombinedOrder = true + } + next.CombinedOrder.Expressions = append(next.CombinedOrder.Expressions, m()) + case dialect.LimitCombined: + next.CombinedLimit.Count = m.Count + case dialect.OffsetCombined: + next.CombinedOffset.Count = m.Count + case dialect.FetchCombined: + next.CombinedFetch.Count = m.Count + next.CombinedFetch.WithTies = m.WithTies + case dialect.FromChain[*dialect.SelectQuery]: + next.TableRef = copyTableRef(m()) + default: + return nil, false + } + } + + return &next, true +} + +func deriveUpdate(base *dialect.UpdateQuery, queryMods ...bob.Mod[*dialect.UpdateQuery]) (*dialect.UpdateQuery, bool) { + next := *base + var cloneWith, cloneSet, cloneWhere, cloneReturning, cloneJoins bool + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Recursive[*dialect.UpdateQuery]: + next.With.Recursive = bool(m) + case dialect.CTEChain[*dialect.UpdateQuery]: + if !cloneWith { + next.With.CTEs = copyExpressionSlice(base.With.CTEs) + cloneWith = true + } + next.With.CTEs = append(next.With.CTEs, m()) + case dialect.UpdateOnly: + next.Only = bool(m) + case dialect.UpdateTable: + next.Table = copyTableRef(clause.TableRef(m)) + case dialect.UpdateSet: + if !cloneSet { + next.Set.Set = copyAnySlice(base.Set.Set) + cloneSet = true + } + next.Set.Set = append(next.Set.Set, []any(m)...) + case mods.Where[*dialect.UpdateQuery]: + if !cloneWhere { + next.Where.Conditions = copyAnySlice(base.Where.Conditions) + cloneWhere = true + } + next.Where.Conditions = append(next.Where.Conditions, m.E) + case mods.Returning[*dialect.UpdateQuery]: + if !cloneReturning { + next.Returning.Expressions = copyAnySlice(base.Returning.Expressions) + cloneReturning = true + } + next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + case dialect.FromChain[*dialect.UpdateQuery]: + next.TableRef = copyTableRef(m()) + case mods.Join[*dialect.UpdateQuery]: + if !cloneJoins { + next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) + cloneJoins = true + } + next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) + case dialect.CrossJoinChain[*dialect.UpdateQuery]: + if !cloneJoins { + next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) + cloneJoins = true + } + next.TableRef.Joins = append(next.TableRef.Joins, m()) + default: + return nil, false + } + } + + return &next, true +} + +func deriveDelete(base *dialect.DeleteQuery, queryMods ...bob.Mod[*dialect.DeleteQuery]) (*dialect.DeleteQuery, bool) { + next := *base + var cloneWith, cloneWhere, cloneReturning, cloneJoins bool + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Recursive[*dialect.DeleteQuery]: + next.With.Recursive = bool(m) + case dialect.CTEChain[*dialect.DeleteQuery]: + if !cloneWith { + next.With.CTEs = copyExpressionSlice(base.With.CTEs) + cloneWith = true + } + next.With.CTEs = append(next.With.CTEs, m()) + case dialect.DeleteOnly: + next.Only = bool(m) + case dialect.DeleteTable: + next.Table = copyTableRef(clause.TableRef(m)) + case mods.Where[*dialect.DeleteQuery]: + if !cloneWhere { + next.Where.Conditions = copyAnySlice(base.Where.Conditions) + cloneWhere = true + } + next.Where.Conditions = append(next.Where.Conditions, m.E) + case mods.Returning[*dialect.DeleteQuery]: + if !cloneReturning { + next.Returning.Expressions = copyAnySlice(base.Returning.Expressions) + cloneReturning = true + } + next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + case dialect.FromChain[*dialect.DeleteQuery]: + next.TableRef = copyTableRef(m()) + case mods.Join[*dialect.DeleteQuery]: + if !cloneJoins { + next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) + cloneJoins = true + } + next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) + case dialect.CrossJoinChain[*dialect.DeleteQuery]: + if !cloneJoins { + next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) + cloneJoins = true + } + next.TableRef.Joins = append(next.TableRef.Joins, m()) + default: + return nil, false + } + } + + return &next, true +} + +func deriveInsert(base *dialect.InsertQuery, queryMods ...bob.Mod[*dialect.InsertQuery]) (*dialect.InsertQuery, bool) { + next := *base + var cloneWith, cloneReturning, cloneVals bool + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Recursive[*dialect.InsertQuery]: + next.With.Recursive = bool(m) + case dialect.CTEChain[*dialect.InsertQuery]: + if !cloneWith { + next.With.CTEs = copyExpressionSlice(base.With.CTEs) + cloneWith = true + } + next.With.CTEs = append(next.With.CTEs, m()) + case dialect.InsertTable: + next.TableRef = copyTableRef(clause.TableRef(m)) + case dialect.InsertOverriding: + next.Overriding = string(m) + case dialect.InsertQuerySource: + next.Values.Query = m.Query + case mods.Returning[*dialect.InsertQuery]: + if !cloneReturning { + next.Returning.Expressions = copyAnySlice(base.Returning.Expressions) + cloneReturning = true + } + next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + case mods.Values[*dialect.InsertQuery]: + if !cloneVals { + next.Values.Vals = append([]clause.Value(nil), base.Values.Vals...) + cloneVals = true + } + next.Values.Vals = append(next.Values.Vals, clause.Value(m)) + case mods.Rows[*dialect.InsertQuery]: + if !cloneVals { + next.Values.Vals = append([]clause.Value(nil), base.Values.Vals...) + cloneVals = true + } + for _, row := range m { + next.Values.Vals = append(next.Values.Vals, clause.Value(row)) + } + case mods.Conflict[*dialect.InsertQuery]: + next.Conflict.Expression = m() + default: + return nil, false + } + } + + return &next, true +} diff --git a/dialect/psql/dialect/clone.go b/dialect/psql/dialect/clone.go index a04155d0..aa37b1bd 100644 --- a/dialect/psql/dialect/clone.go +++ b/dialect/psql/dialect/clone.go @@ -11,7 +11,7 @@ func cloneAnySlice(values []any) []any { if values == nil { return nil } - return append([]any(nil), values...) + return append(make([]any, 0, len(values)), values...) } func cloneStringSlice(values []string) []string { diff --git a/dialect/psql/dialect/delete.go b/dialect/psql/dialect/delete.go index c32e7673..49d4da97 100644 --- a/dialect/psql/dialect/delete.go +++ b/dialect/psql/dialect/delete.go @@ -25,51 +25,60 @@ type DeleteQuery struct { func (d DeleteQuery) WriteSQL(ctx context.Context, w io.StringWriter, dl bob.Dialect, start int) ([]any, error) { var err error - var args []any if ctx, err = d.RunContextualMods(ctx, &d); err != nil { return nil, err } - withArgs, err := bob.ExpressIf(ctx, w, dl, start+len(args), d.With, - len(d.With.CTEs) > 0, "\n", "") - if err != nil { - return nil, err + writer := queryWriter{ + ctx: ctx, + w: w, + start: start, } - args = append(args, withArgs...) - w.WriteString("DELETE FROM ") + if len(d.With.CTEs) > 0 { + args, err := d.With.WriteSQL(ctx, w, dl, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(args) + _, _ = w.WriteString("\n") + } + + _, _ = w.WriteString("DELETE FROM ") if d.Only { - w.WriteString("ONLY ") + _, _ = w.WriteString("ONLY ") } - tableArgs, err := bob.ExpressIf(ctx, w, dl, start+len(args), d.Table, true, "", "") + tableArgs, err := d.Table.WriteSQL(ctx, w, dl, writer.argPos()) if err != nil { return nil, err } - args = append(args, tableArgs...) + writer.appendArgs(tableArgs) - usingArgs, err := bob.ExpressIf(ctx, w, dl, start+len(args), d.TableRef, - d.TableRef.Expression != nil, "\nUSING ", "") - if err != nil { - return nil, err + if d.TableRef.Expression != nil { + _, _ = w.WriteString("\nUSING ") + usingArgs, err := d.TableRef.WriteSQL(ctx, w, dl, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(usingArgs) } - args = append(args, usingArgs...) - whereArgs, err := bob.ExpressIf(ctx, w, dl, start+len(args), d.Where, - len(d.Where.Conditions) > 0, "\n", "") - if err != nil { - return nil, err + if len(d.Where.Conditions) > 0 { + _, _ = w.WriteString("\nWHERE ") + if err := writer.writeSliceAny(d.Where.Conditions, " AND "); err != nil { + return nil, err + } } - args = append(args, whereArgs...) - retArgs, err := bob.ExpressIf(ctx, w, dl, start+len(args), d.Returning, - len(d.Returning.Expressions) > 0, "\n", "") - if err != nil { - return nil, err + if len(d.Returning.Expressions) > 0 { + _, _ = w.WriteString("\nRETURNING ") + if err := writer.writeSliceAny(d.Returning.Expressions, ", "); err != nil { + return nil, err + } } - args = append(args, retArgs...) - return args, nil + return writer.args, nil } diff --git a/dialect/psql/dialect/insert.go b/dialect/psql/dialect/insert.go index 477060b3..85892b5b 100644 --- a/dialect/psql/dialect/insert.go +++ b/dialect/psql/dialect/insert.go @@ -34,52 +34,85 @@ type InsertQuery struct { func (i InsertQuery) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { var err error - var args []any if ctx, err = i.RunContextualMods(ctx, &i); err != nil { return nil, err } - withArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), i.With, - len(i.With.CTEs) > 0, "", "\n") - if err != nil { - return nil, err + writer := queryWriter{ + ctx: ctx, + w: w, + start: start, } - args = append(args, withArgs...) - tableArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), i.TableRef, - true, "INSERT INTO ", "") - if err != nil { - return nil, err + if len(i.With.CTEs) > 0 { + args, err := i.With.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(args) + _, _ = w.WriteString("\n") } - args = append(args, tableArgs...) - _, err = bob.ExpressIf(ctx, w, d, start+len(args), i.Overriding, - i.Overriding != "", "\nOVERRIDING ", " VALUE") - if err != nil { - return nil, err + _, _ = w.WriteString("INSERT INTO ") + if isSimpleTableRef(i.TableRef) { + if err := writer.writeAny(i.TableRef.Expression); err != nil { + return nil, err + } + } else { + tableArgs, err := i.TableRef.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(tableArgs) } - valArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), i.Values, true, "\n", "") - if err != nil { - return nil, err + if i.Overriding != "" { + _, _ = w.WriteString("\nOVERRIDING ") + _, _ = w.WriteString(i.Overriding) + _, _ = w.WriteString(" VALUE") } - args = append(args, valArgs...) - conflictArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), i.Conflict.Expression, - i.Conflict.Expression != nil, "\n", "") - if err != nil { - return nil, err + _, _ = w.WriteString("\n") + switch { + case i.Values.Query != nil: + valArgs, err := i.Values.Query.WriteQuery(ctx, w, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(valArgs) + case len(i.Values.Vals) > 0: + _, _ = w.WriteString("VALUES ") + for rowIndex, row := range i.Values.Vals { + if rowIndex > 0 { + _, _ = w.WriteString(", ") + } + _, _ = w.WriteString("(") + if err := writer.writeSliceExpr(row, ", "); err != nil { + return nil, err + } + _, _ = w.WriteString(")") + } + default: + _, _ = w.WriteString("DEFAULT VALUES") } - args = append(args, conflictArgs...) - retArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), i.Returning, - len(i.Returning.Expressions) > 0, "\n", "") - if err != nil { - return nil, err + if i.Conflict.Expression != nil { + _, _ = w.WriteString("\n") + conflictArgs, err := i.Conflict.Expression.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(conflictArgs) + } + + if len(i.Returning.Expressions) > 0 { + _, _ = w.WriteString("\nRETURNING ") + if err := writer.writeSliceAny(i.Returning.Expressions, ", "); err != nil { + return nil, err + } } - args = append(args, retArgs...) - w.WriteString("\n") - return args, nil + _, _ = w.WriteString("\n") + return writer.args, nil } diff --git a/dialect/psql/dialect/select.go b/dialect/psql/dialect/select.go index 7642294d..1b6a1f22 100644 --- a/dialect/psql/dialect/select.go +++ b/dialect/psql/dialect/select.go @@ -38,18 +38,25 @@ type SelectQuery struct { func (s SelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { var err error - var args []any if ctx, err = s.RunContextualMods(ctx, &s); err != nil { return nil, err } - withArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.With, - len(s.With.CTEs) > 0, "\n", "") - if err != nil { - return nil, err + writer := queryWriter{ + ctx: ctx, + w: w, + start: start, + } + + if len(s.With.CTEs) > 0 { + args, err := s.With.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(args) + _, _ = w.WriteString("\n") } - args = append(args, withArgs...) needsParens := false if len(s.Combines.Queries) > 0 && @@ -58,133 +65,163 @@ func (s SelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dial s.Offset.Count != nil || s.Fetch.Count != nil || len(s.Locks.Locks) > 0) { - w.WriteString("(") + _, _ = w.WriteString("(") needsParens = true } - w.WriteString("SELECT ") + _, _ = w.WriteString("SELECT ") - distinctArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.Distinct, - s.Distinct.On != nil, "", " ") - if err != nil { - return nil, err + if s.Distinct.On != nil { + _, _ = w.WriteString("DISTINCT") + if len(s.Distinct.On) > 0 { + _, _ = w.WriteString(" ON (") + if err := writer.writeSliceAny(s.Distinct.On, ", "); err != nil { + return nil, err + } + _, _ = w.WriteString(")") + } + _, _ = w.WriteString(" ") } - args = append(args, distinctArgs...) - selArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.SelectList, true, "\n", "") - if err != nil { + _, _ = w.WriteString("\n") + allCols := append([]any(nil), s.SelectList.Columns...) + allCols = append(allCols, s.SelectList.PreloadColumns...) + if len(allCols) == 0 { + _, _ = w.WriteString("*") + } else if err := writer.writeSliceAny(allCols, ", "); err != nil { return nil, err } - args = append(args, selArgs...) - fromArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.TableRef, s.TableRef.Expression != nil, "\nFROM ", "") - if err != nil { - return nil, err + if s.TableRef.Expression != nil { + _, _ = w.WriteString("\nFROM ") + args, err := s.TableRef.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(args) } - args = append(args, fromArgs...) - whereArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.Where, - len(s.Where.Conditions) > 0, "\n", "") - if err != nil { - return nil, err + if len(s.Where.Conditions) > 0 { + _, _ = w.WriteString("\nWHERE ") + if err := writer.writeSliceAny(s.Where.Conditions, " AND "); err != nil { + return nil, err + } } - args = append(args, whereArgs...) - groupByArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.GroupBy, - len(s.GroupBy.Groups) > 0, "\n", "") - if err != nil { - return nil, err + if len(s.GroupBy.Groups) > 0 { + _, _ = w.WriteString("\nGROUP BY ") + if s.GroupBy.Distinct { + _, _ = w.WriteString("DISTINCT ") + } + if err := writer.writeSliceAny(s.GroupBy.Groups, ", "); err != nil { + return nil, err + } + if s.GroupBy.With != "" { + _, _ = w.WriteString(" WITH ") + _, _ = w.WriteString(s.GroupBy.With) + } } - args = append(args, groupByArgs...) - havingArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.Having, - len(s.Having.Conditions) > 0, "\n", "") - if err != nil { - return nil, err + if len(s.Having.Conditions) > 0 { + _, _ = w.WriteString("\nHAVING ") + if err := writer.writeSliceAny(s.Having.Conditions, " AND "); err != nil { + return nil, err + } } - args = append(args, havingArgs...) - windowArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.Windows, - len(s.Windows.Windows) > 0, "\n", "") - if err != nil { - return nil, err + if len(s.Windows.Windows) > 0 { + _, _ = w.WriteString("\nWINDOW ") + if err := writer.writeSliceExpr(s.Windows.Windows, ", "); err != nil { + return nil, err + } } - args = append(args, windowArgs...) - orderArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.OrderBy, - len(s.OrderBy.Expressions) > 0, "\n", "") - if err != nil { - return nil, err + if len(s.OrderBy.Expressions) > 0 { + _, _ = w.WriteString("\nORDER BY ") + if err := writer.writeOrderExprs(s.OrderBy.Expressions); err != nil { + return nil, err + } } - args = append(args, orderArgs...) - limitArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.Limit, - s.Limit.Count != nil, "\n", "") - if err != nil { - return nil, err + if s.Limit.Count != nil { + _, _ = w.WriteString("\nLIMIT ") + if err := writer.writeAny(s.Limit.Count); err != nil { + return nil, err + } } - args = append(args, limitArgs...) - offsetArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.Offset, - s.Offset.Count != nil, "\n", "") - if err != nil { - return nil, err + if s.Offset.Count != nil { + _, _ = w.WriteString("\nOFFSET ") + if err := writer.writeAny(s.Offset.Count); err != nil { + return nil, err + } } - args = append(args, offsetArgs...) - fetchArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.Fetch, - s.Fetch.Count != nil, "\n", "") - if err != nil { - return nil, err + if s.Fetch.Count != nil { + _, _ = w.WriteString("\nFETCH NEXT ") + if err := writer.writeAny(s.Fetch.Count); err != nil { + return nil, err + } + if s.Fetch.WithTies { + _, _ = w.WriteString(" ROWS WITH TIES") + } else { + _, _ = w.WriteString(" ROWS ONLY") + } } - args = append(args, fetchArgs...) - lockArgs, err := bob.ExpressSlice(ctx, w, d, start+len(args), s.Locks.Locks, - "\n", "\n", "") - if err != nil { - return nil, err + for _, lock := range s.Locks.Locks { + _, _ = w.WriteString("\n") + if err := writer.writeExpression(lock); err != nil { + return nil, err + } } - args = append(args, lockArgs...) if needsParens { - w.WriteString(")") + _, _ = w.WriteString(")") } - combineArgs, err := bob.ExpressSlice(ctx, w, d, start+len(args), - s.Combines.Queries, "\n", "\n", "") - if err != nil { - return nil, err + for _, combine := range s.Combines.Queries { + _, _ = w.WriteString("\n") + args, err := combine.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(args) } - args = append(args, combineArgs...) - combinedOrderArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.CombinedOrder, - len(s.CombinedOrder.Expressions) > 0, "\n", "") - if err != nil { - return nil, err + if len(s.CombinedOrder.Expressions) > 0 { + _, _ = w.WriteString("\nORDER BY ") + if err := writer.writeOrderExprs(s.CombinedOrder.Expressions); err != nil { + return nil, err + } } - args = append(args, combinedOrderArgs...) - combinedLimitArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.CombinedLimit, - s.CombinedLimit.Count != nil, "\n", "") - if err != nil { - return nil, err + if s.CombinedLimit.Count != nil { + _, _ = w.WriteString("\nLIMIT ") + if err := writer.writeAny(s.CombinedLimit.Count); err != nil { + return nil, err + } } - args = append(args, combinedLimitArgs...) - combinedOffsetArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.CombinedOffset, - s.CombinedOffset.Count != nil, "\n", "") - if err != nil { - return nil, err + if s.CombinedOffset.Count != nil { + _, _ = w.WriteString("\nOFFSET ") + if err := writer.writeAny(s.CombinedOffset.Count); err != nil { + return nil, err + } } - args = append(args, combinedOffsetArgs...) - combinedFetchArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), s.CombinedFetch, - s.CombinedFetch.Count != nil, "\n", "") - if err != nil { - return nil, err + if s.CombinedFetch.Count != nil { + _, _ = w.WriteString("\nFETCH NEXT ") + if err := writer.writeAny(s.CombinedFetch.Count); err != nil { + return nil, err + } + if s.CombinedFetch.WithTies { + _, _ = w.WriteString(" ROWS WITH TIES") + } else { + _, _ = w.WriteString(" ROWS ONLY") + } } - args = append(args, combinedFetchArgs...) - w.WriteString("\n") - return args, nil + _, _ = w.WriteString("\n") + return writer.args, nil } diff --git a/dialect/psql/dialect/update.go b/dialect/psql/dialect/update.go index c392fa45..885f3892 100644 --- a/dialect/psql/dialect/update.go +++ b/dialect/psql/dialect/update.go @@ -26,57 +26,71 @@ type UpdateQuery struct { func (u UpdateQuery) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { var err error - var args []any if ctx, err = u.RunContextualMods(ctx, &u); err != nil { return nil, err } - withArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), u.With, - len(u.With.CTEs) > 0, "\n", "") - if err != nil { - return nil, err + writer := queryWriter{ + ctx: ctx, + w: w, + start: start, } - args = append(args, withArgs...) - w.WriteString("UPDATE ") + if len(u.With.CTEs) > 0 { + args, err := u.With.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(args) + _, _ = w.WriteString("\n") + } + + _, _ = w.WriteString("UPDATE ") if u.Only { - w.WriteString("ONLY ") + _, _ = w.WriteString("ONLY ") } - tableArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), u.Table, true, "", "") + tableArgs, err := u.Table.WriteSQL(ctx, w, d, writer.argPos()) if err != nil { return nil, err } - args = append(args, tableArgs...) + writer.appendArgs(tableArgs) - setArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), u.Set, true, " SET\n", "") + _, _ = w.WriteString(" SET\n") + setArgs, err := u.Set.WriteSQL(ctx, w, d, writer.argPos()) if err != nil { return nil, err } - args = append(args, setArgs...) + writer.appendArgs(setArgs) - fromArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), u.TableRef, - u.TableRef.Expression != nil, "\nFROM ", "") - if err != nil { - return nil, err + if u.TableRef.Expression != nil { + _, _ = w.WriteString("\nFROM ") + fromArgs, err := u.TableRef.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(fromArgs) } - args = append(args, fromArgs...) - whereArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), u.Where, - len(u.Where.Conditions) > 0, "\n", "") - if err != nil { - return nil, err + if len(u.Where.Conditions) > 0 { + _, _ = w.WriteString("\n") + whereArgs, err := u.Where.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(whereArgs) } - args = append(args, whereArgs...) - retArgs, err := bob.ExpressIf(ctx, w, d, start+len(args), u.Returning, - len(u.Returning.Expressions) > 0, "\n", "") - if err != nil { - return nil, err + if len(u.Returning.Expressions) > 0 { + _, _ = w.WriteString("\n") + retArgs, err := u.Returning.WriteSQL(ctx, w, d, writer.argPos()) + if err != nil { + return nil, err + } + writer.appendArgs(retArgs) } - args = append(args, retArgs...) - return args, nil + return writer.args, nil } diff --git a/dialect/psql/dialect/writer.go b/dialect/psql/dialect/writer.go new file mode 100644 index 00000000..612669da --- /dev/null +++ b/dialect/psql/dialect/writer.go @@ -0,0 +1,144 @@ +package dialect + +import ( + "context" + "database/sql" + "fmt" + "io" + "strconv" + + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/clause" +) + +func isSimpleTableRef(ref clause.TableRef) bool { + return ref.Expression != nil && + ref.Alias == "" && + len(ref.Columns) == 0 && + !ref.Only && + !ref.Lateral && + !ref.WithOrdinality && + ref.IndexedBy == nil && + len(ref.Partitions) == 0 && + len(ref.IndexHints) == 0 && + len(ref.Joins) == 0 +} + +type queryWriter struct { + ctx context.Context + w io.StringWriter + args []any + start int +} + +func (w *queryWriter) argPos() int { + return w.start + len(w.args) +} + +func (w *queryWriter) appendArgs(args []any) { + w.args = append(w.args, args...) +} + +func (w *queryWriter) writeExpression(value bob.Expression) error { + args, err := value.WriteSQL(w.ctx, w.w, Dialect, w.argPos()) + if err != nil { + return err + } + w.appendArgs(args) + return nil +} + +func (w *queryWriter) writeAny(value any) error { + switch v := value.(type) { + case nil: + _, _ = w.w.WriteString("NULL") + case string: + _, _ = w.w.WriteString(v) + case []byte: + _, _ = w.w.WriteString(string(v)) + case int: + _, _ = w.w.WriteString(strconv.Itoa(v)) + case int8: + _, _ = w.w.WriteString(strconv.FormatInt(int64(v), 10)) + case int16: + _, _ = w.w.WriteString(strconv.FormatInt(int64(v), 10)) + case int32: + _, _ = w.w.WriteString(strconv.FormatInt(int64(v), 10)) + case int64: + _, _ = w.w.WriteString(strconv.FormatInt(v, 10)) + case uint: + _, _ = w.w.WriteString(strconv.FormatUint(uint64(v), 10)) + case uint8: + _, _ = w.w.WriteString(strconv.FormatUint(uint64(v), 10)) + case uint16: + _, _ = w.w.WriteString(strconv.FormatUint(uint64(v), 10)) + case uint32: + _, _ = w.w.WriteString(strconv.FormatUint(uint64(v), 10)) + case uint64: + _, _ = w.w.WriteString(strconv.FormatUint(v, 10)) + case sql.NamedArg: + return fmt.Errorf("named args are not supported by psql dialect") + case bob.Expression: + return w.writeExpression(v) + default: + _, _ = w.w.WriteString(fmt.Sprint(v)) + } + + return nil +} + +func (w *queryWriter) writeSliceAny(values []any, sep string) error { + for i, value := range values { + if i > 0 { + _, _ = w.w.WriteString(sep) + } + if err := w.writeAny(value); err != nil { + return err + } + } + return nil +} + +func (w *queryWriter) writeSliceExpr(values []bob.Expression, sep string) error { + for i, value := range values { + if i > 0 { + _, _ = w.w.WriteString(sep) + } + if err := w.writeExpression(value); err != nil { + return err + } + } + return nil +} + +func (w *queryWriter) writeOrderExprs(values []bob.Expression) error { + for i, value := range values { + if i > 0 { + _, _ = w.w.WriteString(", ") + } + + switch order := value.(type) { + case clause.OrderDef: + if err := w.writeAny(order.Expression); err != nil { + return err + } + if order.Collation != "" { + _, _ = w.w.WriteString(" COLLATE ") + Dialect.WriteQuoted(w.w, order.Collation) + } + if order.Direction != "" { + _, _ = w.w.WriteString(" ") + _, _ = w.w.WriteString(order.Direction) + } + if order.Nulls != "" { + _, _ = w.w.WriteString(" NULLS ") + _, _ = w.w.WriteString(order.Nulls) + } + default: + if err := w.writeExpression(value); err != nil { + return err + } + } + } + return nil +} diff --git a/dialect/psql/immutable_select.go b/dialect/psql/immutable_select.go deleted file mode 100644 index 4f0b4675..00000000 --- a/dialect/psql/immutable_select.go +++ /dev/null @@ -1,688 +0,0 @@ -package psql - -import ( - "context" - "database/sql" - "fmt" - "io" - "strconv" - "strings" - - "github.com/stephenafamo/bob" - "github.com/stephenafamo/bob/clause" - psqldialect "github.com/stephenafamo/bob/dialect/psql/dialect" - "github.com/stephenafamo/bob/mods" - "github.com/stephenafamo/scan" -) - -type derivedSelectQuery struct { - state immutableSelectState - load bob.Load - hooks bob.EmbeddedHook -} - -type immutableSelectState struct { - DefaultSelectColumns []any - With clause.With - SelectColumns []any - PreloadColumns []any - Distinct psqldialect.Distinct - TableRef clause.TableRef - Where clause.Where - GroupBy clause.GroupBy - Having clause.Having - Windows clause.Windows - Combines clause.Combines - OrderBy clause.OrderBy - Limit clause.Limit - Offset clause.Offset - Fetch clause.Fetch - Locks clause.Locks - - CombinedOrder clause.OrderBy - CombinedLimit clause.Limit - CombinedFetch clause.Fetch - CombinedOffset clause.Offset -} - -func asImmutable(q bob.BaseQuery[*psqldialect.SelectQuery]) derivedSelectQuery { - return derivedSelectQuery{ - state: immutableStateFromMutable(q.Expression), - load: q.Expression.Load, - hooks: q.Expression.EmbeddedHook, - } -} - -func (q derivedSelectQuery) Type() bob.QueryType { - return bob.QueryTypeSelect -} - -func (q derivedSelectQuery) With(queryMods ...bob.Mod[*psqldialect.SelectQuery]) derivedSelectQuery { - next, ok := q.state.withMods(queryMods...) - if ok { - q.state = next - return q - } - - base := q.mutableBase() - mutable := base.Expression - for _, mod := range queryMods { - mod.Apply(mutable) - } - - return asImmutable(base) -} - -func (q derivedSelectQuery) AsCount() derivedSelectQuery { - next := q.state - next.SelectColumns = []any{"count(1)"} - next.DefaultSelectColumns = nil - next.PreloadColumns = nil - next.OrderBy.Expressions = nil - next.GroupBy.Groups = nil - next.GroupBy.With = "" - next.GroupBy.Distinct = false - next.Offset.Count = nil - next.Limit.Count = 1 - - q.state = next - return q -} - -func (q derivedSelectQuery) Build(ctx context.Context) (string, []any, error) { - return q.BuildN(ctx, 1) -} - -func (q derivedSelectQuery) BuildN(ctx context.Context, start int) (string, []any, error) { - var sb strings.Builder - args, err := q.WriteQuery(ctx, &sb, start) - if err != nil { - return "", nil, err - } - - return sb.String(), args, nil -} - -func (q derivedSelectQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - writer := immutableSelectWriter{ - ctx: ctx, - w: w, - start: start, - } - - if err := writer.writeQuery(q.state); err != nil { - return nil, err - } - - return writer.args, nil -} - -func (q derivedSelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { - w.WriteString("(") - args, err := q.WriteQuery(ctx, w, start) - if err != nil { - return nil, err - } - w.WriteString(")") - return args, nil -} - -func (q derivedSelectQuery) Exec(ctx context.Context, exec bob.Executor) (sql.Result, error) { - return bob.Exec(ctx, exec, q) -} - -func (q derivedSelectQuery) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { - return q.hooks.RunHooks(ctx, exec) -} - -func (q derivedSelectQuery) GetLoaders() []bob.Loader { - return q.load.GetLoaders() -} - -func (q derivedSelectQuery) GetMapperMods() []scan.MapperMod { - return q.load.GetMapperMods() -} - -func (q derivedSelectQuery) mutableBase() bob.BaseQuery[*psqldialect.SelectQuery] { - mutable := &psqldialect.SelectQuery{ - With: q.state.With, - SelectList: clause.SelectList{Columns: q.state.selectColumns(), PreloadColumns: q.state.PreloadColumns}, - Distinct: q.state.Distinct, - TableRef: q.state.TableRef, - Where: q.state.Where, - GroupBy: q.state.GroupBy, - Having: q.state.Having, - Windows: q.state.Windows, - Combines: q.state.Combines, - OrderBy: q.state.OrderBy, - Limit: q.state.Limit, - Offset: q.state.Offset, - Fetch: q.state.Fetch, - Locks: q.state.Locks, - - Load: q.load, - EmbeddedHook: q.hooks, - - CombinedOrder: q.state.CombinedOrder, - CombinedLimit: q.state.CombinedLimit, - CombinedFetch: q.state.CombinedFetch, - CombinedOffset: q.state.CombinedOffset, - } - - return bob.BaseQuery[*psqldialect.SelectQuery]{ - Expression: mutable, - Dialect: psqldialect.Dialect, - QueryType: bob.QueryTypeSelect, - } -} - -func immutableStateFromMutable(q *psqldialect.SelectQuery) immutableSelectState { - return immutableSelectState{ - DefaultSelectColumns: nil, - With: clause.With{ - Recursive: q.With.Recursive, - CTEs: append([]bob.Expression(nil), q.With.CTEs...), - }, - SelectColumns: append([]any(nil), q.SelectList.Columns...), - PreloadColumns: append([]any(nil), q.SelectList.PreloadColumns...), - Distinct: psqldialect.Distinct{On: cloneAnySlice(q.Distinct.On)}, - TableRef: cloneTableRef(q.TableRef), - Where: clause.Where{Conditions: append([]any(nil), q.Where.Conditions...)}, - GroupBy: clause.GroupBy{ - Groups: append([]any(nil), q.GroupBy.Groups...), - Distinct: q.GroupBy.Distinct, - With: q.GroupBy.With, - }, - Having: clause.Having{ - Conditions: append([]any(nil), q.Having.Conditions...), - }, - Windows: clause.Windows{ - Windows: append([]bob.Expression(nil), q.Windows.Windows...), - }, - Combines: clause.Combines{ - Queries: append([]clause.Combine(nil), q.Combines.Queries...), - }, - OrderBy: clause.OrderBy{ - Expressions: append([]bob.Expression(nil), q.OrderBy.Expressions...), - }, - Limit: clause.Limit{Count: q.Limit.Count}, - Offset: clause.Offset{ - Count: q.Offset.Count, - }, - Fetch: clause.Fetch{ - Count: q.Fetch.Count, - WithTies: q.Fetch.WithTies, - }, - Locks: clause.Locks{ - Locks: append([]bob.Expression(nil), q.Locks.Locks...), - }, - CombinedOrder: clause.OrderBy{ - Expressions: append([]bob.Expression(nil), q.CombinedOrder.Expressions...), - }, - CombinedLimit: clause.Limit{Count: q.CombinedLimit.Count}, - CombinedFetch: clause.Fetch{ - Count: q.CombinedFetch.Count, - WithTies: q.CombinedFetch.WithTies, - }, - CombinedOffset: clause.Offset{Count: q.CombinedOffset.Count}, - } -} - -func (s immutableSelectState) selectColumns() []any { - if len(s.SelectColumns) > 0 { - return s.SelectColumns - } - return s.DefaultSelectColumns -} - -func cloneAnySlice(values []any) []any { - if values == nil { - return nil - } - return append(make([]any, 0, len(values)), values...) -} - -func (s immutableSelectState) toMutable() psqldialect.SelectQuery { - return psqldialect.SelectQuery{ - With: s.With, - SelectList: clause.SelectList{Columns: s.selectColumns(), PreloadColumns: s.PreloadColumns}, - Distinct: s.Distinct, - TableRef: s.TableRef, - Where: s.Where, - GroupBy: s.GroupBy, - Having: s.Having, - Windows: s.Windows, - Combines: s.Combines, - OrderBy: s.OrderBy, - Limit: s.Limit, - Offset: s.Offset, - Fetch: s.Fetch, - Locks: s.Locks, - CombinedOrder: s.CombinedOrder, - CombinedLimit: s.CombinedLimit, - CombinedFetch: s.CombinedFetch, - CombinedOffset: s.CombinedOffset, - } -} - -func (s immutableSelectState) withMods(queryMods ...bob.Mod[*psqldialect.SelectQuery]) (immutableSelectState, bool) { - next := s - var cloneWith, cloneSelect, cloneWhere, cloneGroup, cloneHaving, cloneOrder, cloneWindows, cloneLocks, cloneJoins, cloneCombines, cloneCombinedOrder, clonePreload bool - - for _, mod := range queryMods { - switch m := mod.(type) { - case mods.Recursive[*psqldialect.SelectQuery]: - next.With.Recursive = bool(m) - case psqldialect.CTEChain[*psqldialect.SelectQuery]: - if !cloneWith { - next.With.CTEs = append([]bob.Expression(nil), s.With.CTEs...) - cloneWith = true - } - next.With.CTEs = append(next.With.CTEs, m()) - case psqldialect.DistinctMod: - next.Distinct.On = cloneAnySlice(m.On) - case mods.Select[*psqldialect.SelectQuery]: - if !cloneSelect { - next.SelectColumns = append([]any(nil), s.SelectColumns...) - cloneSelect = true - } - next.SelectColumns = append(next.SelectColumns, []any(m)...) - case mods.Preload[*psqldialect.SelectQuery]: - if !clonePreload { - next.PreloadColumns = append([]any(nil), s.PreloadColumns...) - clonePreload = true - } - next.PreloadColumns = append(next.PreloadColumns, []any(m)...) - case mods.Where[*psqldialect.SelectQuery]: - if !cloneWhere { - next.Where.Conditions = append([]any(nil), s.Where.Conditions...) - cloneWhere = true - } - next.Where.Conditions = append(next.Where.Conditions, m.E) - case mods.GroupBy[*psqldialect.SelectQuery]: - if !cloneGroup { - next.GroupBy.Groups = append([]any(nil), s.GroupBy.Groups...) - cloneGroup = true - } - next.GroupBy.Groups = append(next.GroupBy.Groups, m.E) - case mods.GroupByDistinct[*psqldialect.SelectQuery]: - next.GroupBy.Distinct = bool(m) - case mods.GroupWith[*psqldialect.SelectQuery]: - next.GroupBy.With = string(m) - case mods.Having[*psqldialect.SelectQuery]: - if !cloneHaving { - next.Having.Conditions = append([]any(nil), s.Having.Conditions...) - cloneHaving = true - } - next.Having.Conditions = append(next.Having.Conditions, []any(m)...) - case mods.Limit[*psqldialect.SelectQuery]: - next.Limit.Count = m.Count - case mods.Offset[*psqldialect.SelectQuery]: - next.Offset.Count = m.Count - case mods.Fetch[*psqldialect.SelectQuery]: - next.Fetch = clause.Fetch(m) - case psqldialect.OrderBy[*psqldialect.SelectQuery]: - if !cloneOrder { - next.OrderBy.Expressions = append([]bob.Expression(nil), s.OrderBy.Expressions...) - cloneOrder = true - } - next.OrderBy.Expressions = append(next.OrderBy.Expressions, m()) - case mods.Join[*psqldialect.SelectQuery]: - if !cloneJoins { - next.TableRef.Joins = append([]clause.Join(nil), s.TableRef.Joins...) - cloneJoins = true - } - next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) - case psqldialect.CrossJoinChain[*psqldialect.SelectQuery]: - if !cloneJoins { - next.TableRef.Joins = append([]clause.Join(nil), s.TableRef.Joins...) - cloneJoins = true - } - next.TableRef.Joins = append(next.TableRef.Joins, m()) - case mods.NamedWindow[*psqldialect.SelectQuery]: - if !cloneWindows { - next.Windows.Windows = append([]bob.Expression(nil), s.Windows.Windows...) - cloneWindows = true - } - next.Windows.Windows = append(next.Windows.Windows, clause.NamedWindow(m)) - case psqldialect.LockChain[*psqldialect.SelectQuery]: - if !cloneLocks { - next.Locks.Locks = append([]bob.Expression(nil), s.Locks.Locks...) - cloneLocks = true - } - next.Locks.Locks = append(next.Locks.Locks, m()) - case mods.Combine[*psqldialect.SelectQuery]: - if !cloneCombines { - next.Combines.Queries = append([]clause.Combine(nil), s.Combines.Queries...) - cloneCombines = true - } - next.Combines.Queries = append(next.Combines.Queries, clause.Combine(m)) - case psqldialect.OrderCombined: - if !cloneCombinedOrder { - next.CombinedOrder.Expressions = append([]bob.Expression(nil), s.CombinedOrder.Expressions...) - cloneCombinedOrder = true - } - next.CombinedOrder.Expressions = append(next.CombinedOrder.Expressions, m()) - case psqldialect.LimitCombined: - next.CombinedLimit.Count = m.Count - case psqldialect.OffsetCombined: - next.CombinedOffset.Count = m.Count - case psqldialect.FetchCombined: - next.CombinedFetch.Count = m.Count - next.CombinedFetch.WithTies = m.WithTies - case psqldialect.FromChain[*psqldialect.SelectQuery]: - next.TableRef = cloneTableRef(m()) - default: - return next, false - } - } - - return next, true -} - -func cloneTableRef(from clause.TableRef) clause.TableRef { - from.Columns = append([]string(nil), from.Columns...) - from.Partitions = append([]string(nil), from.Partitions...) - from.IndexHints = append([]clause.IndexHint(nil), from.IndexHints...) - from.Joins = append([]clause.Join(nil), from.Joins...) - for i := range from.Joins { - from.Joins[i].On = append([]bob.Expression(nil), from.Joins[i].On...) - from.Joins[i].Using = append([]string(nil), from.Joins[i].Using...) - from.Joins[i].To = cloneTableRef(from.Joins[i].To) - } - return from -} - -type immutableSelectWriter struct { - ctx context.Context - w io.StringWriter - args []any - start int -} - -func (w *immutableSelectWriter) writeQuery(q immutableSelectState) error { - if len(q.With.CTEs) > 0 { - if _, err := q.With.WriteSQL(w.ctx, w.w, psqldialect.Dialect, w.argPos()); err != nil { - return err - } - w.w.WriteString("\n") - } - - needsParens := len(q.Combines.Queries) > 0 && - (len(q.OrderBy.Expressions) > 0 || - q.Limit.Count != nil || - q.Offset.Count != nil || - q.Fetch.Count != nil || - len(q.Locks.Locks) > 0) - - if needsParens { - w.w.WriteString("(") - } - - w.w.WriteString("SELECT ") - - if q.Distinct.On != nil { - w.w.WriteString("DISTINCT") - if len(q.Distinct.On) > 0 { - w.w.WriteString(" ON (") - if err := w.writeSliceAny(q.Distinct.On, ", "); err != nil { - return err - } - w.w.WriteString(")") - } - w.w.WriteString(" ") - } - - w.w.WriteString("\n") - selectColumns := q.selectColumns() - if len(selectColumns) == 0 && len(q.PreloadColumns) == 0 { - w.w.WriteString("*") - } else { - allCols := append([]any(nil), selectColumns...) - allCols = append(allCols, q.PreloadColumns...) - if err := w.writeSliceAny(allCols, ", "); err != nil { - return err - } - } - - if q.TableRef.Expression != nil { - w.w.WriteString("\nFROM ") - args, err := q.TableRef.WriteSQL(w.ctx, w.w, psqldialect.Dialect, w.argPos()) - if err != nil { - return err - } - w.args = append(w.args, args...) - } - - if len(q.Where.Conditions) > 0 { - w.w.WriteString("\nWHERE ") - if err := w.writeSliceAny(q.Where.Conditions, " AND "); err != nil { - return err - } - } - - if len(q.GroupBy.Groups) > 0 { - w.w.WriteString("\nGROUP BY ") - if q.GroupBy.Distinct { - w.w.WriteString("DISTINCT ") - } - if err := w.writeSliceAny(q.GroupBy.Groups, ", "); err != nil { - return err - } - if q.GroupBy.With != "" { - w.w.WriteString(" WITH ") - w.w.WriteString(q.GroupBy.With) - } - } - - if len(q.Having.Conditions) > 0 { - w.w.WriteString("\nHAVING ") - if err := w.writeSliceAny(q.Having.Conditions, " AND "); err != nil { - return err - } - } - - if len(q.Windows.Windows) > 0 { - w.w.WriteString("\nWINDOW ") - if err := w.writeSliceExpr(q.Windows.Windows, ", "); err != nil { - return err - } - } - - if len(q.OrderBy.Expressions) > 0 { - w.w.WriteString("\nORDER BY ") - if err := w.writeOrderExprs(q.OrderBy.Expressions); err != nil { - return err - } - } - - if q.Limit.Count != nil { - w.w.WriteString("\nLIMIT ") - if err := w.writeAny(q.Limit.Count); err != nil { - return err - } - } - - if q.Offset.Count != nil { - w.w.WriteString("\nOFFSET ") - if err := w.writeAny(q.Offset.Count); err != nil { - return err - } - } - - if q.Fetch.Count != nil { - w.w.WriteString("\nFETCH NEXT ") - if err := w.writeAny(q.Fetch.Count); err != nil { - return err - } - if q.Fetch.WithTies { - w.w.WriteString(" ROWS WITH TIES") - } else { - w.w.WriteString(" ROWS ONLY") - } - } - - for _, lock := range q.Locks.Locks { - w.w.WriteString("\n") - if err := w.writeAny(lock); err != nil { - return err - } - } - - if needsParens { - w.w.WriteString(")") - } - - for _, combine := range q.Combines.Queries { - w.w.WriteString("\n") - args, err := combine.WriteSQL(w.ctx, w.w, psqldialect.Dialect, w.argPos()) - if err != nil { - return err - } - w.args = append(w.args, args...) - } - - if len(q.CombinedOrder.Expressions) > 0 { - w.w.WriteString("\nORDER BY ") - if err := w.writeOrderExprs(q.CombinedOrder.Expressions); err != nil { - return err - } - } - - if q.CombinedLimit.Count != nil { - w.w.WriteString("\nLIMIT ") - if err := w.writeAny(q.CombinedLimit.Count); err != nil { - return err - } - } - - if q.CombinedOffset.Count != nil { - w.w.WriteString("\nOFFSET ") - if err := w.writeAny(q.CombinedOffset.Count); err != nil { - return err - } - } - - if q.CombinedFetch.Count != nil { - w.w.WriteString("\nFETCH NEXT ") - if err := w.writeAny(q.CombinedFetch.Count); err != nil { - return err - } - if q.CombinedFetch.WithTies { - w.w.WriteString(" ROWS WITH TIES") - } else { - w.w.WriteString(" ROWS ONLY") - } - } - - w.w.WriteString("\n") - return nil -} - -func (w *immutableSelectWriter) argPos() int { - return w.start + len(w.args) -} - -func (w *immutableSelectWriter) writeSliceAny(values []any, sep string) error { - for i, value := range values { - if i > 0 { - w.w.WriteString(sep) - } - if err := w.writeAny(value); err != nil { - return err - } - } - return nil -} - -func (w *immutableSelectWriter) writeSliceExpr(values []bob.Expression, sep string) error { - for i, value := range values { - if i > 0 { - w.w.WriteString(sep) - } - if err := w.writeExpression(value); err != nil { - return err - } - } - return nil -} - -func (w *immutableSelectWriter) writeOrderExprs(values []bob.Expression) error { - for i, value := range values { - if i > 0 { - w.w.WriteString(", ") - } - - switch order := value.(type) { - case clause.OrderDef: - if err := w.writeAny(order.Expression); err != nil { - return err - } - if order.Collation != "" { - w.w.WriteString(" COLLATE ") - psqldialect.Dialect.WriteQuoted(w.w, order.Collation) - } - if order.Direction != "" { - w.w.WriteString(" ") - w.w.WriteString(order.Direction) - } - if order.Nulls != "" { - w.w.WriteString(" NULLS ") - w.w.WriteString(order.Nulls) - } - default: - if err := w.writeExpression(value); err != nil { - return err - } - } - } - return nil -} - -func (w *immutableSelectWriter) writeExpression(value bob.Expression) error { - args, err := value.WriteSQL(w.ctx, w.w, psqldialect.Dialect, w.argPos()) - if err != nil { - return err - } - w.args = append(w.args, args...) - return nil -} - -func (w *immutableSelectWriter) writeAny(value any) error { - switch v := value.(type) { - case nil: - w.w.WriteString("NULL") - case string: - w.w.WriteString(v) - case []byte: - w.w.WriteString(string(v)) - case int: - w.w.WriteString(strconv.Itoa(v)) - case int8: - w.w.WriteString(strconv.FormatInt(int64(v), 10)) - case int16: - w.w.WriteString(strconv.FormatInt(int64(v), 10)) - case int32: - w.w.WriteString(strconv.FormatInt(int64(v), 10)) - case int64: - w.w.WriteString(strconv.FormatInt(v, 10)) - case uint: - w.w.WriteString(strconv.FormatUint(uint64(v), 10)) - case uint8: - w.w.WriteString(strconv.FormatUint(uint64(v), 10)) - case uint16: - w.w.WriteString(strconv.FormatUint(uint64(v), 10)) - case uint32: - w.w.WriteString(strconv.FormatUint(uint64(v), 10)) - case uint64: - w.w.WriteString(strconv.FormatUint(v, 10)) - case sql.NamedArg: - return fmt.Errorf("named args are not supported by psql dialect") - case bob.Expression: - return w.writeExpression(v) - default: - w.w.WriteString(fmt.Sprint(v)) - } - - return nil -} diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index c379fae1..6bcd815c 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -401,7 +401,7 @@ func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := q.Query.derivedSelectQuery.AsCount().Build(ctx); err != nil { + if _, _, err := q.Query.AsCount().Build(ctx); err != nil { b.Fatal(err) } @@ -426,7 +426,7 @@ func BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := q.Query.derivedSelectQuery.AsCount().Build(ctx); err != nil { + if _, _, err := q.Query.AsCount().Build(ctx); err != nil { b.Fatal(err) } diff --git a/dialect/psql/immutable_write.go b/dialect/psql/immutable_write.go deleted file mode 100644 index 17237263..00000000 --- a/dialect/psql/immutable_write.go +++ /dev/null @@ -1,661 +0,0 @@ -package psql - -import ( - "context" - "database/sql" - "io" - "strings" - - "github.com/stephenafamo/bob" - "github.com/stephenafamo/bob/clause" - psqldialect "github.com/stephenafamo/bob/dialect/psql/dialect" - "github.com/stephenafamo/bob/mods" - "github.com/stephenafamo/scan" -) - -type derivedUpdateQuery struct { - state immutableUpdateState - load bob.Load - hooks bob.EmbeddedHook -} - -type immutableUpdateState struct { - With clause.With - Only bool - Table clause.TableRef - From clause.TableRef - Set clause.Set - Where clause.Where - Returning clause.Returning -} - -func asImmutableUpdate(q bob.BaseQuery[*psqldialect.UpdateQuery]) derivedUpdateQuery { - return derivedUpdateQuery{ - state: immutableUpdateState{ - With: clause.With{ - Recursive: q.Expression.With.Recursive, - CTEs: append([]bob.Expression(nil), q.Expression.With.CTEs...), - }, - Only: q.Expression.Only, - Table: cloneTableRef(q.Expression.Table), - From: cloneTableRef(q.Expression.TableRef), - Set: clause.Set{ - Set: append([]any(nil), q.Expression.Set.Set...), - }, - Where: clause.Where{ - Conditions: append([]any(nil), q.Expression.Where.Conditions...), - }, - Returning: clause.Returning{ - Expressions: append([]any(nil), q.Expression.Returning.Expressions...), - }, - }, - load: q.Expression.Load, - hooks: q.Expression.EmbeddedHook, - } -} - -func (q derivedUpdateQuery) Type() bob.QueryType { return bob.QueryTypeUpdate } - -func (q derivedUpdateQuery) With(queryMods ...bob.Mod[*psqldialect.UpdateQuery]) derivedUpdateQuery { - next, ok := q.state.withMods(queryMods...) - if ok { - q.state = next - return q - } - - base := q.mutableBase() - mutable := base.Expression - for _, mod := range queryMods { - mod.Apply(mutable) - } - - return asImmutableUpdate(base) -} - -func (q derivedUpdateQuery) Exec(ctx context.Context, exec bob.Executor) (sql.Result, error) { - return bob.Exec(ctx, exec, q) -} - -func (q derivedUpdateQuery) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { - return q.hooks.RunHooks(ctx, exec) -} - -func (q derivedUpdateQuery) GetLoaders() []bob.Loader { - return q.load.GetLoaders() -} - -func (q derivedUpdateQuery) GetMapperMods() []scan.MapperMod { - return q.load.GetMapperMods() -} - -func (q derivedUpdateQuery) Build(ctx context.Context) (string, []any, error) { - return q.BuildN(ctx, 1) -} - -func (q derivedUpdateQuery) BuildN(ctx context.Context, start int) (string, []any, error) { - var sb strings.Builder - args, err := q.WriteQuery(ctx, &sb, start) - if err != nil { - return "", nil, err - } - return sb.String(), args, nil -} - -func (q derivedUpdateQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - var args []any - - if len(q.state.With.CTEs) > 0 { - withArgs, err := q.state.With.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, withArgs...) - w.WriteString("\n") - } - - w.WriteString("UPDATE ") - if q.state.Only { - w.WriteString("ONLY ") - } - - tableArgs, err := q.state.Table.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, tableArgs...) - - w.WriteString(" SET\n") - setArgs, err := q.state.Set.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, setArgs...) - - if q.state.From.Expression != nil { - w.WriteString("\nFROM ") - fromArgs, err := q.state.From.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, fromArgs...) - } - - if len(q.state.Where.Conditions) > 0 { - w.WriteString("\n") - whereArgs, err := q.state.Where.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, whereArgs...) - } - - if len(q.state.Returning.Expressions) > 0 { - w.WriteString("\n") - retArgs, err := q.state.Returning.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, retArgs...) - } - - return args, nil -} - -func (q derivedUpdateQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { - return q.WriteQuery(ctx, w, start) -} - -func (q derivedUpdateQuery) mutableBase() bob.BaseQuery[*psqldialect.UpdateQuery] { - mutable := &psqldialect.UpdateQuery{ - With: q.state.With, - Only: q.state.Only, - Table: q.state.Table, - Set: q.state.Set, - TableRef: q.state.From, - Where: q.state.Where, - Returning: q.state.Returning, - Load: q.load, - EmbeddedHook: q.hooks, - } - - return bob.BaseQuery[*psqldialect.UpdateQuery]{ - Expression: mutable, - Dialect: psqldialect.Dialect, - QueryType: bob.QueryTypeUpdate, - } -} - -func (s immutableUpdateState) withMods(queryMods ...bob.Mod[*psqldialect.UpdateQuery]) (immutableUpdateState, bool) { - next := s - var cloneWith, cloneSet, cloneWhere, cloneReturning, cloneJoins bool - - for _, mod := range queryMods { - switch m := mod.(type) { - case mods.Recursive[*psqldialect.UpdateQuery]: - next.With.Recursive = bool(m) - case psqldialect.CTEChain[*psqldialect.UpdateQuery]: - if !cloneWith { - next.With.CTEs = append([]bob.Expression(nil), s.With.CTEs...) - cloneWith = true - } - next.With.CTEs = append(next.With.CTEs, m()) - case psqldialect.UpdateOnly: - next.Only = bool(m) - case psqldialect.UpdateTable: - next.Table = cloneTableRef(clause.TableRef(m)) - case psqldialect.UpdateSet: - if !cloneSet { - next.Set.Set = append([]any(nil), s.Set.Set...) - cloneSet = true - } - next.Set.Set = append(next.Set.Set, []any(m)...) - case mods.Where[*psqldialect.UpdateQuery]: - if !cloneWhere { - next.Where.Conditions = append([]any(nil), s.Where.Conditions...) - cloneWhere = true - } - next.Where.Conditions = append(next.Where.Conditions, m.E) - case mods.Returning[*psqldialect.UpdateQuery]: - if !cloneReturning { - next.Returning.Expressions = append([]any(nil), s.Returning.Expressions...) - cloneReturning = true - } - next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) - case psqldialect.FromChain[*psqldialect.UpdateQuery]: - next.From = cloneTableRef(m()) - case mods.Join[*psqldialect.UpdateQuery]: - if !cloneJoins { - next.From.Joins = append([]clause.Join(nil), s.From.Joins...) - cloneJoins = true - } - next.From.Joins = append(next.From.Joins, clause.Join(m)) - case psqldialect.CrossJoinChain[*psqldialect.UpdateQuery]: - if !cloneJoins { - next.From.Joins = append([]clause.Join(nil), s.From.Joins...) - cloneJoins = true - } - next.From.Joins = append(next.From.Joins, m()) - default: - return next, false - } - } - - return next, true -} - -type derivedDeleteQuery struct { - state immutableDeleteState - load bob.Load - hooks bob.EmbeddedHook -} - -type immutableDeleteState struct { - With clause.With - Only bool - Table clause.TableRef - Using clause.TableRef - Where clause.Where - Returning clause.Returning -} - -func asImmutableDelete(q bob.BaseQuery[*psqldialect.DeleteQuery]) derivedDeleteQuery { - return derivedDeleteQuery{ - state: immutableDeleteState{ - With: clause.With{ - Recursive: q.Expression.With.Recursive, - CTEs: append([]bob.Expression(nil), q.Expression.With.CTEs...), - }, - Only: q.Expression.Only, - Table: cloneTableRef(q.Expression.Table), - Using: cloneTableRef(q.Expression.TableRef), - Where: clause.Where{Conditions: append([]any(nil), q.Expression.Where.Conditions...)}, - Returning: clause.Returning{ - Expressions: append([]any(nil), q.Expression.Returning.Expressions...), - }, - }, - load: q.Expression.Load, - hooks: q.Expression.EmbeddedHook, - } -} - -func (q derivedDeleteQuery) Type() bob.QueryType { return bob.QueryTypeDelete } - -func (q derivedDeleteQuery) With(queryMods ...bob.Mod[*psqldialect.DeleteQuery]) derivedDeleteQuery { - next, ok := q.state.withMods(queryMods...) - if ok { - q.state = next - return q - } - - base := q.mutableBase() - mutable := base.Expression - for _, mod := range queryMods { - mod.Apply(mutable) - } - - return asImmutableDelete(base) -} - -func (q derivedDeleteQuery) Exec(ctx context.Context, exec bob.Executor) (sql.Result, error) { - return bob.Exec(ctx, exec, q) -} - -func (q derivedDeleteQuery) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { - return q.hooks.RunHooks(ctx, exec) -} - -func (q derivedDeleteQuery) GetLoaders() []bob.Loader { return q.load.GetLoaders() } - -func (q derivedDeleteQuery) GetMapperMods() []scan.MapperMod { return q.load.GetMapperMods() } - -func (q derivedDeleteQuery) Build(ctx context.Context) (string, []any, error) { - return q.BuildN(ctx, 1) -} - -func (q derivedDeleteQuery) BuildN(ctx context.Context, start int) (string, []any, error) { - var sb strings.Builder - args, err := q.WriteQuery(ctx, &sb, start) - if err != nil { - return "", nil, err - } - return sb.String(), args, nil -} - -func (q derivedDeleteQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - var args []any - if len(q.state.With.CTEs) > 0 { - withArgs, err := q.state.With.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, withArgs...) - w.WriteString("\n") - } - - w.WriteString("DELETE FROM ") - if q.state.Only { - w.WriteString("ONLY ") - } - - tableArgs, err := q.state.Table.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, tableArgs...) - - if q.state.Using.Expression != nil { - w.WriteString("\nUSING ") - usingArgs, err := q.state.Using.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, usingArgs...) - } - - if len(q.state.Where.Conditions) > 0 { - w.WriteString("\n") - whereArgs, err := q.state.Where.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, whereArgs...) - } - - if len(q.state.Returning.Expressions) > 0 { - w.WriteString("\n") - retArgs, err := q.state.Returning.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, retArgs...) - } - - return args, nil -} - -func (q derivedDeleteQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { - return q.WriteQuery(ctx, w, start) -} - -func (q derivedDeleteQuery) mutableBase() bob.BaseQuery[*psqldialect.DeleteQuery] { - mutable := &psqldialect.DeleteQuery{ - With: q.state.With, - Only: q.state.Only, - Table: q.state.Table, - TableRef: q.state.Using, - Where: q.state.Where, - Returning: q.state.Returning, - Load: q.load, - EmbeddedHook: q.hooks, - } - - return bob.BaseQuery[*psqldialect.DeleteQuery]{ - Expression: mutable, - Dialect: psqldialect.Dialect, - QueryType: bob.QueryTypeDelete, - } -} - -func (s immutableDeleteState) withMods(queryMods ...bob.Mod[*psqldialect.DeleteQuery]) (immutableDeleteState, bool) { - next := s - var cloneWith, cloneWhere, cloneReturning, cloneJoins bool - - for _, mod := range queryMods { - switch m := mod.(type) { - case mods.Recursive[*psqldialect.DeleteQuery]: - next.With.Recursive = bool(m) - case psqldialect.CTEChain[*psqldialect.DeleteQuery]: - if !cloneWith { - next.With.CTEs = append([]bob.Expression(nil), s.With.CTEs...) - cloneWith = true - } - next.With.CTEs = append(next.With.CTEs, m()) - case psqldialect.DeleteOnly: - next.Only = bool(m) - case psqldialect.DeleteTable: - next.Table = cloneTableRef(clause.TableRef(m)) - case mods.Where[*psqldialect.DeleteQuery]: - if !cloneWhere { - next.Where.Conditions = append([]any(nil), s.Where.Conditions...) - cloneWhere = true - } - next.Where.Conditions = append(next.Where.Conditions, m.E) - case mods.Returning[*psqldialect.DeleteQuery]: - if !cloneReturning { - next.Returning.Expressions = append([]any(nil), s.Returning.Expressions...) - cloneReturning = true - } - next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) - case psqldialect.FromChain[*psqldialect.DeleteQuery]: - next.Using = cloneTableRef(m()) - case mods.Join[*psqldialect.DeleteQuery]: - if !cloneJoins { - next.Using.Joins = append([]clause.Join(nil), s.Using.Joins...) - cloneJoins = true - } - next.Using.Joins = append(next.Using.Joins, clause.Join(m)) - case psqldialect.CrossJoinChain[*psqldialect.DeleteQuery]: - if !cloneJoins { - next.Using.Joins = append([]clause.Join(nil), s.Using.Joins...) - cloneJoins = true - } - next.Using.Joins = append(next.Using.Joins, m()) - default: - return next, false - } - } - - return next, true -} - -type derivedInsertQuery struct { - state immutableInsertState - load bob.Load - hooks bob.EmbeddedHook -} - -type immutableInsertState struct { - With clause.With - Overriding string - Table clause.TableRef - Values clause.Values - Conflict clause.Conflict - Returning clause.Returning -} - -func asImmutableInsert(q bob.BaseQuery[*psqldialect.InsertQuery]) derivedInsertQuery { - return derivedInsertQuery{ - state: immutableInsertState{ - With: clause.With{ - Recursive: q.Expression.With.Recursive, - CTEs: append([]bob.Expression(nil), q.Expression.With.CTEs...), - }, - Overriding: q.Expression.Overriding, - Table: cloneTableRef(q.Expression.TableRef), - Values: clause.Values{ - Query: q.Expression.Values.Query, - Vals: append([]clause.Value(nil), q.Expression.Values.Vals...), - }, - Conflict: clause.Conflict{Expression: q.Expression.Conflict.Expression}, - Returning: clause.Returning{ - Expressions: append([]any(nil), q.Expression.Returning.Expressions...), - }, - }, - load: q.Expression.Load, - hooks: q.Expression.EmbeddedHook, - } -} - -func (q derivedInsertQuery) Type() bob.QueryType { return bob.QueryTypeInsert } - -func (q derivedInsertQuery) With(queryMods ...bob.Mod[*psqldialect.InsertQuery]) derivedInsertQuery { - next, ok := q.state.withMods(queryMods...) - if ok { - q.state = next - return q - } - - base := q.mutableBase() - mutable := base.Expression - for _, mod := range queryMods { - mod.Apply(mutable) - } - - return asImmutableInsert(base) -} - -func (q derivedInsertQuery) Exec(ctx context.Context, exec bob.Executor) (sql.Result, error) { - return bob.Exec(ctx, exec, q) -} - -func (q derivedInsertQuery) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { - return q.hooks.RunHooks(ctx, exec) -} - -func (q derivedInsertQuery) GetLoaders() []bob.Loader { return q.load.GetLoaders() } - -func (q derivedInsertQuery) GetMapperMods() []scan.MapperMod { return q.load.GetMapperMods() } - -func (q derivedInsertQuery) Build(ctx context.Context) (string, []any, error) { - return q.BuildN(ctx, 1) -} - -func (q derivedInsertQuery) BuildN(ctx context.Context, start int) (string, []any, error) { - var sb strings.Builder - args, err := q.WriteQuery(ctx, &sb, start) - if err != nil { - return "", nil, err - } - return sb.String(), args, nil -} - -func (q derivedInsertQuery) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - var args []any - if len(q.state.With.CTEs) > 0 { - withArgs, err := q.state.With.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, withArgs...) - w.WriteString("\n") - } - - w.WriteString("INSERT INTO ") - tableArgs, err := q.state.Table.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, tableArgs...) - - if q.state.Overriding != "" { - w.WriteString("\nOVERRIDING ") - w.WriteString(q.state.Overriding) - w.WriteString(" VALUE") - } - - w.WriteString("\n") - valArgs, err := q.state.Values.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, valArgs...) - - if q.state.Conflict.Expression != nil { - w.WriteString("\n") - conflictArgs, err := q.state.Conflict.Expression.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, conflictArgs...) - } - - if len(q.state.Returning.Expressions) > 0 { - w.WriteString("\n") - retArgs, err := q.state.Returning.WriteSQL(ctx, w, psqldialect.Dialect, start+len(args)) - if err != nil { - return nil, err - } - args = append(args, retArgs...) - } - - w.WriteString("\n") - return args, nil -} - -func (q derivedInsertQuery) WriteSQL(ctx context.Context, w io.StringWriter, _ bob.Dialect, start int) ([]any, error) { - return q.WriteQuery(ctx, w, start) -} - -func (q derivedInsertQuery) mutableBase() bob.BaseQuery[*psqldialect.InsertQuery] { - values := clause.Values{ - Query: q.state.Values.Query, - Vals: append([]clause.Value(nil), q.state.Values.Vals...), - } - - mutable := &psqldialect.InsertQuery{ - With: q.state.With, - Overriding: q.state.Overriding, - TableRef: q.state.Table, - Values: values, - Conflict: q.state.Conflict, - Returning: q.state.Returning, - Load: q.load, - EmbeddedHook: q.hooks, - } - - return bob.BaseQuery[*psqldialect.InsertQuery]{ - Expression: mutable, - Dialect: psqldialect.Dialect, - QueryType: bob.QueryTypeInsert, - } -} - -func (s immutableInsertState) withMods(queryMods ...bob.Mod[*psqldialect.InsertQuery]) (immutableInsertState, bool) { - next := s - var cloneWith, cloneReturning, cloneVals bool - - for _, mod := range queryMods { - switch m := mod.(type) { - case mods.Recursive[*psqldialect.InsertQuery]: - next.With.Recursive = bool(m) - case psqldialect.CTEChain[*psqldialect.InsertQuery]: - if !cloneWith { - next.With.CTEs = append([]bob.Expression(nil), s.With.CTEs...) - cloneWith = true - } - next.With.CTEs = append(next.With.CTEs, m()) - case psqldialect.InsertTable: - next.Table = cloneTableRef(clause.TableRef(m)) - case psqldialect.InsertOverriding: - next.Overriding = string(m) - case psqldialect.InsertQuerySource: - next.Values.Query = m.Query - case mods.Returning[*psqldialect.InsertQuery]: - if !cloneReturning { - next.Returning.Expressions = append([]any(nil), s.Returning.Expressions...) - cloneReturning = true - } - next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) - case mods.Values[*psqldialect.InsertQuery]: - if !cloneVals { - next.Values.Vals = append([]clause.Value(nil), s.Values.Vals...) - cloneVals = true - } - next.Values.Vals = append(next.Values.Vals, clause.Value(m)) - case mods.Rows[*psqldialect.InsertQuery]: - if !cloneVals { - next.Values.Vals = append([]clause.Value(nil), s.Values.Vals...) - cloneVals = true - } - for _, row := range m { - next.Values.Vals = append(next.Values.Vals, clause.Value(row)) - } - case mods.Conflict[*psqldialect.InsertQuery]: - next.Conflict.Expression = m() - default: - return next, false - } - } - - return next, true -} diff --git a/dialect/psql/insert.go b/dialect/psql/insert.go index b5e82651..b1d3c70c 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -6,11 +6,15 @@ import ( ) type InsertQuery struct { - derivedInsertQuery + bob.BaseQuery[*dialect.InsertQuery] } func (q InsertQuery) With(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { - q.derivedInsertQuery = q.derivedInsertQuery.With(queryMods...) + if next, ok := deriveInsert(q.Expression, queryMods...); ok { + q.Expression = next + return q + } + q.BaseQuery = q.BaseQuery.Apply(queryMods...) return q } @@ -19,25 +23,16 @@ func (q InsertQuery) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQue } func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { - state, ok := (immutableInsertState{}).withMods(queryMods...) - if ok { - return InsertQuery{ - derivedInsertQuery: derivedInsertQuery{ - state: state, - }, - } - } - q := &dialect.InsertQuery{} for _, mod := range queryMods { mod.Apply(q) } return InsertQuery{ - derivedInsertQuery: asImmutableInsert(bob.BaseQuery[*dialect.InsertQuery]{ + BaseQuery: bob.BaseQuery[*dialect.InsertQuery]{ Expression: q, Dialect: dialect.Dialect, QueryType: bob.QueryTypeInsert, - }), + }, } } diff --git a/dialect/psql/select.go b/dialect/psql/select.go index 2530a27e..90fdd2f3 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -6,11 +6,15 @@ import ( ) type SelectQuery struct { - derivedSelectQuery + bob.BaseQuery[*dialect.SelectQuery] } func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { - q.derivedSelectQuery = q.derivedSelectQuery.With(queryMods...) + if next, ok := deriveSelect(q.Expression, queryMods...); ok { + q.Expression = next + return q + } + q.BaseQuery = q.BaseQuery.Apply(queryMods...) return q } @@ -18,26 +22,30 @@ func (q SelectQuery) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQue return q.With(queryMods...) } -func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { - state, ok := (immutableSelectState{}).withMods(queryMods...) - if ok { - return SelectQuery{ - derivedSelectQuery: derivedSelectQuery{ - state: state, - }, - } - } +func (q SelectQuery) AsCount() SelectQuery { + next := q.Clone() + next.Expression.SetSelect("count(1)") + next.Expression.SetPreloadSelect() + next.Expression.SetMapperMods() + next.Expression.SetLoaders() + next.Expression.SetLimit(1) + next.Expression.ClearOrderBy() + next.Expression.SetGroups() + next.Expression.SetOffset(0) + return SelectQuery{BaseQuery: next} +} +func Select(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { q := &dialect.SelectQuery{} for _, mod := range queryMods { mod.Apply(q) } return SelectQuery{ - derivedSelectQuery: asImmutable(bob.BaseQuery[*dialect.SelectQuery]{ + BaseQuery: bob.BaseQuery[*dialect.SelectQuery]{ Expression: q, Dialect: dialect.Dialect, QueryType: bob.QueryTypeSelect, - }), + }, } } diff --git a/dialect/psql/update.go b/dialect/psql/update.go index 6109763c..e61818ea 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -6,11 +6,15 @@ import ( ) type UpdateQuery struct { - derivedUpdateQuery + bob.BaseQuery[*dialect.UpdateQuery] } func (q UpdateQuery) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { - q.derivedUpdateQuery = q.derivedUpdateQuery.With(queryMods...) + if next, ok := deriveUpdate(q.Expression, queryMods...); ok { + q.Expression = next + return q + } + q.BaseQuery = q.BaseQuery.Apply(queryMods...) return q } @@ -19,25 +23,16 @@ func (q UpdateQuery) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQue } func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { - state, ok := (immutableUpdateState{}).withMods(queryMods...) - if ok { - return UpdateQuery{ - derivedUpdateQuery: derivedUpdateQuery{ - state: state, - }, - } - } - q := &dialect.UpdateQuery{} for _, mod := range queryMods { mod.Apply(q) } return UpdateQuery{ - derivedUpdateQuery: asImmutableUpdate(bob.BaseQuery[*dialect.UpdateQuery]{ + BaseQuery: bob.BaseQuery[*dialect.UpdateQuery]{ Expression: q, Dialect: dialect.Dialect, QueryType: bob.QueryTypeUpdate, - }), + }, } } diff --git a/dialect/psql/view.go b/dialect/psql/view.go index f1bb3112..65b5395b 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -85,10 +85,16 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * Scanner: v.scanner, Hooks: &v.SelectQueryHooks, } - q.Query.derivedSelectQuery.state.DefaultSelectColumns = []any{v.Columns} - if len(queryMods) == 0 { - return q - } + + q.Query.Expression.AppendContextualModFunc( + func(ctx context.Context, q *dialect.SelectQuery) (context.Context, error) { + if len(q.SelectList.Columns) == 0 { + q.AppendSelect(v.Columns) + } + return ctx, nil + }, + ) + return q.Apply(queryMods...) } @@ -129,12 +135,12 @@ func (q *ViewQuery[T, Ts]) WriteSQL(ctx context.Context, w io.StringWriter, d bo } // Count the number of matching rows -func (v *ViewQuery[T, Tslice]) Count(ctx context.Context, exec bob.Executor) (int64, error) { - ctx, err := v.RunHooks(ctx, exec) +func (q *ViewQuery[T, Tslice]) Count(ctx context.Context, exec bob.Executor) (int64, error) { + ctx, err := q.RunHooks(ctx, exec) if err != nil { return 0, err } - sql, args, err := v.Query.AsCount().Build(ctx) + sql, args, err := q.Query.AsCount().Build(ctx) if err != nil { return 0, err } @@ -142,8 +148,8 @@ func (v *ViewQuery[T, Tslice]) Count(ctx context.Context, exec bob.Executor) (in } // Exists checks if there is any matching row -func (v *ViewQuery[T, Tslice]) Exists(ctx context.Context, exec bob.Executor) (bool, error) { - count, err := v.Count(ctx, exec) +func (q *ViewQuery[T, Tslice]) Exists(ctx context.Context, exec bob.Executor) (bool, error) { + count, err := q.Count(ctx, exec) return count > 0, err } @@ -183,27 +189,3 @@ func (q *ViewQuery[T, Ts]) GetLoaders() []bob.Loader { func (q *ViewQuery[T, Ts]) GetMapperMods() []scan.MapperMod { return q.Query.GetMapperMods() } - -// asCountQuery clones and rewrites an existing query to a count query -func asCountQuery(query bob.BaseQuery[*dialect.SelectQuery]) bob.BaseQuery[*dialect.SelectQuery] { - // clone the original query, so it's not being modified silently - countQuery := query.Clone() - // only select the count - countQuery.Expression.SetSelect("count(1)") - // don't select any preload columns - countQuery.Expression.SetPreloadSelect() - // disable mapper mods - countQuery.Expression.SetMapperMods() - // disable loaders - countQuery.Expression.SetLoaders() - // set the limit to 1 - countQuery.Expression.SetLimit(1) - // remove ordering - countQuery.Expression.ClearOrderBy() - // remove group by - countQuery.Expression.SetGroups() - // remove offset - countQuery.Expression.SetOffset(0) - - return countQuery -} From a64a5009aecc227190b6ce99fcc8f840ba56dbef Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 11:29:09 -0400 Subject: [PATCH 26/33] refactor(psql): simplify table returning defaults --- clause/returning.go | 4 -- dialect/psql/table.go | 113 ++++++++++++++++--------------------- dialect/psql/table_test.go | 64 +++++++++++++++++++++ mods/mods.go | 4 -- 4 files changed, 113 insertions(+), 72 deletions(-) diff --git a/clause/returning.go b/clause/returning.go index ad0a48d3..e8cdb1d5 100644 --- a/clause/returning.go +++ b/clause/returning.go @@ -15,10 +15,6 @@ func (r *Returning) HasReturning() bool { return len(r.Expressions) > 0 } -func (r *Returning) SetReturning(columns ...any) { - r.Expressions = append(r.Expressions[:0], columns...) -} - func (r *Returning) AppendReturning(columns ...any) { r.Expressions = append(r.Expressions, columns...) } diff --git a/dialect/psql/table.go b/dialect/psql/table.go index 97c33dc2..2b3a5c53 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -22,23 +22,23 @@ type ( type ormInsertQuery[T any, Tslice ~[]T] struct { orm.Query[*dialect.InsertQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - hasDefaultReturning bool + defaultReturning bob.Expression } type ormUpdateQuery[T any, Tslice ~[]T] struct { orm.Query[*dialect.UpdateQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - hasDefaultReturning bool + defaultReturning bob.Expression } type ormDeleteQuery[T any, Tslice ~[]T] struct { orm.Query[*dialect.DeleteQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - hasDefaultReturning bool + defaultReturning bob.Expression } func (q ormInsertQuery[T, Tslice]) clone() ormInsertQuery[T, Tslice] { return ormInsertQuery[T, Tslice]{ - Query: q.Query.Clone(), - hasDefaultReturning: q.hasDefaultReturning, + Query: q.Query.Clone(), + defaultReturning: q.defaultReturning, } } @@ -48,7 +48,9 @@ func (q *ormInsertQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.InsertQue } next := q.clone() - applyTableQueryMods(next.Expression, &next.hasDefaultReturning, queryMods...) + applyTableQueryMods(next.Expression, next.defaultReturning, func(query *dialect.InsertQuery) *clause.Returning { + return &query.Returning + }, queryMods...) return &next } @@ -58,8 +60,8 @@ func (q *ormInsertQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.InsertQu func (q ormUpdateQuery[T, Tslice]) clone() ormUpdateQuery[T, Tslice] { return ormUpdateQuery[T, Tslice]{ - Query: q.Query.Clone(), - hasDefaultReturning: q.hasDefaultReturning, + Query: q.Query.Clone(), + defaultReturning: q.defaultReturning, } } @@ -69,7 +71,9 @@ func (q *ormUpdateQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.UpdateQue } next := q.clone() - applyTableQueryMods(next.Expression, &next.hasDefaultReturning, queryMods...) + applyTableQueryMods(next.Expression, next.defaultReturning, func(query *dialect.UpdateQuery) *clause.Returning { + return &query.Returning + }, queryMods...) return &next } @@ -79,8 +83,8 @@ func (q *ormUpdateQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.UpdateQu func (q ormDeleteQuery[T, Tslice]) clone() ormDeleteQuery[T, Tslice] { return ormDeleteQuery[T, Tslice]{ - Query: q.Query.Clone(), - hasDefaultReturning: q.hasDefaultReturning, + Query: q.Query.Clone(), + defaultReturning: q.defaultReturning, } } @@ -90,7 +94,9 @@ func (q *ormDeleteQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.DeleteQue } next := q.clone() - applyTableQueryMods(next.Expression, &next.hasDefaultReturning, queryMods...) + applyTableQueryMods(next.Expression, next.defaultReturning, func(query *dialect.DeleteQuery) *clause.Returning { + return &query.Returning + }, queryMods...) return &next } @@ -98,18 +104,30 @@ func (q *ormDeleteQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.DeleteQu return q.With(queryMods...) } -func applyTableQueryMods[Q interface{ SetReturning(...any) }](query Q, hasDefaultReturning *bool, queryMods ...bob.Mod[Q]) { - for _, mod := range queryMods { - if returning, ok := any(mod).(interface{ ReturningValues() []any }); ok && *hasDefaultReturning { - query.SetReturning(returning.ReturningValues()...) - *hasDefaultReturning = false - continue - } +func applyTableQueryMods[Q interface{ AppendReturning(...any) }](query Q, defaultReturning bob.Expression, getReturning func(Q) *clause.Returning, queryMods ...bob.Mod[Q]) { + if hasExplicitReturning(queryMods...) && hasOnlyDefaultReturning(getReturning(query).Expressions, defaultReturning) { + getReturning(query).Expressions = nil + } + for _, mod := range queryMods { mod.Apply(query) } } +func hasExplicitReturning[Q interface{ AppendReturning(...any) }](queryMods ...bob.Mod[Q]) bool { + for _, mod := range queryMods { + if _, ok := mod.(bobmods.Returning[Q]); ok { + return true + } + } + + return false +} + +func hasOnlyDefaultReturning(expressions []any, defaultReturning bob.Expression) bool { + return len(expressions) == 1 && reflect.DeepEqual(expressions[0], defaultReturning) +} + func NewTable[T any, Tset setter[T], C bob.Expression](schema, tableName string, columns C) *Table[T, []T, Tset, C] { return NewTablex[T, []T, Tset](schema, tableName, columns) } @@ -163,12 +181,12 @@ func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQ q := &ormInsertQuery[T, Tslice]{ Query: orm.Query[*dialect.InsertQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ - BaseQuery: insertTableBaseQuery(t.NameAs(), t.nonGeneratedCols, t.Columns, queryMods), + BaseQuery: insertTableBaseQuery(t.NameAs(), t.nonGeneratedCols, t.Columns), Hooks: &t.InsertQueryHooks, }, Scanner: t.scanner, }, - hasDefaultReturning: !hasInsertReturning(queryMods), + defaultReturning: t.Columns, } return q.Apply(queryMods...) @@ -179,12 +197,12 @@ func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQ q := &ormUpdateQuery[T, Tslice]{ Query: orm.Query[*dialect.UpdateQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ - BaseQuery: updateTableBaseQuery(t.NameAs(), t.Columns, queryMods), + BaseQuery: updateTableBaseQuery(t.NameAs(), t.Columns), Hooks: &t.UpdateQueryHooks, }, Scanner: t.scanner, }, - hasDefaultReturning: !hasUpdateReturning(queryMods), + defaultReturning: t.Columns, } return q.Apply(queryMods...) @@ -195,12 +213,12 @@ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQ q := &ormDeleteQuery[T, Tslice]{ Query: orm.Query[*dialect.DeleteQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ - BaseQuery: deleteTableBaseQuery(t.NameAs(), t.Columns, queryMods), + BaseQuery: deleteTableBaseQuery(t.NameAs(), t.Columns), Hooks: &t.DeleteQueryHooks, }, Scanner: t.scanner, }, - hasDefaultReturning: !hasDeleteReturning(queryMods), + defaultReturning: t.Columns, } return q.Apply(queryMods...) @@ -235,7 +253,7 @@ func (t *Table[T, Tslice, Tset, C]) Merge(queryMods ...bob.Mod[*dialect.MergeQue return q } -func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Expression, queryMods []bob.Mod[*dialect.InsertQuery]) bob.BaseQuery[*dialect.InsertQuery] { +func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Expression) bob.BaseQuery[*dialect.InsertQuery] { base := bob.BaseQuery[*dialect.InsertQuery]{ Expression: &dialect.InsertQuery{ TableRef: clause.TableRef{ @@ -246,13 +264,11 @@ func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Exp Dialect: dialect.Dialect, QueryType: bob.QueryTypeInsert, } - if !hasInsertReturning(queryMods) { - base.Expression.AppendReturning(returning) - } + base.Expression.AppendReturning(returning) return base } -func updateTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mod[*dialect.UpdateQuery]) bob.BaseQuery[*dialect.UpdateQuery] { +func updateTableBaseQuery(name any, returning bob.Expression) bob.BaseQuery[*dialect.UpdateQuery] { base := bob.BaseQuery[*dialect.UpdateQuery]{ Expression: &dialect.UpdateQuery{ Table: clause.TableRef{ @@ -262,13 +278,11 @@ func updateTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mo Dialect: dialect.Dialect, QueryType: bob.QueryTypeUpdate, } - if !hasUpdateReturning(queryMods) { - base.Expression.AppendReturning(returning) - } + base.Expression.AppendReturning(returning) return base } -func deleteTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mod[*dialect.DeleteQuery]) bob.BaseQuery[*dialect.DeleteQuery] { +func deleteTableBaseQuery(name any, returning bob.Expression) bob.BaseQuery[*dialect.DeleteQuery] { base := bob.BaseQuery[*dialect.DeleteQuery]{ Expression: &dialect.DeleteQuery{ Table: clause.TableRef{ @@ -278,35 +292,6 @@ func deleteTableBaseQuery(name any, returning bob.Expression, queryMods []bob.Mo Dialect: dialect.Dialect, QueryType: bob.QueryTypeDelete, } - if !hasDeleteReturning(queryMods) { - base.Expression.AppendReturning(returning) - } + base.Expression.AppendReturning(returning) return base } - -func hasInsertReturning(mods []bob.Mod[*dialect.InsertQuery]) bool { - for _, mod := range mods { - if _, ok := mod.(bobmods.Returning[*dialect.InsertQuery]); ok { - return true - } - } - return false -} - -func hasUpdateReturning(mods []bob.Mod[*dialect.UpdateQuery]) bool { - for _, mod := range mods { - if _, ok := mod.(bobmods.Returning[*dialect.UpdateQuery]); ok { - return true - } - } - return false -} - -func hasDeleteReturning(mods []bob.Mod[*dialect.DeleteQuery]) bool { - for _, mod := range mods { - if _, ok := mod.(bobmods.Returning[*dialect.DeleteQuery]); ok { - return true - } - } - return false -} diff --git a/dialect/psql/table_test.go b/dialect/psql/table_test.go index f951f821..8900e3e4 100644 --- a/dialect/psql/table_test.go +++ b/dialect/psql/table_test.go @@ -174,6 +174,28 @@ func TestTableUpdateExplicitReturningOverridesDefault(t *testing.T) { } } +func TestTableUpdateAdditionalExplicitReturningAppends(t *testing.T) { + base := userTable.Update( + um.SetCol("name").ToArg("Stephen"), + um.Where(Quote("id").EQ(Arg(1))), + ) + + q := base.With(um.Returning("id")).With(um.Returning("email")) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "UPDATE \"users\" AS \"users\" SET\n\"name\" = $1\nWHERE (\"id\" = $2)\nRETURNING id, email" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 2 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + func TestTableInsertDefaultsReturningAllColumns(t *testing.T) { q := userTable.Insert( im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), @@ -213,6 +235,27 @@ func TestTableInsertExplicitReturningOverridesDefault(t *testing.T) { } } +func TestTableInsertAdditionalExplicitReturningAppends(t *testing.T) { + base := userTable.Insert( + im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), + ) + + q := base.With(im.Returning("id")).With(im.Returning("email")) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "INSERT INTO \"users\" AS \"users\"(\"id\", \"name\", \"email\")\nVALUES ($1, $2, $3)\nRETURNING id, email\n" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 3 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + func TestTableDeleteDefaultsReturningAllColumns(t *testing.T) { q := userTable.Delete( dm.Where(Quote("id").EQ(Arg(1))), @@ -252,6 +295,27 @@ func TestTableDeleteExplicitReturningOverridesDefault(t *testing.T) { } } +func TestTableDeleteAdditionalExplicitReturningAppends(t *testing.T) { + base := userTable.Delete( + dm.Where(Quote("id").EQ(Arg(1))), + ) + + q := base.With(dm.Returning("id")).With(dm.Returning("email")) + + sql, args, err := q.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + + expectedSQL := "DELETE FROM \"users\" AS \"users\"\nWHERE (\"id\" = $1)\nRETURNING id, email" + if sql != expectedSQL { + t.Fatalf("unexpected SQL: %#v", sql) + } + if len(args) != 1 { + t.Fatalf("unexpected arg count: %d", len(args)) + } +} + func TestTableUpdateApplyDoesNotMutateOriginal(t *testing.T) { base := userTable.Update( um.SetCol("name").ToArg("Stephen"), diff --git a/mods/mods.go b/mods/mods.go index f95c2434..f870a2c3 100644 --- a/mods/mods.go +++ b/mods/mods.go @@ -161,10 +161,6 @@ func (s Returning[Q]) Apply(q Q) { q.AppendReturning(s...) } -func (s Returning[Q]) ReturningValues() []any { - return []any(s) -} - type Set[Q interface{ AppendSet(clauses ...any) }] []string func (s Set[Q]) To(to any) bob.Mod[Q] { From 744a5b1ad5bbf233aa795ba320fba0db7da9bb2b Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 11:33:52 -0400 Subject: [PATCH 27/33] refactor(psql): move derivation into dialect queries Benchmarks:\n- BenchmarkBaseQueryImmutableNativeHotPath: 2938 ns/op, 3149 B/op, 35 allocs/op\n- BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath: 6092 ns/op, 6316 B/op, 62 allocs/op\n- BenchmarkUpdateQueryImmutableNativeHotPath: 2561 ns/op, 2543 B/op, 41 allocs/op\n- BenchmarkDeleteQueryImmutableNativeHotPath: 1979 ns/op, 2156 B/op, 27 allocs/op\n- BenchmarkInsertQueryImmutableNativeHotPath: 1532 ns/op, 1720 B/op, 18 allocs/op --- dialect/psql/delete.go | 2 +- dialect/psql/{ => dialect}/derive.go | 194 ++++++++++++--------------- dialect/psql/dialect/insert.go | 8 ++ dialect/psql/dialect/mods.go | 28 ---- dialect/psql/dialect/select.go | 9 ++ dialect/psql/im/qm.go | 6 +- dialect/psql/insert.go | 2 +- dialect/psql/select.go | 2 +- dialect/psql/sm/qm.go | 2 +- dialect/psql/um/qm.go | 2 +- dialect/psql/update.go | 2 +- mods/mods.go | 26 ++++ 12 files changed, 135 insertions(+), 148 deletions(-) rename dialect/psql/{ => dialect}/derive.go (52%) diff --git a/dialect/psql/delete.go b/dialect/psql/delete.go index 135589df..870000be 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -10,7 +10,7 @@ type DeleteQuery struct { } func (q DeleteQuery) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { - if next, ok := deriveDelete(q.Expression, queryMods...); ok { + if next, ok := q.Expression.Derive(queryMods...); ok { q.Expression = next return q } diff --git a/dialect/psql/derive.go b/dialect/psql/dialect/derive.go similarity index 52% rename from dialect/psql/derive.go rename to dialect/psql/dialect/derive.go index d0f94ccb..66a35f14 100644 --- a/dialect/psql/derive.go +++ b/dialect/psql/dialect/derive.go @@ -1,146 +1,118 @@ -package psql +package dialect import ( "github.com/stephenafamo/bob" "github.com/stephenafamo/bob/clause" - "github.com/stephenafamo/bob/dialect/psql/dialect" "github.com/stephenafamo/bob/mods" ) -func copyAnySlice(values []any) []any { - if values == nil { - return nil - } - return append([]any(nil), values...) -} - -func copyExpressionSlice(values []bob.Expression) []bob.Expression { - if values == nil { - return nil - } - return append([]bob.Expression(nil), values...) -} - -func copyTableRef(from clause.TableRef) clause.TableRef { - from.Columns = append([]string(nil), from.Columns...) - from.Partitions = append([]string(nil), from.Partitions...) - from.IndexHints = append([]clause.IndexHint(nil), from.IndexHints...) - from.Joins = append([]clause.Join(nil), from.Joins...) - for i := range from.Joins { - from.Joins[i].On = append([]bob.Expression(nil), from.Joins[i].On...) - from.Joins[i].Using = append([]string(nil), from.Joins[i].Using...) - from.Joins[i].To = copyTableRef(from.Joins[i].To) - } - return from -} - -func deriveSelect(base *dialect.SelectQuery, queryMods ...bob.Mod[*dialect.SelectQuery]) (*dialect.SelectQuery, bool) { +func (base *SelectQuery) Derive(queryMods ...bob.Mod[*SelectQuery]) (*SelectQuery, bool) { next := *base var cloneWith, cloneSelect, clonePreload, cloneWhere, cloneGroup, cloneHaving, cloneOrder, cloneWindows, cloneLocks, cloneJoins, cloneCombines, cloneCombinedOrder bool for _, mod := range queryMods { switch m := mod.(type) { - case mods.Recursive[*dialect.SelectQuery]: + case mods.Recursive[*SelectQuery]: next.With.Recursive = bool(m) - case dialect.CTEChain[*dialect.SelectQuery]: + case CTEChain[*SelectQuery]: if !cloneWith { - next.With.CTEs = copyExpressionSlice(base.With.CTEs) + next.With.CTEs = cloneExpressionSlice(base.With.CTEs) cloneWith = true } next.With.CTEs = append(next.With.CTEs, m()) - case dialect.DistinctMod: - next.Distinct.On = append(make([]any, 0, len(m.On)), m.On...) - case mods.Select[*dialect.SelectQuery]: + case mods.Distinct[*SelectQuery]: + next.SetDistinctValues([]any(m)) + case mods.Select[*SelectQuery]: if !cloneSelect { - next.SelectList.Columns = copyAnySlice(base.SelectList.Columns) + next.SelectList.Columns = cloneAnySlice(base.SelectList.Columns) cloneSelect = true } next.SelectList.Columns = append(next.SelectList.Columns, []any(m)...) - case mods.Preload[*dialect.SelectQuery]: + case mods.Preload[*SelectQuery]: if !clonePreload { - next.SelectList.PreloadColumns = copyAnySlice(base.SelectList.PreloadColumns) + next.SelectList.PreloadColumns = cloneAnySlice(base.SelectList.PreloadColumns) clonePreload = true } next.SelectList.PreloadColumns = append(next.SelectList.PreloadColumns, []any(m)...) - case mods.Where[*dialect.SelectQuery]: + case mods.Where[*SelectQuery]: if !cloneWhere { - next.Where.Conditions = copyAnySlice(base.Where.Conditions) + next.Where.Conditions = cloneAnySlice(base.Where.Conditions) cloneWhere = true } next.Where.Conditions = append(next.Where.Conditions, m.E) - case mods.GroupBy[*dialect.SelectQuery]: + case mods.GroupBy[*SelectQuery]: if !cloneGroup { - next.GroupBy.Groups = copyAnySlice(base.GroupBy.Groups) + next.GroupBy.Groups = cloneAnySlice(base.GroupBy.Groups) cloneGroup = true } next.GroupBy.Groups = append(next.GroupBy.Groups, m.E) - case mods.GroupByDistinct[*dialect.SelectQuery]: + case mods.GroupByDistinct[*SelectQuery]: next.GroupBy.Distinct = bool(m) - case mods.GroupWith[*dialect.SelectQuery]: + case mods.GroupWith[*SelectQuery]: next.GroupBy.With = string(m) - case mods.Having[*dialect.SelectQuery]: + case mods.Having[*SelectQuery]: if !cloneHaving { - next.Having.Conditions = copyAnySlice(base.Having.Conditions) + next.Having.Conditions = cloneAnySlice(base.Having.Conditions) cloneHaving = true } next.Having.Conditions = append(next.Having.Conditions, []any(m)...) - case mods.Limit[*dialect.SelectQuery]: + case mods.Limit[*SelectQuery]: next.Limit.Count = m.Count - case mods.Offset[*dialect.SelectQuery]: + case mods.Offset[*SelectQuery]: next.Offset.Count = m.Count - case mods.Fetch[*dialect.SelectQuery]: + case mods.Fetch[*SelectQuery]: next.Fetch = clause.Fetch(m) - case dialect.OrderBy[*dialect.SelectQuery]: + case OrderBy[*SelectQuery]: if !cloneOrder { - next.OrderBy.Expressions = copyExpressionSlice(base.OrderBy.Expressions) + next.OrderBy.Expressions = cloneExpressionSlice(base.OrderBy.Expressions) cloneOrder = true } next.OrderBy.Expressions = append(next.OrderBy.Expressions, m()) - case mods.Join[*dialect.SelectQuery]: + case mods.Join[*SelectQuery]: if !cloneJoins { next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) cloneJoins = true } next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) - case dialect.CrossJoinChain[*dialect.SelectQuery]: + case CrossJoinChain[*SelectQuery]: if !cloneJoins { next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) cloneJoins = true } next.TableRef.Joins = append(next.TableRef.Joins, m()) - case mods.NamedWindow[*dialect.SelectQuery]: + case mods.NamedWindow[*SelectQuery]: if !cloneWindows { - next.Windows.Windows = copyExpressionSlice(base.Windows.Windows) + next.Windows.Windows = cloneExpressionSlice(base.Windows.Windows) cloneWindows = true } next.Windows.Windows = append(next.Windows.Windows, clause.NamedWindow(m)) - case dialect.LockChain[*dialect.SelectQuery]: + case LockChain[*SelectQuery]: if !cloneLocks { - next.Locks.Locks = copyExpressionSlice(base.Locks.Locks) + next.Locks.Locks = cloneExpressionSlice(base.Locks.Locks) cloneLocks = true } next.Locks.Locks = append(next.Locks.Locks, m()) - case mods.Combine[*dialect.SelectQuery]: + case mods.Combine[*SelectQuery]: if !cloneCombines { next.Combines.Queries = append([]clause.Combine(nil), base.Combines.Queries...) cloneCombines = true } next.Combines.Queries = append(next.Combines.Queries, clause.Combine(m)) - case dialect.OrderCombined: + case OrderCombined: if !cloneCombinedOrder { - next.CombinedOrder.Expressions = copyExpressionSlice(base.CombinedOrder.Expressions) + next.CombinedOrder.Expressions = cloneExpressionSlice(base.CombinedOrder.Expressions) cloneCombinedOrder = true } next.CombinedOrder.Expressions = append(next.CombinedOrder.Expressions, m()) - case dialect.LimitCombined: + case LimitCombined: next.CombinedLimit.Count = m.Count - case dialect.OffsetCombined: + case OffsetCombined: next.CombinedOffset.Count = m.Count - case dialect.FetchCombined: + case FetchCombined: next.CombinedFetch.Count = m.Count next.CombinedFetch.WithTies = m.WithTies - case dialect.FromChain[*dialect.SelectQuery]: - next.TableRef = copyTableRef(m()) + case FromChain[*SelectQuery]: + next.TableRef = cloneTableRef(m()) default: return nil, false } @@ -149,51 +121,51 @@ func deriveSelect(base *dialect.SelectQuery, queryMods ...bob.Mod[*dialect.Selec return &next, true } -func deriveUpdate(base *dialect.UpdateQuery, queryMods ...bob.Mod[*dialect.UpdateQuery]) (*dialect.UpdateQuery, bool) { +func (base *UpdateQuery) Derive(queryMods ...bob.Mod[*UpdateQuery]) (*UpdateQuery, bool) { next := *base var cloneWith, cloneSet, cloneWhere, cloneReturning, cloneJoins bool for _, mod := range queryMods { switch m := mod.(type) { - case mods.Recursive[*dialect.UpdateQuery]: + case mods.Recursive[*UpdateQuery]: next.With.Recursive = bool(m) - case dialect.CTEChain[*dialect.UpdateQuery]: + case CTEChain[*UpdateQuery]: if !cloneWith { - next.With.CTEs = copyExpressionSlice(base.With.CTEs) + next.With.CTEs = cloneExpressionSlice(base.With.CTEs) cloneWith = true } next.With.CTEs = append(next.With.CTEs, m()) - case dialect.UpdateOnly: + case UpdateOnly: next.Only = bool(m) - case dialect.UpdateTable: - next.Table = copyTableRef(clause.TableRef(m)) - case dialect.UpdateSet: + case UpdateTable: + next.Table = cloneTableRef(clause.TableRef(m)) + case mods.SetExprs[*UpdateQuery]: if !cloneSet { - next.Set.Set = copyAnySlice(base.Set.Set) + next.Set.Set = cloneAnySlice(base.Set.Set) cloneSet = true } next.Set.Set = append(next.Set.Set, []any(m)...) - case mods.Where[*dialect.UpdateQuery]: + case mods.Where[*UpdateQuery]: if !cloneWhere { - next.Where.Conditions = copyAnySlice(base.Where.Conditions) + next.Where.Conditions = cloneAnySlice(base.Where.Conditions) cloneWhere = true } next.Where.Conditions = append(next.Where.Conditions, m.E) - case mods.Returning[*dialect.UpdateQuery]: + case mods.Returning[*UpdateQuery]: if !cloneReturning { - next.Returning.Expressions = copyAnySlice(base.Returning.Expressions) + next.Returning.Expressions = cloneAnySlice(base.Returning.Expressions) cloneReturning = true } next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) - case dialect.FromChain[*dialect.UpdateQuery]: - next.TableRef = copyTableRef(m()) - case mods.Join[*dialect.UpdateQuery]: + case FromChain[*UpdateQuery]: + next.TableRef = cloneTableRef(m()) + case mods.Join[*UpdateQuery]: if !cloneJoins { next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) cloneJoins = true } next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) - case dialect.CrossJoinChain[*dialect.UpdateQuery]: + case CrossJoinChain[*UpdateQuery]: if !cloneJoins { next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) cloneJoins = true @@ -207,45 +179,45 @@ func deriveUpdate(base *dialect.UpdateQuery, queryMods ...bob.Mod[*dialect.Updat return &next, true } -func deriveDelete(base *dialect.DeleteQuery, queryMods ...bob.Mod[*dialect.DeleteQuery]) (*dialect.DeleteQuery, bool) { +func (base *DeleteQuery) Derive(queryMods ...bob.Mod[*DeleteQuery]) (*DeleteQuery, bool) { next := *base var cloneWith, cloneWhere, cloneReturning, cloneJoins bool for _, mod := range queryMods { switch m := mod.(type) { - case mods.Recursive[*dialect.DeleteQuery]: + case mods.Recursive[*DeleteQuery]: next.With.Recursive = bool(m) - case dialect.CTEChain[*dialect.DeleteQuery]: + case CTEChain[*DeleteQuery]: if !cloneWith { - next.With.CTEs = copyExpressionSlice(base.With.CTEs) + next.With.CTEs = cloneExpressionSlice(base.With.CTEs) cloneWith = true } next.With.CTEs = append(next.With.CTEs, m()) - case dialect.DeleteOnly: + case DeleteOnly: next.Only = bool(m) - case dialect.DeleteTable: - next.Table = copyTableRef(clause.TableRef(m)) - case mods.Where[*dialect.DeleteQuery]: + case DeleteTable: + next.Table = cloneTableRef(clause.TableRef(m)) + case mods.Where[*DeleteQuery]: if !cloneWhere { - next.Where.Conditions = copyAnySlice(base.Where.Conditions) + next.Where.Conditions = cloneAnySlice(base.Where.Conditions) cloneWhere = true } next.Where.Conditions = append(next.Where.Conditions, m.E) - case mods.Returning[*dialect.DeleteQuery]: + case mods.Returning[*DeleteQuery]: if !cloneReturning { - next.Returning.Expressions = copyAnySlice(base.Returning.Expressions) + next.Returning.Expressions = cloneAnySlice(base.Returning.Expressions) cloneReturning = true } next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) - case dialect.FromChain[*dialect.DeleteQuery]: - next.TableRef = copyTableRef(m()) - case mods.Join[*dialect.DeleteQuery]: + case FromChain[*DeleteQuery]: + next.TableRef = cloneTableRef(m()) + case mods.Join[*DeleteQuery]: if !cloneJoins { next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) cloneJoins = true } next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) - case dialect.CrossJoinChain[*dialect.DeleteQuery]: + case CrossJoinChain[*DeleteQuery]: if !cloneJoins { next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) cloneJoins = true @@ -259,39 +231,39 @@ func deriveDelete(base *dialect.DeleteQuery, queryMods ...bob.Mod[*dialect.Delet return &next, true } -func deriveInsert(base *dialect.InsertQuery, queryMods ...bob.Mod[*dialect.InsertQuery]) (*dialect.InsertQuery, bool) { +func (base *InsertQuery) Derive(queryMods ...bob.Mod[*InsertQuery]) (*InsertQuery, bool) { next := *base var cloneWith, cloneReturning, cloneVals bool for _, mod := range queryMods { switch m := mod.(type) { - case mods.Recursive[*dialect.InsertQuery]: + case mods.Recursive[*InsertQuery]: next.With.Recursive = bool(m) - case dialect.CTEChain[*dialect.InsertQuery]: + case CTEChain[*InsertQuery]: if !cloneWith { - next.With.CTEs = copyExpressionSlice(base.With.CTEs) + next.With.CTEs = cloneExpressionSlice(base.With.CTEs) cloneWith = true } next.With.CTEs = append(next.With.CTEs, m()) - case dialect.InsertTable: - next.TableRef = copyTableRef(clause.TableRef(m)) - case dialect.InsertOverriding: + case InsertTable: + next.TableRef = cloneTableRef(clause.TableRef(m)) + case mods.Overriding[*InsertQuery]: next.Overriding = string(m) - case dialect.InsertQuerySource: + case mods.QuerySource[*InsertQuery]: next.Values.Query = m.Query - case mods.Returning[*dialect.InsertQuery]: + case mods.Returning[*InsertQuery]: if !cloneReturning { - next.Returning.Expressions = copyAnySlice(base.Returning.Expressions) + next.Returning.Expressions = cloneAnySlice(base.Returning.Expressions) cloneReturning = true } next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) - case mods.Values[*dialect.InsertQuery]: + case mods.Values[*InsertQuery]: if !cloneVals { next.Values.Vals = append([]clause.Value(nil), base.Values.Vals...) cloneVals = true } next.Values.Vals = append(next.Values.Vals, clause.Value(m)) - case mods.Rows[*dialect.InsertQuery]: + case mods.Rows[*InsertQuery]: if !cloneVals { next.Values.Vals = append([]clause.Value(nil), base.Values.Vals...) cloneVals = true @@ -299,7 +271,7 @@ func deriveInsert(base *dialect.InsertQuery, queryMods ...bob.Mod[*dialect.Inser for _, row := range m { next.Values.Vals = append(next.Values.Vals, clause.Value(row)) } - case mods.Conflict[*dialect.InsertQuery]: + case mods.Conflict[*InsertQuery]: next.Conflict.Expression = m() default: return nil, false diff --git a/dialect/psql/dialect/insert.go b/dialect/psql/dialect/insert.go index 85892b5b..8847cfae 100644 --- a/dialect/psql/dialect/insert.go +++ b/dialect/psql/dialect/insert.go @@ -32,6 +32,14 @@ type InsertQuery struct { bob.ContextualModdable[*InsertQuery] } +func (i *InsertQuery) SetOverriding(overriding string) { + i.Overriding = overriding +} + +func (i *InsertQuery) SetQuery(q bob.Query) { + i.Values.Query = q +} + func (i InsertQuery) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { var err error diff --git a/dialect/psql/dialect/mods.go b/dialect/psql/dialect/mods.go index e573815b..ce235c39 100644 --- a/dialect/psql/dialect/mods.go +++ b/dialect/psql/dialect/mods.go @@ -19,14 +19,6 @@ func (di Distinct) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialec return bob.ExpressSlice(ctx, w, d, start, di.On, " ON (", ", ", ")") } -type DistinctMod struct { - On []any -} - -func (d DistinctMod) Apply(q *SelectQuery) { - q.Distinct.On = d.On -} - type UpdateOnly bool func (o UpdateOnly) Apply(q *UpdateQuery) { @@ -57,26 +49,6 @@ func (t InsertTable) Apply(q *InsertQuery) { q.TableRef = clause.TableRef(t) } -type UpdateSet []any - -func (s UpdateSet) Apply(q *UpdateQuery) { - q.Set.Set = append(q.Set.Set, []any(s)...) -} - -type InsertOverriding string - -func (o InsertOverriding) Apply(q *InsertQuery) { - q.Overriding = string(o) -} - -type InsertQuerySource struct { - Query bob.Query -} - -func (s InsertQuerySource) Apply(q *InsertQuery) { - q.Query = s.Query -} - func With[Q interface{ AppendCTE(bob.Expression) }](name string, columns ...string) CTEChain[Q] { return CTEChain[Q](func() clause.CTE { return clause.CTE{ diff --git a/dialect/psql/dialect/select.go b/dialect/psql/dialect/select.go index 1b6a1f22..5bebc909 100644 --- a/dialect/psql/dialect/select.go +++ b/dialect/psql/dialect/select.go @@ -36,6 +36,15 @@ type SelectQuery struct { CombinedOffset clause.Offset } +func (s *SelectQuery) SetDistinctValues(on []any) { + if on == nil { + s.Distinct.On = nil + return + } + + s.Distinct.On = append(make([]any, 0, len(on)), on...) +} + func (s SelectQuery) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { var err error diff --git a/dialect/psql/im/qm.go b/dialect/psql/im/qm.go index bb6c835b..f673a411 100644 --- a/dialect/psql/im/qm.go +++ b/dialect/psql/im/qm.go @@ -33,11 +33,11 @@ func IntoAs(name any, alias string, columns ...string) bob.Mod[*dialect.InsertQu } func OverridingSystem() bob.Mod[*dialect.InsertQuery] { - return dialect.InsertOverriding("SYSTEM") + return mods.Overriding[*dialect.InsertQuery]("SYSTEM") } func OverridingUser() bob.Mod[*dialect.InsertQuery] { - return dialect.InsertOverriding("USER") + return mods.Overriding[*dialect.InsertQuery]("USER") } func Values(clauses ...bob.Expression) bob.Mod[*dialect.InsertQuery] { @@ -50,7 +50,7 @@ func Rows(rows ...[]bob.Expression) bob.Mod[*dialect.InsertQuery] { // Insert from a query func Query(q bob.Query) bob.Mod[*dialect.InsertQuery] { - return dialect.InsertQuerySource{Query: q} + return mods.QuerySource[*dialect.InsertQuery]{Query: q} } // The column to target. Will auto add brackets diff --git a/dialect/psql/insert.go b/dialect/psql/insert.go index b1d3c70c..d2c10d63 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -10,7 +10,7 @@ type InsertQuery struct { } func (q InsertQuery) With(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { - if next, ok := deriveInsert(q.Expression, queryMods...); ok { + if next, ok := q.Expression.Derive(queryMods...); ok { q.Expression = next return q } diff --git a/dialect/psql/select.go b/dialect/psql/select.go index 90fdd2f3..1e9c5e74 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -10,7 +10,7 @@ type SelectQuery struct { } func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { - if next, ok := deriveSelect(q.Expression, queryMods...); ok { + if next, ok := q.Expression.Derive(queryMods...); ok { q.Expression = next return q } diff --git a/dialect/psql/sm/qm.go b/dialect/psql/sm/qm.go index 8fcaf5f1..2d0af599 100644 --- a/dialect/psql/sm/qm.go +++ b/dialect/psql/sm/qm.go @@ -20,7 +20,7 @@ func Distinct(on ...any) bob.Mod[*dialect.SelectQuery] { on = []any{} // nil means no distinct } - return dialect.DistinctMod{On: on} + return mods.Distinct[*dialect.SelectQuery](on) } func Columns(clauses ...any) bob.Mod[*dialect.SelectQuery] { diff --git a/dialect/psql/um/qm.go b/dialect/psql/um/qm.go index c93f0302..68a3ac95 100644 --- a/dialect/psql/um/qm.go +++ b/dialect/psql/um/qm.go @@ -33,7 +33,7 @@ func TableAs(name any, alias string) bob.Mod[*dialect.UpdateQuery] { } func Set(sets ...bob.Expression) bob.Mod[*dialect.UpdateQuery] { - return dialect.UpdateSet(internal.ToAnySlice(sets)) + return mods.SetExprs[*dialect.UpdateQuery](internal.ToAnySlice(sets)) } func SetCol(from string) mods.Set[*dialect.UpdateQuery] { diff --git a/dialect/psql/update.go b/dialect/psql/update.go index e61818ea..14b281d6 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -10,7 +10,7 @@ type UpdateQuery struct { } func (q UpdateQuery) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { - if next, ok := deriveUpdate(q.Expression, queryMods...); ok { + if next, ok := q.Expression.Derive(queryMods...); ok { q.Expression = next return q } diff --git a/mods/mods.go b/mods/mods.go index f870a2c3..642a04ed 100644 --- a/mods/mods.go +++ b/mods/mods.go @@ -43,6 +43,12 @@ func (s Select[Q]) Apply(q Q) { q.AppendSelect(s...) } +type Distinct[Q interface{ SetDistinctValues([]any) }] []any + +func (d Distinct[Q]) Apply(q Q) { + q.SetDistinctValues([]any(d)) +} + type Preload[Q interface{ AppendPreloadSelect(columns ...any) }] []any func (s Preload[Q]) Apply(q Q) { @@ -155,6 +161,20 @@ func (r Rows[Q]) Apply(q Q) { } } +type QuerySource[Q interface{ SetQuery(bob.Query) }] struct { + Query bob.Query +} + +func (s QuerySource[Q]) Apply(q Q) { + q.SetQuery(s.Query) +} + +type Overriding[Q interface{ SetOverriding(string) }] string + +func (o Overriding[Q]) Apply(q Q) { + q.SetOverriding(string(o)) +} + type Returning[Q interface{ AppendReturning(vals ...any) }] []any func (s Returning[Q]) Apply(q Q) { @@ -171,6 +191,12 @@ func (s Set[Q]) ToArg(to any) bob.Mod[Q] { return set[Q]{expr.OP("=", expr.Quote(s...), expr.Arg(to))} } +type SetExprs[Q interface{ AppendSet(clauses ...any) }] []any + +func (s SetExprs[Q]) Apply(q Q) { + q.AppendSet(s...) +} + type set[Q interface{ AppendSet(clauses ...any) }] []any func (s set[Q]) Apply(q Q) { From d9deaced4b24476990eddccab0068971d7d5c18f Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 11:49:54 -0400 Subject: [PATCH 28/33] refactor(psql): reduce target table mod glue Benchmarks:\n- BenchmarkBaseQueryImmutableNativeHotPath: 2868 ns/op, 3149 B/op, 35 allocs/op\n- BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath: 5948 ns/op, 6316 B/op, 62 allocs/op\n- BenchmarkUpdateQueryImmutableNativeHotPath: 2509 ns/op, 2543 B/op, 41 allocs/op\n- BenchmarkDeleteQueryImmutableNativeHotPath: 1957 ns/op, 2156 B/op, 27 allocs/op\n- BenchmarkInsertQueryImmutableNativeHotPath: 1487 ns/op, 1720 B/op, 18 allocs/op --- dialect/psql/dialect/delete.go | 12 +++ dialect/psql/dialect/derive.go | 190 +++++++++------------------------ dialect/psql/dialect/insert.go | 8 ++ dialect/psql/dialect/mods.go | 30 ------ dialect/psql/dialect/update.go | 12 +++ dialect/psql/dm/qm.go | 6 +- dialect/psql/im/qm.go | 4 +- dialect/psql/um/qm.go | 6 +- mods/mods.go | 18 ++++ 9 files changed, 107 insertions(+), 179 deletions(-) diff --git a/dialect/psql/dialect/delete.go b/dialect/psql/dialect/delete.go index 49d4da97..86df4fa3 100644 --- a/dialect/psql/dialect/delete.go +++ b/dialect/psql/dialect/delete.go @@ -23,6 +23,18 @@ type DeleteQuery struct { bob.ContextualModdable[*DeleteQuery] } +func (d *DeleteQuery) SetTargetOnly(only bool) { + d.Table.SetOnly(only) +} + +func (d *DeleteQuery) SetTargetTable(table any) { + d.Table.SetTable(table) +} + +func (d *DeleteQuery) SetTargetTableAlias(alias string, columns ...string) { + d.Table.SetTableAlias(alias, columns...) +} + func (d DeleteQuery) WriteSQL(ctx context.Context, w io.StringWriter, dl bob.Dialect, start int) ([]any, error) { var err error diff --git a/dialect/psql/dialect/derive.go b/dialect/psql/dialect/derive.go index 66a35f14..c357aea3 100644 --- a/dialect/psql/dialect/derive.go +++ b/dialect/psql/dialect/derive.go @@ -6,6 +6,22 @@ import ( "github.com/stephenafamo/bob/mods" ) +func cloneSlice[T any](values []T) []T { + if values == nil { + return nil + } + return append([]T(nil), values...) +} + +func appendDerived[T any](target *[]T, base []T, cloned *bool, values ...T) { + if !*cloned { + *target = cloneSlice(base) + *cloned = true + } + + *target = append(*target, values...) +} + func (base *SelectQuery) Derive(queryMods ...bob.Mod[*SelectQuery]) (*SelectQuery, bool) { next := *base var cloneWith, cloneSelect, clonePreload, cloneWhere, cloneGroup, cloneHaving, cloneOrder, cloneWindows, cloneLocks, cloneJoins, cloneCombines, cloneCombinedOrder bool @@ -15,47 +31,23 @@ func (base *SelectQuery) Derive(queryMods ...bob.Mod[*SelectQuery]) (*SelectQuer case mods.Recursive[*SelectQuery]: next.With.Recursive = bool(m) case CTEChain[*SelectQuery]: - if !cloneWith { - next.With.CTEs = cloneExpressionSlice(base.With.CTEs) - cloneWith = true - } - next.With.CTEs = append(next.With.CTEs, m()) + appendDerived[bob.Expression](&next.With.CTEs, base.With.CTEs, &cloneWith, m()) case mods.Distinct[*SelectQuery]: next.SetDistinctValues([]any(m)) case mods.Select[*SelectQuery]: - if !cloneSelect { - next.SelectList.Columns = cloneAnySlice(base.SelectList.Columns) - cloneSelect = true - } - next.SelectList.Columns = append(next.SelectList.Columns, []any(m)...) + appendDerived(&next.SelectList.Columns, base.SelectList.Columns, &cloneSelect, []any(m)...) case mods.Preload[*SelectQuery]: - if !clonePreload { - next.SelectList.PreloadColumns = cloneAnySlice(base.SelectList.PreloadColumns) - clonePreload = true - } - next.SelectList.PreloadColumns = append(next.SelectList.PreloadColumns, []any(m)...) + appendDerived(&next.SelectList.PreloadColumns, base.SelectList.PreloadColumns, &clonePreload, []any(m)...) case mods.Where[*SelectQuery]: - if !cloneWhere { - next.Where.Conditions = cloneAnySlice(base.Where.Conditions) - cloneWhere = true - } - next.Where.Conditions = append(next.Where.Conditions, m.E) + appendDerived[any](&next.Where.Conditions, base.Where.Conditions, &cloneWhere, m.E) case mods.GroupBy[*SelectQuery]: - if !cloneGroup { - next.GroupBy.Groups = cloneAnySlice(base.GroupBy.Groups) - cloneGroup = true - } - next.GroupBy.Groups = append(next.GroupBy.Groups, m.E) + appendDerived(&next.GroupBy.Groups, base.GroupBy.Groups, &cloneGroup, m.E) case mods.GroupByDistinct[*SelectQuery]: next.GroupBy.Distinct = bool(m) case mods.GroupWith[*SelectQuery]: next.GroupBy.With = string(m) case mods.Having[*SelectQuery]: - if !cloneHaving { - next.Having.Conditions = cloneAnySlice(base.Having.Conditions) - cloneHaving = true - } - next.Having.Conditions = append(next.Having.Conditions, []any(m)...) + appendDerived(&next.Having.Conditions, base.Having.Conditions, &cloneHaving, []any(m)...) case mods.Limit[*SelectQuery]: next.Limit.Count = m.Count case mods.Offset[*SelectQuery]: @@ -63,47 +55,19 @@ func (base *SelectQuery) Derive(queryMods ...bob.Mod[*SelectQuery]) (*SelectQuer case mods.Fetch[*SelectQuery]: next.Fetch = clause.Fetch(m) case OrderBy[*SelectQuery]: - if !cloneOrder { - next.OrderBy.Expressions = cloneExpressionSlice(base.OrderBy.Expressions) - cloneOrder = true - } - next.OrderBy.Expressions = append(next.OrderBy.Expressions, m()) + appendDerived[bob.Expression](&next.OrderBy.Expressions, base.OrderBy.Expressions, &cloneOrder, m()) case mods.Join[*SelectQuery]: - if !cloneJoins { - next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) - cloneJoins = true - } - next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, clause.Join(m)) case CrossJoinChain[*SelectQuery]: - if !cloneJoins { - next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) - cloneJoins = true - } - next.TableRef.Joins = append(next.TableRef.Joins, m()) + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, m()) case mods.NamedWindow[*SelectQuery]: - if !cloneWindows { - next.Windows.Windows = cloneExpressionSlice(base.Windows.Windows) - cloneWindows = true - } - next.Windows.Windows = append(next.Windows.Windows, clause.NamedWindow(m)) + appendDerived[bob.Expression](&next.Windows.Windows, base.Windows.Windows, &cloneWindows, clause.NamedWindow(m)) case LockChain[*SelectQuery]: - if !cloneLocks { - next.Locks.Locks = cloneExpressionSlice(base.Locks.Locks) - cloneLocks = true - } - next.Locks.Locks = append(next.Locks.Locks, m()) + appendDerived[bob.Expression](&next.Locks.Locks, base.Locks.Locks, &cloneLocks, m()) case mods.Combine[*SelectQuery]: - if !cloneCombines { - next.Combines.Queries = append([]clause.Combine(nil), base.Combines.Queries...) - cloneCombines = true - } - next.Combines.Queries = append(next.Combines.Queries, clause.Combine(m)) + appendDerived(&next.Combines.Queries, base.Combines.Queries, &cloneCombines, clause.Combine(m)) case OrderCombined: - if !cloneCombinedOrder { - next.CombinedOrder.Expressions = cloneExpressionSlice(base.CombinedOrder.Expressions) - cloneCombinedOrder = true - } - next.CombinedOrder.Expressions = append(next.CombinedOrder.Expressions, m()) + appendDerived[bob.Expression](&next.CombinedOrder.Expressions, base.CombinedOrder.Expressions, &cloneCombinedOrder, m()) case LimitCombined: next.CombinedLimit.Count = m.Count case OffsetCombined: @@ -130,47 +94,23 @@ func (base *UpdateQuery) Derive(queryMods ...bob.Mod[*UpdateQuery]) (*UpdateQuer case mods.Recursive[*UpdateQuery]: next.With.Recursive = bool(m) case CTEChain[*UpdateQuery]: - if !cloneWith { - next.With.CTEs = cloneExpressionSlice(base.With.CTEs) - cloneWith = true - } - next.With.CTEs = append(next.With.CTEs, m()) - case UpdateOnly: + appendDerived[bob.Expression](&next.With.CTEs, base.With.CTEs, &cloneWith, m()) + case mods.TargetOnly[*UpdateQuery]: next.Only = bool(m) - case UpdateTable: + case mods.TargetTable[*UpdateQuery]: next.Table = cloneTableRef(clause.TableRef(m)) case mods.SetExprs[*UpdateQuery]: - if !cloneSet { - next.Set.Set = cloneAnySlice(base.Set.Set) - cloneSet = true - } - next.Set.Set = append(next.Set.Set, []any(m)...) + appendDerived(&next.Set.Set, base.Set.Set, &cloneSet, []any(m)...) case mods.Where[*UpdateQuery]: - if !cloneWhere { - next.Where.Conditions = cloneAnySlice(base.Where.Conditions) - cloneWhere = true - } - next.Where.Conditions = append(next.Where.Conditions, m.E) + appendDerived[any](&next.Where.Conditions, base.Where.Conditions, &cloneWhere, m.E) case mods.Returning[*UpdateQuery]: - if !cloneReturning { - next.Returning.Expressions = cloneAnySlice(base.Returning.Expressions) - cloneReturning = true - } - next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + appendDerived(&next.Returning.Expressions, base.Returning.Expressions, &cloneReturning, []any(m)...) case FromChain[*UpdateQuery]: next.TableRef = cloneTableRef(m()) case mods.Join[*UpdateQuery]: - if !cloneJoins { - next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) - cloneJoins = true - } - next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, clause.Join(m)) case CrossJoinChain[*UpdateQuery]: - if !cloneJoins { - next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) - cloneJoins = true - } - next.TableRef.Joins = append(next.TableRef.Joins, m()) + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, m()) default: return nil, false } @@ -188,41 +128,21 @@ func (base *DeleteQuery) Derive(queryMods ...bob.Mod[*DeleteQuery]) (*DeleteQuer case mods.Recursive[*DeleteQuery]: next.With.Recursive = bool(m) case CTEChain[*DeleteQuery]: - if !cloneWith { - next.With.CTEs = cloneExpressionSlice(base.With.CTEs) - cloneWith = true - } - next.With.CTEs = append(next.With.CTEs, m()) - case DeleteOnly: + appendDerived[bob.Expression](&next.With.CTEs, base.With.CTEs, &cloneWith, m()) + case mods.TargetOnly[*DeleteQuery]: next.Only = bool(m) - case DeleteTable: + case mods.TargetTable[*DeleteQuery]: next.Table = cloneTableRef(clause.TableRef(m)) case mods.Where[*DeleteQuery]: - if !cloneWhere { - next.Where.Conditions = cloneAnySlice(base.Where.Conditions) - cloneWhere = true - } - next.Where.Conditions = append(next.Where.Conditions, m.E) + appendDerived[any](&next.Where.Conditions, base.Where.Conditions, &cloneWhere, m.E) case mods.Returning[*DeleteQuery]: - if !cloneReturning { - next.Returning.Expressions = cloneAnySlice(base.Returning.Expressions) - cloneReturning = true - } - next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + appendDerived(&next.Returning.Expressions, base.Returning.Expressions, &cloneReturning, []any(m)...) case FromChain[*DeleteQuery]: next.TableRef = cloneTableRef(m()) case mods.Join[*DeleteQuery]: - if !cloneJoins { - next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) - cloneJoins = true - } - next.TableRef.Joins = append(next.TableRef.Joins, clause.Join(m)) + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, clause.Join(m)) case CrossJoinChain[*DeleteQuery]: - if !cloneJoins { - next.TableRef.Joins = append([]clause.Join(nil), base.TableRef.Joins...) - cloneJoins = true - } - next.TableRef.Joins = append(next.TableRef.Joins, m()) + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, m()) default: return nil, false } @@ -240,32 +160,20 @@ func (base *InsertQuery) Derive(queryMods ...bob.Mod[*InsertQuery]) (*InsertQuer case mods.Recursive[*InsertQuery]: next.With.Recursive = bool(m) case CTEChain[*InsertQuery]: - if !cloneWith { - next.With.CTEs = cloneExpressionSlice(base.With.CTEs) - cloneWith = true - } - next.With.CTEs = append(next.With.CTEs, m()) - case InsertTable: + appendDerived[bob.Expression](&next.With.CTEs, base.With.CTEs, &cloneWith, m()) + case mods.TargetTable[*InsertQuery]: next.TableRef = cloneTableRef(clause.TableRef(m)) case mods.Overriding[*InsertQuery]: next.Overriding = string(m) case mods.QuerySource[*InsertQuery]: next.Values.Query = m.Query case mods.Returning[*InsertQuery]: - if !cloneReturning { - next.Returning.Expressions = cloneAnySlice(base.Returning.Expressions) - cloneReturning = true - } - next.Returning.Expressions = append(next.Returning.Expressions, []any(m)...) + appendDerived(&next.Returning.Expressions, base.Returning.Expressions, &cloneReturning, []any(m)...) case mods.Values[*InsertQuery]: - if !cloneVals { - next.Values.Vals = append([]clause.Value(nil), base.Values.Vals...) - cloneVals = true - } - next.Values.Vals = append(next.Values.Vals, clause.Value(m)) + appendDerived(&next.Values.Vals, base.Values.Vals, &cloneVals, clause.Value(m)) case mods.Rows[*InsertQuery]: if !cloneVals { - next.Values.Vals = append([]clause.Value(nil), base.Values.Vals...) + next.Values.Vals = cloneSlice(base.Values.Vals) cloneVals = true } for _, row := range m { diff --git a/dialect/psql/dialect/insert.go b/dialect/psql/dialect/insert.go index 8847cfae..6558b4fc 100644 --- a/dialect/psql/dialect/insert.go +++ b/dialect/psql/dialect/insert.go @@ -32,6 +32,14 @@ type InsertQuery struct { bob.ContextualModdable[*InsertQuery] } +func (i *InsertQuery) SetTargetTable(table any) { + i.TableRef.SetTable(table) +} + +func (i *InsertQuery) SetTargetTableAlias(alias string, columns ...string) { + i.TableRef.SetTableAlias(alias, columns...) +} + func (i *InsertQuery) SetOverriding(overriding string) { i.Overriding = overriding } diff --git a/dialect/psql/dialect/mods.go b/dialect/psql/dialect/mods.go index ce235c39..8fcd80a4 100644 --- a/dialect/psql/dialect/mods.go +++ b/dialect/psql/dialect/mods.go @@ -19,36 +19,6 @@ func (di Distinct) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialec return bob.ExpressSlice(ctx, w, d, start, di.On, " ON (", ", ", ")") } -type UpdateOnly bool - -func (o UpdateOnly) Apply(q *UpdateQuery) { - q.Only = bool(o) -} - -type DeleteOnly bool - -func (o DeleteOnly) Apply(q *DeleteQuery) { - q.Only = bool(o) -} - -type UpdateTable clause.TableRef - -func (t UpdateTable) Apply(q *UpdateQuery) { - q.Table = clause.TableRef(t) -} - -type DeleteTable clause.TableRef - -func (t DeleteTable) Apply(q *DeleteQuery) { - q.Table = clause.TableRef(t) -} - -type InsertTable clause.TableRef - -func (t InsertTable) Apply(q *InsertQuery) { - q.TableRef = clause.TableRef(t) -} - func With[Q interface{ AppendCTE(bob.Expression) }](name string, columns ...string) CTEChain[Q] { return CTEChain[Q](func() clause.CTE { return clause.CTE{ diff --git a/dialect/psql/dialect/update.go b/dialect/psql/dialect/update.go index 885f3892..2f02cc57 100644 --- a/dialect/psql/dialect/update.go +++ b/dialect/psql/dialect/update.go @@ -24,6 +24,18 @@ type UpdateQuery struct { bob.ContextualModdable[*UpdateQuery] } +func (u *UpdateQuery) SetTargetOnly(only bool) { + u.Table.SetOnly(only) +} + +func (u *UpdateQuery) SetTargetTable(table any) { + u.Table.SetTable(table) +} + +func (u *UpdateQuery) SetTargetTableAlias(alias string, columns ...string) { + u.Table.SetTableAlias(alias, columns...) +} + func (u UpdateQuery) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { var err error diff --git a/dialect/psql/dm/qm.go b/dialect/psql/dm/qm.go index 4635356d..8dff1b96 100644 --- a/dialect/psql/dm/qm.go +++ b/dialect/psql/dm/qm.go @@ -15,17 +15,17 @@ func Recursive(r bool) bob.Mod[*dialect.DeleteQuery] { } func Only() bob.Mod[*dialect.DeleteQuery] { - return dialect.DeleteOnly(true) + return mods.TargetOnly[*dialect.DeleteQuery](true) } func From(name any) bob.Mod[*dialect.DeleteQuery] { - return dialect.DeleteTable{ + return mods.TargetTable[*dialect.DeleteQuery]{ Expression: name, } } func FromAs(name any, alias string) bob.Mod[*dialect.DeleteQuery] { - return dialect.DeleteTable{ + return mods.TargetTable[*dialect.DeleteQuery]{ Expression: name, Alias: alias, } diff --git a/dialect/psql/im/qm.go b/dialect/psql/im/qm.go index f673a411..e573396f 100644 --- a/dialect/psql/im/qm.go +++ b/dialect/psql/im/qm.go @@ -18,14 +18,14 @@ func Recursive(r bool) bob.Mod[*dialect.InsertQuery] { } func Into(name any, columns ...string) bob.Mod[*dialect.InsertQuery] { - return dialect.InsertTable{ + return mods.TargetTable[*dialect.InsertQuery]{ Expression: name, Columns: columns, } } func IntoAs(name any, alias string, columns ...string) bob.Mod[*dialect.InsertQuery] { - return dialect.InsertTable{ + return mods.TargetTable[*dialect.InsertQuery]{ Expression: name, Alias: alias, Columns: columns, diff --git a/dialect/psql/um/qm.go b/dialect/psql/um/qm.go index 68a3ac95..58a25822 100644 --- a/dialect/psql/um/qm.go +++ b/dialect/psql/um/qm.go @@ -16,17 +16,17 @@ func Recursive(r bool) bob.Mod[*dialect.UpdateQuery] { } func Only() bob.Mod[*dialect.UpdateQuery] { - return dialect.UpdateOnly(true) + return mods.TargetOnly[*dialect.UpdateQuery](true) } func Table(name any) bob.Mod[*dialect.UpdateQuery] { - return dialect.UpdateTable{ + return mods.TargetTable[*dialect.UpdateQuery]{ Expression: name, } } func TableAs(name any, alias string) bob.Mod[*dialect.UpdateQuery] { - return dialect.UpdateTable{ + return mods.TargetTable[*dialect.UpdateQuery]{ Expression: name, Alias: alias, } diff --git a/mods/mods.go b/mods/mods.go index 642a04ed..3f67be1a 100644 --- a/mods/mods.go +++ b/mods/mods.go @@ -37,6 +37,24 @@ func (r Recursive[Q]) Apply(q Q) { q.SetRecursive(bool(r)) } +type TargetOnly[Q interface{ SetTargetOnly(bool) }] bool + +func (o TargetOnly[Q]) Apply(q Q) { + q.SetTargetOnly(bool(o)) +} + +type TargetTable[Q interface { + SetTargetTable(any) + SetTargetTableAlias(alias string, columns ...string) +}] clause.TableRef + +func (t TargetTable[Q]) Apply(q Q) { + q.SetTargetTable(t.Expression) + if t.Alias != "" { + q.SetTargetTableAlias(t.Alias, t.Columns...) + } +} + type Select[Q interface{ AppendSelect(columns ...any) }] []any func (s Select[Q]) Apply(q Q) { From 5caedca467f529d8b1f434b5c73226d84381d5e1 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 11:58:17 -0400 Subject: [PATCH 29/33] refactor(psql): simplify update derivation Benchmarks:\n- BenchmarkBaseQueryImmutableNativeHotPath: 2857 ns/op, 3149 B/op, 35 allocs/op\n- BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath: 5767 ns/op, 6316 B/op, 62 allocs/op\n- BenchmarkUpdateQueryImmutableNativeHotPath: 2532 ns/op, 2576 B/op, 43 allocs/op\n- BenchmarkDeleteQueryImmutableNativeHotPath: 1872 ns/op, 2156 B/op, 27 allocs/op\n- BenchmarkInsertQueryImmutableNativeHotPath: 1422 ns/op, 1720 B/op, 18 allocs/op --- dialect/psql/dialect/derive.go | 34 ---------------------------------- dialect/psql/update.go | 4 ---- 2 files changed, 38 deletions(-) diff --git a/dialect/psql/dialect/derive.go b/dialect/psql/dialect/derive.go index c357aea3..6e3683a9 100644 --- a/dialect/psql/dialect/derive.go +++ b/dialect/psql/dialect/derive.go @@ -85,40 +85,6 @@ func (base *SelectQuery) Derive(queryMods ...bob.Mod[*SelectQuery]) (*SelectQuer return &next, true } -func (base *UpdateQuery) Derive(queryMods ...bob.Mod[*UpdateQuery]) (*UpdateQuery, bool) { - next := *base - var cloneWith, cloneSet, cloneWhere, cloneReturning, cloneJoins bool - - for _, mod := range queryMods { - switch m := mod.(type) { - case mods.Recursive[*UpdateQuery]: - next.With.Recursive = bool(m) - case CTEChain[*UpdateQuery]: - appendDerived[bob.Expression](&next.With.CTEs, base.With.CTEs, &cloneWith, m()) - case mods.TargetOnly[*UpdateQuery]: - next.Only = bool(m) - case mods.TargetTable[*UpdateQuery]: - next.Table = cloneTableRef(clause.TableRef(m)) - case mods.SetExprs[*UpdateQuery]: - appendDerived(&next.Set.Set, base.Set.Set, &cloneSet, []any(m)...) - case mods.Where[*UpdateQuery]: - appendDerived[any](&next.Where.Conditions, base.Where.Conditions, &cloneWhere, m.E) - case mods.Returning[*UpdateQuery]: - appendDerived(&next.Returning.Expressions, base.Returning.Expressions, &cloneReturning, []any(m)...) - case FromChain[*UpdateQuery]: - next.TableRef = cloneTableRef(m()) - case mods.Join[*UpdateQuery]: - appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, clause.Join(m)) - case CrossJoinChain[*UpdateQuery]: - appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, m()) - default: - return nil, false - } - } - - return &next, true -} - func (base *DeleteQuery) Derive(queryMods ...bob.Mod[*DeleteQuery]) (*DeleteQuery, bool) { next := *base var cloneWith, cloneWhere, cloneReturning, cloneJoins bool diff --git a/dialect/psql/update.go b/dialect/psql/update.go index 14b281d6..3e519b65 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -10,10 +10,6 @@ type UpdateQuery struct { } func (q UpdateQuery) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { - if next, ok := q.Expression.Derive(queryMods...); ok { - q.Expression = next - return q - } q.BaseQuery = q.BaseQuery.Apply(queryMods...) return q } From 6446a55d789415f16cfa8f693984ae35f81ee65d Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 12:02:35 -0400 Subject: [PATCH 30/33] refactor(psql): simplify view query wrapper Benchmarks:\n- BenchmarkBaseQueryImmutableNativeHotPath: 2726 ns/op, 3149 B/op, 35 allocs/op\n- BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath: 5769 ns/op, 6315 B/op, 62 allocs/op\n- BenchmarkUpdateQueryImmutableNativeHotPath: 2477 ns/op, 2576 B/op, 43 allocs/op\n- BenchmarkDeleteQueryImmutableNativeHotPath: 1856 ns/op, 2156 B/op, 27 allocs/op\n- BenchmarkInsertQueryImmutableNativeHotPath: 1375 ns/op, 1720 B/op, 18 allocs/op --- dialect/psql/immutable_select_test.go | 4 +-- dialect/psql/view.go | 47 +++++---------------------- dialect/psql/with_regression_test.go | 2 +- 3 files changed, 12 insertions(+), 41 deletions(-) diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 6bcd815c..f1b3012b 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -401,7 +401,7 @@ func BenchmarkViewQueryCountThenPaginateApplyMain(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := q.Query.AsCount().Build(ctx); err != nil { + if _, _, err := q.AsCount().Build(ctx); err != nil { b.Fatal(err) } @@ -426,7 +426,7 @@ func BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath(b *testing.B) { sm.Where(Quote("id").GT(Arg(0))), ) - if _, _, err := q.Query.AsCount().Build(ctx); err != nil { + if _, _, err := q.AsCount().Build(ctx); err != nil { b.Fatal(err) } diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 65b5395b..fc2958ab 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -3,7 +3,6 @@ package psql import ( "context" "fmt" - "io" "reflect" "github.com/stephenafamo/bob" @@ -81,12 +80,12 @@ func (v *View[T, Tslice, C]) Alias() string { // Starts a select query func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Tslice] { q := &ViewQuery[T, Tslice]{ - Query: Select(sm.From(v.NameAs())), - Scanner: v.scanner, - Hooks: &v.SelectQueryHooks, + SelectQuery: Select(sm.From(v.NameAs())), + Scanner: v.scanner, + Hooks: &v.SelectQueryHooks, } - q.Query.Expression.AppendContextualModFunc( + q.SelectQuery.Expression.AppendContextualModFunc( func(ctx context.Context, q *dialect.SelectQuery) (context.Context, error) { if len(q.SelectList.Columns) == 0 { q.AppendSelect(v.Columns) @@ -99,14 +98,14 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * } type ViewQuery[T any, Ts ~[]T] struct { - Query SelectQuery + SelectQuery Scanner scan.Mapper[T] Hooks *bob.Hooks[*SelectQuery, bob.SkipQueryHooksKey] } func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { next := *q - next.Query = next.Query.Apply(queryMods...) + next.SelectQuery = next.SelectQuery.Apply(queryMods...) return &next } @@ -114,33 +113,13 @@ func (q *ViewQuery[T, Ts]) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) *Vi return q.With(queryMods...) } -func (q *ViewQuery[T, Ts]) Type() bob.QueryType { - return q.Query.Type() -} - -func (q *ViewQuery[T, Ts]) Build(ctx context.Context) (string, []any, error) { - return q.Query.Build(ctx) -} - -func (q *ViewQuery[T, Ts]) BuildN(ctx context.Context, start int) (string, []any, error) { - return q.Query.BuildN(ctx, start) -} - -func (q *ViewQuery[T, Ts]) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { - return q.Query.WriteQuery(ctx, w, start) -} - -func (q *ViewQuery[T, Ts]) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dialect, start int) ([]any, error) { - return q.Query.WriteSQL(ctx, w, d, start) -} - // Count the number of matching rows func (q *ViewQuery[T, Tslice]) Count(ctx context.Context, exec bob.Executor) (int64, error) { ctx, err := q.RunHooks(ctx, exec) if err != nil { return 0, err } - sql, args, err := q.Query.AsCount().Build(ctx) + sql, args, err := q.AsCount().Build(ctx) if err != nil { return 0, err } @@ -170,7 +149,7 @@ func (q *ViewQuery[T, Ts]) Each(ctx context.Context, exec bob.Executor) (func(fu } func (q *ViewQuery[T, Ts]) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { - ctx, err := q.Query.RunHooks(ctx, exec) + ctx, err := q.SelectQuery.RunHooks(ctx, exec) if err != nil { return ctx, err } @@ -179,13 +158,5 @@ func (q *ViewQuery[T, Ts]) RunHooks(ctx context.Context, exec bob.Executor) (con return ctx, nil } - return q.Hooks.RunHooks(ctx, exec, &q.Query) -} - -func (q *ViewQuery[T, Ts]) GetLoaders() []bob.Loader { - return q.Query.GetLoaders() -} - -func (q *ViewQuery[T, Ts]) GetMapperMods() []scan.MapperMod { - return q.Query.GetMapperMods() + return q.Hooks.RunHooks(ctx, exec, &q.SelectQuery) } diff --git a/dialect/psql/with_regression_test.go b/dialect/psql/with_regression_test.go index d743dc18..159a5089 100644 --- a/dialect/psql/with_regression_test.go +++ b/dialect/psql/with_regression_test.go @@ -92,7 +92,7 @@ func TestViewQueryWithRegression(t *testing.T) { assertQueriesEqual(t, base, withTestStructView.Query( sm.Where(psql.Quote("id").GT(psql.Arg(0))), )) - assertQueriesEqual(t, derived.Query, withTestStructView.Query( + assertQueriesEqual(t, derived, withTestStructView.Query( sm.Where(psql.Quote("id").GT(psql.Arg(0))), sm.OrderBy("id").Desc(), sm.Limit(10), From a8ea6d87480d9e709109f2701061516777738118 Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 12:05:51 -0400 Subject: [PATCH 31/33] refactor(psql): simplify table returning behavior Benchmarks:\n- BenchmarkBaseQueryImmutableNativeHotPath: 2931 ns/op, 3149 B/op, 35 allocs/op\n- BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath: 5941 ns/op, 6316 B/op, 62 allocs/op\n- BenchmarkUpdateQueryImmutableNativeHotPath: 2601 ns/op, 2575 B/op, 43 allocs/op\n- BenchmarkDeleteQueryImmutableNativeHotPath: 1962 ns/op, 2156 B/op, 27 allocs/op\n- BenchmarkInsertQueryImmutableNativeHotPath: 1507 ns/op, 1720 B/op, 18 allocs/op --- dialect/psql/table.go | 223 +++++++++--------------------------------- mods/mods.go | 2 +- 2 files changed, 48 insertions(+), 177 deletions(-) diff --git a/dialect/psql/table.go b/dialect/psql/table.go index 2b3a5c53..5f51e2c4 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -5,129 +5,25 @@ import ( "reflect" "github.com/stephenafamo/bob" - "github.com/stephenafamo/bob/clause" "github.com/stephenafamo/bob/dialect/psql/dialect" + "github.com/stephenafamo/bob/dialect/psql/dm" + "github.com/stephenafamo/bob/dialect/psql/im" "github.com/stephenafamo/bob/dialect/psql/mm" + "github.com/stephenafamo/bob/dialect/psql/um" "github.com/stephenafamo/bob/expr" "github.com/stephenafamo/bob/internal" "github.com/stephenafamo/bob/internal/mappings" - bobmods "github.com/stephenafamo/bob/mods" "github.com/stephenafamo/bob/orm" ) type ( - setter[T any] = orm.Setter[T, *dialect.InsertQuery, *dialect.UpdateQuery] - ormMergeQuery[T any, Tslice ~[]T] = orm.Query[*dialect.MergeQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] + setter[T any] = orm.Setter[T, *dialect.InsertQuery, *dialect.UpdateQuery] + ormInsertQuery[T any, Tslice ~[]T] = orm.Query[*dialect.InsertQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] + ormUpdateQuery[T any, Tslice ~[]T] = orm.Query[*dialect.UpdateQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] + ormDeleteQuery[T any, Tslice ~[]T] = orm.Query[*dialect.DeleteQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] + ormMergeQuery[T any, Tslice ~[]T] = orm.Query[*dialect.MergeQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] ) -type ormInsertQuery[T any, Tslice ~[]T] struct { - orm.Query[*dialect.InsertQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - defaultReturning bob.Expression -} - -type ormUpdateQuery[T any, Tslice ~[]T] struct { - orm.Query[*dialect.UpdateQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - defaultReturning bob.Expression -} - -type ormDeleteQuery[T any, Tslice ~[]T] struct { - orm.Query[*dialect.DeleteQuery, T, Tslice, bob.SliceTransformer[T, Tslice]] - defaultReturning bob.Expression -} - -func (q ormInsertQuery[T, Tslice]) clone() ormInsertQuery[T, Tslice] { - return ormInsertQuery[T, Tslice]{ - Query: q.Query.Clone(), - defaultReturning: q.defaultReturning, - } -} - -func (q *ormInsertQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { - if q == nil { - return nil - } - - next := q.clone() - applyTableQueryMods(next.Expression, next.defaultReturning, func(query *dialect.InsertQuery) *clause.Returning { - return &query.Returning - }, queryMods...) - return &next -} - -func (q *ormInsertQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { - return q.With(queryMods...) -} - -func (q ormUpdateQuery[T, Tslice]) clone() ormUpdateQuery[T, Tslice] { - return ormUpdateQuery[T, Tslice]{ - Query: q.Query.Clone(), - defaultReturning: q.defaultReturning, - } -} - -func (q *ormUpdateQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { - if q == nil { - return nil - } - - next := q.clone() - applyTableQueryMods(next.Expression, next.defaultReturning, func(query *dialect.UpdateQuery) *clause.Returning { - return &query.Returning - }, queryMods...) - return &next -} - -func (q *ormUpdateQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { - return q.With(queryMods...) -} - -func (q ormDeleteQuery[T, Tslice]) clone() ormDeleteQuery[T, Tslice] { - return ormDeleteQuery[T, Tslice]{ - Query: q.Query.Clone(), - defaultReturning: q.defaultReturning, - } -} - -func (q *ormDeleteQuery[T, Tslice]) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { - if q == nil { - return nil - } - - next := q.clone() - applyTableQueryMods(next.Expression, next.defaultReturning, func(query *dialect.DeleteQuery) *clause.Returning { - return &query.Returning - }, queryMods...) - return &next -} - -func (q *ormDeleteQuery[T, Tslice]) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { - return q.With(queryMods...) -} - -func applyTableQueryMods[Q interface{ AppendReturning(...any) }](query Q, defaultReturning bob.Expression, getReturning func(Q) *clause.Returning, queryMods ...bob.Mod[Q]) { - if hasExplicitReturning(queryMods...) && hasOnlyDefaultReturning(getReturning(query).Expressions, defaultReturning) { - getReturning(query).Expressions = nil - } - - for _, mod := range queryMods { - mod.Apply(query) - } -} - -func hasExplicitReturning[Q interface{ AppendReturning(...any) }](queryMods ...bob.Mod[Q]) bool { - for _, mod := range queryMods { - if _, ok := mod.(bobmods.Returning[Q]); ok { - return true - } - } - - return false -} - -func hasOnlyDefaultReturning(expressions []any, defaultReturning bob.Expression) bool { - return len(expressions) == 1 && reflect.DeepEqual(expressions[0], defaultReturning) -} - func NewTable[T any, Tset setter[T], C bob.Expression](schema, tableName string, columns C) *Table[T, []T, Tset, C] { return NewTablex[T, []T, Tset](schema, tableName, columns) } @@ -179,48 +75,66 @@ func (t *Table[T, Tslice, Tset, C]) PrimaryKey() expr.ColumnsExpr { // Starts an insert query for this table func (t *Table[T, Tslice, Tset, C]) Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) *ormInsertQuery[T, Tslice] { q := &ormInsertQuery[T, Tslice]{ - Query: orm.Query[*dialect.InsertQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ - ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ - BaseQuery: insertTableBaseQuery(t.NameAs(), t.nonGeneratedCols, t.Columns), - Hooks: &t.InsertQueryHooks, - }, - Scanner: t.scanner, + ExecQuery: orm.ExecQuery[*dialect.InsertQuery]{ + BaseQuery: Insert(im.Into(t.NameAs(), t.nonGeneratedCols...)).BaseQuery, + Hooks: &t.InsertQueryHooks, }, - defaultReturning: t.Columns, + Scanner: t.scanner, } + q.Expression.AppendContextualModFunc( + func(ctx context.Context, q *dialect.InsertQuery) (context.Context, error) { + if !q.HasReturning() { + q.AppendReturning(t.Columns) + } + return ctx, nil + }, + ) + return q.Apply(queryMods...) } // Starts an Update query for this table func (t *Table[T, Tslice, Tset, C]) Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) *ormUpdateQuery[T, Tslice] { q := &ormUpdateQuery[T, Tslice]{ - Query: orm.Query[*dialect.UpdateQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ - ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ - BaseQuery: updateTableBaseQuery(t.NameAs(), t.Columns), - Hooks: &t.UpdateQueryHooks, - }, - Scanner: t.scanner, + ExecQuery: orm.ExecQuery[*dialect.UpdateQuery]{ + BaseQuery: Update(um.Table(t.NameAs())).BaseQuery, + Hooks: &t.UpdateQueryHooks, }, - defaultReturning: t.Columns, + Scanner: t.scanner, } + q.Expression.AppendContextualModFunc( + func(ctx context.Context, q *dialect.UpdateQuery) (context.Context, error) { + if !q.HasReturning() { + q.AppendReturning(t.Columns) + } + return ctx, nil + }, + ) + return q.Apply(queryMods...) } // Starts a Delete query for this table func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) *ormDeleteQuery[T, Tslice] { q := &ormDeleteQuery[T, Tslice]{ - Query: orm.Query[*dialect.DeleteQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ - ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ - BaseQuery: deleteTableBaseQuery(t.NameAs(), t.Columns), - Hooks: &t.DeleteQueryHooks, - }, - Scanner: t.scanner, + ExecQuery: orm.ExecQuery[*dialect.DeleteQuery]{ + BaseQuery: Delete(dm.From(t.NameAs())).BaseQuery, + Hooks: &t.DeleteQueryHooks, }, - defaultReturning: t.Columns, + Scanner: t.scanner, } + q.Expression.AppendContextualModFunc( + func(ctx context.Context, q *dialect.DeleteQuery) (context.Context, error) { + if !q.HasReturning() { + q.AppendReturning(t.Columns) + } + return ctx, nil + }, + ) + return q.Apply(queryMods...) } @@ -252,46 +166,3 @@ func (t *Table[T, Tslice, Tset, C]) Merge(queryMods ...bob.Mod[*dialect.MergeQue return q } - -func insertTableBaseQuery(name any, nonGeneratedCols []string, returning bob.Expression) bob.BaseQuery[*dialect.InsertQuery] { - base := bob.BaseQuery[*dialect.InsertQuery]{ - Expression: &dialect.InsertQuery{ - TableRef: clause.TableRef{ - Expression: name, - Columns: nonGeneratedCols, - }, - }, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeInsert, - } - base.Expression.AppendReturning(returning) - return base -} - -func updateTableBaseQuery(name any, returning bob.Expression) bob.BaseQuery[*dialect.UpdateQuery] { - base := bob.BaseQuery[*dialect.UpdateQuery]{ - Expression: &dialect.UpdateQuery{ - Table: clause.TableRef{ - Expression: name, - }, - }, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeUpdate, - } - base.Expression.AppendReturning(returning) - return base -} - -func deleteTableBaseQuery(name any, returning bob.Expression) bob.BaseQuery[*dialect.DeleteQuery] { - base := bob.BaseQuery[*dialect.DeleteQuery]{ - Expression: &dialect.DeleteQuery{ - Table: clause.TableRef{ - Expression: name, - }, - }, - Dialect: dialect.Dialect, - QueryType: bob.QueryTypeDelete, - } - base.Expression.AppendReturning(returning) - return base -} diff --git a/mods/mods.go b/mods/mods.go index 3f67be1a..badb9fdd 100644 --- a/mods/mods.go +++ b/mods/mods.go @@ -50,7 +50,7 @@ type TargetTable[Q interface { func (t TargetTable[Q]) Apply(q Q) { q.SetTargetTable(t.Expression) - if t.Alias != "" { + if t.Alias != "" || len(t.Columns) > 0 { q.SetTargetTableAlias(t.Alias, t.Columns...) } } From 46ec3ae855f5539578f51faf8214c5ebe473a13b Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Tue, 21 Apr 2026 12:32:34 -0400 Subject: [PATCH 32/33] fix(psql): address immutable query review feedback --- dialect/psql/dialect/delete.go | 3 ++- dialect/psql/dialect/derive.go | 1 + dialect/psql/dialect/update.go | 3 ++- dialect/psql/dialect/writer.go | 2 +- dialect/psql/immutable_select_test.go | 32 ++++++++++++++++++++++-- dialect/psql/immutable_write_test.go | 35 +++++++++++++++++++++++++++ dialect/psql/view.go | 10 +++++--- orm/query.go | 8 ++++-- 8 files changed, 84 insertions(+), 10 deletions(-) diff --git a/dialect/psql/dialect/delete.go b/dialect/psql/dialect/delete.go index 86df4fa3..3cc1cd3d 100644 --- a/dialect/psql/dialect/delete.go +++ b/dialect/psql/dialect/delete.go @@ -24,7 +24,8 @@ type DeleteQuery struct { } func (d *DeleteQuery) SetTargetOnly(only bool) { - d.Table.SetOnly(only) + d.Only = only + d.Table.SetOnly(false) } func (d *DeleteQuery) SetTargetTable(table any) { diff --git a/dialect/psql/dialect/derive.go b/dialect/psql/dialect/derive.go index 6e3683a9..c9435069 100644 --- a/dialect/psql/dialect/derive.go +++ b/dialect/psql/dialect/derive.go @@ -97,6 +97,7 @@ func (base *DeleteQuery) Derive(queryMods ...bob.Mod[*DeleteQuery]) (*DeleteQuer appendDerived[bob.Expression](&next.With.CTEs, base.With.CTEs, &cloneWith, m()) case mods.TargetOnly[*DeleteQuery]: next.Only = bool(m) + next.Table.Only = false case mods.TargetTable[*DeleteQuery]: next.Table = cloneTableRef(clause.TableRef(m)) case mods.Where[*DeleteQuery]: diff --git a/dialect/psql/dialect/update.go b/dialect/psql/dialect/update.go index 2f02cc57..03c330b3 100644 --- a/dialect/psql/dialect/update.go +++ b/dialect/psql/dialect/update.go @@ -25,7 +25,8 @@ type UpdateQuery struct { } func (u *UpdateQuery) SetTargetOnly(only bool) { - u.Table.SetOnly(only) + u.Only = only + u.Table.SetOnly(false) } func (u *UpdateQuery) SetTargetTable(table any) { diff --git a/dialect/psql/dialect/writer.go b/dialect/psql/dialect/writer.go index 612669da..ba34fc49 100644 --- a/dialect/psql/dialect/writer.go +++ b/dialect/psql/dialect/writer.go @@ -77,7 +77,7 @@ func (w *queryWriter) writeAny(value any) error { case uint64: _, _ = w.w.WriteString(strconv.FormatUint(v, 10)) case sql.NamedArg: - return fmt.Errorf("named args are not supported by psql dialect") + return bob.ErrNoNamedArgs case bob.Expression: return w.writeExpression(v) default: diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index f1b3012b..86f8da8c 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -2,9 +2,12 @@ package psql import ( "context" + "database/sql" + "errors" "testing" "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/dialect/psql/dialect" "github.com/stephenafamo/bob/dialect/psql/sm" "github.com/stephenafamo/bob/expr" ) @@ -74,6 +77,14 @@ func TestImmutableViewQueryWithDoesNotMutateOriginal(t *testing.T) { } } +func TestViewQueryWithNilReceiver(t *testing.T) { + var q *ViewQuery[*someStruct, []*someStruct] + + if got := q.With(sm.Where(Quote("id").EQ(Arg(1)))); got != nil { + t.Fatalf("expected nil view query, got %#v", got) + } +} + func TestImmutableSelectQueryApplyDoesNotMutateOriginal(t *testing.T) { base := Select( sm.Columns("id"), @@ -243,8 +254,12 @@ func TestViewSelectQueryHooksUseImmutableSelectQuery(t *testing.T) { view := NewView[*someStruct, bob.Expression]("public", "some_struct", expr.ColsForStruct[someStruct]("some_struct")) var hookSQL string - view.SelectQueryHooks.AppendHooks(func(ctx context.Context, exec bob.Executor, q *SelectQuery) (context.Context, error) { - sql, _, err := q.Build(ctx) + view.SelectQueryHooks.AppendHooks(func(ctx context.Context, exec bob.Executor, q *dialect.SelectQuery) (context.Context, error) { + sql, _, err := bob.BaseQuery[*dialect.SelectQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeSelect, + }.Build(ctx) if err != nil { return ctx, err } @@ -272,6 +287,19 @@ func TestViewSelectQueryHooksUseImmutableSelectQuery(t *testing.T) { } } +func TestImmutableSelectQueryBuildReturnsErrNoNamedArgs(t *testing.T) { + _, _, err := Select( + sm.Columns(sql.Named("id", 1)), + sm.From("users"), + ).Build(t.Context()) + if err == nil { + t.Fatal("expected named arg error") + } + if !errors.Is(err, bob.ErrNoNamedArgs) { + t.Fatalf("expected ErrNoNamedArgs, got %v", err) + } +} + func TestImmutableSelectQueryApplySupportsCommonDerivedMods(t *testing.T) { base := Select( sm.Columns("users.id", "users.name"), diff --git a/dialect/psql/immutable_write_test.go b/dialect/psql/immutable_write_test.go index 5e8afa7c..eba01b8c 100644 --- a/dialect/psql/immutable_write_test.go +++ b/dialect/psql/immutable_write_test.go @@ -147,6 +147,41 @@ func TestDeleteApplyDoesNotMutateOriginal(t *testing.T) { } } +func TestUpdateApplyDoesNotDuplicateOnly(t *testing.T) { + base := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Drama"), + um.Only(), + ) + + derived := base.Apply(um.Only()) + + sql, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if sql != "UPDATE ONLY films SET\n\"kind\" = $1" { + t.Fatalf("unexpected update SQL: %#v", sql) + } +} + +func TestDeleteApplyDoesNotDuplicateOnly(t *testing.T) { + base := Delete( + dm.From("films"), + dm.Only(), + ) + + derived := base.Apply(dm.Only()) + + sql, _, err := derived.Build(t.Context()) + if err != nil { + t.Fatal(err) + } + if sql != "DELETE FROM ONLY films" { + t.Fatalf("unexpected delete SQL: %#v", sql) + } +} + func TestInsertApplyDoesNotMutateOriginal(t *testing.T) { base := Insert( im.Into("films"), diff --git a/dialect/psql/view.go b/dialect/psql/view.go index fc2958ab..17d9ab0e 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -57,7 +57,7 @@ type View[T any, Tslice ~[]T, C bob.Expression] struct { Columns C AfterSelectHooks bob.Hooks[Tslice, bob.SkipModelHooksKey] - SelectQueryHooks bob.Hooks[*SelectQuery, bob.SkipQueryHooksKey] + SelectQueryHooks bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] } func (v *View[T, Tslice, C]) Name() Expression { @@ -100,10 +100,14 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * type ViewQuery[T any, Ts ~[]T] struct { SelectQuery Scanner scan.Mapper[T] - Hooks *bob.Hooks[*SelectQuery, bob.SkipQueryHooksKey] + Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] } func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { + if q == nil { + return nil + } + next := *q next.SelectQuery = next.SelectQuery.Apply(queryMods...) return &next @@ -158,5 +162,5 @@ func (q *ViewQuery[T, Ts]) RunHooks(ctx context.Context, exec bob.Executor) (con return ctx, nil } - return q.Hooks.RunHooks(ctx, exec, &q.SelectQuery) + return q.Hooks.RunHooks(ctx, exec, q.SelectQuery.Expression) } diff --git a/orm/query.go b/orm/query.go index c6ab7f80..3909b652 100644 --- a/orm/query.go +++ b/orm/query.go @@ -28,7 +28,9 @@ func (q *ExecQuery[Q]) With(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { } next := q.Clone() - next.BaseQuery = next.BaseQuery.Apply(queryMods...) + for _, mod := range queryMods { + mod.Apply(next.BaseQuery.Expression) + } return &next } @@ -79,7 +81,9 @@ func (q *Query[Q, T, Ts, Tr]) With(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] } next := q.Clone() - next.BaseQuery = next.BaseQuery.Apply(queryMods...) + for _, mod := range queryMods { + mod.Apply(next.BaseQuery.Expression) + } return &next } From d6e79df730120f6607d8677094c73db6c585d15f Mon Sep 17 00:00:00 2001 From: Jay Patel <36803168+jay-babu@users.noreply.github.com> Date: Sun, 3 May 2026 17:23:12 +0000 Subject: [PATCH 33/33] refactor: remove query With methods --- dialect/mysql/table.go | 8 ++------ dialect/psql/delete.go | 6 +----- dialect/psql/dialect/derive.go | 2 +- dialect/psql/dialect/insert.go | 4 ++-- dialect/psql/immutable_select_test.go | 20 ++++++++++---------- dialect/psql/immutable_write_test.go | 18 +++++++++--------- dialect/psql/insert.go | 6 +----- dialect/psql/select.go | 6 +----- dialect/psql/table.go | 4 +--- dialect/psql/table_test.go | 10 +++++----- dialect/psql/update.go | 6 +----- dialect/psql/view.go | 6 +----- dialect/psql/with_regression_test.go | 24 ++++++++++++------------ orm/query.go | 12 ++---------- query.go | 6 +----- query_immutable_test.go | 4 ++-- 16 files changed, 52 insertions(+), 90 deletions(-) diff --git a/dialect/mysql/table.go b/dialect/mysql/table.go index b6861987..27fe4edb 100644 --- a/dialect/mysql/table.go +++ b/dialect/mysql/table.go @@ -120,20 +120,16 @@ type insertQuery[T any, Ts ~[]T, Tset setter[T], C bob.Expression] struct { table *Table[T, Ts, Tset, C] } -func (t *insertQuery[T, Ts, Tset, C]) With(queryMods ...bob.Mod[*dialect.InsertQuery]) *insertQuery[T, Ts, Tset, C] { +func (t *insertQuery[T, Ts, Tset, C]) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) *insertQuery[T, Ts, Tset, C] { if t == nil { return nil } next := *t - next.ExecQuery = *t.ExecQuery.With(queryMods...) + next.ExecQuery = *t.ExecQuery.Apply(queryMods...) return &next } -func (t *insertQuery[T, Ts, Tset, C]) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) *insertQuery[T, Ts, Tset, C] { - return t.With(queryMods...) -} - // Insert One Row // NOTE: Because MySQL does not support RETURNING, this will insert the row and then run a SELECT query // to retrieve the row. diff --git a/dialect/psql/delete.go b/dialect/psql/delete.go index 870000be..d540bce5 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -9,7 +9,7 @@ type DeleteQuery struct { bob.BaseQuery[*dialect.DeleteQuery] } -func (q DeleteQuery) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { +func (q DeleteQuery) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { if next, ok := q.Expression.Derive(queryMods...); ok { q.Expression = next return q @@ -18,10 +18,6 @@ func (q DeleteQuery) With(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuer return q } -func (q DeleteQuery) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { - return q.With(queryMods...) -} - func Delete(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { q := &dialect.DeleteQuery{} for _, mod := range queryMods { diff --git a/dialect/psql/dialect/derive.go b/dialect/psql/dialect/derive.go index c9435069..3b5d91ff 100644 --- a/dialect/psql/dialect/derive.go +++ b/dialect/psql/dialect/derive.go @@ -131,7 +131,7 @@ func (base *InsertQuery) Derive(queryMods ...bob.Mod[*InsertQuery]) (*InsertQuer case mods.TargetTable[*InsertQuery]: next.TableRef = cloneTableRef(clause.TableRef(m)) case mods.Overriding[*InsertQuery]: - next.Overriding = string(m) + next.Overriding = OverridingType(m) case mods.QuerySource[*InsertQuery]: next.Values.Query = m.Query case mods.Returning[*InsertQuery]: diff --git a/dialect/psql/dialect/insert.go b/dialect/psql/dialect/insert.go index 6558b4fc..9eda6eac 100644 --- a/dialect/psql/dialect/insert.go +++ b/dialect/psql/dialect/insert.go @@ -41,7 +41,7 @@ func (i *InsertQuery) SetTargetTableAlias(alias string, columns ...string) { } func (i *InsertQuery) SetOverriding(overriding string) { - i.Overriding = overriding + i.Overriding = OverridingType(overriding) } func (i *InsertQuery) SetQuery(q bob.Query) { @@ -85,7 +85,7 @@ func (i InsertQuery) WriteSQL(ctx context.Context, w io.StringWriter, d bob.Dial if i.Overriding != "" { _, _ = w.WriteString("\nOVERRIDING ") - _, _ = w.WriteString(i.Overriding) + _, _ = w.WriteString(string(i.Overriding)) _, _ = w.WriteString(" VALUE") } diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go index 86f8da8c..c5c0c279 100644 --- a/dialect/psql/immutable_select_test.go +++ b/dialect/psql/immutable_select_test.go @@ -12,13 +12,13 @@ import ( "github.com/stephenafamo/bob/expr" ) -func TestImmutableSelectQueryWithDoesNotMutateOriginal(t *testing.T) { +func TestImmutableSelectQueryApplyDoesNotMutateOriginalFromLegacyWithCase(t *testing.T) { base := Select( sm.Columns("id"), sm.From("users"), - ).With() + ).Apply() - derived := base.With( + derived := base.Apply( sm.OrderBy("id").Desc(), sm.Limit(10), sm.Offset(20), @@ -45,12 +45,12 @@ func TestImmutableSelectQueryWithDoesNotMutateOriginal(t *testing.T) { } } -func TestImmutableViewQueryWithDoesNotMutateOriginal(t *testing.T) { +func TestImmutableViewQueryApplyDoesNotMutateOriginalFromLegacyWithCase(t *testing.T) { base := someStructView.Query( sm.Where(Quote("id").GT(Arg(0))), - ).With() + ).Apply() - derived := base.With( + derived := base.Apply( sm.OrderBy("id").Desc(), sm.Limit(10), sm.Offset(20), @@ -77,10 +77,10 @@ func TestImmutableViewQueryWithDoesNotMutateOriginal(t *testing.T) { } } -func TestViewQueryWithNilReceiver(t *testing.T) { +func TestViewQueryApplyNilReceiver(t *testing.T) { var q *ViewQuery[*someStruct, []*someStruct] - if got := q.With(sm.Where(Quote("id").EQ(Arg(1)))); got != nil { + if got := q.Apply(sm.Where(Quote("id").EQ(Arg(1)))); got != nil { t.Fatalf("expected nil view query, got %#v", got) } } @@ -408,7 +408,7 @@ func BenchmarkBaseQueryImmutableNativeHotPath(b *testing.B) { sm.From("users"), sm.Where(Quote("tenant_id").EQ(Arg(42))), ) - derived := q.With( + derived := q.Apply( sm.OrderBy("id").Desc(), sm.Limit(10), sm.Offset(20), @@ -458,7 +458,7 @@ func BenchmarkViewQueryCountThenPaginateImmutableNativeHotPath(b *testing.B) { b.Fatal(err) } - derived := q.With( + derived := q.Apply( sm.OrderBy("id").Desc(), sm.Limit(10), sm.Offset(20), diff --git a/dialect/psql/immutable_write_test.go b/dialect/psql/immutable_write_test.go index eba01b8c..b68f62a9 100644 --- a/dialect/psql/immutable_write_test.go +++ b/dialect/psql/immutable_write_test.go @@ -10,13 +10,13 @@ import ( "github.com/stephenafamo/bob/dialect/psql/um" ) -func TestUpdateWithDoesNotMutateOriginal(t *testing.T) { +func TestUpdateApplyDoesNotMutateOriginalFromLegacyWithCase(t *testing.T) { base := Update( um.Table("films"), um.SetCol("kind").ToArg("Dramatic"), ) - derived := base.With( + derived := base.Apply( um.Where(Quote("kind").EQ(Arg("Drama"))), um.Returning("id"), ) @@ -38,12 +38,12 @@ func TestUpdateWithDoesNotMutateOriginal(t *testing.T) { } } -func TestDeleteWithDoesNotMutateOriginal(t *testing.T) { +func TestDeleteApplyDoesNotMutateOriginalFromLegacyWithCase(t *testing.T) { base := Delete( dm.From("films"), ) - derived := base.With( + derived := base.Apply( dm.Where(Quote("kind").EQ(Arg("Drama"))), dm.Returning("id"), ) @@ -65,13 +65,13 @@ func TestDeleteWithDoesNotMutateOriginal(t *testing.T) { } } -func TestInsertWithDoesNotMutateOriginal(t *testing.T) { +func TestInsertApplyDoesNotMutateOriginalFromLegacyWithCase(t *testing.T) { base := Insert( im.Into("films"), im.Values(Arg("UA502", "Bananas")), ) - derived := base.With( + derived := base.Apply( im.Returning("id"), ) @@ -427,7 +427,7 @@ func BenchmarkUpdateQueryImmutableNativeHotPath(b *testing.B) { um.Table("films"), um.SetCol("kind").ToArg("Dramatic"), ) - derived := q.With( + derived := q.Apply( um.Where(Quote("kind").EQ(Arg("Drama"))), um.Returning("id"), ) @@ -465,7 +465,7 @@ func BenchmarkDeleteQueryImmutableNativeHotPath(b *testing.B) { q := Delete( dm.From("films"), ) - derived := q.With( + derived := q.Apply( dm.Where(Quote("kind").EQ(Arg("Drama"))), dm.Returning("id"), ) @@ -504,7 +504,7 @@ func BenchmarkInsertQueryImmutableNativeHotPath(b *testing.B) { im.Into("films"), im.Values(Arg("UA502", "Bananas")), ) - derived := q.With( + derived := q.Apply( im.Returning("id"), ) diff --git a/dialect/psql/insert.go b/dialect/psql/insert.go index d2c10d63..099419da 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -9,7 +9,7 @@ type InsertQuery struct { bob.BaseQuery[*dialect.InsertQuery] } -func (q InsertQuery) With(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { +func (q InsertQuery) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { if next, ok := q.Expression.Derive(queryMods...); ok { q.Expression = next return q @@ -18,10 +18,6 @@ func (q InsertQuery) With(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuer return q } -func (q InsertQuery) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { - return q.With(queryMods...) -} - func Insert(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { q := &dialect.InsertQuery{} for _, mod := range queryMods { diff --git a/dialect/psql/select.go b/dialect/psql/select.go index 1e9c5e74..cf1a4844 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -9,7 +9,7 @@ type SelectQuery struct { bob.BaseQuery[*dialect.SelectQuery] } -func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { +func (q SelectQuery) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { if next, ok := q.Expression.Derive(queryMods...); ok { q.Expression = next return q @@ -18,10 +18,6 @@ func (q SelectQuery) With(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuer return q } -func (q SelectQuery) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { - return q.With(queryMods...) -} - func (q SelectQuery) AsCount() SelectQuery { next := q.Clone() next.Expression.SetSelect("count(1)") diff --git a/dialect/psql/table.go b/dialect/psql/table.go index 5f51e2c4..72a6b488 100644 --- a/dialect/psql/table.go +++ b/dialect/psql/table.go @@ -162,7 +162,5 @@ func (t *Table[T, Tslice, Tset, C]) Merge(queryMods ...bob.Mod[*dialect.MergeQue }, ) - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } diff --git a/dialect/psql/table_test.go b/dialect/psql/table_test.go index 8900e3e4..ca63e397 100644 --- a/dialect/psql/table_test.go +++ b/dialect/psql/table_test.go @@ -180,7 +180,7 @@ func TestTableUpdateAdditionalExplicitReturningAppends(t *testing.T) { um.Where(Quote("id").EQ(Arg(1))), ) - q := base.With(um.Returning("id")).With(um.Returning("email")) + q := base.Apply(um.Returning("id")).Apply(um.Returning("email")) sql, args, err := q.Build(t.Context()) if err != nil { @@ -240,7 +240,7 @@ func TestTableInsertAdditionalExplicitReturningAppends(t *testing.T) { im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), ) - q := base.With(im.Returning("id")).With(im.Returning("email")) + q := base.Apply(im.Returning("id")).Apply(im.Returning("email")) sql, args, err := q.Build(t.Context()) if err != nil { @@ -300,7 +300,7 @@ func TestTableDeleteAdditionalExplicitReturningAppends(t *testing.T) { dm.Where(Quote("id").EQ(Arg(1))), ) - q := base.With(dm.Returning("id")).With(dm.Returning("email")) + q := base.Apply(dm.Returning("id")).Apply(dm.Returning("email")) sql, args, err := q.Build(t.Context()) if err != nil { @@ -346,12 +346,12 @@ func TestTableUpdateApplyDoesNotMutateOriginal(t *testing.T) { } } -func TestTableInsertWithDoesNotMutateOriginal(t *testing.T) { +func TestTableInsertApplyDoesNotMutateOriginal(t *testing.T) { base := userTable.Insert( im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), ) - derived := base.With(im.Returning("id")) + derived := base.Apply(im.Returning("id")) baseSQL, _, err := base.Build(t.Context()) if err != nil { diff --git a/dialect/psql/update.go b/dialect/psql/update.go index 3e519b65..ee63521b 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -9,15 +9,11 @@ type UpdateQuery struct { bob.BaseQuery[*dialect.UpdateQuery] } -func (q UpdateQuery) With(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { +func (q UpdateQuery) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { q.BaseQuery = q.BaseQuery.Apply(queryMods...) return q } -func (q UpdateQuery) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { - return q.With(queryMods...) -} - func Update(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { q := &dialect.UpdateQuery{} for _, mod := range queryMods { diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 17d9ab0e..8a4270dd 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -103,7 +103,7 @@ type ViewQuery[T any, Ts ~[]T] struct { Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] } -func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { +func (q *ViewQuery[T, Ts]) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { if q == nil { return nil } @@ -113,10 +113,6 @@ func (q *ViewQuery[T, Ts]) With(queryMods ...bob.Mod[*dialect.SelectQuery]) *Vie return &next } -func (q *ViewQuery[T, Ts]) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { - return q.With(queryMods...) -} - // Count the number of matching rows func (q *ViewQuery[T, Tslice]) Count(ctx context.Context, exec bob.Executor) (int64, error) { ctx, err := q.RunHooks(ctx, exec) diff --git a/dialect/psql/with_regression_test.go b/dialect/psql/with_regression_test.go index 159a5089..192d7624 100644 --- a/dialect/psql/with_regression_test.go +++ b/dialect/psql/with_regression_test.go @@ -25,7 +25,7 @@ var withTestStructView = psql.NewView[*withTestStruct, bob.Expression]( expr.ColsForStruct[withTestStruct]("with_test_struct"), ) -func TestSelectWithRegression(t *testing.T) { +func TestSelectApplyRegression(t *testing.T) { t.Run("native path matches direct construction", func(t *testing.T) { base := psql.Select( sm.Columns("id", "name"), @@ -33,7 +33,7 @@ func TestSelectWithRegression(t *testing.T) { sm.Where(psql.Quote("tenant_id").EQ(psql.Arg(42))), ) - derived := base.With( + derived := base.Apply( sm.OrderBy("id").Desc(), sm.Limit(10), sm.Offset(20), @@ -60,7 +60,7 @@ func TestSelectWithRegression(t *testing.T) { sm.From("users"), ) - derived := base.With( + derived := base.Apply( sm.LeftJoin("teams").Using("id"), sm.ForUpdate("users").SkipLocked(), ) @@ -78,12 +78,12 @@ func TestSelectWithRegression(t *testing.T) { }) } -func TestViewQueryWithRegression(t *testing.T) { +func TestViewQueryApplyRegression(t *testing.T) { base := withTestStructView.Query( sm.Where(psql.Quote("id").GT(psql.Arg(0))), ) - derived := base.With( + derived := base.Apply( sm.OrderBy("id").Desc(), sm.Limit(10), sm.Offset(20), @@ -100,13 +100,13 @@ func TestViewQueryWithRegression(t *testing.T) { )) } -func TestUpdateWithRegression(t *testing.T) { +func TestUpdateApplyRegression(t *testing.T) { base := psql.Update( um.Table("films"), um.SetCol("kind").ToArg("Dramatic"), ) - derived := base.With( + derived := base.Apply( um.SetCol("updated_at").To("NOW()"), um.Where(psql.Quote("kind").EQ(psql.Arg("Drama"))), um.Returning("id"), @@ -125,12 +125,12 @@ func TestUpdateWithRegression(t *testing.T) { )) } -func TestDeleteWithRegression(t *testing.T) { +func TestDeleteApplyRegression(t *testing.T) { base := psql.Delete( dm.From("employees"), ) - derived := base.With( + derived := base.Apply( dm.Using("accounts"), dm.Where(psql.Quote("accounts", "name").EQ(psql.Arg("Acme Corporation"))), dm.Where(psql.Quote("employees", "id").EQ(psql.Quote("accounts", "sales_person"))), @@ -149,14 +149,14 @@ func TestDeleteWithRegression(t *testing.T) { )) } -func TestInsertWithRegression(t *testing.T) { +func TestInsertApplyRegression(t *testing.T) { t.Run("native path matches direct construction", func(t *testing.T) { base := psql.Insert( im.Into("films"), im.Values(psql.Arg("UA502", "Bananas")), ) - derived := base.With( + derived := base.Apply( im.Returning("id"), ) @@ -177,7 +177,7 @@ func TestInsertWithRegression(t *testing.T) { im.Values(psql.Arg(8, "Anvil Distribution")), ) - derived := base.With( + derived := base.Apply( im.OnConflict("did").DoUpdate( im.SetExcluded("dname"), im.Where(psql.Quote("d", "zipcode").NE(psql.S("21201"))), diff --git a/orm/query.go b/orm/query.go index 3909b652..8678213d 100644 --- a/orm/query.go +++ b/orm/query.go @@ -22,7 +22,7 @@ func (q ExecQuery[Q]) Clone() ExecQuery[Q] { } } -func (q *ExecQuery[Q]) With(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { +func (q *ExecQuery[Q]) Apply(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { if q == nil { return nil } @@ -34,10 +34,6 @@ func (q *ExecQuery[Q]) With(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { return &next } -func (q *ExecQuery[Q]) Apply(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { - return q.With(queryMods...) -} - func (q ExecQuery[Q]) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { var err error @@ -75,7 +71,7 @@ func (q Query[Q, T, Ts, Tr]) Clone() Query[Q, T, Ts, Tr] { } } -func (q *Query[Q, T, Ts, Tr]) With(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] { +func (q *Query[Q, T, Ts, Tr]) Apply(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] { if q == nil { return nil } @@ -87,10 +83,6 @@ func (q *Query[Q, T, Ts, Tr]) With(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] return &next } -func (q *Query[Q, T, Ts, Tr]) Apply(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] { - return q.With(queryMods...) -} - // First matching row func (q Query[Q, T, Ts, Tr]) One(ctx context.Context, exec bob.Executor) (T, error) { return bob.One(ctx, exec, q, q.Scanner) diff --git a/query.go b/query.go index 767d3df4..59c53deb 100644 --- a/query.go +++ b/query.go @@ -117,7 +117,7 @@ func (b BaseQuery[E]) GetMapperMods() []scan.MapperMod { return nil } -func (b BaseQuery[E]) With(mods ...Mod[E]) BaseQuery[E] { +func (b BaseQuery[E]) Apply(mods ...Mod[E]) BaseQuery[E] { next := b.Clone() for _, mod := range mods { mod.Apply(next.Expression) @@ -125,10 +125,6 @@ func (b BaseQuery[E]) With(mods ...Mod[E]) BaseQuery[E] { return next } -func (b BaseQuery[E]) Apply(mods ...Mod[E]) BaseQuery[E] { - return b.With(mods...) -} - func (b BaseQuery[E]) WriteQuery(ctx context.Context, w io.StringWriter, start int) ([]any, error) { // If it a query, just call its WriteQuery method if e, ok := any(b.Expression).(interface { diff --git a/query_immutable_test.go b/query_immutable_test.go index 44acac0f..07ee5b8a 100644 --- a/query_immutable_test.go +++ b/query_immutable_test.go @@ -27,13 +27,13 @@ func (m appendExprMod) Apply(e *cloneableExpr) { e.parts = append(e.parts, string(m)) } -func TestBaseQueryWithDoesNotMutateOriginal(t *testing.T) { +func TestBaseQueryApplyDoesNotMutateOriginalFromLegacyWithCase(t *testing.T) { base := BaseQuery[*cloneableExpr]{ Expression: &cloneableExpr{parts: []string{"base"}}, QueryType: QueryTypeSelect, } - derived := base.With(appendExprMod("derived")) + derived := base.Apply(appendExprMod("derived")) if !slices.Equal(base.Expression.parts, []string{"base"}) { t.Fatalf("base query changed unexpectedly: %#v", base.Expression.parts)