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
8 changes: 8 additions & 0 deletions generator/generator.go
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,14 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default
}
return false
},
"hasAnyLog": func() bool {
for _, l := range defaultLogs {
if l != "none" {
return true
}
}
return false
},
})
tmpl, err := tmpl.ParseFS(templatesFS, "templates/*.gotpl")
if err != nil {
Expand Down
52 changes: 49 additions & 3 deletions generator/templates/builders_create.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,53 @@ type CreateOmitBuilder[M any, I any, S any, O any] struct {
func (b *CreateOmitBuilder[M, I, S, O]) Exec(ctx context.Context) (*M, error) {
return b.builder.execFunc(ctx, b.builder.input, nil, &b.omits)
}

type CreateManyBuilder[M any, I any] struct {
client *Queries
inputs []I
execFunc func(ctx context.Context, inputs []I) (int64, error)
}

func (b *CreateManyBuilder[M, I]) Exec(ctx context.Context) (int64, error) {
return b.execFunc(ctx, b.inputs)
}

type CreateManyAndReturnBuilder[M any, I any, S any, O any] struct {
client *Queries
inputs []I
execFunc func(ctx context.Context, inputs []I, s *S, o *O) ([]*M, error)
}

func (b *CreateManyAndReturnBuilder[M, I, S, O]) Select(s S) *CreateManyAndReturnSelectBuilder[M, I, S, O] {
return &CreateManyAndReturnSelectBuilder[M, I, S, O]{builder: b, selects: s}
}

func (b *CreateManyAndReturnBuilder[M, I, S, O]) Omit(o O) *CreateManyAndReturnOmitBuilder[M, I, S, O] {
return &CreateManyAndReturnOmitBuilder[M, I, S, O]{builder: b, omits: o}
}

func (b *CreateManyAndReturnBuilder[M, I, S, O]) Exec(ctx context.Context) ([]*M, error) {
return b.execFunc(ctx, b.inputs, nil, nil)
}

type CreateManyAndReturnSelectBuilder[M any, I any, S any, O any] struct {
builder *CreateManyAndReturnBuilder[M, I, S, O]
selects S
}

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

type CreateManyAndReturnOmitBuilder[M any, I any, S any, O any] struct {
builder *CreateManyAndReturnBuilder[M, I, S, O]
omits O
}

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

func executeInsert[M any](
ctx context.Context,
q *Queries,
Expand All @@ -53,7 +100,7 @@ func executeInsert[M any](
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(col)
sb.WriteString(q.dialect.Quote(col))
}
sb.WriteString(") VALUES (")
for i := range cols {
Expand Down Expand Up @@ -93,9 +140,8 @@ func executeInsert[M any](
}

var idVal any
quotedIdCol := q.dialect.Quote(idCol)
for i, c := range cols {
if c == quotedIdCol {
if c == idCol {
idVal = vals[i]
break
}
Expand Down
3 changes: 3 additions & 0 deletions generator/templates/client.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -2,20 +2,23 @@ type Dialect interface {
Quote(ident string) string
BindVar(idx int) string
SupportsReturning() bool
SupportsBulkInsert() bool
}

{{- if or (eq .Schema.Datasource.Provider "postgres") (eq .Schema.Datasource.Provider "postgresql") }}
type postgresDialect struct{}
func (postgresDialect) Quote(ident string) string { return `"` + ident + `"` }
func (postgresDialect) BindVar(idx int) string { return fmt.Sprintf("$%d", idx) }
func (postgresDialect) SupportsReturning() bool { return true }
func (postgresDialect) SupportsBulkInsert() bool { return true }
{{- end }}

{{- if or (eq .Schema.Datasource.Provider "sqlite") (eq .Schema.Datasource.Provider "sqlite3") }}
type sqliteDialect struct{}
func (sqliteDialect) Quote(ident string) string { return `"` + ident + `"` }
func (sqliteDialect) BindVar(idx int) string { return "?" }
func (sqliteDialect) SupportsReturning() bool { return true }
func (sqliteDialect) SupportsBulkInsert() bool { return false }
{{- end }}

type DBTX interface {
Expand Down
2 changes: 2 additions & 0 deletions generator/templates/header.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@ import (
"embed"
{{- end }}
"fmt"
{{- if hasAnyLog }}
"log"
{{- end }}
"strconv"
"strings"
"time"
Expand Down
203 changes: 162 additions & 41 deletions generator/templates/model_create.gotpl
Original file line number Diff line number Diff line change
@@ -1,3 +1,20 @@
var {{ .Model.Name }}ColOrder = []string{
{{- range $field := .Model.ScalarFields }}
"{{ $field.EffectiveColName }}",
{{- end }}
}

func (s *{{ .Model.Name }}Select) hasAnyRelation() bool {
if s == nil {
return false
}
{{- if .Model.RelationFields }}
return {{ range $i, $rel := .Model.RelationFields }}{{ if $i }} || {{ end }}s.{{ capitalize $rel.Name }} != nil{{ end }}
{{- else }}
return false
{{- end }}
}

func (d *{{ .Model.Name }}Delegate) Create(input {{ .Model.Name }}CreateInput) *CreateBuilder[{{ .Model.Name }}, {{ .Model.Name }}CreateInput, {{ .Model.Name }}Select, {{ .Model.Name }}Omit] {
return &CreateBuilder[{{ .Model.Name }}, {{ .Model.Name }}CreateInput, {{ .Model.Name }}Select, {{ .Model.Name }}Omit]{
client: d.client,
Expand All @@ -7,96 +24,200 @@ func (d *{{ .Model.Name }}Delegate) Create(input {{ .Model.Name }}CreateInput) *
}

func (q *Queries) execute{{ .Model.Name }}Create(ctx context.Context, input {{ .Model.Name }}CreateInput, selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit) (*{{ .Model.Name }}, error) {
var cols []string
var vals []any
m := q.{{ .Model.Name }}InputToMap(input)
cols, vals := mapToColsVals(m, {{ .Model.Name }}ColOrder)

returningCols := q.select{{ .Model.Name }}Cols(selects, omits)

scanFunc := func(res *{{ .Model.Name }}, cols []string) []any {
return res.ScanFields(cols)
}

idCol := "{{ range $field := .Model.ScalarFields }}{{ if $field.IsID }}{{ $field.EffectiveColName }}{{ end }}{{ end }}"

hasRelations := selects.hasAnyRelation()

var res *{{ .Model.Name }}
var err error
if hasRelations {
err = q.transaction(ctx, func(txQ *Queries) error {
var err error
res, err = executeInsert(ctx, txQ, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, idCol, scanFunc)
if err != nil {
return err
}
return txQ.load{{ .Model.Name }}Relations(ctx, []*{{ .Model.Name }}{res}, selects)
})
} else {
res, err = executeInsert(ctx, q, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, idCol, scanFunc)
}
if err != nil {
return nil, err
}

return res, nil
}

func (q *Queries) {{ .Model.Name }}InputToMap(input {{ .Model.Name }}CreateInput) map[string]any {
m := make(map[string]any)
{{- range $field := .Model.ScalarFields }}
{{- if $field.EnumRef }}
{{- if $field.IsArray }}
if input.{{ capitalize $field.Name }} != nil {
cols = append(cols, q.dialect.Quote("{{ $field.EffectiveColName }}"))
vals = append(vals, input.{{ capitalize $field.Name }})
m["{{ $field.EffectiveColName }}"] = input.{{ capitalize $field.Name }}
}
{{- else }}
if input.{{ capitalize $field.Name }} != nil {
cols = append(cols, q.dialect.Quote("{{ $field.EffectiveColName }}"))
vals = append(vals, *input.{{ capitalize $field.Name }})
m["{{ $field.EffectiveColName }}"] = *input.{{ capitalize $field.Name }}
}
{{- end }}
{{- else }}
{{- if $field.IsArray }}
if input.{{ capitalize $field.Name }} != nil {
cols = append(cols, q.dialect.Quote("{{ $field.EffectiveColName }}"))
vals = append(vals, input.{{ capitalize $field.Name }})
m["{{ $field.EffectiveColName }}"] = input.{{ capitalize $field.Name }}
}
{{- else }}
{{- if and $field.Default (eq $field.Default.Kind.String "Func") }}
if input.{{ capitalize $field.Name }} != nil {
cols = append(cols, q.dialect.Quote("{{ $field.EffectiveColName }}"))
vals = append(vals, *input.{{ capitalize $field.Name }})
m["{{ $field.EffectiveColName }}"] = *input.{{ capitalize $field.Name }}
} else {
cols = append(cols, q.dialect.Quote("{{ $field.EffectiveColName }}"))
{{- if eq $field.Default.FuncName "cuid" }}
vals = append(vals, generateCUID())
m["{{ $field.EffectiveColName }}"] = generateCUID()
{{- else if eq $field.Default.FuncName "uuid" }}
vals = append(vals, generateUUID())
m["{{ $field.EffectiveColName }}"] = generateUUID()
{{- else if eq $field.Default.FuncName "now" }}
vals = append(vals, time.Now())
m["{{ $field.EffectiveColName }}"] = time.Now()
{{- end }}
}
{{- else if or $field.Optional (ne $field.Default nil) }}
if input.{{ capitalize $field.Name }} != nil {
cols = append(cols, q.dialect.Quote("{{ $field.EffectiveColName }}"))
vals = append(vals, *input.{{ capitalize $field.Name }})
m["{{ $field.EffectiveColName }}"] = *input.{{ capitalize $field.Name }}
}
{{- else }}
{{- if and $field.IsID (eq $field.GoType "string") }}
cols = append(cols, q.dialect.Quote("{{ $field.EffectiveColName }}"))
if input.{{ capitalize $field.Name }} != "" {
vals = append(vals, input.{{ capitalize $field.Name }})
m["{{ $field.EffectiveColName }}"] = input.{{ capitalize $field.Name }}
} else {
vals = append(vals, generateCUID())
m["{{ $field.EffectiveColName }}"] = generateCUID()
}
{{- else }}
cols = append(cols, q.dialect.Quote("{{ $field.EffectiveColName }}"))
vals = append(vals, input.{{ capitalize $field.Name }})
m["{{ $field.EffectiveColName }}"] = input.{{ capitalize $field.Name }}
{{- end }}
{{- end }}
{{- end }}
{{- end }}
{{- end }}
return m
}

returningCols := q.select{{ .Model.Name }}Cols(selects, omits)
func (d *{{ .Model.Name }}Delegate) CreateMany(inputs []{{ .Model.Name }}CreateInput) *CreateManyBuilder[{{ .Model.Name }}, {{ .Model.Name }}CreateInput] {
return &CreateManyBuilder[{{ .Model.Name }}, {{ .Model.Name }}CreateInput]{
client: d.client,
inputs: inputs,
execFunc: d.client.execute{{ .Model.Name }}CreateMany,
}
}

scanFunc := func(res *{{ .Model.Name }}, cols []string) []any {
return res.ScanFields(cols)
func (d *{{ .Model.Name }}Delegate) CreateManyAndReturn(inputs []{{ .Model.Name }}CreateInput) *CreateManyAndReturnBuilder[{{ .Model.Name }}, {{ .Model.Name }}CreateInput, {{ .Model.Name }}Select, {{ .Model.Name }}Omit] {
return &CreateManyAndReturnBuilder[{{ .Model.Name }}, {{ .Model.Name }}CreateInput, {{ .Model.Name }}Select, {{ .Model.Name }}Omit]{
client: d.client,
inputs: inputs,
execFunc: d.client.execute{{ .Model.Name }}CreateManyAndReturn,
}
}

idCol := "{{ range $field := .Model.ScalarFields }}{{ if $field.IsID }}{{ $field.EffectiveColName }}{{ end }}{{ end }}"
func (q *Queries) execute{{ .Model.Name }}CreateMany(ctx context.Context, inputs []{{ .Model.Name }}CreateInput) (int64, error) {
if len(inputs) == 0 {
return 0, nil
}

{{- if .Model.RelationFields }}
hasRelations := selects != nil && ({{ range $i, $rel := .Model.RelationFields }}{{ if $i }} || {{ end }}selects.{{ capitalize $rel.Name }} != nil{{ end }})
{{- else }}
hasRelations := false
{{- end }}
if q.dialect.SupportsBulkInsert() {
rowMaps := make([]map[string]any, len(inputs))
for i, input := range inputs {
rowMaps[i] = q.{{ .Model.Name }}InputToMap(input)
}
query, vals := buildBulkInsertSQL(q.dialect, "{{ .Model.EffectiveTableName }}", rowMaps, {{ .Model.Name }}ColOrder, nil)
res, err := q.exec(ctx, query, vals...)
if err != nil {
return 0, err
}
return res.RowsAffected()
}

var res *{{ .Model.Name }}
var err error
if hasRelations {
err = q.transaction(ctx, func(txQ *Queries) error {
var err error
res, err = executeInsert(ctx, txQ, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, idCol, scanFunc)
var count int64
err := q.transaction(ctx, func(txQ *Queries) error {
for _, input := range inputs {
_, err := txQ.execute{{ .Model.Name }}Create(ctx, input, nil, nil)
if err != nil {
return err
}
return txQ.load{{ .Model.Name }}Relations(ctx, []*{{ .Model.Name }}{res}, selects)
count++
}
return nil
})
return count, err
}

func (q *Queries) execute{{ .Model.Name }}CreateManyAndReturn(ctx context.Context, inputs []{{ .Model.Name }}CreateInput, selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit) ([]*{{ .Model.Name }}, error) {
if len(inputs) == 0 {
return nil, nil
}

hasRelations := selects.hasAnyRelation()
returningCols := q.select{{ .Model.Name }}Cols(selects, omits)

if q.dialect.SupportsBulkInsert() {
rowMaps := make([]map[string]any, len(inputs))
for i, input := range inputs {
rowMaps[i] = q.{{ .Model.Name }}InputToMap(input)
}
query, vals := buildBulkInsertSQL(q.dialect, "{{ .Model.EffectiveTableName }}", rowMaps, {{ .Model.Name }}ColOrder, returningCols)
var records []*{{ .Model.Name }}
err := q.transaction(ctx, func(txQ *Queries) error {
rows, err := txQ.query(ctx, query, vals...)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var record {{ .Model.Name }}
if err := rows.Scan(record.ScanFields(returningCols)...); err != nil {
return err
}
records = append(records, &record)
}
if err := rows.Err(); err != nil {
return err
}
if hasRelations {
return txQ.load{{ .Model.Name }}Relations(ctx, records, selects)
}
return nil
})
} else {
res, err = executeInsert(ctx, q, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, idCol, scanFunc)
if err != nil {
return nil, err
}
return records, nil
}

// Fallback to loop inside transaction
var records []*{{ .Model.Name }}
err := q.transaction(ctx, func(txQ *Queries) error {
for _, input := range inputs {
res, err := txQ.execute{{ .Model.Name }}Create(ctx, input, nil, nil)
if err != nil {
return err
}
records = append(records, res)
}

if hasRelations {
return txQ.load{{ .Model.Name }}Relations(ctx, records, selects)
}
return nil
})
if err != nil {
return nil, err
}

return res, nil
return records, nil
}
Loading
Loading