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
9 changes: 3 additions & 6 deletions generator/templates/builders_create.gotpl
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
type CreateBuilder[M any, S any, O any] struct {
client *Queries
assignments []FieldAssignment
execFunc func(ctx context.Context, assignments []FieldAssignment, s *S, o *O, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) (*M, error)
conflictAction *ConflictAction
Expand Down Expand Up @@ -37,7 +36,6 @@ func (b *CreateOmitBuilder[M, S, O]) Exec(ctx context.Context) (*M, error) {
}

type CreateManyBuilder[M any] struct {
client *Queries
records []RecordInput
execFunc func(ctx context.Context, records []RecordInput, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) (int64, error)
conflictAction *ConflictAction
Expand All @@ -54,7 +52,6 @@ func (b *CreateManyBuilder[M]) Exec(ctx context.Context) (int64, error) {
}

type CreateManyAndReturnBuilder[M any, S any, O any] struct {
client *Queries
records []RecordInput
execFunc func(ctx context.Context, records []RecordInput, s *S, o *O, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) ([]*M, error)
conflictAction *ConflictAction
Expand Down Expand Up @@ -288,7 +285,7 @@ func executeCreateManyAndReturn[M any, S any, O any](
selects *S,
omits *O,
selectColsFn func(*S, *O, ...string) []string,
loadRelationsFn func(context.Context, []*M, *S) error,
loadRelationsFn func(context.Context, *Queries, []*M, *S) error,
scanFunc func(*M, []string) []any,
hasRelationsFn func(*S) bool,
pkCols []string,
Expand All @@ -314,7 +311,7 @@ func executeCreateManyAndReturn[M any, S any, O any](
recordsOut = append(recordsOut, res)
}
if hasRelations {
return loadRelationsFn(ctx, recordsOut, selects)
return loadRelationsFn(ctx, txQ, recordsOut, selects)
}
return nil
})
Expand Down Expand Up @@ -369,7 +366,7 @@ func executeCreateManyAndReturn[M any, S any, O any](
}
}
if hasRelations {
return loadRelationsFn(ctx, recordsOut, selects)
return loadRelationsFn(ctx, txQ, recordsOut, selects)
}
return nil
})
Expand Down
3 changes: 0 additions & 3 deletions generator/templates/builders_query.gotpl
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
type FindUniqueBuilder[M any, S any, O any] struct {
client *Queries
where UniquePredicate[M]
additional []PredicateOf[M]
execFunc func(ctx context.Context, where UniquePredicate[M], additional []PredicateOf[M], s *S, o *O) (*M, error)
Expand Down Expand Up @@ -36,7 +35,6 @@ func (b *FindUniqueOmitBuilder[M, S, O]) Exec(ctx context.Context) (*M, error) {
}

type FindFirstBuilder[M any, S any, O any] struct {
client *Queries
where []PredicateOf[M]
skip *int
execFunc func(ctx context.Context, params QueryParams[M], s *S, o *O) (*M, error)
Expand Down Expand Up @@ -90,7 +88,6 @@ func (b *FindFirstOmitBuilder[M, S, O]) Exec(ctx context.Context) (*M, error) {
}

type FindManyBuilder[M any, S any, O any] struct {
client *Queries
where []PredicateOf[M]
take *int
skip *int
Expand Down
39 changes: 19 additions & 20 deletions generator/templates/model_create.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -48,9 +48,8 @@ func (b *{{ $.Model.Name }}CreateBuilder) Set{{ capitalize $field.Name }}(v {{ $
func (d *{{ .Model.Name }}Delegate) Create(assignments ...FieldAssignment) *{{ .Model.Name }}CreateBuilder {
return &{{ .Model.Name }}CreateBuilder{
CreateBuilder: &CreateBuilder[{{ .Model.Name }}, {{ .Model.Name }}Select, {{ .Model.Name }}Omit]{
client: d.client,
assignments: assignments,
execFunc: d.client.execute{{ .Model.Name }}Create,
execFunc: d.executeCreate,
},
}
}
Expand Down Expand Up @@ -256,7 +255,7 @@ func (s *{{ .Model.Name }}Create) ToRowMap() map[string]any {
return m
}

func (q *Queries) execute{{ .Model.Name }}Create(ctx context.Context, assignments []FieldAssignment, selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) (*{{ .Model.Name }}, error) {
func (d *{{ .Model.Name }}Delegate) executeCreate(ctx context.Context, assignments []FieldAssignment, selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) (*{{ .Model.Name }}, error) {
input, err := assignmentsTo{{ .Model.Name }}Create(assignments)
if err != nil {
return nil, err
Expand All @@ -265,7 +264,7 @@ func (q *Queries) execute{{ .Model.Name }}Create(ctx context.Context, assignment
curr := func(c context.Context, args *{{ .Model.Name }}Create) (*{{ .Model.Name }}, error) {
cols, vals := args.ToColsVals()

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

scanFunc := func(res *{{ .Model.Name }}, cols []string) []any {
return res.ScanFields(cols)
Expand All @@ -290,24 +289,24 @@ func (q *Queries) execute{{ .Model.Name }}Create(ctx context.Context, assignment
var res *{{ .Model.Name }}
var err error
if hasRelations {
err = q.transaction(c, func(txQ *Queries) error {
err = d.client.transaction(c, func(txQ *Queries) error {
var err error
res, err = executeInsert(c, txQ, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, pkCols, scanFunc, conflictTarget, conflictAction)
if err != nil {
return err
}
return txQ.load{{ .Model.Name }}Relations(c, []*{{ .Model.Name }}{res}, selects)
return txQ.{{ .Model.Name }}.loadRelations(c, []*{{ .Model.Name }}{res}, selects)
})
} else {
res, err = executeInsert(c, q, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, pkCols, scanFunc, conflictTarget, conflictAction)
res, err = executeInsert(c, d.client, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, pkCols, scanFunc, conflictTarget, conflictAction)
}
if err != nil {
return nil, err
}
return res, nil
}

for _, ext := range slices.Backward(q.{{ .Model.Name }}.extensions) {
for _, ext := range slices.Backward(d.extensions) {
if ext.Create != nil {
next, hook := curr, ext.Create
curr = func(c context.Context, input *{{ .Model.Name }}Create) (*{{ .Model.Name }}, error) {
Expand Down Expand Up @@ -356,9 +355,8 @@ func (d *{{ .Model.Name }}Delegate) CreateMany(builders ...*{{ .Model.Name }}Cre
}
return &{{ .Model.Name }}CreateManyBuilder{
CreateManyBuilder: &CreateManyBuilder[{{ .Model.Name }}]{
client: d.client,
records: records,
execFunc: d.client.execute{{ .Model.Name }}CreateMany,
execFunc: d.executeCreateMany,
},
}
}
Expand All @@ -370,14 +368,13 @@ func (d *{{ .Model.Name }}Delegate) CreateManyAndReturn(builders ...*{{ .Model.N
}
return &{{ .Model.Name }}CreateManyAndReturnBuilder{
CreateManyAndReturnBuilder: &CreateManyAndReturnBuilder[{{ .Model.Name }}, {{ .Model.Name }}Select, {{ .Model.Name }}Omit]{
client: d.client,
records: records,
execFunc: d.client.execute{{ .Model.Name }}CreateManyAndReturn,
execFunc: d.executeCreateManyAndReturn,
},
}
}

func (q *Queries) execute{{ .Model.Name }}CreateMany(ctx context.Context, records []RecordInput, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) (int64, error) {
func (d *{{ .Model.Name }}Delegate) executeCreateMany(ctx context.Context, records []RecordInput, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) (int64, error) {
inputs := make([]*{{ .Model.Name }}Create, len(records))
for i, rec := range records {
input, err := assignmentsTo{{ .Model.Name }}Create(rec.Assignments)
Expand Down Expand Up @@ -407,10 +404,10 @@ func (q *Queries) execute{{ .Model.Name }}CreateMany(ctx context.Context, record
{{- end }}
}

return executeCreateMany(c, q, rowMaps, "{{ .Model.EffectiveTableName }}", {{ .Model.Name }}ColOrder, pkCols, conflictTarget, conflictAction)
return executeCreateMany(c, d.client, rowMaps, "{{ .Model.EffectiveTableName }}", {{ .Model.Name }}ColOrder, pkCols, conflictTarget, conflictAction)
}

for _, ext := range slices.Backward(q.{{ .Model.Name }}.extensions) {
for _, ext := range slices.Backward(d.extensions) {
if ext.CreateMany != nil {
next, hook := curr, ext.CreateMany
curr = func(c context.Context, inputs []*{{ .Model.Name }}Create) (int64, error) {
Expand All @@ -422,7 +419,7 @@ func (q *Queries) execute{{ .Model.Name }}CreateMany(ctx context.Context, record
return curr(ctx, inputs)
}

func (q *Queries) execute{{ .Model.Name }}CreateManyAndReturn(ctx context.Context, records []RecordInput, selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) ([]*{{ .Model.Name }}, error) {
func (d *{{ .Model.Name }}Delegate) executeCreateManyAndReturn(ctx context.Context, records []RecordInput, selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit, conflictTarget UniqueConstraintTarget, conflictAction *ConflictAction) ([]*{{ .Model.Name }}, error) {
inputs := make([]*{{ .Model.Name }}Create, len(records))
for i, rec := range records {
input, err := assignmentsTo{{ .Model.Name }}Create(rec.Assignments)
Expand Down Expand Up @@ -452,9 +449,11 @@ func (q *Queries) execute{{ .Model.Name }}CreateManyAndReturn(ctx context.Contex
{{- end }}
}

return executeCreateManyAndReturn(c, q, rowMaps, "{{ .Model.EffectiveTableName }}", {{ .Model.Name }}ColOrder, selects, omits,
q.select{{ .Model.Name }}Cols,
q.load{{ .Model.Name }}Relations,
return executeCreateManyAndReturn(c, d.client, rowMaps, "{{ .Model.EffectiveTableName }}", {{ .Model.Name }}ColOrder, selects, omits,
d.selectCols,
func(ctx context.Context, txQ *Queries, results []*{{ .Model.Name }}, sel *{{ .Model.Name }}Select) error {
return txQ.{{ .Model.Name }}.loadRelations(ctx, results, sel)
},
(*{{ .Model.Name }}).ScanFields,
(*{{ .Model.Name }}Select).hasAnyRelation,
pkCols,
Expand All @@ -463,7 +462,7 @@ func (q *Queries) execute{{ .Model.Name }}CreateManyAndReturn(ctx context.Contex
)
}

for _, ext := range slices.Backward(q.{{ .Model.Name }}.extensions) {
for _, ext := range slices.Backward(d.extensions) {
if ext.CreateManyAndReturn != nil {
next, hook := curr, ext.CreateManyAndReturn
curr = func(c context.Context, inputs []*{{ .Model.Name }}Create) ([]*{{ .Model.Name }}, error) {
Expand Down
45 changes: 21 additions & 24 deletions generator/templates/model_query.gotpl
Original file line number Diff line number Diff line change
@@ -1,29 +1,26 @@
func (d *{{ .Model.Name }}Delegate) FindUnique(where UniquePredicate[{{ .Model.Name }}], additional ...PredicateOf[{{ .Model.Name }}]) *FindUniqueBuilder[{{ .Model.Name }}, {{ .Model.Name }}Select, {{ .Model.Name }}Omit] {
return &FindUniqueBuilder[{{ .Model.Name }}, {{ .Model.Name }}Select, {{ .Model.Name }}Omit]{
client: d.client,
where: where,
additional: additional,
execFunc: d.client.execute{{ .Model.Name }}FindUnique,
execFunc: d.executeFindUnique,
}
}

func (d *{{ .Model.Name }}Delegate) FindFirst(preds ...PredicateOf[{{ .Model.Name }}]) *FindFirstBuilder[{{ .Model.Name }}, {{ .Model.Name }}Select, {{ .Model.Name }}Omit] {
return &FindFirstBuilder[{{ .Model.Name }}, {{ .Model.Name }}Select, {{ .Model.Name }}Omit]{
client: d.client,
where: preds,
execFunc: d.client.execute{{ .Model.Name }}FindFirst,
execFunc: d.executeFindFirst,
}
}

func (d *{{ .Model.Name }}Delegate) FindMany(preds ...PredicateOf[{{ .Model.Name }}]) *FindManyBuilder[{{ .Model.Name }}, {{ .Model.Name }}Select, {{ .Model.Name }}Omit] {
return &FindManyBuilder[{{ .Model.Name }}, {{ .Model.Name }}Select, {{ .Model.Name }}Omit]{
client: d.client,
where: preds,
execFunc: d.client.execute{{ .Model.Name }}FindMany,
execFunc: d.executeFindMany,
}
}

func (q *Queries) execute{{ .Model.Name }}FindUnique(ctx context.Context, where UniquePredicate[{{ .Model.Name }}], additional []PredicateOf[{{ .Model.Name }}], selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit) (*{{ .Model.Name }}, error) {
func (d *{{ .Model.Name }}Delegate) executeFindUnique(ctx context.Context, where UniquePredicate[{{ .Model.Name }}], additional []PredicateOf[{{ .Model.Name }}], selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit) (*{{ .Model.Name }}, error) {
curr := func(c context.Context, w UniquePredicate[{{ .Model.Name }}], add []PredicateOf[{{ .Model.Name }}], sel *{{ .Model.Name }}Select, o *{{ .Model.Name }}Omit) (*{{ .Model.Name }}, error) {
if err := w.Validate(); err != nil {
return nil, err
Expand All @@ -36,22 +33,22 @@ func (q *Queries) execute{{ .Model.Name }}FindUnique(ctx context.Context, where
}
}
allPreds := append([]PredicateOf[{{ .Model.Name }}]{w}, add...)
whereClause, vals := CompilePredicates(q.dialect, allPreds)
whereClause, vals := CompilePredicates(d.client.dialect, allPreds)
if whereClause != "" {
whereClause = " WHERE " + whereClause
}
returningCols := q.select{{ .Model.Name }}Cols(sel, o)
return executeSingleWithRelations(c, q, "{{ .Model.EffectiveTableName }}", whereClause, vals, returningCols,
returningCols := d.selectCols(sel, o)
return executeSingleWithRelations(c, d.client, "{{ .Model.EffectiveTableName }}", whereClause, vals, returningCols,
func(res *{{ .Model.Name }}, cols []string) []any { return res.ScanFields(cols) },
sel.hasAnyRelation(),
func(ctx context.Context, txQ *Queries, results []*{{ .Model.Name }}) error {
return txQ.load{{ .Model.Name }}Relations(ctx, results, sel)
return txQ.{{ .Model.Name }}.loadRelations(ctx, results, sel)
},
nil,
)
}

for _, ext := range slices.Backward(q.{{ .Model.Name }}.extensions) {
for _, ext := range slices.Backward(d.extensions) {
if ext.FindUnique != nil {
next, hook := curr, ext.FindUnique
curr = func(c context.Context, w UniquePredicate[{{ .Model.Name }}], add []PredicateOf[{{ .Model.Name }}], sel *{{ .Model.Name }}Select, o *{{ .Model.Name }}Omit) (*{{ .Model.Name }}, error) {
Expand All @@ -63,7 +60,7 @@ func (q *Queries) execute{{ .Model.Name }}FindUnique(ctx context.Context, where
return curr(ctx, where, additional, selects, omits)
}

func (q *Queries) execute{{ .Model.Name }}FindFirst(
func (d *{{ .Model.Name }}Delegate) executeFindFirst(
ctx context.Context,
params QueryParams[{{ .Model.Name }}],
selects *{{ .Model.Name }}Select,
Expand All @@ -77,22 +74,22 @@ func (q *Queries) execute{{ .Model.Name }}FindFirst(
}
}
}
whereClause, vals := CompilePredicates(q.dialect, p.Where)
whereClause, vals := CompilePredicates(d.client.dialect, p.Where)
if whereClause != "" {
whereClause = " WHERE " + whereClause
}
returningCols := q.select{{ .Model.Name }}Cols(sel, o)
return executeSingleWithRelations(c, q, "{{ .Model.EffectiveTableName }}", whereClause, vals, returningCols,
returningCols := d.selectCols(sel, o)
return executeSingleWithRelations(c, d.client, "{{ .Model.EffectiveTableName }}", whereClause, vals, returningCols,
func(res *{{ .Model.Name }}, cols []string) []any { return res.ScanFields(cols) },
sel.hasAnyRelation(),
func(ctx context.Context, txQ *Queries, results []*{{ .Model.Name }}) error {
return txQ.load{{ .Model.Name }}Relations(ctx, results, sel)
return txQ.{{ .Model.Name }}.loadRelations(ctx, results, sel)
},
p.Skip,
)
}

for _, ext := range slices.Backward(q.{{ .Model.Name }}.extensions) {
for _, ext := range slices.Backward(d.extensions) {
if ext.FindFirst != nil {
next, hook := curr, ext.FindFirst
curr = func(c context.Context, p QueryParams[{{ .Model.Name }}], sel *{{ .Model.Name }}Select, o *{{ .Model.Name }}Omit) (*{{ .Model.Name }}, error) {
Expand All @@ -104,7 +101,7 @@ func (q *Queries) execute{{ .Model.Name }}FindFirst(
return curr(ctx, params, selects, omits)
}

func (q *Queries) execute{{ .Model.Name }}FindMany(
func (d *{{ .Model.Name }}Delegate) executeFindMany(
ctx context.Context,
params QueryParams[{{ .Model.Name }}],
selects *{{ .Model.Name }}Select,
Expand All @@ -118,23 +115,23 @@ func (q *Queries) execute{{ .Model.Name }}FindMany(
}
}
}
whereClause, vals := CompilePredicates(q.dialect, p.Where)
whereClause, vals := CompilePredicates(d.client.dialect, p.Where)
if whereClause != "" {
whereClause = " WHERE " + whereClause
}
returningCols := q.select{{ .Model.Name }}Cols(sel, o)
return executeManyWithRelations(c, q, "{{ .Model.EffectiveTableName }}", whereClause, vals, returningCols,
returningCols := d.selectCols(sel, o)
return executeManyWithRelations(c, d.client, "{{ .Model.EffectiveTableName }}", whereClause, vals, returningCols,
func(res *{{ .Model.Name }}, cols []string) []any { return res.ScanFields(cols) },
sel.hasAnyRelation(),
func(ctx context.Context, txQ *Queries, results []*{{ .Model.Name }}) error {
return txQ.load{{ .Model.Name }}Relations(ctx, results, sel)
return txQ.{{ .Model.Name }}.loadRelations(ctx, results, sel)
},
p.Take,
p.Skip,
)
}

for _, ext := range slices.Backward(q.{{ .Model.Name }}.extensions) {
for _, ext := range slices.Backward(d.extensions) {
if ext.FindMany != nil {
next, hook := curr, ext.FindMany
curr = func(c context.Context, p QueryParams[{{ .Model.Name }}], sel *{{ .Model.Name }}Select, o *{{ .Model.Name }}Omit) ([]*{{ .Model.Name }}, error) {
Expand Down
8 changes: 4 additions & 4 deletions generator/templates/model_relations.gotpl
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
func (q *Queries) load{{ .Model.Name }}Relations(ctx context.Context, records []*{{ .Model.Name }}, selects *{{ .Model.Name }}Select) error {
func (d *{{ .Model.Name }}Delegate) loadRelations(ctx context.Context, records []*{{ .Model.Name }}, selects *{{ .Model.Name }}Select) error {
if selects == nil || len(records) == 0 {
return nil
}
Expand All @@ -12,15 +12,15 @@ func (q *Queries) load{{ .Model.Name }}Relations(ctx context.Context, records []
{{- $forceCol = (index $relation.Inverse.FKFields 0).EffectiveColName }}
{{- end }}
relationSelects, relationOmits, relationParams := selects.{{ capitalize $relation.Name }}.GetRelationParams()
returningCols := q.select{{ $relation.TargetModelName }}Cols(relationSelects, relationOmits, "{{ $forceCol }}")
returningCols := d.client.{{ $relation.TargetModelName }}.selectCols(relationSelects, relationOmits, "{{ $forceCol }}")

{{- if gt (len $relation.FKFields) 0 }}
// Current model holds the FK: {{ $.Model.Name }}.{{ (index $relation.FKFields 0).Name }}
{{- else }}
// Inverse holds the FK: {{ $relation.TargetModelName }}.{{ (index $relation.Inverse.FKFields 0).Name }}
{{- end }}
allChildren, err := loadRelation(
ctx, q, records,
ctx, d.client, records,
{{- if gt (len $relation.FKFields) 0 }}
{{- if (index $relation.FKFields 0).Optional }}
optionalKey(func(p *{{ $.Model.Name }}) {{ (index $relation.FKFields 0).GoType }} { return p.{{ capitalize (index $relation.FKFields 0).Name }} }),
Expand Down Expand Up @@ -57,7 +57,7 @@ func (q *Queries) load{{ .Model.Name }}Relations(ctx context.Context, records []
if err != nil {
return fmt.Errorf("loading {{ $relation.Name }}: %w", err)
}
if err := q.load{{ $relation.TargetModelName }}Relations(ctx, allChildren, relationSelects); err != nil {
if err := d.client.{{ $relation.TargetModelName }}.loadRelations(ctx, allChildren, relationSelects); err != nil {
return err
}
}
Expand Down
2 changes: 1 addition & 1 deletion generator/templates/model_structs.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -140,7 +140,7 @@ var {{ lowercase .Model.Name }}DefaultCols = []string{
{{- end }}
}

func (q *Queries) select{{ .Model.Name }}Cols(selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit, forceCols ...string) []string {
func (d *{{ .Model.Name }}Delegate) selectCols(selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit, forceCols ...string) []string {
if selects == nil && omits == nil && len(forceCols) == 0 {
return {{ lowercase .Model.Name }}DefaultCols
}
Expand Down
4 changes: 2 additions & 2 deletions makefile
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
.PHONY: build build-prod run test install db-up db-down db-clean bi fmt fmt-check tidy tidy-check vulncheck vet integration-gen integration-test bench race lint test-sqlite test-pg test-dbs ci-local
.PHONY: build build-prod run test install db-up db-down db-clean bi fmt fmt-check tidy tidy-check vulncheck vet integration-gen integration-test bench race lint test-sqlite test-pg test-dbs ci-local bench

bi: build install

Expand All @@ -8,7 +8,7 @@ race:
go test -race ./... && cd integration && go test -race ./...

bench:
cd integration && go test -bench=. -benchmem -benchtime=2s -count=3
cd benchmark && make bench && cd ..

build:
go build -o bin/valk
Expand Down
Loading