Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
145 changes: 97 additions & 48 deletions generator/templates/builders_create.gotpl
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
type CreateBuilder[M any, S any, O any] struct {
client *Queries
assignments []FieldAssignment
execFunc func(ctx context.Context, assignments []FieldAssignment, s *S, o *O) (*M, error)
client *Queries
assignments []FieldAssignment
execFunc func(ctx context.Context, assignments []FieldAssignment, s *S, o *O, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) (*M, error)
conflictAction *ConflictAction
conflictTarget UniqueConstraintTarget
}

func (b *CreateBuilder[M, S, O]) Select(s S) *CreateSelectBuilder[M, S, O] {
Expand All @@ -13,7 +15,7 @@ func (b *CreateBuilder[M, S, O]) Omit(o O) *CreateOmitBuilder[M, S, O] {
}

func (b *CreateBuilder[M, S, O]) Exec(ctx context.Context) (*M, error) {
return b.execFunc(ctx, b.assignments, nil, nil)
return b.execFunc(ctx, b.assignments, nil, nil, b.conflictTarget, b.conflictAction)
}

type CreateSelectBuilder[M any, S any, O any] struct {
Expand All @@ -22,7 +24,7 @@ type CreateSelectBuilder[M any, S any, O any] struct {
}

func (b *CreateSelectBuilder[M, S, O]) Exec(ctx context.Context) (*M, error) {
return b.builder.execFunc(ctx, b.builder.assignments, &b.selects, nil)
return b.builder.execFunc(ctx, b.builder.assignments, &b.selects, nil, b.builder.conflictTarget, b.builder.conflictAction)
}

type CreateOmitBuilder[M any, S any, O any] struct {
Expand All @@ -31,34 +33,36 @@ type CreateOmitBuilder[M any, S any, O any] struct {
}

func (b *CreateOmitBuilder[M, S, O]) Exec(ctx context.Context) (*M, error) {
return b.builder.execFunc(ctx, b.builder.assignments, nil, &b.omits)
return b.builder.execFunc(ctx, b.builder.assignments, nil, &b.omits, b.builder.conflictTarget, b.builder.conflictAction)
}

type CreateManyBuilder[M any] struct {
client *Queries
records []RecordInput
execFunc func(ctx context.Context, records []RecordInput, skipDuplicates bool) (int64, error)
skipDuplicates bool
execFunc func(ctx context.Context, records []RecordInput, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) (int64, error)
conflictAction *ConflictAction
conflictTarget UniqueConstraintTarget
}

func (b *CreateManyBuilder[M]) SkipDuplicates() *CreateManyBuilder[M] {
b.skipDuplicates = true
b.conflictAction = &ConflictAction{Type: ConflictActionIgnore}
return b
}

func (b *CreateManyBuilder[M]) Exec(ctx context.Context) (int64, error) {
return b.execFunc(ctx, b.records, b.skipDuplicates)
return b.execFunc(ctx, b.records, b.conflictTarget, b.conflictAction)
}

type CreateManyAndReturnBuilder[M any, S any, O any] struct {
client *Queries
records []RecordInput
execFunc func(ctx context.Context, records []RecordInput, s *S, o *O, skipDuplicates bool) ([]*M, error)
skipDuplicates bool
execFunc func(ctx context.Context, records []RecordInput, s *S, o *O, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) ([]*M, error)
conflictAction *ConflictAction
conflictTarget UniqueConstraintTarget
}

func (b *CreateManyAndReturnBuilder[M, S, O]) SkipDuplicates() *CreateManyAndReturnBuilder[M, S, O] {
b.skipDuplicates = true
b.conflictAction = &ConflictAction{Type: ConflictActionIgnore}
return b
}

Expand All @@ -71,7 +75,7 @@ func (b *CreateManyAndReturnBuilder[M, S, O]) Omit(o O) *CreateManyAndReturnOmit
}

func (b *CreateManyAndReturnBuilder[M, S, O]) Exec(ctx context.Context) ([]*M, error) {
return b.execFunc(ctx, b.records, nil, nil, b.skipDuplicates)
return b.execFunc(ctx, b.records, nil, nil, b.conflictTarget, b.conflictAction)
}

type CreateManyAndReturnSelectBuilder[M any, S any, O any] struct {
Expand All @@ -80,7 +84,7 @@ type CreateManyAndReturnSelectBuilder[M any, S any, O any] struct {
}

func (b *CreateManyAndReturnSelectBuilder[M, S, O]) Exec(ctx context.Context) ([]*M, error) {
return b.builder.execFunc(ctx, b.builder.records, &b.selects, nil, b.builder.skipDuplicates)
return b.builder.execFunc(ctx, b.builder.records, &b.selects, nil, b.builder.conflictTarget, b.builder.conflictAction)
}

type CreateManyAndReturnOmitBuilder[M any, S any, O any] struct {
Expand All @@ -89,7 +93,7 @@ type CreateManyAndReturnOmitBuilder[M any, S any, O any] struct {
}

func (b *CreateManyAndReturnOmitBuilder[M, S, O]) Exec(ctx context.Context) ([]*M, error) {
return b.builder.execFunc(ctx, b.builder.records, nil, &b.omits, b.builder.skipDuplicates)
return b.builder.execFunc(ctx, b.builder.records, nil, &b.omits, b.builder.conflictTarget, b.builder.conflictAction)
}

func executeInsert[M any](
Expand All @@ -99,8 +103,10 @@ func executeInsert[M any](
cols []string,
vals []any,
returningCols []string,
idCol string,
pkCols []string,
scanFunc func(record *M, cols []string) []any,
conflictTarget UniqueConstraintTarget,
conflictAction *ConflictAction,
) (*M, error) {
var sb strings.Builder
sb.Grow(128 + len(table) + len(cols)*15 + len(returningCols)*15)
Expand All @@ -123,6 +129,22 @@ func executeInsert[M any](
}
sb.WriteString(")")

var conflictCols []string
if conflictTarget != nil {
conflictCols = conflictTarget.UniqueColumns()
}

var nonConflictCols []string
if conflictAction != nil && conflictAction.Type == ConflictActionUpdateNewValues {
nonConflictCols = computeNonConflictCols(cols, conflictCols, pkCols)
}

clause, clauseArgs := q.dialect.ConflictClause(conflictCols, conflictAction, nonConflictCols, len(vals)+1)
sb.WriteString(clause)
if len(clauseArgs) > 0 {
vals = append(vals, clauseArgs...)
}

if q.dialect.SupportsReturning() && len(returningCols) > 0 {
sb.WriteString(" RETURNING ")
for i, col := range returningCols {
Expand All @@ -136,13 +158,20 @@ func executeInsert[M any](

var res M
if q.dialect.SupportsReturning() {
row := q.queryRow(ctx, query, vals...)

scanTargets := scanFunc(&res, returningCols)
if err := row.Scan(scanTargets...); err != nil {
rows, err := q.query(ctx, query, vals...)
if err != nil {
return nil, err
}
return &res, nil
defer rows.Close()

if rows.Next() {
scanTargets := scanFunc(&res, returningCols)
if err := rows.Scan(scanTargets...); err != nil {
return nil, err
}
return &res, nil
}
return nil, rows.Err()
}

// Fallback for dialects without RETURNING (MySQL)
Expand All @@ -151,23 +180,27 @@ func executeInsert[M any](
return nil, err
}

var idVal any
for i, c := range cols {
if c == idCol {
idVal = vals[i]
break
var pkVals []any
for _, pkCol := range pkCols {
var val any
for i, c := range cols {
if c == pkCol {
val = vals[i]
break
}
}
}
if idVal == nil {
lastID, err := result.LastInsertId()
if err != nil {
return nil, err
if val == nil && len(pkCols) == 1 {
lastID, err := result.LastInsertId()
if err != nil {
return nil, err
}
val = lastID
}
idVal = lastID
pkVals = append(pkVals, val)
}

var selectSb strings.Builder
selectSb.Grow(64 + len(returningCols)*15 + len(table) + len(idCol))
selectSb.Grow(64 + len(returningCols)*15 + len(table) + len(pkCols)*15)
selectSb.WriteString("SELECT ")
for i, col := range returningCols {
if i > 0 {
Expand All @@ -178,15 +211,28 @@ func executeInsert[M any](
selectSb.WriteString(" FROM ")
selectSb.WriteString(q.dialect.Quote(table))
selectSb.WriteString(" WHERE ")
selectSb.WriteString(q.dialect.Quote(idCol))
selectSb.WriteString(" = ?")
for i, pkCol := range pkCols {
if i > 0 {
selectSb.WriteString(" AND ")
}
selectSb.WriteString(q.dialect.Quote(pkCol))
selectSb.WriteString(" = ?")
}

row := q.queryRow(ctx, selectSb.String(), idVal)
scanTargets := scanFunc(&res, returningCols)
if err := row.Scan(scanTargets...); err != nil {
rows, err := q.query(ctx, selectSb.String(), pkVals...)
if err != nil {
return nil, err
}
return &res, nil
defer rows.Close()

if rows.Next() {
scanTargets := scanFunc(&res, returningCols)
if err := rows.Scan(scanTargets...); err != nil {
return nil, err
}
return &res, nil
}
return nil, rows.Err()
}


Expand All @@ -196,15 +242,17 @@ func executeCreateMany(
rowMaps []map[string]any,
tableName string,
colOrder []string,
skipDuplicates bool,
pkCols []string,
conflictTarget UniqueConstraintTarget,
conflictAction *ConflictAction,
) (int64, error) {
if len(rowMaps) == 0 {
return 0, nil
}

batches := partitionRowMaps(q.dialect, rowMaps)
if len(batches) == 1 {
query, vals := buildBulkInsertSQL(q.dialect, tableName, batches[0], colOrder, nil, skipDuplicates)
query, vals := buildBulkInsertSQL(q.dialect, tableName, batches[0], colOrder, nil, pkCols, conflictTarget, conflictAction)
res, err := q.exec(ctx, query, vals...)
if err != nil {
return 0, err
Expand All @@ -215,7 +263,7 @@ func executeCreateMany(
var count int64
err := q.transaction(ctx, func(txQ *Queries) error {
for _, batch := range batches {
query, vals := buildBulkInsertSQL(txQ.dialect, tableName, batch, colOrder, nil, skipDuplicates)
query, vals := buildBulkInsertSQL(txQ.dialect, tableName, batch, colOrder, nil, pkCols, conflictTarget, conflictAction)
res, err := txQ.exec(ctx, query, vals...)
if err != nil {
return err
Expand Down Expand Up @@ -243,8 +291,9 @@ func executeCreateManyAndReturn[M any, S any, O any](
loadRelationsFn func(context.Context, []*M, *S) error,
scanFunc func(*M, []string) []any,
hasRelationsFn func(*S) bool,
idCol string,
skipDuplicates bool,
pkCols []string,
conflictTarget UniqueConstraintTarget,
conflictAction *ConflictAction,
) ([]*M, error) {
if len(rowMaps) == 0 {
return nil, nil
Expand All @@ -258,7 +307,7 @@ func executeCreateManyAndReturn[M any, S any, O any](
err := q.transaction(ctx, func(txQ *Queries) error {
for _, rowMap := range rowMaps {
cols, vals := mapToColsVals(rowMap, colOrder)
res, err := executeInsert(ctx, txQ, tableName, cols, vals, returningCols, idCol, scanFunc)
res, err := executeInsert(ctx, txQ, tableName, cols, vals, returningCols, pkCols, scanFunc, conflictTarget, conflictAction)
if err != nil {
return err
}
Expand All @@ -278,7 +327,7 @@ func executeCreateManyAndReturn[M any, S any, O any](
batches := partitionRowMaps(q.dialect, rowMaps)
recordsOut := make([]*M, 0)
if len(batches) == 1 && !hasRelations {
query, vals := buildBulkInsertSQL(q.dialect, tableName, batches[0], colOrder, returningCols, skipDuplicates)
query, vals := buildBulkInsertSQL(q.dialect, tableName, batches[0], colOrder, returningCols, pkCols, conflictTarget, conflictAction)
rows, err := q.query(ctx, query, vals...)
if err != nil {
return nil, err
Expand All @@ -299,7 +348,7 @@ func executeCreateManyAndReturn[M any, S any, O any](

err := q.transaction(ctx, func(txQ *Queries) error {
for _, batch := range batches {
query, vals := buildBulkInsertSQL(txQ.dialect, tableName, batch, colOrder, returningCols, skipDuplicates)
query, vals := buildBulkInsertSQL(txQ.dialect, tableName, batch, colOrder, returningCols, pkCols, conflictTarget, conflictAction)
err := func() error {
rows, err := txQ.query(ctx, query, vals...)
if err != nil {
Expand Down
24 changes: 12 additions & 12 deletions generator/templates/client.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -17,12 +17,8 @@ func (postgresDialect) FormatLimitOffset(take *int, skip *int) string {
return ""
}
func (postgresDialect) SupportsDefaultKeyword() bool { return true }
func (postgresDialect) InsertPrefix(skipDuplicates bool) string { return "INSERT INTO" }
func (postgresDialect) ConflictClause(skipDuplicates bool) string {
if skipDuplicates {
return " ON CONFLICT DO NOTHING"
}
return ""
func (postgresDialect) ConflictClause(conflictCols []string, action *ConflictAction, nonConflictCols []string, startParamIndex int) (string, []any) {
return buildConflictClause(postgresDialect{}.Quote, postgresDialect{}.BindVar, conflictCols, action, nonConflictCols, startParamIndex)
}

{{- end }}
Expand All @@ -46,13 +42,9 @@ func (sqliteDialect) FormatLimitOffset(take *int, skip *int) string {
return ""
}
func (sqliteDialect) SupportsDefaultKeyword() bool { return false }
func (sqliteDialect) InsertPrefix(skipDuplicates bool) string {
if skipDuplicates {
return "INSERT OR IGNORE INTO"
}
return "INSERT INTO"
func (sqliteDialect) ConflictClause(conflictCols []string, action *ConflictAction, nonConflictCols []string, startParamIndex int) (string, []any) {
return buildConflictClause(sqliteDialect{}.Quote, sqliteDialect{}.BindVar, conflictCols, action, nonConflictCols, startParamIndex)
}
func (sqliteDialect) ConflictClause(skipDuplicates bool) string { return "" }

{{- end }}

Expand Down Expand Up @@ -529,6 +521,10 @@ type UniqueField[M any, T any] struct {
Column string
}

func (f UniqueField[M, T]) UniqueColumns() []string {
return []string{f.Column}
}

func (f UniqueField[M, T]) Set(val T) FieldAssignment {
return FieldAssignment{Col: f.Column, Val: val}
}
Expand Down Expand Up @@ -757,6 +753,10 @@ type StringUniqueField[M any] struct {
Column string
}

func (f StringUniqueField[M]) UniqueColumns() []string {
return []string{f.Column}
}

func (f StringUniqueField[M]) Set(val string) FieldAssignment {
return FieldAssignment{Col: f.Column, Val: val}
}
Expand Down
Loading
Loading