diff --git a/dialect/mysql/table.go b/dialect/mysql/table.go index 3d32a296..27fe4edb 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,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]) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) *insertQuery[T, Ts, Tset, C] { + if t == nil { + return nil + } + + next := *t + next.ExecQuery = *t.ExecQuery.Apply(queryMods...) + return &next +} + // 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. @@ -287,8 +291,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/psql/delete.go b/dialect/psql/delete.go index cdbdd214..d540bce5 100644 --- a/dialect/psql/delete.go +++ b/dialect/psql/delete.go @@ -5,15 +5,30 @@ 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) Apply(queryMods ...bob.Mod[*dialect.DeleteQuery]) DeleteQuery { + if next, ok := q.Expression.Derive(queryMods...); ok { + q.Expression = next + return q + } + q.BaseQuery = q.BaseQuery.Apply(queryMods...) + return q +} + +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/dialect/clone.go b/dialect/psql/dialect/clone.go new file mode 100644 index 00000000..aa37b1bd --- /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(make([]any, 0, len(values)), 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())) + } +} diff --git a/dialect/psql/dialect/delete.go b/dialect/psql/dialect/delete.go index c32e7673..3cc1cd3d 100644 --- a/dialect/psql/dialect/delete.go +++ b/dialect/psql/dialect/delete.go @@ -23,53 +23,75 @@ type DeleteQuery struct { bob.ContextualModdable[*DeleteQuery] } +func (d *DeleteQuery) SetTargetOnly(only bool) { + d.Only = only + d.Table.SetOnly(false) +} + +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 - 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/derive.go b/dialect/psql/dialect/derive.go new file mode 100644 index 00000000..3b5d91ff --- /dev/null +++ b/dialect/psql/dialect/derive.go @@ -0,0 +1,157 @@ +package dialect + +import ( + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/clause" + "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 + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Recursive[*SelectQuery]: + next.With.Recursive = bool(m) + case CTEChain[*SelectQuery]: + appendDerived[bob.Expression](&next.With.CTEs, base.With.CTEs, &cloneWith, m()) + case mods.Distinct[*SelectQuery]: + next.SetDistinctValues([]any(m)) + case mods.Select[*SelectQuery]: + appendDerived(&next.SelectList.Columns, base.SelectList.Columns, &cloneSelect, []any(m)...) + case mods.Preload[*SelectQuery]: + appendDerived(&next.SelectList.PreloadColumns, base.SelectList.PreloadColumns, &clonePreload, []any(m)...) + case mods.Where[*SelectQuery]: + appendDerived[any](&next.Where.Conditions, base.Where.Conditions, &cloneWhere, m.E) + case mods.GroupBy[*SelectQuery]: + 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]: + appendDerived(&next.Having.Conditions, base.Having.Conditions, &cloneHaving, []any(m)...) + case mods.Limit[*SelectQuery]: + next.Limit.Count = m.Count + case mods.Offset[*SelectQuery]: + next.Offset.Count = m.Count + case mods.Fetch[*SelectQuery]: + next.Fetch = clause.Fetch(m) + case OrderBy[*SelectQuery]: + appendDerived[bob.Expression](&next.OrderBy.Expressions, base.OrderBy.Expressions, &cloneOrder, m()) + case mods.Join[*SelectQuery]: + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, clause.Join(m)) + case CrossJoinChain[*SelectQuery]: + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, m()) + case mods.NamedWindow[*SelectQuery]: + appendDerived[bob.Expression](&next.Windows.Windows, base.Windows.Windows, &cloneWindows, clause.NamedWindow(m)) + case LockChain[*SelectQuery]: + appendDerived[bob.Expression](&next.Locks.Locks, base.Locks.Locks, &cloneLocks, m()) + case mods.Combine[*SelectQuery]: + appendDerived(&next.Combines.Queries, base.Combines.Queries, &cloneCombines, clause.Combine(m)) + case OrderCombined: + appendDerived[bob.Expression](&next.CombinedOrder.Expressions, base.CombinedOrder.Expressions, &cloneCombinedOrder, m()) + case LimitCombined: + next.CombinedLimit.Count = m.Count + case OffsetCombined: + next.CombinedOffset.Count = m.Count + case FetchCombined: + next.CombinedFetch.Count = m.Count + next.CombinedFetch.WithTies = m.WithTies + case FromChain[*SelectQuery]: + next.TableRef = cloneTableRef(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 + + for _, mod := range queryMods { + switch m := mod.(type) { + case mods.Recursive[*DeleteQuery]: + next.With.Recursive = bool(m) + case CTEChain[*DeleteQuery]: + 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]: + appendDerived[any](&next.Where.Conditions, base.Where.Conditions, &cloneWhere, m.E) + case mods.Returning[*DeleteQuery]: + appendDerived(&next.Returning.Expressions, base.Returning.Expressions, &cloneReturning, []any(m)...) + case FromChain[*DeleteQuery]: + next.TableRef = cloneTableRef(m()) + case mods.Join[*DeleteQuery]: + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, clause.Join(m)) + case CrossJoinChain[*DeleteQuery]: + appendDerived(&next.TableRef.Joins, base.TableRef.Joins, &cloneJoins, m()) + default: + return nil, false + } + } + + return &next, true +} + +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[*InsertQuery]: + next.With.Recursive = bool(m) + case CTEChain[*InsertQuery]: + 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 = OverridingType(m) + case mods.QuerySource[*InsertQuery]: + next.Values.Query = m.Query + case mods.Returning[*InsertQuery]: + appendDerived(&next.Returning.Expressions, base.Returning.Expressions, &cloneReturning, []any(m)...) + case mods.Values[*InsertQuery]: + appendDerived(&next.Values.Vals, base.Values.Vals, &cloneVals, clause.Value(m)) + case mods.Rows[*InsertQuery]: + if !cloneVals { + next.Values.Vals = cloneSlice(base.Values.Vals) + cloneVals = true + } + for _, row := range m { + next.Values.Vals = append(next.Values.Vals, clause.Value(row)) + } + case mods.Conflict[*InsertQuery]: + next.Conflict.Expression = m() + default: + return nil, false + } + } + + return &next, true +} diff --git a/dialect/psql/dialect/insert.go b/dialect/psql/dialect/insert.go index 477060b3..9eda6eac 100644 --- a/dialect/psql/dialect/insert.go +++ b/dialect/psql/dialect/insert.go @@ -32,54 +32,103 @@ 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 = OverridingType(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 - 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(string(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/mods.go b/dialect/psql/dialect/mods.go index 223dccae..8fcd80a4 100644 --- a/dialect/psql/dialect/mods.go +++ b/dialect/psql/dialect/mods.go @@ -223,6 +223,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/dialect/select.go b/dialect/psql/dialect/select.go index 7642294d..5bebc909 100644 --- a/dialect/psql/dialect/select.go +++ b/dialect/psql/dialect/select.go @@ -36,20 +36,36 @@ 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 - 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 +74,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..03c330b3 100644 --- a/dialect/psql/dialect/update.go +++ b/dialect/psql/dialect/update.go @@ -24,59 +24,86 @@ type UpdateQuery struct { bob.ContextualModdable[*UpdateQuery] } +func (u *UpdateQuery) SetTargetOnly(only bool) { + u.Only = only + u.Table.SetOnly(false) +} + +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 - 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..ba34fc49 --- /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 bob.ErrNoNamedArgs + 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/dm/qm.go b/dialect/psql/dm/qm.go index b282f86e..8dff1b96 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 mods.TargetOnly[*dialect.DeleteQuery](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 mods.TargetTable[*dialect.DeleteQuery]{ + 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 mods.TargetTable[*dialect.DeleteQuery]{ + 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..e573396f 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 mods.TargetTable[*dialect.InsertQuery]{ + 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 mods.TargetTable[*dialect.InsertQuery]{ + 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 mods.Overriding[*dialect.InsertQuery]("SYSTEM") } func OverridingUser() bob.Mod[*dialect.InsertQuery] { - return bob.ModFunc[*dialect.InsertQuery](func(i *dialect.InsertQuery) { - i.Overriding = dialect.OverridingUser - }) + return mods.Overriding[*dialect.InsertQuery]("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 mods.QuerySource[*dialect.InsertQuery]{Query: q} } // The column to target. Will auto add brackets diff --git a/dialect/psql/immutable_select_test.go b/dialect/psql/immutable_select_test.go new file mode 100644 index 00000000..c5c0c279 --- /dev/null +++ b/dialect/psql/immutable_select_test.go @@ -0,0 +1,471 @@ +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" +) + +func TestImmutableSelectQueryApplyDoesNotMutateOriginalFromLegacyWithCase(t *testing.T) { + base := Select( + sm.Columns("id"), + sm.From("users"), + ).Apply() + + derived := base.Apply( + 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(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 TestImmutableViewQueryApplyDoesNotMutateOriginalFromLegacyWithCase(t *testing.T) { + base := someStructView.Query( + sm.Where(Quote("id").GT(Arg(0))), + ).Apply() + + derived := base.Apply( + 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(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 TestViewQueryApplyNilReceiver(t *testing.T) { + var q *ViewQuery[*someStruct, []*someStruct] + + if got := q.Apply(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"), + 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 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), + ) + + 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))), + ) + + 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 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 *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 + } + 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 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"), + 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() + + 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 = 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 := b.Context() + + 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))), + ) + derived := q.Apply( + 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 := b.Context() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := someStructView.Query( + sm.Where(Quote("id").GT(Arg(0))), + ) + + if _, _, err := q.AsCount().Build(ctx); err != nil { + b.Fatal(err) + } + + q = 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 := b.Context() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := someStructView.Query( + sm.Where(Quote("id").GT(Arg(0))), + ) + + if _, _, err := q.AsCount().Build(ctx); err != nil { + b.Fatal(err) + } + + derived := q.Apply( + 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/immutable_write_test.go b/dialect/psql/immutable_write_test.go new file mode 100644 index 00000000..b68f62a9 --- /dev/null +++ b/dialect/psql/immutable_write_test.go @@ -0,0 +1,515 @@ +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" +) + +func TestUpdateApplyDoesNotMutateOriginalFromLegacyWithCase(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 TestDeleteApplyDoesNotMutateOriginalFromLegacyWithCase(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 TestInsertApplyDoesNotMutateOriginalFromLegacyWithCase(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 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 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"), + 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 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() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + ) + q = 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 := b.Context() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + ) + derived := q.Apply( + 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 := b.Context() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Delete( + dm.From("films"), + ) + q = 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 := b.Context() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Delete( + dm.From("films"), + ) + derived := q.Apply( + 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 := b.Context() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Insert( + im.Into("films"), + im.Values(Arg("UA502", "Bananas")), + ) + q = q.Apply( + im.Returning("id"), + ) + + if _, _, err := q.Build(ctx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkInsertQueryImmutableNativeHotPath(b *testing.B) { + ctx := b.Context() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + q := Insert( + im.Into("films"), + im.Values(Arg("UA502", "Bananas")), + ) + derived := q.Apply( + 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..099419da 100644 --- a/dialect/psql/insert.go +++ b/dialect/psql/insert.go @@ -5,15 +5,30 @@ 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) Apply(queryMods ...bob.Mod[*dialect.InsertQuery]) InsertQuery { + if next, ok := q.Expression.Derive(queryMods...); ok { + q.Expression = next + return q + } + q.BaseQuery = q.BaseQuery.Apply(queryMods...) + return q +} + +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/select.go b/dialect/psql/select.go index e177663d..cf1a4844 100644 --- a/dialect/psql/select.go +++ b/dialect/psql/select.go @@ -5,15 +5,43 @@ 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) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) SelectQuery { + if next, ok := q.Expression.Derive(queryMods...); ok { + q.Expression = next + return q + } + q.BaseQuery = q.BaseQuery.Apply(queryMods...) + return q +} + +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 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/sm/qm.go b/dialect/psql/sm/qm.go index 81a81dd9..2d0af599 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 mods.Distinct[*dialect.SelectQuery](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} } diff --git a/dialect/psql/table.go b/dialect/psql/table.go index 63cc5f03..72a6b488 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, @@ -91,16 +91,14 @@ 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 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, @@ -115,16 +113,14 @@ 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 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, @@ -139,9 +135,7 @@ func (t *Table[T, Tslice, Tset, C]) Delete(queryMods ...bob.Mod[*dialect.DeleteQ }, ) - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } // Starts a Merge query for this table @@ -168,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 c7303733..ca63e397 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,277 @@ 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 TestTableUpdateAdditionalExplicitReturningAppends(t *testing.T) { + base := userTable.Update( + um.SetCol("name").ToArg("Stephen"), + um.Where(Quote("id").EQ(Arg(1))), + ) + + q := base.Apply(um.Returning("id")).Apply(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")}), + ) + + 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 TestTableInsertAdditionalExplicitReturningAppends(t *testing.T) { + base := userTable.Insert( + im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), + ) + + q := base.Apply(im.Returning("id")).Apply(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))), + ) + + 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 TestTableDeleteAdditionalExplicitReturningAppends(t *testing.T) { + base := userTable.Delete( + dm.Where(Quote("id").EQ(Arg(1))), + ) + + q := base.Apply(dm.Returning("id")).Apply(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"), + ) + + 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 TestTableInsertApplyDoesNotMutateOriginal(t *testing.T) { + base := userTable.Insert( + im.Rows([]bob.Expression{Arg(int64(1)), Arg("Stephen"), Arg("stephen@example.com")}), + ) + + derived := base.Apply(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/psql/um/qm.go b/dialect/psql/um/qm.go index 85abc22b..58a25822 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 mods.TargetOnly[*dialect.UpdateQuery](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 mods.TargetTable[*dialect.UpdateQuery]{ + 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 mods.TargetTable[*dialect.UpdateQuery]{ + 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 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 8df715df..ee63521b 100644 --- a/dialect/psql/update.go +++ b/dialect/psql/update.go @@ -5,15 +5,26 @@ 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) Apply(queryMods ...bob.Mod[*dialect.UpdateQuery]) UpdateQuery { + q.BaseQuery = q.BaseQuery.Apply(queryMods...) + return q +} + +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, + }, } } diff --git a/dialect/psql/view.go b/dialect/psql/view.go index 104ca0de..8a4270dd 100644 --- a/dialect/psql/view.go +++ b/dialect/psql/view.go @@ -80,16 +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: orm.Query[*dialect.SelectQuery, T, Tslice, bob.SliceTransformer[T, Tslice]]{ - ExecQuery: orm.ExecQuery[*dialect.SelectQuery]{ - BaseQuery: Select(sm.From(v.NameAs())), - Hooks: &v.SelectQueryHooks, - }, - Scanner: v.scanner, - }, + SelectQuery: Select(sm.From(v.NameAs())), + Scanner: v.scanner, + Hooks: &v.SelectQueryHooks, } - q.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) @@ -98,50 +94,69 @@ func (v *View[T, Tslice, C]) Query(queryMods ...bob.Mod[*dialect.SelectQuery]) * }, ) - q.Apply(queryMods...) - - return q + return q.Apply(queryMods...) } type ViewQuery[T any, Ts ~[]T] struct { - orm.Query[*dialect.SelectQuery, T, Ts, bob.SliceTransformer[T, Ts]] + SelectQuery + Scanner scan.Mapper[T] + Hooks *bob.Hooks[*dialect.SelectQuery, bob.SkipQueryHooksKey] +} + +func (q *ViewQuery[T, Ts]) Apply(queryMods ...bob.Mod[*dialect.SelectQuery]) *ViewQuery[T, Ts] { + if q == nil { + return nil + } + + next := *q + next.SelectQuery = next.SelectQuery.Apply(queryMods...) + return &next } // 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 } - return bob.One(ctx, exec, asCountQuery(v.BaseQuery), scan.SingleColumnMapper[int64]) + sql, args, err := q.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 -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 } -// 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 +func (q *ViewQuery[T, Ts]) One(ctx context.Context, exec bob.Executor) (T, error) { + return bob.One(ctx, exec, q, q.Scanner) +} + +func (q *ViewQuery[T, Ts]) All(ctx context.Context, exec bob.Executor) (Ts, error) { + 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 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 bob.Each(ctx, exec, q, q.Scanner) +} + +func (q *ViewQuery[T, Ts]) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { + ctx, err := q.SelectQuery.RunHooks(ctx, exec) + if err != nil { + return ctx, err + } + + if q.Hooks == nil { + return ctx, nil + } + + return q.Hooks.RunHooks(ctx, exec, q.SelectQuery.Expression) } diff --git a/dialect/psql/view_test.go b/dialect/psql/view_test.go index 9e68a9d3..af4e21f5 100644 --- a/dialect/psql/view_test.go +++ b/dialect/psql/view_test.go @@ -2,12 +2,10 @@ package psql import ( "bytes" - "context" "testing" _ "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,9 +54,23 @@ func TestSomeViewQuery(t *testing.T) { } } -func selectToString(t *testing.T, query bob.BaseQuery[*dialect.SelectQuery], argsLen int) string { +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 := context.Background() + ctx := t.Context() buf := new(bytes.Buffer) args, err := query.WriteQuery(ctx, buf, 0) if err != nil { @@ -74,5 +86,5 @@ func selectToString(t *testing.T, query bob.BaseQuery[*dialect.SelectQuery], arg func viewToString(t *testing.T, query *ViewQuery[*someStruct, []*someStruct]) string { t.Helper() - return selectToString(t, query.BaseQuery, 3) + return selectToString(t, query, 3) } diff --git a/dialect/psql/with_regression_test.go b/dialect/psql/with_regression_test.go new file mode 100644 index 00000000..192d7624 --- /dev/null +++ b/dialect/psql/with_regression_test.go @@ -0,0 +1,226 @@ +package psql_test + +import ( + "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 TestSelectApplyRegression(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.Apply( + 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.Apply( + 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 TestViewQueryApplyRegression(t *testing.T) { + base := withTestStructView.Query( + sm.Where(psql.Quote("id").GT(psql.Arg(0))), + ) + + derived := base.Apply( + 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, withTestStructView.Query( + sm.Where(psql.Quote("id").GT(psql.Arg(0))), + sm.OrderBy("id").Desc(), + sm.Limit(10), + sm.Offset(20), + )) +} + +func TestUpdateApplyRegression(t *testing.T) { + base := psql.Update( + um.Table("films"), + um.SetCol("kind").ToArg("Dramatic"), + ) + + derived := base.Apply( + 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 TestDeleteApplyRegression(t *testing.T) { + base := psql.Delete( + dm.From("employees"), + ) + + 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"))), + 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 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.Apply( + 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.Apply( + 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(t.Context(), got) + if err != nil { + t.Fatalf("build got: %v", err) + } + + wantSQL, wantArgs, err := bob.Build(t.Context(), 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) + } +} 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/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/mods/mods.go b/mods/mods.go index f870a2c3..badb9fdd 100644 --- a/mods/mods.go +++ b/mods/mods.go @@ -37,12 +37,36 @@ 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 != "" || len(t.Columns) > 0 { + q.SetTargetTableAlias(t.Alias, t.Columns...) + } +} + type Select[Q interface{ AppendSelect(columns ...any) }] []any 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 +179,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 +209,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) { diff --git a/orm/query.go b/orm/query.go index 5ed7729d..8678213d 100644 --- a/orm/query.go +++ b/orm/query.go @@ -22,6 +22,18 @@ func (q ExecQuery[Q]) Clone() ExecQuery[Q] { } } +func (q *ExecQuery[Q]) Apply(queryMods ...bob.Mod[Q]) *ExecQuery[Q] { + if q == nil { + return nil + } + + next := q.Clone() + for _, mod := range queryMods { + mod.Apply(next.BaseQuery.Expression) + } + return &next +} + func (q ExecQuery[Q]) RunHooks(ctx context.Context, exec bob.Executor) (context.Context, error) { var err error @@ -55,7 +67,20 @@ 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]) Apply(queryMods ...bob.Mod[Q]) *Query[Q, T, Ts, Tr] { + if q == nil { + return nil + } + + next := q.Clone() + for _, mod := range queryMods { + mod.Apply(next.BaseQuery.Expression) } + return &next } // First matching row diff --git a/query.go b/query.go index eb227f8c..59c53deb 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, } } @@ -115,10 +117,12 @@ func (b BaseQuery[E]) GetMapperMods() []scan.MapperMod { return nil } -func (b BaseQuery[E]) Apply(mods ...Mod[E]) { +func (b BaseQuery[E]) Apply(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]) 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..07ee5b8a --- /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 TestBaseQueryApplyDoesNotMutateOriginalFromLegacyWithCase(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) + } +} + +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) + } +}