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
1 change: 1 addition & 0 deletions generator/generator.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default
"client.gotpl",
"tx.gotpl",
"builders_create.gotpl",
"relations_runtime.gotpl",
}
for _, file := range files {
if err := tmpl.ExecuteTemplate(&buf, file, data); err != nil {
Expand Down
56 changes: 20 additions & 36 deletions generator/templates/model_relations.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -20,53 +20,37 @@ func (q *Queries) load{{ .Model.Name }}Relations(ctx context.Context, records []
{{- end }}
allChildren, err := loadRelation(
ctx, q, records,
func(p *{{ $.Model.Name }}) (string, bool) {
{{- if gt (len $relation.FKFields) 0 }}
{{- if (index $relation.FKFields 0).Optional }}
if p.{{ capitalize (index $relation.FKFields 0).Name }} == nil {
return "", false
}
return fmt.Sprint(*p.{{ capitalize (index $relation.FKFields 0).Name }}), true
{{- else }}
return fmt.Sprint(p.{{ capitalize (index $relation.FKFields 0).Name }}), true
{{- end }}
{{- 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 }} }),
{{- else }}
return fmt.Sprint(p.{{ capitalize (index $relation.Inverse.RefFields 0).Name }}), true
directKey(func(p *{{ $.Model.Name }}) {{ (index $relation.FKFields 0).GoType }} { return p.{{ capitalize (index $relation.FKFields 0).Name }} }),
{{- end }}
},
{{- else }}
directKey(func(p *{{ $.Model.Name }}) {{ (index $relation.Inverse.RefFields 0).GoType }} { return p.{{ capitalize (index $relation.Inverse.RefFields 0).Name }} }),
{{- end }}
"{{ $relation.TargetModel.EffectiveTableName }}",
{{- if gt (len $relation.FKFields) 0 }}
"{{ (index $relation.RefFields 0).EffectiveColName }}",
{{- else }}
"{{ (index $relation.Inverse.FKFields 0).EffectiveColName }}",
{{- end }}
returningCols,
func(rows *sql.Rows, child *{{ $relation.TargetModelName }}) error {
return rows.Scan(child.ScanFields(returningCols)...)
},
func(child *{{ $relation.TargetModelName }}) (string, bool) {
{{- if gt (len $relation.FKFields) 0 }}
return fmt.Sprint(child.{{ capitalize (index $relation.RefFields 0).Name }}), true
{{- else }}
{{- if (index $relation.Inverse.FKFields 0).Optional }}
if child.{{ capitalize (index $relation.Inverse.FKFields 0).Name }} == nil {
return "", false
}
return fmt.Sprint(*child.{{ capitalize (index $relation.Inverse.FKFields 0).Name }}), true
{{- else }}
return fmt.Sprint(child.{{ capitalize (index $relation.Inverse.FKFields 0).Name }}), true
{{- end }}
{{- end }}
},
func(p *{{ $.Model.Name }}, children []*{{ $relation.TargetModelName }}) {
{{- if and (eq (len $relation.FKFields) 0) $relation.IsArray }}
p.{{ capitalize $relation.Name }} = append(p.{{ capitalize $relation.Name }}, children...)
scanInto(returningCols, (*{{ $relation.TargetModelName }}).ScanFields),
{{- if gt (len $relation.FKFields) 0 }}
directKey(func(c *{{ $relation.TargetModelName }}) {{ (index $relation.RefFields 0).GoType }} { return c.{{ capitalize (index $relation.RefFields 0).Name }} }),
{{- else }}
{{- if (index $relation.Inverse.FKFields 0).Optional }}
optionalKey(func(c *{{ $relation.TargetModelName }}) {{ (index $relation.Inverse.FKFields 0).GoType }} { return c.{{ capitalize (index $relation.Inverse.FKFields 0).Name }} }),
{{- else }}
if len(children) > 0 {
p.{{ capitalize $relation.Name }} = children[0]
}
directKey(func(c *{{ $relation.TargetModelName }}) {{ (index $relation.Inverse.FKFields 0).GoType }} { return c.{{ capitalize (index $relation.Inverse.FKFields 0).Name }} }),
{{- end }}
},
{{- end }}
{{- if and (eq (len $relation.FKFields) 0) $relation.IsArray }}
appendMany(func(p *{{ $.Model.Name }}) *[]*{{ $relation.TargetModelName }} { return &p.{{ capitalize $relation.Name }} }),
{{- else }}
setOne(func(p *{{ $.Model.Name }}, c *{{ $relation.TargetModelName }}) { p.{{ capitalize $relation.Name }} = c }),
{{- end }}
)
if err != nil {
return fmt.Errorf("loading {{ $relation.Name }}: %w", err)
Expand Down
52 changes: 7 additions & 45 deletions generator/templates/model_structs.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -63,56 +63,18 @@ func (q *Queries) select{{ .Model.Name }}Cols(selects *{{ .Model.Name }}Select,
return {{ lowercase .Model.Name }}DefaultCols
}

var cols []string
var anySelected bool
if selects != nil {
{{- range $field := .Model.ScalarFields }}
if selects.{{ capitalize $field.Name }} {
anySelected = true
}
{{- end }}
{{- range $relation := .Model.RelationFields }}
if selects.{{ capitalize $relation.Name }} != nil {
anySelected = true
}
{{- end }}
}
anySelected := selects != nil && ({{ range $i, $f := .Model.ScalarFields }}{{ if $i }} || {{ end }}selects.{{ capitalize $f.Name }}{{ end }}{{ range .Model.RelationFields }} || selects.{{ capitalize .Name }} != nil{{ end }})

{{- range $field := .Model.ScalarFields }}
{
include := true
if selects != nil {
include = false
if !anySelected {
include = true
} else if selects.{{ capitalize $field.Name }} {
include = true
}
{{- $relName := fkForRelation $.Model $field }}
{{- if ne $relName "" }}
// Force-include FK when its relation is selected
if selects.{{ capitalize $relName }} != nil {
include = true
}
{{- end }}
} else if omits != nil {
if omits.{{ capitalize $field.Name }} {
include = false
}
}
if include {
cols = append(cols, "{{ $field.EffectiveColName }}")
}
}
{{- end }}

if len(cols) == 0 {
specs := []colSpec{
{{- range $field := .Model.ScalarFields }}
cols = append(cols, "{{ $field.EffectiveColName }}")
{"{{ $field.EffectiveColName }}", selects != nil && selects.{{ capitalize $field.Name }}, omits != nil && omits.{{ capitalize $field.Name }},
{{- $relName := fkForRelation $.Model $field }}
{{- if ne $relName "" }} selects != nil && selects.{{ capitalize $relName }} != nil{{ else }} false{{ end }}},
{{- end }}
}

// Force-include any requested columns
cols := computeCols(specs, selects != nil, anySelected)

for _, f := range forceCols {
if !slices.Contains(cols, f) {
cols = append(cols, f)
Expand Down
65 changes: 65 additions & 0 deletions generator/templates/relations_runtime.gotpl
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
func directKey[T any, K any](get func(*T) K) func(*T) (string, bool) {
return func(t *T) (string, bool) {
return fmt.Sprint(get(t)), true
}
}

func optionalKey[T any, K any](get func(*T) *K) func(*T) (string, bool) {
return func(t *T) (string, bool) {
if p := get(t); p != nil {
return fmt.Sprint(*p), true
}
return "", false
}
}

func setOne[P any, C any](set func(*P, *C)) func(*P, []*C) {
return func(p *P, children []*C) {
if len(children) > 0 {
set(p, children[0])
}
}
}

func appendMany[P any, C any](get func(*P) *[]*C) func(*P, []*C) {
return func(p *P, children []*C) {
if s := get(p); s != nil {
*s = append(*s, children...)
}
}
}

func scanInto[C any](cols []string, scan func(*C, []string) []any) func(*sql.Rows, *C) error {
return func(rows *sql.Rows, c *C) error {
return rows.Scan(scan(c, cols)...)
}
}

type colSpec struct {
col string
selected bool
omitted bool
forceIn bool
}

func computeCols(specs []colSpec, hasSelects, anySelected bool) []string {
var cols []string
for _, s := range specs {
include := true
if hasSelects {
include = !anySelected || s.selected || s.forceIn
} else if s.omitted {
include = false
}
if include {
cols = append(cols, s.col)
}
}
if len(cols) == 0 {
for _, s := range specs {
cols = append(cols, s.col)
}
}
return cols
}

Loading
Loading